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
+12 -1
View File
@@ -36,7 +36,18 @@ class ModelSettingsInput(Contract):
class PageContext(Contract):
page: Literal["alphas", "account", "datasets", "backtests", "operators", "templates", "features", "variants"] = "alphas"
page: Literal[
"alphas",
"account",
"datasets",
"backtests",
"operators",
"templates",
"features",
"variants",
"pipeline",
] = "alphas"
research_run_id: str | None = Field(default=None, max_length=36)
research_asset_id: str | None = Field(default=None, max_length=36)
research_experiment_id: str | None = Field(default=None, max_length=36)
catalog_scope: Scope | None = None
+6
View File
@@ -26,6 +26,7 @@ from .db import create_database
from .jobs import AUTH_KINDS, Runner, create_job
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession
from .research.routes import router as research_router
from .research.runtime import ResearchRuntime
from .schemas import (
AccountOutput,
AlphaDetail,
@@ -89,6 +90,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
engine, sessions = create_database(settings.database_url)
runner = Runner(sessions, settings, client=wq_client)
ai_runtime = AIRuntime(sessions, settings, runner, ai_model_factory)
research_runtime = ResearchRuntime(sessions, ai_runtime, runner)
@asynccontextmanager
async def lifespan(app):
@@ -100,7 +102,10 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
await ai_runtime.start()
if settings.enable_runner:
await runner.start()
await research_runtime.start()
yield
if settings.enable_runner:
await research_runtime.stop()
await ai_runtime.stop()
if settings.enable_runner:
await runner.stop()
@@ -117,6 +122,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
app.state.engine, app.state.sessions, app.state.runner = engine, sessions, runner
app.state.settings = settings
app.state.ai = ai_runtime
app.state.research = research_runtime
login_failures = defaultdict(list)
@app.exception_handler(RequestValidationError)
+35
View File
@@ -458,3 +458,38 @@ class ResearchParent(Base):
child_id: Mapped[str] = mapped_column(ForeignKey("research_experiments.id"), primary_key=True)
parent_kind: Mapped[str] = mapped_column(String(30), primary_key=True)
parent_id: Mapped[str] = mapped_column(String(100), primary_key=True, index=True)
class ResearchFlowRun(Base):
"""A user's finite authorization plus an immutable workflow and scope snapshot."""
__tablename__ = "research_flow_runs"
id: Mapped[str] = mapped_column(String(36), primary_key=True)
request_id: Mapped[str] = mapped_column(String(100), unique=True)
name: Mapped[str] = mapped_column(String(200))
definition: Mapped[dict] = mapped_column(JSON)
authorization: Mapped[dict] = mapped_column(JSON)
model_revision: Mapped[int | None] = mapped_column(Integer)
status: Mapped[str] = mapped_column(String(30), default="queued", index=True)
version: Mapped[int] = mapped_column(Integer, default=1)
round: Mapped[int] = mapped_column(Integer, default=1)
simulations_used: Mapped[int] = mapped_column(Integer, default=0)
model_calls_used: Mapped[int] = mapped_column(Integer, default=0)
error: Mapped[str | None] = mapped_column(Text)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
class ResearchStepRun(Base):
__tablename__ = "research_step_runs"
id: Mapped[str] = mapped_column(String(36), primary_key=True)
run_id: Mapped[str] = mapped_column(ForeignKey("research_flow_runs.id"), index=True)
node_id: Mapped[str] = mapped_column(String(100))
round: Mapped[int] = mapped_column(Integer)
status: Mapped[str] = mapped_column(String(30), default="running")
output: Mapped[dict] = mapped_column(JSON, default=dict)
backtest_run_id: Mapped[str | None] = mapped_column(ForeignKey("backtest_runs.id"))
error: Mapped[str | None] = mapped_column(Text)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
__table_args__ = (UniqueConstraint("run_id", "node_id", "round"),)
+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 不授予自动研究执行权限。"
@@ -0,0 +1,50 @@
"""Stage three: finite research authorizations and durable steps."""
import sqlalchemy as sa
from alembic import op
revision = "0008"
down_revision = "0007"
branch_labels = None
depends_on = None
def upgrade():
op.create_table(
"research_flow_runs",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("request_id", sa.String(100), nullable=False, unique=True),
sa.Column("name", sa.String(200), nullable=False),
sa.Column("definition", sa.JSON(), nullable=False),
sa.Column("authorization", sa.JSON(), nullable=False),
sa.Column("model_revision", sa.Integer(), nullable=True),
sa.Column("status", sa.String(30), nullable=False),
sa.Column("version", sa.Integer(), nullable=False),
sa.Column("round", sa.Integer(), nullable=False),
sa.Column("simulations_used", sa.Integer(), nullable=False),
sa.Column("model_calls_used", sa.Integer(), nullable=False),
sa.Column("error", sa.Text(), nullable=True),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
)
op.create_index("ix_research_flow_runs_status", "research_flow_runs", ["status"])
op.create_table(
"research_step_runs",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("run_id", sa.String(36), sa.ForeignKey("research_flow_runs.id"), nullable=False),
sa.Column("node_id", sa.String(100), nullable=False),
sa.Column("round", sa.Integer(), nullable=False),
sa.Column("status", sa.String(30), nullable=False),
sa.Column("output", sa.JSON(), nullable=False),
sa.Column("backtest_run_id", sa.String(36), sa.ForeignKey("backtest_runs.id"), nullable=True),
sa.Column("error", sa.Text(), nullable=True),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.UniqueConstraint("run_id", "node_id", "round"),
)
op.create_index("ix_research_step_runs_run_id", "research_step_runs", ["run_id"])
def downgrade():
op.drop_table("research_step_runs")
op.drop_table("research_flow_runs")
+35 -2
View File
@@ -5,7 +5,7 @@ import json
from contextlib import asynccontextmanager
from uuid import uuid4
from pydantic_ai.messages import ToolReturnPart, UserPromptPart
from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart, ToolReturnPart, UserPromptPart
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
from tests.research_fake import research_step
@@ -95,6 +95,39 @@ async def fake_stream(messages, info):
yield {0: DeltaToolCall(name=name, json_args=json.dumps(args), tool_call_id=uuid4().hex)}
def fake_structured(messages, info):
if not info.output_tools:
return ModelResponse(parts=[TextPart("READY")])
tool = info.output_tools[0]
context = next(
(json.loads(p.content) for m in reversed(messages) for p in m.parts if isinstance(p, UserPromptPart)),
{},
)
fields = [name for name, kind in context.get("fields", {}).items() if kind == "MATRIX"][:2]
template = {
"name": "合成流水线模板",
"description": "合成模型研究假设",
"expression": "rank({field})",
"variables": {
"field": {"kind": "field", "field_type": "MATRIX", "values": fields or ["TEST_FIN_001"]}
},
}
properties = tool.parameters_json_schema.get("properties", {})
if "summary" in properties:
data = {"summary": "合成评估建议", "risks": ["仅供验收"], "suggestions": ["继续核实缺失证据"]}
elif "input_ids" in properties:
data = {
"name": "合成特征方案",
"hypothesis": context.get("hypothesis", "合成假设"),
"input_ids": [i["id"] for i in context.get("inputs", [])],
"steps": [],
"template": template,
}
else:
data = template
return ModelResponse(parts=[ToolCallPart(tool.name, data)])
@asynccontextmanager
async def fake_model(config, settings):
yield FunctionModel(stream_function=fake_stream, model_name="test-model")
yield FunctionModel(function=fake_structured, stream_function=fake_stream, model_name="test-model")
+121
View File
@@ -0,0 +1,121 @@
"""Stage-three PostgreSQL acceptance in dedicated databases only."""
import asyncio
import os
import subprocess
from pathlib import Path
from alembic import command
from alembic.config import Config
from cryptography.fernet import Fernet
NAME = "wq_research_stage3_test"
RESTORE = "wq_research_restore_stage3"
os.environ.update(
DATABASE_URL=f"postgresql+asyncpg://postgres:research-test-only@127.0.0.1:18436/{NAME}",
ADMIN_PASSWORD="research-acceptance-only",
ENCRYPTION_KEY=Fernet.generate_key().decode(),
)
def docker(*args, **kwargs):
return subprocess.run(["docker", "exec", "-i", "wq-research-acceptance-pg", *args], check=True, **kwargs)
async def acceptance():
from unittest.mock import patch
import httpx
from sqlalchemy import select
from app.config import Settings
from app.main import create_app
from app.models import TemplateInput
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
from tests.test_backtests import setup
from tests.test_research_flows import begin, get, test_fixed_two_rounds_and_idempotent_start
from tests.test_research_workspace import template
app = create_app(Settings(_env_file=None, enable_runner=False, public_origin="http://testserver"))
async with app.router.lifespan_context(app):
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app), base_url="http://testserver", headers={"X-WQ-Request": "1"}
) as client:
assert (
await client.post(
"/api/v1/auth/login", json={"username": "admin", "password": "research-acceptance-only"}
)
).status_code == 200
# Configure a deterministic model; no provider or real platform network.
from tests.ai_fake import fake_model
app.state.ai.model_factory = fake_model
await configure(app, client)
platform, lane = await setup(app)
calls = []
async def model(ai, context, output_type, revision):
calls.append(context)
value = template()
value["expression"] = f"rank({{field}}) + {len(calls)}"
return TemplateSpec.model_validate(value), {
"model": "fixture",
"revision": revision,
"usage": {"requests": 1},
}
async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput))
body = {
"request_id": "finite-run",
"name": "PG 有限研究",
"input_ids": [fixed.id],
"hypothesis": "排名稳定性",
"settings": {"region": "USA", "universe": "TOP3000", "delay": 1},
"budget": {"max_rounds": 2, "max_simulations": 4, "max_model_calls": 3},
"batch_candidates": 2,
}
with patch("app.research.runtime.request_model", model):
await test_fixed_two_rounds_and_idempotent_start(app, client, (body, platform, lane, calls))
short = {
**body,
"request_id": "concurrent-budget",
"budget": {"max_rounds": 1, "max_simulations": 1, "max_model_calls": 2},
}
first, second = await asyncio.gather(begin(client, short), begin(client, short))
assert first["id"] == second["id"]
from app.research.runtime import ResearchRuntime
another = ResearchRuntime(app.state.sessions, app.state.ai, app.state.runner)
for _ in range(7):
await asyncio.gather(
app.state.research.advance(first["id"]), another.advance(first["id"])
)
result = await get(client, first["id"])
assert result["status"] == "budget_exhausted" and result["simulations_used"] == 0
assert result["model_calls_used"] == 1
print("PASS PostgreSQL: two-round execution, idempotent starts, concurrent reservations and budget gate")
if __name__ == "__main__":
docker("createdb", "-U", "postgres", NAME)
with Path("/tmp/wq-research-stage2.dump").open("rb") as source:
docker("pg_restore", "-U", "postgres", "-d", NAME, stdin=source)
config = Config("alembic.ini")
command.upgrade(config, "0008")
command.check(config)
asyncio.run(acceptance())
dump = Path("/tmp/wq-research-stage3.dump")
with dump.open("wb") as output:
docker("pg_dump", "-U", "postgres", "-Fc", NAME, stdout=output)
docker("createdb", "-U", "postgres", RESTORE)
with dump.open("rb") as source:
docker("pg_restore", "-U", "postgres", "-d", RESTORE, stdin=source)
query = "SELECT (SELECT count(*) FROM research_revisions),(SELECT count(*) FROM research_flow_runs),(SELECT count(*) FROM research_step_runs),(SELECT note FROM research WHERE alpha_id='OLD_RESEARCH')"
a = docker("psql", "-U", "postgres", "-d", NAME, "-Atc", query, capture_output=True).stdout
b = docker("psql", "-U", "postgres", "-d", RESTORE, "-Atc", query, capture_output=True).stdout
assert a == b
print(
"PASS PostgreSQL 17: 0007 → 0008 and pg_dump/pg_restore preserve flow budgets, steps and old research notes"
)
+257
View File
@@ -0,0 +1,257 @@
"""Fixed research: durable previews, budget grants, pause/stop and recovery."""
import asyncio
import pytest
from sqlalchemy import func, select
from app.models import Account, BacktestRun
from app.research.runtime import ResearchRuntime
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
from tests.test_backtests import execute, setup
from tests.test_research_workspace import catalog, research_input, template
__all__ = ["catalog", "research_input"]
@pytest.fixture
async def flow_setup(app, logged_in, research_input, monkeypatch):
await configure(app, logged_in)
platform, lane = await setup(app)
calls = []
async def model(ai, context, output_type, revision):
calls.append(context)
value = template()
value["expression"] = f"rank({{field}}) + {len(calls)}"
return TemplateSpec.model_validate(value), {
"model": "fixture",
"revision": revision,
"usage": {"requests": 1},
}
monkeypatch.setattr("app.research.runtime.request_model", model)
body = {
"request_id": "finite-run",
"name": "有限研究",
"input_ids": [research_input["id"]],
"hypothesis": "排名稳定性",
"settings": {"region": "USA", "universe": "TOP3000", "delay": 1},
"budget": {"max_rounds": 2, "max_simulations": 4, "max_model_calls": 3},
"batch_candidates": 2,
}
return body, platform, lane, calls
async def begin(client, body):
result = await client.post("/api/v1/research/flows/runs", json=body)
assert result.status_code == 201, result.text
return result.json()
async def get(client, run_id):
return (await client.get(f"/api/v1/research/flows/runs/{run_id}")).json()
async def drive(app, client, run_id, lane, ticks=45):
for _ in range(ticks):
await app.state.research.advance(run_id)
run = await get(client, run_id)
for step in run["steps"]:
if step["status"] == "waiting" and step["backtest_run_id"]:
await execute(app, lane, step["backtest_run_id"])
if run["status"] not in ("queued", "running"):
return run
raise AssertionError(await get(client, run_id))
async def test_fixed_two_rounds_and_idempotent_start(app, logged_in, flow_setup):
body, platform, lane, calls = flow_setup
first = await begin(logged_in, body)
assert (await begin(logged_in, body))["id"] == first["id"]
assert (
await logged_in.post("/api/v1/research/flows/runs", json={**body, "name": "different"})
).status_code == 409
result = await drive(app, logged_in, first["id"], lane)
assert result["status"] == "completed", result
assert result["round"] == 2 and result["simulations_used"] == 4 and result["model_calls_used"] == 3
assert len(result["steps"]) == 16 and len(calls) == 3 and len(platform.posts) == 2
assert all(s["output"].get("digest") for s in result["steps"] if s["node_id"] == "simulate")
assert len(calls[1]["evaluation"]["records"]) == 2
async def test_preview_commits_before_budget_gate_and_parallel_ticks(app, logged_in, flow_setup):
body, platform, _, _ = flow_setup
body["budget"]["max_simulations"] = 1
run = await begin(logged_in, body)
for _ in range(6):
await asyncio.gather(app.state.research.advance(run["id"]), app.state.research.advance(run["id"]))
result = await get(logged_in, run["id"])
assert result["status"] == "budget_exhausted", result
step = next(s for s in result["steps"] if s["node_id"] == "simulate")
assert step["status"] == "previewed" and step["output"]["preview_id"]
assert result["simulations_used"] == 0 and not platform.posts
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
async def test_pause_stop_keep_known_simulation_and_block_next_steps(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
run = await begin(logged_in, body)
for _ in range(5):
await app.state.research.advance(run["id"])
current = await get(logged_in, run["id"])
simulation = next(s for s in current["steps"] if s["node_id"] == "simulate")
assert simulation["backtest_run_id"]
# Issue remote simulation first. Stop must continue collecting it.
async with app.state.sessions() as db:
from app.models import SimulationAttempt
aid = await db.scalar(
select(SimulationAttempt.id).where(SimulationAttempt.run_id == simulation["backtest_run_id"])
)
await lane.step(aid)
for action in ("pause", "stop"):
current = await get(logged_in, run["id"])
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": action, "version": current["version"]},
)
assert response.status_code == 200, response.text
await app.state.research.advance(run["id"])
await lane.step(aid)
result = (await logged_in.get(f"/api/v1/backtests/runs/{simulation['backtest_run_id']}/results")).json()
assert all(i["persistence_status"] == "saved" for i in result["items"])
final = await get(logged_in, run["id"])
assert final["status"] == "stopped" and len(final["steps"]) == 4 and len(platform.posts) == 1
async def test_recovery_marks_model_interrupted_without_refund(app, logged_in, flow_setup):
body, _, _, calls = flow_setup
run = await begin(logged_in, body)
await app.state.research.advance(run["id"])
work = await app.state.research.prepare(run["id"])
assert work and not calls
recovered = ResearchRuntime(app.state.sessions, app.state.ai, app.state.runner)
await recovered.recover()
current = await get(logged_in, run["id"])
assert current["status"] == "interrupted" and current["model_calls_used"] == 1
await recovered.advance(run["id"])
assert not calls
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": "resume", "version": current["version"]},
)
assert response.status_code == 200
await recovered.advance(run["id"])
current = await get(logged_in, run["id"])
assert current["model_calls_used"] == 2 and len(calls) == 1
assert [a["status"] for a in current["steps"][-1]["output"]["model_attempts"]] == [
"interrupted",
"completed",
]
@pytest.mark.parametrize(
"key,value",
[
("max_rounds", 0),
("max_simulations", -1),
("max_model_calls", 0),
("max_rounds", None),
("max_rounds", True),
("max_model_calls", 1.5),
],
)
async def test_finite_positive_budgets(logged_in, flow_setup, key, value):
body, _, _, _ = flow_setup
body["budget"][key] = value
assert (await logged_in.post("/api/v1/research/flows/runs", json=body)).status_code == 422
async def test_changed_account_stops_authorized_execution(app, logged_in, flow_setup):
body, platform, _, calls = flow_setup
run = await begin(logged_in, body)
async with app.state.sessions.begin() as db:
account = await db.get(Account, 1)
account.wq_user_id = "different"
await app.state.research.advance(run["id"])
assert (await get(logged_in, run["id"]))["status"] == "interrupted"
assert not platform.posts and not calls
async def test_unknown_submission_retains_budget_and_is_never_reposted(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
platform.reject = "unknown"
run = await begin(logged_in, body)
result = await drive(app, logged_in, run["id"], lane)
assert result["status"] == "needs_review" and result["simulations_used"] == 2
assert len(platform.posts) == 1
for _ in range(3):
await app.state.research.advance(run["id"])
await app.state.research.recover()
assert len(platform.posts) == 1
async def test_source_fields_do_not_create_automatic_grant(app, logged_in, research_input):
from app.ai.tools import CAPABILITIES
from app.models import ResearchFlowRun
from tests.test_backtests import candidate, preview
await preview(logged_in, [candidate()])
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(ResearchFlowRun)) == 0
assert CAPABILITIES["start_backtest"].requires_confirmation
assert not any("start_flow" in key for key in CAPABILITIES)
async def test_invalid_final_candidates_cannot_resume_into_completion(
app, logged_in, flow_setup, monkeypatch
):
body, _, lane, _ = flow_setup
body["budget"]["max_rounds"] = 1
run = await begin(logged_in, body)
for _ in range(6):
await app.state.research.advance(run["id"])
current = await get(logged_in, run["id"])
for step in current["steps"]:
if step["status"] == "waiting":
await execute(app, lane, step["backtest_run_id"])
async def invalid(ai, context, output_type, revision):
value = template()
value["expression"] = "unknown_operator({field})"
return TemplateSpec.model_validate(value), {"model": "fixture", "revision": revision}
monkeypatch.setattr("app.research.runtime.request_model", invalid)
current = await drive(app, logged_in, run["id"], lane)
assert current["status"] == "needs_review", current
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": "resume", "version": current["version"]},
)
assert response.status_code == 409
await app.state.research.advance(run["id"])
assert (await get(logged_in, run["id"]))["status"] == "needs_review"
async def test_restart_reuses_preview_and_known_backtest(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
run = await begin(logged_in, body)
for _ in range(4):
await app.state.research.advance(run["id"])
before = await get(logged_in, run["id"])
preview_id = before["steps"][-1]["output"]["preview_id"]
runtime = ResearchRuntime(app.state.sessions, app.state.ai, app.state.runner)
await runtime.recover()
await runtime.advance(run["id"])
current = await get(logged_in, run["id"])
backtest_id = current["steps"][-1]["backtest_run_id"]
assert current["steps"][-1]["output"]["preview_id"] == preview_id
await runtime.recover()
await execute(app, lane, backtest_id)
await runtime.advance(run["id"])
after = await get(logged_in, run["id"])
assert after["simulations_used"] == 2 and len(platform.posts) == 1
assert next(s for s in after["steps"] if s["node_id"] == "simulate")["backtest_run_id"] == backtest_id