feat: compose native research workflows in QuantFlow

This commit is contained in:
yuxuanhui
2026-09-08 22:42:29 +08:00
parent 7860434b04
commit a6e36e50ec
22 changed files with 2274 additions and 85 deletions
+67 -8
View File
@@ -5,7 +5,7 @@ from collections import defaultdict
from fastapi import HTTPException
from sqlalchemy import func, select
from ..backtests.contracts import fingerprint
from ..backtests.contracts import SimulationSettings, fingerprint
from ..backtests.service import uid
from ..models import Account, ResearchFlowRun, ResearchStepRun
from .assets import Assets
@@ -64,8 +64,12 @@ def validate_graph(graph):
if node.type != "summarize" and len(incoming[node.id]) > 1:
raise HTTPException(422, "仅汇总节点接受多个上游;其他节点需要唯一输入")
allowed = (
{"prompt"}
if node.type in ("feature", "generate")
{"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"}
@@ -76,6 +80,15 @@ def validate_graph(graph):
)
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
):
@@ -84,7 +97,7 @@ def validate_graph(graph):
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"})
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 (
@@ -169,8 +182,48 @@ class Workflows:
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))
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 = SimulationSettings.model_validate(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))
@@ -187,6 +240,7 @@ class Workflows:
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
@@ -212,6 +266,8 @@ class Workflows:
"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",
@@ -272,13 +328,16 @@ class Workflows:
}
)
async def list(self, limit=25, offset=0):
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(
select(ResearchFlowRun).order_by(ResearchFlowRun.created_at.desc()).limit(limit).offset(offset)
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(ResearchFlowRun)),
"total": await self.db.scalar(select(func.count()).select_from(query.subquery())),
"limit": limit,
"offset": offset,
}