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
+1
View File
@@ -46,6 +46,7 @@ class PageContext(Contract):
"features",
"variants",
"pipeline",
"quantflow",
] = "alphas"
research_run_id: str | None = Field(default=None, max_length=36)
research_asset_id: str | None = Field(default=None, max_length=36)
+6 -3
View File
@@ -310,8 +310,10 @@ class Experiments:
preserve_source=True,
)
async def setting_variants(self, body):
parents = await self.parents([body.alpha_id], [])
async def setting_variants(self, body, *, parent_snapshot=None, extra_evidence=None, kind="variant"):
parents = (
[parent_snapshot] if parent_snapshot is not None else await self.parents([body.alpha_id], [])
)
original = parents[0]
base = SimulationSettings.model_validate(original["settings"])
expression = original["expression"]
@@ -362,12 +364,13 @@ class Experiments:
)
return await self.save(
f"{body.alpha_id} · 设置变体",
"variant",
kind,
body.hypothesis,
snapshots,
parents,
candidates,
{
**(extra_evidence or {}),
"method": "settings",
"rejected": rejected,
"operators_snapshot": operators_snapshot,
+45 -9
View File
@@ -22,6 +22,7 @@ from .workspace_contracts import (
ImportCommit,
ImportPreview,
SettingVariants,
WorkflowSpec,
)
router = APIRouter(prefix="/api/v1/research", tags=["research"], dependencies=[Depends(require_auth)])
@@ -41,7 +42,7 @@ async def assets(
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
if kind not in ("template", "feature", "view"):
if kind not in ("template", "feature", "view", "workflow"):
raise HTTPException(422, "当前素材类型尚未开放")
async with request.app.state.sessions() as db:
return await Assets(db).list(kind, q, limit, offset)
@@ -49,7 +50,7 @@ async def assets(
@router.post("/assets", status_code=201)
async def save_asset(body: AssetWrite, request: Request):
if body.kind not in ("template", "feature", "view"):
if body.kind not in ("template", "feature", "view", "workflow"):
raise HTTPException(422, "当前素材类型尚未开放")
async with request.app.state.sessions.begin() as db:
return await Assets(db).save(body)
@@ -63,7 +64,7 @@ async def asset(asset_id: str, request: Request, version: int | None = Query(Non
@router.put("/assets/{asset_id}")
async def update_asset(asset_id: str, body: AssetWrite, request: Request):
if body.kind not in ("template", "feature", "view"):
if body.kind not in ("template", "feature", "view", "workflow"):
raise HTTPException(422, "当前素材类型尚未开放")
async with request.app.state.sessions.begin() as db:
return await Assets(db).save(body, asset_id)
@@ -240,21 +241,42 @@ async def fixed_recipe():
async def start_flow(body: FlowStart, request: Request):
from .workflows import Workflows
if body.workflow_id:
raise HTTPException(422, "自定义流程将在 QuantFlow 阶段开放")
async with request.app.state.sessions.begin() as db:
config = await request.app.state.ai.config(db)
result = await Workflows(db).start(body, config.revision)
from .workflows import fixed_workflow
from .workspace_contracts import WorkflowSpec
graph = (
WorkflowSpec.model_validate(
(await Assets(db).get(body.workflow_id, body.workflow_version, "workflow"))["content"]
)
if body.workflow_id
else fixed_workflow(body.budget.max_rounds)
)
needs_model = any(
n.type == "generate"
or (n.type == "feature" and not n.config.get("asset_id"))
or (n.type == "variant" and n.config.get("method", "structure") == "structure")
for n in graph.nodes
)
config = await request.app.state.ai.config(db) if needs_model else None
result = await Workflows(db).start(body, config.revision if config else None)
request.app.state.research.wake.set()
return result
@router.get("/flows/runs")
async def flow_runs(request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)):
async def flow_runs(
request: Request,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
kind: str | None = None,
):
from .workflows import Workflows
async with request.app.state.sessions() as db:
return await Workflows(db).list(limit, offset)
if kind not in (None, "pipeline", "quantflow"):
raise HTTPException(422, "未知研究运行类型")
return await Workflows(db).list(limit, offset, kind)
@router.get("/flows/runs/{run_id}")
@@ -276,3 +298,17 @@ async def control_flow(run_id: str, body: FlowControl, request: Request):
request.app.state.research.wake.set()
request.app.state.runner.backtests.wake.set()
return result
@router.get("/flows/nodes")
async def flow_nodes():
from .workflows import NODE_TYPES
return {"items": [{"type": key, **value} for key, value in NODE_TYPES.items()]}
@router.post("/flows/validate")
async def validate_flow(body: WorkflowSpec):
from .workflows import validate_graph
return {"valid": True, "order": validate_graph(body)}
+214 -56
View File
@@ -6,21 +6,32 @@ Backtest source fields are provenance, never a grant to execute automatically.
import asyncio
import logging
import random
from collections import defaultdict
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError
from ..backtests.contracts import ControlInput, StartInput
from ..backtests.contracts import ControlInput, SimulationSettings, StartInput
from ..backtests.service import Backtests, uid
from ..models import Account, ResearchFlowRun, ResearchStepRun, now
from .assets import Assets
from .evaluations import Evaluations
from .experiments import Experiments
from .experiments import Experiments, scope_of
from .features import Features
from .model import request_model
from .workflows import validate_graph
from .workspace_contracts import AssetWrite, EvaluateInput, Expansion, Generation, TemplateSpec, WorkflowSpec
from .workspace_contracts import (
AssetWrite,
EvaluateInput,
Expansion,
FeatureSpec,
Generation,
SettingVariants,
TemplateSpec,
WorkflowSpec,
)
logger = logging.getLogger(__name__)
ACTIVE = ("queued", "running")
@@ -162,7 +173,8 @@ class ResearchRuntime:
model_work = await self.prepare(run_id)
if model_work:
step_id, context, revision = model_work
result, evidence = await request_model(self.ai, context, TemplateSpec, revision)
output_type = FeatureSpec if context.get("method") == "feature" else TemplateSpec
result, evidence = await request_model(self.ai, context, output_type, revision)
await self.finish_model(run_id, step_id, result, evidence)
except HTTPException as exc:
await self.fail(run_id, str(exc.detail))
@@ -271,7 +283,7 @@ class ResearchRuntime:
else None
)
step.output = {
**(previous.output if previous else {}),
**(previous.output if previous else {"template": run.authorization.get("template")}),
"type": "context",
"input_ids": run.authorization["input_ids"],
}
@@ -287,50 +299,45 @@ class ResearchRuntime:
}
else:
return await self.reserve_model(db, run, step, data, node)
elif node.type == "feature":
reference = run.authorization.get("node_assets", {}).get(key)
if step.output.get("feature"):
ref = step.output["feature"]
reference = await Assets(db).get(ref["id"], ref["version"], "feature")
if reference:
template = await Features(db).to_template(reference["id"], reference["version"])
step.output = {
**data,
**step.output,
"type": "context",
"feature": {"id": reference["id"], "version": reference["version"]},
"input_ids": reference["content"]["input_ids"],
"template": {"id": template["id"], "version": template["version"]},
}
else:
return await self.reserve_model(db, run, step, data, node)
elif node.type == "expand":
if not data.get("template"):
raise HTTPException(422, "展开节点没有固定模板版本")
body = Expansion(
asset_id=data["template"]["id"],
version=data["template"]["version"],
input_ids=run.authorization["input_ids"],
hypothesis=run.authorization["hypothesis"],
settings=run.authorization["settings"],
mode="random",
limit=run.authorization["batch_candidates"],
seed=run.authorization["seed"] + run.round,
parent_alpha_ids=run.authorization["parent_alpha_ids"],
parent_experiment_ids=data.get("parent_experiment_ids", []),
reference = run.authorization.get("node_assets", {}).get(key) or data.get("template")
await self.expand(db, run, step, {**data, "template": reference})
return
elif node.type == "variant":
if node.config.get("method", "structure") == "structure":
if step.output.get("template"):
await self.expand(db, run, step, step.output)
return
return await self.reserve_model(db, run, step, data, node)
experiment = await Experiments(db).setting_variants(
SettingVariants(
alpha_id=run.authorization["parent_alpha_ids"][0],
input_ids=run.authorization["input_ids"],
hypothesis=run.authorization["hypothesis"],
),
parent_snapshot=run.authorization["parents"][0],
extra_evidence={"flow_run_id": run.id, "node_id": key, "round": run.round},
kind=run.authorization["kind"],
)
frozen_parents = run.authorization["parents"] + await Experiments(db).parents(
[], body.parent_experiment_ids
)
experiment = await Experiments(db).create(
body,
"pipeline",
{
"flow_run_id": run.id,
"node_id": key,
"round": run.round,
"method": "pipeline",
"authorized_seed_snapshots": run.authorization["parents"],
},
parent_snapshots=frozen_parents,
)
step.output = {
"type": "candidates",
"experiment_id": experiment["id"],
"candidate_ids": [
c["client_item_id"]
for c in experiment["candidates"]
if c["validation"]["status"] == "valid"
],
"template": data["template"],
}
if not step.output["candidate_ids"]:
step.status, step.error = "blocked", "候选均未通过本地校验,请核实字段、算子和设置"
halt(run, "needs_review", step.error)
return
self.candidates(run, step, experiment, {})
return
elif node.type == "backtest":
await self.backtest(db, run, step, data)
return
@@ -349,6 +356,55 @@ class ResearchRuntime:
"experiment_id": data["experiment_id"],
"backtest_run_id": data["backtest_run_id"],
}
elif node.type == "condition":
step.output = {**data, "type": "evaluation"}
elif node.type == "filter":
report = (await Evaluations(db).get(data["evaluation_id"]))["report"]
ids = [
r["client_item_id"]
for r in report["records"]
if r["verdict"] in node.config.get("verdicts", ["pass"])
]
experiment = await Experiments(db).get(data["experiment_id"])
template = experiment["evidence"].get("template")
step.output = {
**data,
"type": "candidates",
"candidate_ids": ids,
"template": {"id": template["id"], "version": template["version"]}
if template
else None,
}
if not ids:
step.status, step.updated_at = "skipped", now()
return
elif node.type == "summarize":
# References keep joins bounded, without recursively copying the upstream graph.
step.output = {
"type": "summary",
"artifacts": [
{
"step_id": steps[e.source].id,
"node_id": e.source,
**{
k: steps[e.source].output[k]
for k in (
"type",
"template",
"feature",
"experiment_id",
"evaluation_id",
"backtest_run_id",
"verdict",
)
if k in steps[e.source].output
},
}
for e in upstream_edges
if steps[e.source].status != "skipped"
and (not e.branch or steps[e.source].output.get("verdict") == e.branch)
],
}
elif node.type == "iterate":
step.output = {
**data,
@@ -370,10 +426,16 @@ class ResearchRuntime:
generation = Generation(
name=run.name,
hypothesis=run.authorization["hypothesis"],
input_ids=run.authorization["input_ids"],
input_ids=data.get("input_ids", run.authorization["input_ids"]),
method="feature"
if node.type == "feature"
else "structure"
if node.type == "variant"
else "template",
parent_experiment_ids=[data["experiment_id"]] if data.get("experiment_id") else [],
)
context = await Experiments(db).generation_context(generation)
context["method"] = generation.method
context["parents"] = run.authorization["parents"] + context["parents"]
context["operators"] = run.authorization["operators_snapshot"]["content"]["items"][:100]
context["node_prompt"] = node.config.get("prompt", "")
@@ -414,9 +476,16 @@ class ResearchRuntime:
step = await db.get(ResearchStepRun, step_id)
if not step or step.status != "running":
return
if isinstance(result, FeatureSpec) and set(result.input_ids) != set(
[i["id"] for i in step.output["context"]["inputs"]]
):
raise HTTPException(422, "模型不能改变已固定的输入范围")
# A paused/stopped run may collect this already-issued model output, but cannot advance.
asset = await Assets(db).save(
AssetWrite(kind="template", content=result.model_dump(mode="json")),
AssetWrite(
kind="feature" if isinstance(result, FeatureSpec) else "template",
content=result.model_dump(mode="json"),
),
provenance={
"flow_run_id": run_id,
"step_id": step_id,
@@ -424,18 +493,92 @@ class ResearchRuntime:
"context": step.output["context"],
},
)
feature = (
{"id": asset["id"], "version": asset["version"]} if isinstance(result, FeatureSpec) else None
)
mark_model_attempt(step, "completed")
step.output = {
"model_attempts": step.output.get("model_attempts", []),
"type": "template",
"template": {"id": asset["id"], "version": asset["version"]},
"type": "context" if feature else "template",
"feature": feature,
"input_ids": result.input_ids
if feature
else [i["id"] for i in step.output["context"]["inputs"]],
"template": None if feature else {"id": asset["id"], "version": asset["version"]},
"parent_experiment_ids": step.output.get("parent_experiment_ids", []),
"generation": evidence,
"reserved_call": step.output["reserved_call"],
}
step.status, step.updated_at = "completed", now()
node = next(n for n in run.definition["nodes"] if n["id"] == step.node_id)
# Keep post-processing durable and separate: pause/account changes are checked again
# before converting features or expanding variants on the next active tick.
step.status = "generated" if node["type"] in ("variant", "feature") else "completed"
step.updated_at = now()
changed(run)
def candidates(self, run, step, experiment, data):
step.output = {
**data,
"type": "candidates",
"experiment_id": experiment["id"],
"candidate_ids": [
c["client_item_id"] for c in experiment["candidates"] if c["validation"]["status"] == "valid"
],
}
ids = step.output["candidate_ids"]
maximum = run.authorization["batch_candidates"]
if len(ids) > maximum:
step.output = {
**step.output,
"candidate_ids": random.Random(run.authorization["seed"] + run.round).sample(ids, maximum),
}
step.status, step.updated_at = "completed", now()
if not step.output["candidate_ids"]:
step.status, step.error = "blocked", "候选均未通过本地校验,请核实字段、算子和设置"
if run.status in ACTIVE:
halt(run, "needs_review", step.error)
async def expand(self, db, run, step, data):
reference = data.get("template")
if not reference:
raise HTTPException(422, "展开节点没有固定模板版本")
scope = scope_of(SimulationSettings.model_validate(run.authorization["settings"]))
ids = [
i["id"]
for i in run.authorization["inputs"]
if i["scope"] == scope and i["id"] in data.get("input_ids", run.authorization["input_ids"])
]
if not ids:
raise HTTPException(422, "模板展开需要与基础设置匹配的固定输入")
body = Expansion(
asset_id=reference["id"],
version=reference["version"],
input_ids=ids,
hypothesis=run.authorization["hypothesis"],
settings=run.authorization["settings"],
mode="random",
limit=run.authorization["batch_candidates"],
seed=run.authorization["seed"] + run.round,
parent_alpha_ids=run.authorization["parent_alpha_ids"],
parent_experiment_ids=data.get("parent_experiment_ids", []),
)
frozen_parents = run.authorization["parents"] + await Experiments(db).parents(
[], body.parent_experiment_ids
)
experiment = await Experiments(db).create(
body,
run.authorization["kind"],
{
"flow_run_id": run.id,
"node_id": step.node_id,
"round": run.round,
"method": run.authorization["kind"],
"authorized_seed_snapshots": run.authorization["parents"],
},
parent_snapshots=frozen_parents,
)
self.candidates(run, step, experiment, data)
async def backtest(self, db, run, step, data):
service = Backtests(db)
if step.backtest_run_id:
@@ -470,12 +613,27 @@ class ResearchRuntime:
):
raise HTTPException(409, "研究候选预览不再匹配保存的授权步骤")
experiment = await Experiments(db).get(step.output["experiment_id"])
if experiment["evidence"].get("flow_run_id") != run.id or {
s["id"] for s in experiment["inputs"]
} != set(run.authorization["input_ids"]):
if (
experiment["evidence"].get("flow_run_id") != run.id
or {s["id"] for s in experiment["inputs"]}.issubset(set(run.authorization["input_ids"])) is False
):
raise HTTPException(403, "候选不属于此研究运行的固定输入范围")
if "backtest" not in run.authorization["methods"] or any(
c["settings"] != run.authorization["settings"]
c["settings"]
not in (
run.authorization.get("allowed_settings", [run.authorization["settings"]])
if experiment["evidence"].get("method") == "settings"
else [run.authorization["settings"]]
)
or not c.get("input_ids")
or any(
not any(
i["id"] == input_id
and i["scope"] == scope_of(SimulationSettings.model_validate(c["settings"]))
for i in run.authorization["inputs"]
)
for input_id in c["input_ids"]
)
for c in experiment["candidates"]
if c["client_item_id"] in step.output["candidate_ids"]
):
+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,
}
+18
View File
@@ -184,6 +184,24 @@ async def list_flows(ctx, args):
CAPABILITIES += (
Capability(
name="search_research_workflows",
schema=AssetQuery,
description="分页查阅 QuantFlow 原生流程与版本,不启动运行。",
label="搜索研究流程",
renderer="research",
effect="query",
handler=lambda ctx, args: Assets(ctx.business.db).list("workflow", **args.model_dump()),
),
Capability(
name="get_research_workflow",
schema=FixedAssetReference,
description="读取指定研究流程版本及原生节点连接,配合 get_research_run 解释执行产物。",
label="读取流程版本",
renderer="research",
effect="query",
handler=lambda ctx, args: Assets(ctx.business.db).get(args.asset_id, args.version, "workflow"),
),
Capability(
name="get_research_run",
schema=FlowReference,