2026-09-08 22:16:05 +08:00
|
|
|
|
"""Typed native research graphs and immutable, finite run authorizations."""
|
|
|
|
|
|
|
|
|
|
|
|
from collections import defaultdict
|
|
|
|
|
|
|
|
|
|
|
|
from fastapi import HTTPException
|
|
|
|
|
|
from sqlalchemy import func, select
|
|
|
|
|
|
|
2026-09-08 22:42:29 +08:00
|
|
|
|
from ..backtests.contracts import SimulationSettings, fingerprint
|
2026-09-08 22:16:05 +08:00
|
|
|
|
from ..backtests.service import uid
|
|
|
|
|
|
from ..models import Account, ResearchFlowRun, ResearchStepRun
|
|
|
|
|
|
from .assets import Assets
|
2026-09-08 23:32:52 +08:00
|
|
|
|
from .experiments import Experiments, scope_of, seed_settings
|
2026-09-08 22:16:05 +08:00
|
|
|
|
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 = (
|
2026-09-08 22:42:29 +08:00
|
|
|
|
{"prompt", "asset_id", "version"}
|
|
|
|
|
|
if node.type == "feature"
|
|
|
|
|
|
else {"asset_id", "version"}
|
|
|
|
|
|
if node.type == "expand"
|
|
|
|
|
|
else {"prompt"}
|
|
|
|
|
|
if node.type == "generate"
|
2026-09-08 22:16:05 +08:00
|
|
|
|
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} 包含不支持的配置")
|
2026-09-08 22:42:29 +08:00
|
|
|
|
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 和正整数版本")
|
2026-09-08 22:16:05 +08:00
|
|
|
|
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)
|
2026-09-08 22:42:29 +08:00
|
|
|
|
or any(v not in ("pass", "review", "block") for v in node.config.get("verdicts", ["pass"]))
|
2026-09-08 22:16:05 +08:00
|
|
|
|
):
|
|
|
|
|
|
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)
|
2026-09-08 22:42:29 +08:00
|
|
|
|
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)
|
|
|
|
|
|
)
|
2026-09-08 22:16:05 +08:00
|
|
|
|
parents = await experiments.parents(body.parent_alpha_ids, [])
|
2026-09-08 22:42:29 +08:00
|
|
|
|
allowed_settings = [body.settings.model_dump(mode="json")]
|
|
|
|
|
|
if settings_variant:
|
2026-09-08 23:32:52 +08:00
|
|
|
|
base = seed_settings(parents[0]["settings"])
|
2026-09-08 22:42:29 +08:00
|
|
|
|
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
|
2026-09-08 22:16:05 +08:00
|
|
|
|
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"
|
2026-09-08 22:42:29 +08:00
|
|
|
|
and not node.config.get("asset_id")
|
2026-09-08 22:16:05 +08:00
|
|
|
|
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,
|
2026-09-08 22:42:29 +08:00
|
|
|
|
"node_assets": node_assets,
|
|
|
|
|
|
"allowed_settings": allowed_settings,
|
2026-09-08 22:16:05 +08:00
|
|
|
|
"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
|
|
|
|
|
|
],
|
|
|
|
|
|
}
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-09-08 22:42:29 +08:00
|
|
|
|
async def list(self, limit=25, offset=0, kind=None):
|
|
|
|
|
|
query = select(ResearchFlowRun)
|
|
|
|
|
|
if kind:
|
|
|
|
|
|
query = query.where(ResearchFlowRun.authorization["kind"].as_string() == kind)
|
2026-09-08 22:16:05 +08:00
|
|
|
|
rows = await self.db.scalars(
|
2026-09-08 22:42:29 +08:00
|
|
|
|
query.order_by(ResearchFlowRun.created_at.desc()).limit(limit).offset(offset)
|
2026-09-08 22:16:05 +08:00
|
|
|
|
)
|
|
|
|
|
|
return {
|
|
|
|
|
|
"items": [await self.get(row.id) for row in rows],
|
2026-09-08 22:42:29 +08:00
|
|
|
|
"total": await self.db.scalar(select(func.count()).select_from(query.subquery())),
|
2026-09-08 22:16:05 +08:00
|
|
|
|
"limit": limit,
|
|
|
|
|
|
"offset": offset,
|
|
|
|
|
|
}
|