"""Validated public API contracts. Platform state is intentionally not a closed enum.""" import re from datetime import date, datetime, timezone from typing import Literal from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator ResearchState = Literal["inbox", "candidate", "optimizing", "archived"] Submission = Literal["UNSUBMITTED", "SUBMITTED"] CheckType = Literal["PENDING", "PRE_CHECK", "PASS", "FAIL_1", "FAIL_2"] SortField = Literal[ "id", "name", "sharpe", "fitness", "returns", "turnover", "margin", "drawdown", "sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl", "date_created", "date_submitted", "synced_at", ] class Contract(BaseModel): model_config = ConfigDict(extra="forbid", allow_inf_nan=False) class LoginInput(Contract): username: str = Field(min_length=1, max_length=100) password: str = Field(min_length=1, max_length=1024) class CredentialsInput(Contract): email: str = Field(max_length=254) password: str = Field(min_length=1, max_length=1024) @field_validator("email") @classmethod def valid_email(cls, value): value = value.strip() if not re.fullmatch(r"[^\s@]+@[^\s@]+\.[^\s@]+", value): raise ValueError("邮箱格式不正确") return value class PreferencesInput(Contract): display_name: str = Field(min_length=1, max_length=100) theme: Literal["light", "dark"] timezone: str page_size: Literal[25, 50, 100] @field_validator("timezone") @classmethod def valid_timezone(cls, value): try: ZoneInfo(value) except (ZoneInfoNotFoundError, ValueError): raise ValueError("未知时区") from None return value class AlphaFilters(Contract): submission: Submission | None = None source: str | None = Field(default=None, max_length=100) source_reference: str | None = Field(default=None, max_length=200) research_id: str | None = Field(default=None, max_length=200) backtest_run_id: str | None = Field(default=None, max_length=36) q: str | None = Field(default=None, max_length=300) region: str | None = None universe: str | None = None alpha_type: str | None = None language: str | None = None status: str | None = None stage: str | None = None hidden: bool | None = None check_type: CheckType | None = None neutralization: str | None = None research_state: ResearchState | None = None favorite: bool | None = None tag: str | None = Field(default=None, max_length=60) created_from: datetime | None = None created_to: datetime | None = None sharpe_min: float | None = None sharpe_max: float | None = None fitness_min: float | None = None fitness_max: float | None = None returns_min: float | None = None returns_max: float | None = None turnover_min: float | None = None turnover_max: float | None = None margin_min: float | None = None margin_max: float | None = None drawdown_min: float | None = None drawdown_max: float | None = None sub_universe_sharpe_min: float | None = None sub_universe_sharpe_max: float | None = None robust_universe_sharpe_min: float | None = None robust_universe_sharpe_max: float | None = None two_year_sharpe_min: float | None = None two_year_sharpe_max: float | None = None prod_correlation_min: float | None = None prod_correlation_max: float | None = None pnl_min: float | None = None pnl_max: float | None = None sort: SortField = "date_created" direction: Literal["asc", "desc"] = "desc" limit: int = Field(default=25, ge=1, le=100) offset: int = Field(default=0, ge=0) @field_validator("created_from", "created_to") @classmethod def utc_dates(cls, value): return value.replace(tzinfo=timezone.utc) if value and value.tzinfo is None else value @model_validator(mode="after") def range_order(self): for key in ( "sharpe", "fitness", "returns", "turnover", "margin", "drawdown", "sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl", ): lo, hi = getattr(self, f"{key}_min"), getattr(self, f"{key}_max") if lo is not None and hi is not None and lo > hi: raise ValueError(f"{key} 最小值不能大于最大值") if self.created_from and self.created_to and self.created_from > self.created_to: raise ValueError("开始时间不能晚于结束时间") return self def normalize_tags(values): result = sorted(set(v.strip() for v in values if v.strip())) if len(result) > 30 or any(len(v) > 60 for v in result): raise ValueError("最多 30 个标签,每个标签最多 60 字符") return result class ResearchInput(Contract): note: str = Field(default="", max_length=20000) tags: list[str] = Field(default_factory=list) favorite: bool = False state: ResearchState = "inbox" _tags = field_validator("tags")(normalize_tags) class ResearchUpdate(ResearchInput): version: int = Field(ge=1) def valid_ids(values): result = list(dict.fromkeys(values)) if not result or len(result) > 100 or any(not re.fullmatch(r"[A-Za-z0-9_-]{1,100}", v) for v in result): raise ValueError("请输入 1–100 个合法 Alpha ID(字母、数字、下划线或短横线)") return result class BulkInput(Contract): alpha_ids: list[str] add_tags: list[str] = Field(default_factory=list) remove_tags: list[str] = Field(default_factory=list) state: ResearchState | None = None _ids = field_validator("alpha_ids")(valid_ids) _tags = field_validator("add_tags", "remove_tags")(normalize_tags) @model_validator(mode="after") def disjoint(self): if set(self.add_tags) & set(self.remove_tags): raise ValueError("同一标签不能同时添加和移除") return self class BulkUpdate(BulkInput): versions: dict[str, int] @model_validator(mode="after") def check_versions(self): if set(self.versions) != set(self.alpha_ids) or any(v < 1 for v in self.versions.values()): raise ValueError("每个目标 Alpha 都必须提供当前版本") return self class JobInput(Contract): kind: Literal["full_sync", "daily_sync", "alpha_refresh", "pnl_refresh", "self_correlation"] alpha_ids: list[str] = Field(default_factory=list) submission: Submission | None = None date_from: date | None = None date_to: date | None = None @model_validator(mode="after") def validate_ids(self): if self.kind in ("full_sync", "daily_sync"): if self.alpha_ids: raise ValueError("列表同步不接受 Alpha ID") if self.kind == "full_sync": if self.submission == "UNSUBMITTED": raise ValueError("待提交 Alpha 必须选择日期逐天同步") self.submission = "SUBMITTED" if self.date_from is not None or self.date_to is not None: raise ValueError("全量同步不接受日期范围") else: if not self.submission or not self.date_from or not self.date_to: raise ValueError("按天同步必须选择待提交/已提交和起止日期") if self.date_from > self.date_to: raise ValueError("开始日期不能晚于结束日期") if self.date_to > datetime.now(timezone.utc).date(): raise ValueError("同步日期不能晚于今天(UTC)") else: self.alpha_ids = valid_ids(self.alpha_ids) if self.submission is not None or self.date_from is not None or self.date_to is not None: raise ValueError("按 ID 操作不接受分组或日期范围") return self class ResearchOutput(ResearchInput): updated_at: datetime version: int class AlphaSummary(BaseModel): id: str name: str | None expression_preview: str alpha_type: str | None language: str | None stage: str | None status: str | None hidden: bool region: str | None universe: str | None sharpe: float | None fitness: float | None returns: float | None turnover: float | None margin: float | None drawdown: float | None date_created: datetime | None date_submitted: datetime | None synced_at: datetime check_type: CheckType = "PENDING" failed_checks: list[str] = Field(default_factory=list) neutralization: str | None = None sub_universe_sharpe: float | None = None robust_universe_sharpe: float | None = None two_year_sharpe: float | None = None prod_correlation: float | None = None pnl: float | None = None research: ResearchOutput local_correlation: dict | None = None source_kinds: list[str] = Field(default_factory=list) class AlphaSourceOutput(BaseModel): backtest_run_id: str name: str source: dict item_id: str client_item_id: str observed_at: datetime class AlphaSourcePage(BaseModel): alpha_id: str items: list[AlphaSourceOutput] total: int limit: int offset: int class AlphaDetail(AlphaSummary): expression: str | None selection: str | None combo: str | None settings: dict is_metrics: dict os_metrics: dict checks: list class AlphaPage(BaseModel): items: list[AlphaSummary] total: int limit: int offset: int class JobOutput(BaseModel): model_config = ConfigDict(from_attributes=True) id: str kind: str status: str payload: dict = Field(default_factory=dict) processed: int failed: int total: int | None error: str | None next_retry_at: datetime | None created_at: datetime updated_at: datetime checkpoint: dict = Field(default_factory=dict) class PlatformSessionOutput(BaseModel): authenticated: bool expires_at: datetime | None remaining_seconds: int | None total_seconds: float | None class AccountOutput(BaseModel): email: str | None configured: bool credentials_source: Literal["environment", "database"] wq_user_id: str | None profile: dict connection_status: str connection_error: str | None verification_url: str | None last_synced_at: datetime | None display_name: str theme: Literal["light", "dark"] timezone: str page_size: int session: PlatformSessionOutput class SessionOutput(BaseModel): username: str class OkOutput(BaseModel): ok: bool class ErrorOutput(BaseModel): detail: str class HealthOutput(BaseModel): status: Literal["ok"] class BulkOutput(BaseModel): updated: int class PnlPoint(BaseModel): date: str value: float | None class PnlOutput(BaseModel): cached: bool points: list[PnlPoint] fetched_at: datetime | None class SelfCorrelationOutput(BaseModel): cached: bool result: dict | None class FacetsOutput(BaseModel): region: list[str] universe: list[str] alpha_type: list[str] language: list[str] status: list[str] stage: list[str] tags: list[str] total: int favorites: int last_sync: datetime | None source: list[str] = Field(default_factory=list) class JobErrorOutput(BaseModel): alpha_id: str error: str