333 lines
9.5 KiB
Python
333 lines
9.5 KiB
Python
"""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
|
||
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
|
||
|
||
|
||
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
|
||
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
|
||
|
||
|
||
class JobErrorOutput(BaseModel):
|
||
alpha_id: str
|
||
error: str
|