feat: run fixed research pipelines within durable budgets

This commit is contained in:
yuxuanhui
2026-09-08 22:16:05 +08:00
parent eb4850a003
commit 7860434b04
25 changed files with 2357 additions and 37 deletions
+6 -2
View File
@@ -121,14 +121,18 @@ class Experiments:
validation["status"] = "needs_review"
return validation
async def create(self, body, kind="template", extra_evidence=None):
async def create(self, body, kind="template", extra_evidence=None, *, parent_snapshots=None):
asset = await self.assets.get(body.asset_id, body.version, "template") if body.asset_id else None
template = TemplateSpec.model_validate(asset["content"]) if asset else body.template
scope = scope_of(body.settings)
if template.scope and template.scope.model_dump() != scope:
raise HTTPException(422, "模板适用范围与候选设置不同")
snapshots, fields = await self.inputs(body.input_ids, scope)
parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
parents = (
parent_snapshots
if parent_snapshots is not None
else await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
)
variables = {}
for name, variable in template.variables.items():
if variable.kind == "field":
+51
View File
@@ -16,6 +16,8 @@ from .workspace_contracts import (
Expansion,
ExperimentPreview,
FeatureConversion,
FlowControl,
FlowStart,
Generation,
ImportCommit,
ImportPreview,
@@ -225,3 +227,52 @@ async def research_lineage(
async with request.app.state.sessions() as db:
return await lineage(db, alpha_id, experiment_id, limit, offset)
@router.get("/flows/recipe")
async def fixed_recipe():
from .workflows import fixed_workflow
return fixed_workflow().model_dump(mode="json")
@router.post("/flows/runs", status_code=201)
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)
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)):
from .workflows import Workflows
async with request.app.state.sessions() as db:
return await Workflows(db).list(limit, offset)
@router.get("/flows/runs/{run_id}")
async def flow_run(run_id: str, request: Request):
from .workflows import Workflows
async with request.app.state.sessions() as db:
return await Workflows(db).get(run_id)
@router.post("/flows/runs/{run_id}/control")
async def control_flow(run_id: str, body: FlowControl, request: Request):
from .runtime import control
from .workflows import Workflows
async with request.app.state.sessions.begin() as db:
await control(db, run_id, body)
result = await Workflows(db).get(run_id)
request.app.state.research.wake.set()
request.app.state.runner.backtests.wake.set()
return result
+497
View File
@@ -0,0 +1,497 @@
"""Native research execution: durable intent, finite reservations, existing simulation lane.
Only this server-owned runner may use a saved flow authorization. Caller-supplied
Backtest source fields are provenance, never a grant to execute automatically.
"""
import asyncio
import logging
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.service import Backtests, uid
from ..models import Account, ResearchFlowRun, ResearchStepRun, now
from .assets import Assets
from .evaluations import Evaluations
from .experiments import Experiments
from .model import request_model
from .workflows import validate_graph
from .workspace_contracts import AssetWrite, EvaluateInput, Expansion, Generation, TemplateSpec, WorkflowSpec
logger = logging.getLogger(__name__)
ACTIVE = ("queued", "running")
DONE = ("completed", "skipped")
def changed(run):
run.updated_at = now()
run.version += 1
def halt(run, status, message):
run.status, run.error = status, message
changed(run)
def mark_model_attempt(step, status, error=None):
attempts = [dict(item) for item in step.output.get("model_attempts", [])]
if attempts:
attempts[-1] = {**attempts[-1], "status": status, "error": error}
step.output = {**step.output, "model_attempts": attempts}
async def locked_flow(db, run_id):
# All paths that may start/control a Backtest acquire these locks in this order.
await db.scalar(select(Account).where(Account.id == 1).with_for_update())
row = await db.scalar(select(ResearchFlowRun).where(ResearchFlowRun.id == run_id).with_for_update())
if not row:
raise HTTPException(404, "研究运行不存在")
return row
async def control(db, run_id, body):
run = await locked_flow(db, run_id)
if run.version != body.version:
raise HTTPException(409, "研究运行已变化,请刷新后重试")
if run.status in ("completed", "stopped"):
raise HTTPException(409, "已结束的研究不能恢复;请重新确认并启动新研究")
if body.action == "resume" and run.status == "budget_exhausted":
raise HTTPException(409, "预算已用尽;扩大预算需要重新确认并启动新研究")
steps = list(await db.scalars(select(ResearchStepRun).where(ResearchStepRun.run_id == run_id)))
if body.action == "resume" and any(step.status == "blocked" for step in steps):
raise HTTPException(409, "候选校验未通过,不能跳过此步骤;请修正输入或模板后重新确认研究")
for step in steps:
if step.backtest_run_id and step.status not in DONE:
backtest = await Backtests(db).run(step.backtest_run_id)
if backtest["status"] not in ("completed", "completed_with_errors", "stopped"):
if body.action != "resume" or backtest["control"] == "paused":
await Backtests(db).control(
step.backtest_run_id, ControlInput(action=body.action, version=backtest["version"])
)
if body.action == "resume" and step.status == "interrupted":
# Reservations already spent are never refunded; retry needs a fresh call budget.
step.status = "queued"
run.status = {"pause": "paused", "stop": "stopped", "resume": "queued"}[body.action]
run.error = None
changed(run)
class ResearchRuntime:
def __init__(self, sessions, ai, runner):
self.sessions, self.ai, self.runner = sessions, ai, runner
self.tasks = {}
self.locks = defaultdict(asyncio.Lock)
self.wake = asyncio.Event()
self.stopping = False
self.loop_task = None
async def recover(self):
async with self.sessions.begin() as db:
steps = list(await db.scalars(select(ResearchStepRun).where(ResearchStepRun.status == "running")))
for step in steps:
step.status, step.error = (
"interrupted",
"服务在模型步骤期间中断,已预留调用不退回;恢复后需要新的调用预算",
)
mark_model_attempt(step, "interrupted", step.error)
run = await db.get(ResearchFlowRun, step.run_id)
if run.status in ACTIVE:
halt(run, "interrupted", step.error)
# Durable preview/waiting steps reconcile their existing Backtest on the next tick.
async def start(self):
self.stopping = False
await self.recover()
self.loop_task = asyncio.create_task(self.loop())
async def stop(self):
self.stopping = True
self.wake.set()
if self.loop_task:
await self.loop_task
for task in self.tasks.values():
task.cancel()
await asyncio.gather(*self.tasks.values(), return_exceptions=True)
self.tasks.clear()
await self.recover()
async def loop(self):
while not self.stopping:
try:
await self.tick()
except (SQLAlchemyError, OSError):
logger.warning("Research runner waiting for database recovery")
self.wake.clear()
try:
await asyncio.wait_for(self.wake.wait(), timeout=0.5)
except TimeoutError:
pass
async def tick(self):
for key, task in list(self.tasks.items()):
if task.done():
self.tasks.pop(key)
try:
task.result()
except asyncio.CancelledError:
pass
except Exception:
logger.warning("Research step interrupted; durable state retained")
if self.stopping:
return
async with self.sessions() as db:
ids = list(
await db.scalars(
select(ResearchFlowRun.id)
.where(ResearchFlowRun.status.in_(ACTIVE))
.order_by(ResearchFlowRun.created_at)
.limit(20)
)
)
for run_id in ids:
if run_id not in self.tasks:
self.tasks[run_id] = asyncio.create_task(self.advance(run_id))
async def advance(self, run_id):
async with self.locks[run_id]:
try:
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)
await self.finish_model(run_id, step_id, result, evidence)
except HTTPException as exc:
await self.fail(run_id, str(exc.detail))
except asyncio.CancelledError:
raise
except Exception:
await self.fail(run_id, "研究步骤执行中断,产物已保留;请检查配置并恢复")
finally:
self.runner.backtests.wake.set()
async def fail(self, run_id, message):
async with self.sessions.begin() as db:
run = await locked_flow(db, run_id)
steps = list(
await db.scalars(
select(ResearchStepRun).where(
ResearchStepRun.run_id == run_id, ResearchStepRun.status == "running"
)
)
)
for step in steps:
step.status, step.error = "interrupted", message
mark_model_attempt(step, "interrupted", message)
if run.status in ACTIVE:
halt(run, "interrupted", message)
async def prepare(self, run_id):
async with self.sessions.begin() as db:
run = await locked_flow(db, run_id)
if run.status not in ACTIVE:
return
account = await db.get(Account, 1)
if account.wq_user_id != run.authorization["account_id"]:
raise HTTPException(409, "账户身份已变化,需重新确认研究授权")
graph = WorkflowSpec.model_validate(run.definition)
order = validate_graph(graph)
by_id = {n.id: n for n in graph.nodes}
rows = list(
await db.scalars(
select(ResearchStepRun).where(
ResearchStepRun.run_id == run.id, ResearchStepRun.round == run.round
)
)
)
steps = {s.node_id: s for s in rows}
if all(key in steps and steps[key].status in DONE for key in order):
loop = next(
(s for s in rows if by_id[s.node_id].type == "iterate" and s.status == "completed"), None
)
maximum = (
min(run.authorization["budget"]["max_rounds"], by_id[loop.node_id].config["max_rounds"])
if loop
else 1
)
if loop and run.round < maximum:
run.round += 1
changed(run)
else:
halt(run, "completed", None)
return
for key in order:
step = steps.get(key)
if step and step.status in DONE:
continue
node = by_id[key]
upstream_edges = [e for e in graph.edges if e.target == key]
if any(e.source not in steps or steps[e.source].status not in DONE for e in upstream_edges):
continue
if step and step.status == "running":
# An in-flight model step must only be completed by its owning worker.
return
if step is None:
step = ResearchStepRun(
id=uid(), run_id=run.id, node_id=key, round=run.round, status="queued", output={}
)
db.add(step)
await db.flush()
if step.status in ("interrupted", "blocked"):
halt(run, "needs_review" if step.status == "blocked" else "interrupted", step.error)
return
run.status = "running"
if step.status != "waiting":
changed(run)
upstream = [
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)
]
if upstream_edges and not upstream:
step.status = "skipped"
return
data = upstream[0] if upstream else {}
if node.type == "input":
previous = (
await db.scalar(
select(ResearchStepRun).where(
ResearchStepRun.run_id == run.id,
ResearchStepRun.round == run.round - 1,
ResearchStepRun.node_id.in_(
[n.id for n in graph.nodes if n.type == "iterate"]
),
)
)
if run.round > 1
else None
)
step.output = {
**(previous.output if previous else {}),
"type": "context",
"input_ids": run.authorization["input_ids"],
}
elif node.type == "generate":
reference = data.get("template") or (
run.authorization.get("template") if run.round == 1 and key == "generate" else None
)
if reference:
step.output = {
**data,
"type": "template",
"template": {"id": reference["id"], "version": reference["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", []),
)
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
elif node.type == "backtest":
await self.backtest(db, run, step, data)
return
elif node.type == "evaluate":
result = await Evaluations(db).create(
EvaluateInput(
experiment_id=data["experiment_id"],
backtest_run_id=data["backtest_run_id"],
rules=run.authorization["rules"],
)
)
step.output = {
"type": "evaluation",
"evaluation_id": result["id"],
"verdict": result["report"]["verdict"],
"experiment_id": data["experiment_id"],
"backtest_run_id": data["backtest_run_id"],
}
elif node.type == "iterate":
step.output = {
**data,
"type": "iteration",
"parent_experiment_ids": [data["experiment_id"]] if data.get("experiment_id") else [],
}
else:
raise HTTPException(422, "此研究节点尚未开放")
step.status, step.updated_at = "completed", now()
return
async def reserve_model(self, db, run, step, data, node):
if run.model_calls_used >= run.authorization["budget"]["max_model_calls"]:
halt(run, "budget_exhausted", "剩余模型调用预算不足,停止推进")
return
config = await self.ai.config(db)
if config.revision != run.model_revision:
raise HTTPException(409, "模型配置已变化,需重新确认研究运行")
generation = Generation(
name=run.name,
hypothesis=run.authorization["hypothesis"],
input_ids=run.authorization["input_ids"],
parent_experiment_ids=[data["experiment_id"]] if data.get("experiment_id") else [],
)
context = await Experiments(db).generation_context(generation)
context["parents"] = run.authorization["parents"] + context["parents"]
context["operators"] = run.authorization["operators_snapshot"]["content"]["items"][:100]
context["node_prompt"] = node.config.get("prompt", "")
if data.get("evaluation_id"):
report = (await Evaluations(db).get(data["evaluation_id"]))["report"]
context["evaluation"] = {
"rules": report["rules"],
"verdict": report["verdict"],
"records": [
{k: row.get(k) for k in ("client_item_id", "evidence", "missing", "failed")}
for row in report["records"][:100]
],
}
run.model_calls_used += 1
step.status = "running"
step.output = {
"model_attempts": step.output.get("model_attempts", [])
+ [
{
"reservation": run.model_calls_used,
"status": "running",
"reserved_at": now().isoformat(),
"model_revision": run.model_revision,
}
],
"type": "model_request",
"context": context,
"parent_experiment_ids": generation.parent_experiment_ids,
"reserved_call": run.model_calls_used,
"previous_error": step.error,
}
step.error = None
return step.id, context, run.model_revision
async def finish_model(self, run_id, step_id, result, evidence):
async with self.sessions.begin() as db:
run = await locked_flow(db, run_id)
step = await db.get(ResearchStepRun, step_id)
if not step or step.status != "running":
return
# 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")),
provenance={
"flow_run_id": run_id,
"step_id": step_id,
"generation": evidence,
"context": step.output["context"],
},
)
mark_model_attempt(step, "completed")
step.output = {
"model_attempts": step.output.get("model_attempts", []),
"type": "template",
"template": {"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()
changed(run)
async def backtest(self, db, run, step, data):
service = Backtests(db)
if step.backtest_run_id:
current = await service.run(step.backtest_run_id)
if current["status"] in ("needs_review", "stopped", "stopping"):
halt(run, "needs_review", "关联回测需要人工处理;未知提交不会重提")
elif current["status"] in ("completed", "completed_with_errors"):
step.status = "completed"
changed(run)
step.output = {**step.output, "type": "results"}
return
if not step.output.get("preview_id"):
ids = data.get("candidate_ids")
if not ids:
halt(run, "needs_review", "没有通过校验的候选可供回测")
return
preview = await Experiments(db).preview(data["experiment_id"], ids)
step.output = {
"type": "preview",
"preview_id": preview["preview_id"],
"version": preview["version"],
"digest": preview["digest"],
"experiment_id": data["experiment_id"],
"candidate_ids": ids,
}
step.status = "previewed"
return # The immutable preview commits before authorization and execution.
preview = await service.get_preview(step.output["preview_id"])
if (
preview["digest"] != step.output["digest"]
or preview["source"].get("research_id") != step.output["experiment_id"]
):
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"]):
raise HTTPException(403, "候选不属于此研究运行的固定输入范围")
if "backtest" not in run.authorization["methods"] or any(
c["settings"] != run.authorization["settings"]
for c in experiment["candidates"]
if c["client_item_id"] in step.output["candidate_ids"]
):
raise HTTPException(403, "候选方法或设置超出研究授权")
count = preview["total"]
if run.simulations_used + count > run.authorization["budget"]["max_simulations"]:
halt(run, "budget_exhausted", "剩余模拟条目预算不足以执行此固定预览")
return
run.simulations_used += count
current = await service.start(
StartInput(
preview_id=preview["preview_id"],
version=preview["version"],
idempotency_key=f"research:{step.id}",
)
)
step.backtest_run_id = current["backtest_run_id"]
step.output = {**step.output, "backtest_run_id": current["backtest_run_id"]}
step.status = "waiting"
+284
View File
@@ -0,0 +1,284 @@
"""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 fingerprint
from ..backtests.service import uid
from ..models import Account, ResearchFlowRun, ResearchStepRun
from .assets import Assets
from .experiments import Experiments, scope_of
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"}
if node.type in ("feature", "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 "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 not set(node.config.get("verdicts", ["pass"])).issubset({"pass", "review", "block"})
):
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)
inputs, _ = await experiments.inputs(body.input_ids, scope_of(body.settings))
parents = await experiments.parents(body.parent_alpha_ids, [])
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 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,
"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):
rows = await self.db.scalars(
select(ResearchFlowRun).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)),
"limit": limit,
"offset": offset,
}
+3 -3
View File
@@ -219,9 +219,9 @@ class WorkflowSpec(Contract):
class Budget(Contract):
max_rounds: int = Field(ge=1, le=100)
max_simulations: int = Field(ge=1, le=10000)
max_model_calls: int = Field(ge=1, le=1000)
max_rounds: int = Field(ge=1, le=100, strict=True)
max_simulations: int = Field(ge=1, le=10000, strict=True)
max_model_calls: int = Field(ge=1, le=1000, strict=True)
class FlowStart(Contract):
+44
View File
@@ -160,3 +160,47 @@ CAPABILITIES += (
handler=lambda ctx, args: Evaluations(ctx.business.db).create(args),
),
)
class FlowReference(Contract):
run_id: str = Field(min_length=1, max_length=36)
class FlowQuery(Contract):
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
async def read_flow(ctx, args):
from .workflows import Workflows
return await Workflows(ctx.business.db).get(args.run_id)
async def list_flows(ctx, args):
from .workflows import Workflows
return await Workflows(ctx.business.db).list(args.limit, args.offset)
CAPABILITIES += (
Capability(
name="get_research_run",
schema=FlowReference,
description="读取研究运行的固定授权、预算、阶段和产物,不能启动或扩大研究。",
label="读取研究运行",
renderer="research",
effect="query",
handler=read_flow,
),
Capability(
name="list_research_runs",
schema=FlowQuery,
description="分页查看已有研究运行。",
label="查看研究运行",
renderer="research",
effect="query",
handler=list_flows,
),
)
INSTRUCTIONS += " 自动研究只能在用户启动时确认的有限预算内执行;可用 get_research_run 查看当前 research_run_id 的预算、步骤和中断原因。普通 Chatbox 不授予自动研究执行权限。"