"""Bounded direct research inputs; unknown properties are rejected at the interface.""" from datetime import date, datetime from typing import Annotated, Literal from pydantic import Field, model_validator from ..backtests.contracts import Candidate, SimulationSettings from ..catalog.contracts import CatalogFilters, Scope from ..schemas import Contract Identifier = Annotated[str, Field(min_length=1, max_length=100)] RunId = Annotated[str, Field(min_length=1, max_length=36)] AlphaId = Annotated[str, Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")] class Empty(Contract): pass class Authentication(Contract): action: Literal["connect", "verify"] = "connect" class ConnectionReference(Contract): job_id: RunId | None = None class Page(Contract): limit: int = Field(default=25, ge=1, le=100) offset: int = Field(default=0, ge=0) class CompleteSettings(SimulationSettings): @model_validator(mode="before") @classmethod def complete(cls, value): if isinstance(value, dict) and set(cls.model_fields) - value.keys(): raise ValueError("必须提供每项完整设置;先读取 get_research_capabilities") return value model_config = {"json_schema_extra": {"required": list(SimulationSettings.model_fields)}} class DirectCandidate(Candidate): settings: CompleteSettings class Provenance(Contract): reference: str | None = Field(default=None, max_length=200) batch_id: str | None = Field(default=None, max_length=200) hypothesis: str | None = Field(default=None, max_length=2000) parent_run_id: RunId | None = None class Submit(Contract): name: str = Field(min_length=1, max_length=200) candidates: list[DirectCandidate] = Field(min_length=1, max_length=100) idempotency_key: Identifier duplicate_policy: Literal["reject", "rerun"] = "reject" source: Provenance = Field(default_factory=Provenance) @model_validator(mode="after") def unique_ids(self): if len({c.client_item_id for c in self.candidates}) != len(self.candidates): raise ValueError("client_item_id 必须唯一") return self class Control(Contract): run_id: RunId action: Literal["pause", "resume", "stop", "recover"] expected_version: int = Field(ge=1) idempotency_key: Identifier class CatalogSearch(Contract): filters: CatalogFilters dataset_id: str | None = Field(default=None, min_length=1, max_length=200) class Scopes(Contract): kind: Literal["scopes"] class SettingOptions(Page): kind: Literal["settings"] class Operators(Page): kind: Literal["operators"] q: str = Field(default="", max_length=300) category: str | None = None class Availability(Contract): kind: Literal["field_availability"] field_id: Identifier scope: Scope class Metadata(Contract): query: Annotated[Scopes | SettingOptions | Operators | Availability, Field(discriminator="kind")] class CatalogRefresh(Contract): kind: Literal["catalog"] scope: Scope dataset_id: str | None = Field(default=None, min_length=1, max_length=200) class OperatorsRefresh(Contract): kind: Literal["operators"] class SettingsRefresh(Contract): kind: Literal["settings"] class PnlRefresh(Contract): kind: Literal["pnl"] alpha_ids: list[Identifier] = Field(min_length=1, max_length=100) class Refresh(Contract): query: Annotated[ CatalogRefresh | OperatorsRefresh | SettingsRefresh | Availability | PnlRefresh, Field(discriminator="kind"), ] class JobReference(Page): job_id: RunId class SelfCorrelationCheck(Contract): alpha_ids: list[AlphaId] = Field(min_length=1, max_length=100) class SelfCorrelationReference(Contract): alpha_id: AlphaId class History(Page): source: str | None = Field(default=None, max_length=100) reference: str | None = Field(default=None, max_length=200) status: str | None = Field(default=None, max_length=30) created_from: datetime | None = None created_to: datetime | None = None scope: Scope | None = None q: str = Field(default="", max_length=300) candidates: list[DirectCandidate] | None = Field(default=None, min_length=1, max_length=100) @model_validator(mode="after") def dates(self): for value in (self.created_from, self.created_to): if value and not value.tzinfo: raise ValueError("时间须包含时区") if self.created_from and self.created_to and self.created_from > self.created_to: raise ValueError("起始时间不能晚于结束时间") return self class RunReference(Contract): run_id: RunId after: int | None = Field(default=None, ge=0) event_limit: int = Field(default=25, ge=1, le=100) class Results(Page): run_id: RunId item_ids: list[RunId] | None = Field(default=None, min_length=1, max_length=100) class Artifact(Page): item_id: RunId kind: Literal["snapshot", "pnl"] date_from: date | None = None date_to: date | None = None @model_validator(mode="after") def dates(self): if self.kind == "snapshot" and (self.date_from or self.date_to): raise ValueError("日期筛选仅用于 PnL") if self.date_from and self.date_to and self.date_from > self.date_to: raise ValueError("起始日期不能晚于结束日期") return self