344 lines
15 KiB
Python
344 lines
15 KiB
Python
"""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 SimulationSettings, fingerprint
|
||
from ..backtests.service import uid
|
||
from ..models import Account, ResearchFlowRun, ResearchStepRun
|
||
from .assets import Assets
|
||
from .experiments import Experiments, scope_of, seed_settings
|
||
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", "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"}
|
||
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 "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
|
||
):
|
||
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 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 (
|
||
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)
|
||
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 = seed_settings(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))
|
||
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 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
|
||
)
|
||
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,
|
||
"node_assets": node_assets,
|
||
"allowed_settings": allowed_settings,
|
||
"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, kind=None):
|
||
query = select(ResearchFlowRun)
|
||
if kind:
|
||
query = query.where(ResearchFlowRun.authorization["kind"].as_string() == kind)
|
||
rows = await self.db.scalars(
|
||
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(query.subquery())),
|
||
"limit": limit,
|
||
"offset": offset,
|
||
}
|