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

297 lines
7.8 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 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"]
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):
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", "alpha_refresh", "pnl_refresh"]
alpha_ids: list[str] = Field(default_factory=list)
@model_validator(mode="after")
def validate_ids(self):
if self.kind == "full_sync":
if self.alpha_ids:
raise ValueError("全量同步不接受 Alpha ID")
else:
self.alpha_ids = valid_ids(self.alpha_ids)
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
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
processed: int
failed: int
total: int | None
error: str | None
next_retry_at: datetime | None
created_at: datetime
updated_at: datetime
class AccountOutput(BaseModel):
email: str | None
configured: bool
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
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 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
class JobErrorOutput(BaseModel):
alpha_id: str
error: str