Files
worldquant-alpha-system/backend/app/backtests/contracts.py
T

220 lines
6.1 KiB
Python
Raw Normal View History

"""Fixed, typed inputs shared by HTTP, AI and research producers."""
import hashlib
import json
from typing import Literal
from pydantic import Field, field_validator, model_validator
from ..schemas import Contract
class SimulationSettings(Contract):
instrumentType: Literal["EQUITY"] = "EQUITY"
region: str = Field(min_length=1, max_length=50, pattern=r"^[A-Z0-9_]+$")
universe: str = Field(min_length=1, max_length=100, pattern=r"^[A-Z0-9_]+$")
delay: Literal[0, 1]
decay: int = Field(default=0, ge=0, le=10000)
neutralization: str = Field(default="INDUSTRY", min_length=1, max_length=50, pattern=r"^[A-Z_]+$")
truncation: float = Field(default=0.08, ge=0, le=1)
pasteurization: Literal["ON", "OFF"] = "ON"
unitHandling: Literal["VERIFY"] = "VERIFY"
nanHandling: Literal["ON", "OFF"] = "OFF"
language: Literal["FASTEXPR"] = "FASTEXPR"
visualization: bool = False
maxTrade: Literal["ON", "OFF"] = "OFF"
class Candidate(Contract):
client_item_id: str = Field(min_length=1, max_length=100)
expression: str = Field(min_length=1, max_length=20000)
settings: SimulationSettings
alpha_type: Literal["REGULAR"] = "REGULAR"
@field_validator("expression")
@classmethod
def nonempty(cls, value):
value = value.strip()
if not value:
raise ValueError("表达式不能为空")
return value
def platform_input(self):
return {"type": self.alpha_type, "regular": self.expression, "settings": self.settings.model_dump()}
class Source(Contract):
kind: str = Field(default="manual", min_length=1, max_length=100)
reference: str | None = Field(default=None, max_length=200)
batch_id: str | None = Field(default=None, max_length=200)
template_input_id: str | None = Field(default=None, max_length=200)
research_id: str | None = Field(default=None, max_length=200)
parent_run_id: str | None = Field(default=None, max_length=36)
hypothesis: str | None = Field(default=None, max_length=2000)
class DraftInput(Contract):
name: str = Field(min_length=1, max_length=200)
source: Source = Field(default_factory=Source)
candidates: list[Candidate] = Field(min_length=1, max_length=10000)
@model_validator(mode="after")
def unique_ids(self):
if len({c.client_item_id for c in self.candidates}) != len(self.candidates):
raise ValueError("client_item_id 在候选集合内必须唯一")
return self
class DraftUpdate(DraftInput):
version: int = Field(ge=1)
class PreviewInput(Contract):
draft_id: str | None = Field(default=None, max_length=36)
draft_version: int | None = Field(default=None, ge=1)
selection: list[str] | None = Field(default=None, min_length=1, max_length=10000)
inline: DraftInput | None = None
@model_validator(mode="after")
def one_input(self):
if (self.inline is None) == (self.draft_id is None):
raise ValueError("必须提供 inline 或 draft_id 之一")
if self.draft_id and self.draft_version is None:
raise ValueError("引用草稿时必须提供 draft_version")
if self.inline and (self.draft_version is not None or self.selection is not None):
raise ValueError("inline 已经是完整固定集合")
return self
class StartInput(Contract):
preview_id: str = Field(min_length=1, max_length=36)
version: int = Field(default=1, ge=1)
idempotency_key: str = Field(min_length=1, max_length=100)
class ControlInput(Contract):
action: Literal["pause", "resume", "stop", "recover"]
version: int = Field(ge=1)
class RerunInput(Contract):
item_ids: list[str] = Field(min_length=1, max_length=10000)
class SchedulerInput(Contract):
concurrency: int = Field(default=3, ge=1, le=8)
batch_size: int = Field(default=8, ge=1, le=10)
version: int = Field(ge=1)
def fingerprint(payload: dict) -> str:
return hashlib.sha256(json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
def group_key(candidate: dict):
settings = candidate["settings"]
return tuple(settings[k] for k in ("region", "delay", "language", "instrumentType"))
class ReferenceInput(Contract):
progress_url: str = Field(min_length=1, max_length=2000)
version: int = Field(ge=1)
class SubsetInput(Contract):
exclude_ids: list[str] = Field(min_length=1, max_length=10000)
# OpenAPI outputs deliberately keep platform snapshots as extensible objects.
class SchedulerOutput(Contract):
concurrency: int
batch_size: int
version: int
blocked_reason: str | None
blocked_until: str | None
class PreviewOutput(Contract):
preview_id: str
version: int
name: str
source: Source
digest: str
total: int
batch_count: int
batch_size: int
duplicate_count: int
duplicates: list[dict]
items: list[Candidate]
limit: int
offset: int
has_more: bool
created_at: str
class RunOutput(Contract):
backtest_run_id: str
preview_id: str
name: str
source: Source
ai_context: dict
control: Literal["active", "paused", "stopped"]
status: str
version: int
total: int
batch_size: int
created_at: str
updated_at: str
counts: dict[str, dict[str, int]]
cursor: int
scheduler: SchedulerOutput
class RunPage(Contract):
items: list[RunOutput]
total: int
limit: int
offset: int
class ResultSnapshot(Contract):
snapshot: dict
observed_at: str
complete: bool
class ItemOutput(Contract):
id: str
client_item_id: str
expression: str
settings: SimulationSettings
attempt_id: str
platform_status: str
collection_status: str
persistence_status: str
simulation_id: str | None
alpha_id: str | None
error: str | None
result: ResultSnapshot | None
class ResultPage(Contract):
backtest_run_id: str
total: int
limit: int
offset: int
items: list[ItemOutput]
class EventOutput(Contract):
seq: int
kind: str
payload: dict
created_at: str
class EventPage(Contract):
items: list[EventOutput]
next_cursor: int
has_more: bool