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

346 lines
15 KiB
Python
Raw Normal View History

"""Typed native research graphs and immutable, finite run authorizations."""
from collections import defaultdict
from fastapi import HTTPException
from sqlalchemy import func, select
from ..backtests.contracts import SimulationSettings, fingerprint
from ..backtests.service import uid
from ..models import Account, ResearchFlowRun, ResearchStepRun
from .assets import Assets
from .experiments import Experiments, scope_of, seed_settings
from .serialization import encode_snapshot as jsonable_encoder
from .workspace_contracts import WorkflowSpec
NODE_TYPES = {
"input": {"label": "固定输入", "accepts": [], "produces": "context"},
"feature": {"label": "特征方案", "accepts": ["context"], "produces": "context"},
"generate": {
"label": "模板生成 / 增强",
"accepts": ["context", "evaluation", "candidates"],
"produces": "template",
},
"expand": {"label": "校验与展开", "accepts": ["template", "context"], "produces": "candidates"},
"variant": {"label": "Alpha 变体", "accepts": ["context"], "produces": "candidates"},
"backtest": {"label": "回测", "accepts": ["candidates"], "produces": "results"},
"evaluate": {"label": "评估决策", "accepts": ["results"], "produces": "evaluation"},
"filter": {"label": "候选筛选", "accepts": ["evaluation"], "produces": "candidates"},
"condition": {"label": "条件分支", "accepts": ["evaluation"], "produces": "evaluation"},
"summarize": {
"label": "研究汇总",
"accepts": ["context", "template", "candidates", "results", "evaluation"],
"produces": "summary",
},
"iterate": {"label": "有界迭代", "accepts": ["evaluation", "candidates"], "produces": "iteration"},
}
def validate_graph(graph):
"""Require one input, compatible ports, acyclic edges and a single terminal loop."""
nodes = {node.id: node for node in graph.nodes}
if len(nodes) != len(graph.nodes):
raise HTTPException(422, "节点 ID 不能重复")
if sum(node.type == "input" for node in graph.nodes) != 1:
raise HTTPException(422, "流程需要且只能有一个固定输入节点")
incoming, outgoing = defaultdict(list), defaultdict(list)
seen = set()
for edge in graph.edges:
if edge.source not in nodes or edge.target not in nodes or edge.source == edge.target:
raise HTTPException(422, "连线端点不存在或连接自身")
if (edge.source, edge.target, edge.branch) in seen:
raise HTTPException(422, "连线重复")
seen.add((edge.source, edge.target, edge.branch))
source, target = nodes[edge.source], nodes[edge.target]
if NODE_TYPES[source.type]["produces"] not in NODE_TYPES[target.type]["accepts"]:
raise HTTPException(422, f"{source.id} → {target.id} 的输入输出类型不兼容")
if edge.branch and source.type != "condition":
raise HTTPException(422, "分支条件只能设置在条件节点的出边")
incoming[edge.target].append(edge)
outgoing[edge.source].append(edge)
for node in graph.nodes:
if node.type != "input" and not incoming[node.id]:
raise HTTPException(422, f"节点 {node.id} 未连接上游")
if node.type != "summarize" and len(incoming[node.id]) > 1:
raise HTTPException(422, "仅汇总节点接受多个上游;其他节点需要唯一输入")
allowed = (
{"prompt", "asset_id", "version"}
if node.type == "feature"
else {"asset_id", "version"}
if node.type == "expand"
else {"prompt"}
if node.type == "generate"
else {"method"}
if node.type == "variant"
else {"verdicts"}
if node.type == "filter"
else {"max_rounds"}
if node.type == "iterate"
else set()
)
if set(node.config) - allowed:
raise HTTPException(422, f"节点 {node.id} 包含不支持的配置")
if "asset_id" in node.config or "version" in node.config:
if (
not isinstance(node.config.get("asset_id"), str)
or not node.config["asset_id"]
or len(node.config["asset_id"]) > 36
or type(node.config.get("version")) is not int
or node.config["version"] < 1
):
raise HTTPException(422, "节点素材引用需要 ID 和正整数版本")
if "prompt" in node.config and (
not isinstance(node.config["prompt"], str) or len(node.config["prompt"]) > 10000
):
raise HTTPException(422, "节点提示词格式或长度无效")
if node.type == "variant" and node.config.get("method", "structure") not in ("structure", "settings"):
raise HTTPException(422, "未知变体方法")
if node.type == "filter" and (
not isinstance(node.config.get("verdicts", ["pass"]), list)
or any(v not in ("pass", "review", "block") for v in node.config.get("verdicts", ["pass"]))
):
raise HTTPException(422, "筛选结果必须为 pass/review/block")
if node.type == "iterate" and (
outgoing[node.id]
or type(node.config.get("max_rounds")) is not int
or not 1 <= node.config["max_rounds"] <= 100
):
raise HTTPException(422, "迭代必须是终点且明确 1–100 轮上限")
if sum(node.type == "iterate" for node in graph.nodes) > 1:
raise HTTPException(422, "首版每个流程只支持一个有界迭代节点")
remaining, order = {node.id: len(incoming[node.id]) for node in graph.nodes}, []
ready = sorted(key for key, count in remaining.items() if count == 0)
while ready:
key = ready.pop(0)
order.append(key)
for edge in outgoing[key]:
remaining[edge.target] -= 1
if remaining[edge.target] == 0:
ready.append(edge.target)
if len(order) != len(nodes):
raise HTTPException(422, "普通连线不能形成循环,请使用有界迭代节点")
if not any(node.type in ("summarize", "iterate", "evaluate") for node in graph.nodes):
raise HTTPException(422, "流程需要评估、汇总或迭代产物")
return order
def fixed_workflow(max_rounds=3):
steps = [
("input", "input", "固定输入"),
("generate", "generate", "生成研究模板"),
("inspect", "expand", "校验与设参"),
("simulate", "backtest", "回测"),
("decide", "evaluate", "评估决策"),
("enhance", "generate", "增强模板"),
("implement", "expand", "重新展开"),
("iterate", "iterate", "下一轮"),
]
graph = WorkflowSpec.model_validate(
{
"name": "固定研究流水线",
"nodes": [
{
"id": key,
"type": kind,
"label": label,
"x": 40 + (i % 4) * 240,
"y": 50 + (i // 4) * 180,
"config": {"max_rounds": max_rounds} if kind == "iterate" else {},
}
for i, (key, kind, label) in enumerate(steps)
],
"edges": [{"source": steps[i][0], "target": steps[i + 1][0]} for i in range(len(steps) - 1)],
}
)
return graph
class Workflows:
def __init__(self, db):
self.db = db
async def start(self, body, model_revision):
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
if not account or account.connection_status != "connected" or not account.wq_user_id:
raise HTTPException(409, "启动研究前请连接并确认账户身份")
previous = await self.db.scalar(
select(ResearchFlowRun).where(ResearchFlowRun.request_id == body.request_id)
)
request_digest = fingerprint(body.model_dump(mode="json"))
if previous:
if previous.authorization["request_digest"] != request_digest:
raise HTTPException(409, "启动请求标识已用于其他研究")
return await self.get(previous.id)
if body.workflow_id:
asset = await Assets(self.db).get(body.workflow_id, body.workflow_version, "workflow")
graph = WorkflowSpec.model_validate(asset["content"])
else:
asset = None
graph = fixed_workflow(body.budget.max_rounds)
validate_graph(graph)
for node in graph.nodes:
if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds:
raise HTTPException(422, "流程迭代上限超过本次授权轮数")
from ..preparations.service import Preparations
await Preparations(self.db).bind(body)
experiments = Experiments(self.db)
settings_variant = any(
n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes
)
if settings_variant and len(body.parent_alpha_ids) != 1:
raise HTTPException(422, "设置变体需要且只能选择一个种子 Alpha")
if any(n.type == "variant" for n in graph.nodes) and not body.parent_alpha_ids:
raise HTTPException(422, "变体节点需要种子 Alpha")
inputs, _ = await experiments.inputs(
body.input_ids, None if settings_variant else scope_of(body.settings)
)
parents = await experiments.parents(body.parent_alpha_ids, [])
allowed_settings = [body.settings.model_dump(mode="json")]
if settings_variant:
base = seed_settings(parents[0]["settings"])
for snapshot in inputs:
scope = snapshot["scope"]
target = SimulationSettings.model_validate(
{
**base.model_dump(),
"instrumentType": scope["instrument_type"],
**{key: scope[key] for key in ("region", "universe", "delay")},
}
)
errors, _ = await experiments.settings_check(target)
if errors:
raise HTTPException(422, "目标范围:" + ";".join(errors))
if target.model_dump(mode="json") not in allowed_settings:
allowed_settings.append(target.model_dump(mode="json"))
node_assets = {}
for node in graph.nodes:
if node.config.get("asset_id"):
ref = await Assets(self.db).get(
node.config["asset_id"],
node.config["version"],
"feature" if node.type == "feature" else "template",
)
if node.type == "feature":
if not set(ref["content"]["input_ids"]).issubset(body.input_ids):
raise HTTPException(422, "特征方案输入超出本次固定范围")
if not ref["content"].get("template"):
raise HTTPException(422, "特征方案需要输出模板")
node_assets[node.id] = ref
errors, settings_snapshot = await experiments.settings_check(body.settings)
if errors:
raise HTTPException(422, ";".join(errors))
from ..catalog.research_metadata import ResearchMetadata
operators_snapshot = await ResearchMetadata(self.db).get("operators")
if not operators_snapshot["content"].get("items"):
raise HTTPException(422, "启动前需要同步算子目录")
template = (
await Assets(self.db).get(body.template_id, body.template_version, "template")
if body.template_id
else None
)
if (
any(
node.type == "expand"
and not node.config.get("asset_id")
and any(
e.target == node.id and next(n for n in graph.nodes if n.id == e.source).type == "input"
for e in graph.edges
)
for node in graph.nodes
)
and template is None
):
raise HTTPException(422, "直接展开固定输入时需要选择模板版本")
methods = sorted({node.type for node in graph.nodes})
row = ResearchFlowRun(
id=uid(),
request_id=body.request_id,
name=body.name,
definition=graph.model_dump(mode="json"),
authorization=jsonable_encoder(
{
**body.model_dump(mode="json"),
"request_digest": request_digest,
"account_id": account.wq_user_id,
"inputs": inputs,
"parents": parents,
"template": template,
"workflow": asset,
"methods": methods,
"node_assets": node_assets,
"allowed_settings": allowed_settings,
"settings_snapshot": settings_snapshot,
"operators_snapshot": operators_snapshot,
"kind": "quantflow" if asset else "pipeline",
}
),
model_revision=model_revision,
)
self.db.add(row)
await self.db.flush()
return await self.get(row.id)
async def get(self, run_id):
row = await self.db.get(ResearchFlowRun, run_id)
if not row:
raise HTTPException(404, "研究运行不存在")
steps = await self.db.scalars(
select(ResearchStepRun)
.where(ResearchStepRun.run_id == run_id)
.order_by(ResearchStepRun.round, ResearchStepRun.created_at)
)
return jsonable_encoder(
{
**{
key: getattr(row, key)
for key in (
"id",
"name",
"definition",
"authorization",
"model_revision",
"status",
"version",
"round",
"simulations_used",
"model_calls_used",
"error",
"created_at",
"updated_at",
)
},
"steps": [
{
key: getattr(step, key)
for key in (
"id",
"node_id",
"round",
"status",
"output",
"backtest_run_id",
"error",
"created_at",
"updated_at",
)
}
for step in steps
],
}
)
async def list(self, limit=25, offset=0, kind=None):
query = select(ResearchFlowRun)
if kind:
query = query.where(ResearchFlowRun.authorization["kind"].as_string() == kind)
rows = await self.db.scalars(
query.order_by(ResearchFlowRun.created_at.desc()).limit(limit).offset(offset)
)
return {
"items": [await self.get(row.id) for row in rows],
"total": await self.db.scalar(select(func.count()).select_from(query.subquery())),
"limit": limit,
"offset": offset,
}