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

285 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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 fingerprint
from ..backtests.service import uid
from ..models import Account, ResearchFlowRun, ResearchStepRun
from .assets import Assets
from .experiments import Experiments, scope_of
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"}
if node.type in ("feature", "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 "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 not set(node.config.get("verdicts", ["pass"])).issubset({"pass", "review", "block"})
):
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, "流程迭代上限超过本次授权轮数")
experiments = Experiments(self.db)
inputs, _ = await experiments.inputs(body.input_ids, scope_of(body.settings))
parents = await experiments.parents(body.parent_alpha_ids, [])
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 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,
"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):
rows = await self.db.scalars(
select(ResearchFlowRun).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(ResearchFlowRun)),
"limit": limit,
"offset": offset,
}