feat: compose native research workflows in QuantFlow
This commit is contained in:
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user