Files
worldquant-alpha-system/backend/app/schemas.py
T

357 lines
10 KiB
Python
Raw Normal View History

"""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"]
SortField = Literal[
"id",
"name",
"sharpe",
"fitness",
"returns",
"turnover",
"margin",
"drawdown",
"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
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
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"):
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
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