Files
worldquant-alpha-system/backend/app/schemas.py
T
yuxuanhui b8429efa3d
Deploy production / deploy (push) Has been cancelled
feat: configure Gitea deployment and environment-managed WorldQuant credentials
2026-09-09 09:58:09 +08:00

357 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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