"""Strict output schemas for every LLM task (spec §24: structured JSON validated against schemas). The gateway validates model output against these models and retries once with the validation error; anything that still fails is recorded as a failed job — never stored as an event. Schema versions are written to `events.schema_version` / `llm_jobs.result`. """ from __future__ import annotations from typing import Any from pydantic import BaseModel, ConfigDict, Field, field_validator from companyatlas.taxonomy import EVENT_SUBTYPES, FORBIDDEN_WORDING SCHEMA_VERSIONS = {"ChangeClassification": "classify-v1", "EventSummary": "summary-v1", "LegalDiffSummary": "legal-v1", "IndustryTags": "industry-v1", "AskRoute": "ask-v1", "CompanyProfileText": "profile-v1"} def _clean_text(value: str | None, limit: int) -> str | None: if value is None: return None text = " ".join(str(value).split()) for bad in FORBIDDEN_WORDING: if bad in text.lower(): raise ValueError(f"forbidden wording: {bad!r} — use 'no longer listed' / 'observed' phrasing") return text[:limit] or None class _Strict(BaseModel): model_config = ConfigDict(extra="ignore", str_strip_whitespace=True) class ChangeClassification(_Strict): event_subtype: str = Field(description="One of the Company Atlas event subtypes, or OTHER") importance: float = Field(ge=0, le=1) confidence: float = Field(ge=0, le=1) title: str = Field(min_length=3, max_length=120) summary: str | None = Field(default=None, max_length=400) old_value: str | None = Field(default=None, max_length=200) new_value: str | None = Field(default=None, max_length=200) entities: dict[str, Any] = Field(default_factory=dict) tags: list[str] = Field(default_factory=list, max_length=12) language: str | None = Field(default=None, max_length=8) is_material: bool = True @field_validator("event_subtype", mode="before") @classmethod def _subtype(cls, v: Any) -> str: s = str(v or "OTHER").strip().upper().replace(" ", "_").replace("-", "_") return s if s in EVENT_SUBTYPES else "OTHER" @field_validator("title", mode="before") @classmethod def _title(cls, v: Any) -> str: return _clean_text(str(v or ""), 120) or "" @field_validator("summary", "old_value", "new_value", mode="before") @classmethod def _texts(cls, v: Any) -> str | None: return _clean_text(None if v is None else str(v), 400) @field_validator("tags", mode="before") @classmethod def _tags(cls, v: Any) -> list[str]: if not v: return [] out = [] for t in list(v)[:12]: s = str(t).strip().lower()[:40] if s and s not in out: out.append(s) return out class EventSummary(_Strict): summary: str = Field(min_length=3, max_length=400) key_points: list[str] = Field(default_factory=list, max_length=5) confidence: float = Field(default=0.7, ge=0, le=1) language: str | None = Field(default=None, max_length=8) @field_validator("summary", mode="before") @classmethod def _summary(cls, v: Any) -> str: return _clean_text(str(v or ""), 400) or "" @field_validator("key_points", mode="before") @classmethod def _points(cls, v: Any) -> list[str]: return [p for p in (_clean_text(str(x), 160) for x in (v or [])[:5]) if p] class LegalSection(_Strict): section: str = Field(max_length=120) change: str = Field(max_length=240) class LegalDiffSummary(_Strict): sections_changed: list[LegalSection] = Field(default_factory=list, max_length=12) materiality: str = Field(default="unclear") # editorial | minor | material | unclear summary: str = Field(min_length=3, max_length=400) user_impact: str | None = Field(default=None, max_length=240) confidence: float = Field(default=0.6, ge=0, le=1) @field_validator("materiality", mode="before") @classmethod def _mat(cls, v: Any) -> str: s = str(v or "unclear").strip().lower() return s if s in ("editorial", "minor", "material", "unclear") else "unclear" @field_validator("summary", "user_impact", mode="before") @classmethod def _texts(cls, v: Any) -> str | None: return _clean_text(None if v is None else str(v), 400) class IndustryTags(_Strict): industries: list[str] = Field(default_factory=list, max_length=5) # industry slugs from the registry primary: str | None = None keywords: list[str] = Field(default_factory=list, max_length=10) confidence: float = Field(default=0.6, ge=0, le=1) @field_validator("industries", "keywords", mode="before") @classmethod def _slugs(cls, v: Any) -> list[str]: out: list[str] = [] for x in (v or [])[:10]: s = str(x).strip().lower().replace(" ", "-")[:60] if s and s not in out: out.append(s) return out class CompanyProfileText(_Strict): """Grounded 2–3 sentence company description (services/enrichment.py, job kind `company_profile`). Numbers are additionally checked against the source text by the caller (`enrichment.numbers_grounded`) — any figure absent from the source rejects the output.""" description: str = Field(min_length=40, max_length=700) language: str | None = Field(default=None, max_length=8) confidence: float = Field(default=0.6, ge=0, le=1) @field_validator("description", mode="before") @classmethod def _description(cls, v: Any) -> str: return _clean_text(str(v or ""), 700) or "" class AskRoute(_Strict): """LLM refinement of the deterministic question parser (services/llm/ask.py). It may only tighten filters, never invent results.""" intent: str = Field(default="search") # search | hiring | pricing | ai | launch | leadership | expansion | compare | trend countries: list[str] = Field(default_factory=list, max_length=8) industries: list[str] = Field(default_factory=list, max_length=8) event_types: list[str] = Field(default_factory=list, max_length=8) event_subtypes: list[str] = Field(default_factory=list, max_length=8) companies: list[str] = Field(default_factory=list, max_length=8) window_days: int | None = Field(default=None, ge=1, le=3650) keywords: list[str] = Field(default_factory=list, max_length=8) answer_style: str = Field(default="list") # list | count | compare | timeline confidence: float = Field(default=0.6, ge=0, le=1) @field_validator("event_subtypes", mode="before") @classmethod def _subtypes(cls, v: Any) -> list[str]: return [s for s in (str(x).strip().upper() for x in (v or [])[:8]) if s in EVENT_SUBTYPES] @field_validator("countries", mode="before") @classmethod def _countries(cls, v: Any) -> list[str]: return [s for s in (str(x).strip().upper()[:2] for x in (v or [])[:8]) if len(s) == 2 and s.isalpha()] __all__ = ["SCHEMA_VERSIONS", "AskRoute", "ChangeClassification", "CompanyProfileText", "EventSummary", "IndustryTags", "LegalDiffSummary", "LegalSection"]