Files

141 lines
5.7 KiB
Python
Raw Permalink Normal View History

"""Bounded SUPER authoring inputs, independent from regular data-field preparation."""
from typing import Literal
from pydantic import Field, field_validator, model_validator
from ..backtests.contracts import SuperSimulationSettings
from ..research.expressions import IDENTIFIER, PLACEHOLDER
from ..research.workspace_contracts import Variable
from ..schemas import AlphaFilters, Contract
class PlanSpec(Contract):
name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=2000)
selection: str = Field(min_length=1, max_length=20000)
combo: str = Field(min_length=1, max_length=20000)
variables: dict[str, Variable] = Field(default_factory=dict, max_length=50)
settings: SuperSimulationSettings
setting_variants: dict[str, list[str | int | float | bool]] = Field(default_factory=dict, max_length=20)
include_baseline: bool = False
reference: str = Field(default="", max_length=200)
parent_plan_id: str | None = Field(default=None, max_length=36)
parent_plan_version: int | None = Field(default=None, ge=1)
parent_alpha_id: str | None = Field(default=None, pattern=r"^[A-Za-z0-9_-]{1,100}$")
parent_experiment_id: str | None = Field(default=None, max_length=36)
@field_validator("name", "hypothesis", "selection", "combo")
@classmethod
def text(cls, value):
if not value.strip():
raise ValueError("内容不能为空")
return value.strip()
@model_validator(mode="after")
def bindings(self):
text = self.selection + "\n" + self.combo
if set(PLACEHOLDER.findall(text)) != set(self.variables):
raise ValueError("Selection/Combo 占位符必须与变量逐一对应")
if any(not IDENTIFIER.fullmatch(k) or v.kind == "field" for k, v in self.variables.items()):
raise ValueError("SUPER 变量须使用合法名称,不能使用 REGULAR 数据字段绑定")
if "{" in PLACEHOLDER.sub("", text) or "}" in PLACEHOLDER.sub("", text):
raise ValueError("占位符格式错误")
for key, values in self.setting_variants.items():
if key not in SuperSimulationSettings.model_fields or not 1 <= len(values) <= 100:
raise ValueError("设置变量必须为已支持设置,每项 1–100 个候选值")
for value in values:
SuperSimulationSettings.model_validate({**self.settings.model_dump(), key: value})
if bool(self.parent_plan_id) != bool(self.parent_plan_version):
raise ValueError("父方案必须同时指定 ID 和版本")
return self
class PlanSearch(Contract):
q: str = Field(default="", max_length=200)
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class PlanReference(PlanSearch):
plan_id: str | None = Field(default=None, min_length=1, max_length=36)
experiment_id: str | None = Field(default=None, min_length=1, max_length=36)
version: int | None = Field(default=None, ge=1)
@model_validator(mode="after")
def one_reference(self):
if bool(self.plan_id) == bool(self.experiment_id) or (self.version and not self.plan_id):
raise ValueError("提供 plan_id 或 experiment_id 之一;version 仅用于方案")
return self
class PlanSave(Contract):
plan: PlanSpec
plan_id: str | None = Field(default=None, max_length=36)
version: int | None = Field(default=None, ge=1)
idempotency_key: str = Field(min_length=1, max_length=100)
@model_validator(mode="after")
def reference(self):
if bool(self.plan_id) != bool(self.version):
raise ValueError("更新须同时提供方案 ID 与当前版本")
return self
class SelectionPreview(Contract):
plan_id: str | None = Field(default=None, max_length=36)
version: int | None = Field(default=None, ge=1)
selection: str = Field(min_length=1, max_length=20000)
settings: SuperSimulationSettings
@field_validator("selection")
@classmethod
def concrete(cls, value):
if not value.strip() or "{" in value or "}" in value:
raise ValueError("预览须提供展开后的非空 Selection")
return value.strip()
def platform_query(self):
return {"selection": self.selection, **self.settings.model_dump(include={
"instrumentType", "region", "delay", "selectionLimit", "selectionHandling"})}
class SelectionReference(PlanSearch):
snapshot_id: str | None = Field(default=None, max_length=36)
job_id: str | None = Field(default=None, max_length=36)
@model_validator(mode="after")
def one(self):
if bool(self.snapshot_id) == bool(self.job_id):
raise ValueError("提供 snapshot_id 或 job_id 之一")
return self
class BuildCandidates(Contract):
plan: PlanSpec | None = None
plan_id: str | None = Field(default=None, max_length=36)
version: int | None = Field(default=None, ge=1)
mode: Literal["all", "random"] = "all"
limit: int = Field(default=100, ge=1, le=10000)
seed: int = Field(default=0, ge=0, le=2147483647)
selection_snapshot_ids: list[str] = Field(default_factory=list, max_length=100)
idempotency_key: str = Field(min_length=1, max_length=100)
@model_validator(mode="after")
def one(self):
if bool(self.plan) == bool(self.plan_id) or bool(self.plan_id) != bool(self.version):
raise ValueError("提供内联方案或方案 ID/版本之一")
return self
class SuperAlphaSearch(Contract):
filters: AlphaFilters = Field(default_factory=AlphaFilters)
class AlphaReference(Contract):
alpha_id: str = Field(pattern=r"^[A-Za-z0-9_-]{1,100}$")
class ExperimentPreview(Contract):
candidate_ids: list[str] = Field(min_length=1, max_length=10000)