Files
worldquant-alpha-system/backend/app/research/workspace_contracts.py
T

250 lines
9.0 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 ..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(min_length=1, max_length=20)
steps: list[FeatureStep] = Field(default_factory=list, max_length=30)
template: TemplateSpec | None = None
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
AlphaFilters.model_validate(value)
return value
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(min_length=1, 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(min_length=1, 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(min_length=1, 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 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)
max_simulations: int = Field(ge=1, le=10000)
max_model_calls: int = Field(ge=1, le=1000)
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(min_length=1, 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)