"""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, SuperSimulationSettings from ..catalog.contracts import CatalogFilters, Scope from ..preparations.contracts import PreparationReference from ..research.workspace_contracts import TemplateSpec 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 PyramidQuery(Contract): current_date: date = Field(description="用于确定季度的日期,格式 YYYY-MM-DD;自动查询该季度完整起止范围") region: str = Field(min_length=3, max_length=10, pattern=r"^[A-Z]+$") delay: int = Field(ge=0, le=1, strict=True) 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 CompleteSuperSettings(SuperSimulationSettings): @model_validator(mode="before") @classmethod def complete(cls, value): if isinstance(value, dict) and set(cls.model_fields) - value.keys(): raise ValueError("必须提供每项完整 SUPER 设置;先读取 get_research_capabilities") return value model_config = {"json_schema_extra": {"required": list(SuperSimulationSettings.model_fields)}} class DirectCandidate(Candidate): settings: CompleteSuperSettings | CompleteSettings class Provenance(Contract): research_id: RunId | None = None superalpha_plan_id: RunId | None = None superalpha_plan_version: int | None = Field(default=None, ge=1) selection_snapshot_ids: list[RunId] = Field(default_factory=list, max_length=100) 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 PreparationSearch(Contract): q: str = Field(default="", max_length=300) scope_key: str | None = None limit: int = Field(default=25, ge=1, le=100) offset: int = Field(default=0, ge=0) class PreparationRead(Contract): id: str = Field(min_length=1, max_length=36) version: int = Field(ge=1) q: str = Field(default="", max_length=300) limit: int = Field(default=25, ge=1, le=100) offset: int = Field(default=0, ge=0) class Submit(Contract): name: str = Field(min_length=1, max_length=200) preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20) 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 CreateTemplate(Contract): template: TemplateSpec hypothesis: str = Field(min_length=1, max_length=10000) source_item_ids: list[RunId] = Field(default_factory=list, max_length=20) reference: str | None = Field(default=None, max_length=200) idempotency_key: Identifier @model_validator(mode="after") def research_template(self): self.template.name = self.template.name.strip() self.hypothesis = self.hypothesis.strip() if not self.template.name or not self.hypothesis or not self.template.expression.strip(): raise ValueError("模板名称、表达式和研究假设不能为空") if self.template.category != "template": raise ValueError("此工具仅创建完整模板,不保存表达式片段") if len(set(self.source_item_ids)) != len(self.source_item_ids): raise ValueError("来源候选 ID 不能重复") return self class CreateTemplateVersion(CreateTemplate): template_id: RunId expected_version: int = Field(ge=1) class TemplateRead(Contract): template_id: RunId version: int | None = Field(default=None, ge=1) class TemplateSearch(Page): q: str = Field(default="", max_length=200) class TemplateExpansion(Contract): template_id: RunId version: int = Field(ge=1) preparation_refs: list[PreparationReference] = Field(min_length=1, max_length=20) settings: CompleteSettings mode: Literal["all", "random"] = "all" limit: int = Field(default=100, ge=1, le=10000) seed: int = 0 idempotency_key: Identifier class TemplateCandidates(Page): experiment_id: RunId class SubmitTemplateBacktest(Contract): experiment_id: RunId candidate_ids: list[Identifier] = Field(min_length=1, max_length=10000) 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"] alpha_type: Literal["REGULAR", "SUPER"] = "REGULAR" class Operators(Page): kind: Literal["operators"] q: str = Field(default="", max_length=300) category: str | None = None stage: Literal["REGULAR", "SELECTION", "COMBO"] | None = None class SuperMetadata(Contract): kind: Literal["superalpha"] class Availability(Contract): kind: Literal["field_availability"] field_id: Identifier scope: Scope class Metadata(Contract): query: Annotated[Scopes | SettingOptions | Operators | Availability | SuperMetadata, 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 SubmissionCheck(Contract): alpha_id: AlphaId snapshot: str = Field(pattern=r"^[a-f0-9]{64}$") descriptions: dict[str, str] = Field(min_length=1, max_length=2) class History(Page): research_id: str | None = Field(default=None, max_length=36) alpha_type: Literal["REGULAR", "SUPER"] | None = None 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", "components"] date_from: date | None = None date_to: date | None = None @model_validator(mode="after") def dates(self): if self.kind != "pnl" 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