261 lines
9.9 KiB
Python
261 lines
9.9 KiB
Python
"""Public typed research inputs; arbitrary code, URLs and credentials are not accepted."""
|
|
|
|
import math
|
|
from typing import Literal
|
|
|
|
from pydantic import Field, field_validator, model_validator
|
|
|
|
from ..backtests.contracts import SimulationSettings
|
|
from ..catalog.contracts import Scope
|
|
from ..preparations.contracts import PreparationReference
|
|
from ..schemas import Contract
|
|
from .expressions import IDENTIFIER, PLACEHOLDER, normalize_template
|
|
|
|
AssetKind = Literal["template", "feature", "view", "workflow"]
|
|
|
|
|
|
class Variable(Contract):
|
|
kind: Literal["field", "operator", "integer", "number", "group", "string", "fragment"]
|
|
values: list[str | int | float] = Field(min_length=1, max_length=10000)
|
|
field_type: Literal["MATRIX", "VECTOR", "GROUP"] | None = None
|
|
|
|
@model_validator(mode="after")
|
|
def valid_values(self):
|
|
for value in self.values:
|
|
if isinstance(value, float) and not math.isfinite(value):
|
|
raise ValueError("变量数值必须有限")
|
|
if self.kind in ("field", "operator", "group") and not IDENTIFIER.fullmatch(str(value)):
|
|
raise ValueError("字段、算子和分组值必须是标识符")
|
|
if self.kind == "integer" and (type(value) is not int):
|
|
raise ValueError("整数参数只能包含整数")
|
|
if self.kind == "number" and type(value) not in (int, float):
|
|
raise ValueError("数值参数只能包含数值")
|
|
if self.kind in ("string", "fragment") and not isinstance(value, str):
|
|
raise ValueError("字符串和表达式片段变量必须包含文本")
|
|
if self.kind == "field" and self.field_type is None:
|
|
raise ValueError("字段变量需要明确 MATRIX/VECTOR/GROUP 类型")
|
|
if self.kind != "field" and self.field_type is not None:
|
|
raise ValueError("仅字段变量可以声明字段类型")
|
|
return self
|
|
|
|
|
|
class TemplateSpec(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
description: str = Field(default="", max_length=10000)
|
|
expression: str = Field(min_length=1, max_length=20000)
|
|
variables: dict[str, Variable] = Field(default_factory=dict, max_length=100)
|
|
scope: Scope | None = None
|
|
category: Literal["template", "fragment"] = "template"
|
|
|
|
@field_validator("expression")
|
|
@classmethod
|
|
def normalize(cls, value):
|
|
return normalize_template(value)
|
|
|
|
@model_validator(mode="after")
|
|
def bindings(self):
|
|
if set(PLACEHOLDER.findall(self.expression)) != set(self.variables):
|
|
raise ValueError("模板变量必须与占位符逐一对应")
|
|
remainder = PLACEHOLDER.sub("", self.expression)
|
|
if "{" in remainder or "}" in remainder:
|
|
raise ValueError("模板占位符格式错误")
|
|
return self
|
|
|
|
|
|
class FeatureStep(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
rationale: str = Field(min_length=1, max_length=3000)
|
|
expression: str = Field(default="", max_length=20000)
|
|
|
|
|
|
class FeatureSpec(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
hypothesis: str = Field(min_length=1, max_length=10000)
|
|
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
|
steps: list[FeatureStep] = Field(default_factory=list, max_length=30)
|
|
template: TemplateSpec | None = None
|
|
|
|
|
|
class FeatureConversion(Contract):
|
|
version: int = Field(ge=1)
|
|
|
|
|
|
class ViewSpec(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
filters: dict = Field(default_factory=dict)
|
|
columns: list[str] = Field(default_factory=list, max_length=50)
|
|
|
|
@field_validator("filters")
|
|
@classmethod
|
|
def valid_filters(cls, value):
|
|
from ..schemas import AlphaFilters
|
|
|
|
return AlphaFilters.model_validate(value).model_dump(mode="json", exclude_none=True)
|
|
|
|
|
|
class AssetWrite(Contract):
|
|
kind: AssetKind
|
|
content: dict
|
|
version: int | None = Field(default=None, ge=1)
|
|
|
|
|
|
class Expansion(Contract):
|
|
asset_id: str | None = Field(default=None, max_length=36)
|
|
version: int | None = Field(default=None, ge=1)
|
|
template: TemplateSpec | None = None
|
|
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
|
hypothesis: str = Field(min_length=1, max_length=10000)
|
|
settings: SimulationSettings
|
|
mode: Literal["all", "random"] = "all"
|
|
limit: int = Field(default=100, ge=1, le=10000)
|
|
seed: int = 0
|
|
parent_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
parent_experiment_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
|
|
@model_validator(mode="after")
|
|
def template_reference(self):
|
|
if (self.template is None) == (self.asset_id is None):
|
|
raise ValueError("提供模板版本引用或内联模板之一")
|
|
if self.asset_id and self.version is None:
|
|
raise ValueError("引用模板必须指定版本")
|
|
return self
|
|
|
|
|
|
class Generation(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
hypothesis: str = Field(min_length=1, max_length=10000)
|
|
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
|
parent_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
parent_experiment_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
method: Literal["template", "structure", "feature"] = "template"
|
|
|
|
|
|
class SettingVariants(Contract):
|
|
alpha_id: str = Field(min_length=1, max_length=100)
|
|
input_ids: list[str] = Field(default_factory=list, max_length=100)
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=100)
|
|
hypothesis: str = Field(default="保持表达式,比较字段共同支持的市场与设置", max_length=10000)
|
|
|
|
|
|
class ExperimentPreview(Contract):
|
|
candidate_ids: list[str] | None = Field(default=None, min_length=1, max_length=10000)
|
|
|
|
|
|
class EvaluationRules(Contract):
|
|
version: Literal["research-v1"] = "research-v1"
|
|
sharpe_min: float = Field(default=1.0, allow_inf_nan=False)
|
|
fitness_min: float = Field(default=0.5, allow_inf_nan=False)
|
|
turnover_max: float = Field(default=0.7, ge=0, le=1)
|
|
|
|
|
|
class EvaluateInput(Contract):
|
|
alpha_id: str | None = Field(default=None, max_length=100)
|
|
experiment_id: str | None = Field(default=None, max_length=36)
|
|
backtest_run_id: str | None = Field(default=None, max_length=36)
|
|
rules: EvaluationRules = Field(default_factory=EvaluationRules)
|
|
|
|
@model_validator(mode="after")
|
|
def target(self):
|
|
if self.alpha_id and self.experiment_id:
|
|
raise ValueError("实验评估应选择关联回测运行;Alpha 评估单独保存")
|
|
if bool(self.alpha_id) == bool(self.backtest_run_id):
|
|
raise ValueError("选择 Alpha 或回测运行之一")
|
|
return self
|
|
|
|
|
|
class CompareInput(Contract):
|
|
alpha_ids: list[str] = Field(min_length=2, max_length=20)
|
|
|
|
|
|
class OperatorAnnotation(Contract):
|
|
note: str = Field(default="", max_length=10000)
|
|
favorite: bool = False
|
|
version: int = Field(ge=0)
|
|
|
|
|
|
class FieldAvailabilityInput(Contract):
|
|
field_id: str = Field(pattern=r"^[A-Za-z_][A-Za-z0-9_]*$", max_length=200)
|
|
scope: Scope
|
|
|
|
|
|
class ImportPreview(Contract):
|
|
templates: list[dict] = Field(min_length=1, max_length=100)
|
|
|
|
|
|
class ImportCommit(Contract):
|
|
templates: list[TemplateSpec] = Field(min_length=1, max_length=100)
|
|
digest: str = Field(min_length=64, max_length=64)
|
|
|
|
|
|
class Node(Contract):
|
|
id: str = Field(pattern=r"^[A-Za-z_][A-Za-z0-9_-]*$", max_length=100)
|
|
type: Literal[
|
|
"input",
|
|
"feature",
|
|
"generate",
|
|
"expand",
|
|
"variant",
|
|
"backtest",
|
|
"evaluate",
|
|
"filter",
|
|
"condition",
|
|
"summarize",
|
|
"iterate",
|
|
]
|
|
label: str = Field(default="", max_length=100)
|
|
x: float = Field(default=0, ge=0, le=10000)
|
|
y: float = Field(default=0, ge=0, le=10000)
|
|
config: dict = Field(default_factory=dict)
|
|
|
|
|
|
class Edge(Contract):
|
|
source: str
|
|
target: str
|
|
branch: Literal["pass", "review", "block"] | None = None
|
|
|
|
|
|
class WorkflowSpec(Contract):
|
|
name: str = Field(min_length=1, max_length=200)
|
|
nodes: list[Node] = Field(min_length=1, max_length=50)
|
|
edges: list[Edge] = Field(default_factory=list, max_length=100)
|
|
|
|
|
|
class Budget(Contract):
|
|
max_rounds: int = Field(ge=1, le=100, strict=True)
|
|
max_simulations: int = Field(ge=1, le=10000, strict=True)
|
|
max_model_calls: int = Field(ge=1, le=1000, strict=True)
|
|
|
|
|
|
class FlowStart(Contract):
|
|
request_id: str = Field(min_length=1, max_length=100)
|
|
name: str = Field(min_length=1, max_length=200)
|
|
workflow_id: str | None = None
|
|
workflow_version: int | None = Field(default=None, ge=1)
|
|
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
|
hypothesis: str = Field(min_length=1, max_length=10000)
|
|
settings: SimulationSettings
|
|
budget: Budget
|
|
rules: EvaluationRules = Field(default_factory=EvaluationRules)
|
|
batch_candidates: int = Field(default=8, ge=1, le=100)
|
|
seed: int = 0
|
|
template_id: str | None = None
|
|
template_version: int | None = Field(default=None, ge=1)
|
|
parent_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
|
|
|
|
@model_validator(mode="after")
|
|
def fixed_references(self):
|
|
if bool(self.workflow_id) != (self.workflow_version is not None):
|
|
raise ValueError("流程引用必须同时提供 ID 和版本")
|
|
if bool(self.template_id) != (self.template_version is not None):
|
|
raise ValueError("模板引用必须同时提供 ID 和版本")
|
|
return self
|
|
|
|
|
|
class FlowControl(Contract):
|
|
action: Literal["pause", "resume", "stop"]
|
|
version: int = Field(ge=1)
|