227 lines
6.4 KiB
Python
227 lines
6.4 KiB
Python
"""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 ..preparations.contracts import PreparationReference
|
|
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"
|
|
maxPosition: 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)
|
|
input_snapshot_ids: list[str] = Field(default_factory=list, max_length=100)
|
|
input_snapshot_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):
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
|
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
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
|