From 7860434b04f177d76d06f9fc3bf4d91085c0fb08 Mon Sep 17 00:00:00 2001 From: yuxuanhui Date: Tue, 8 Sep 2026 22:16:05 +0800 Subject: [PATCH] feat: run fixed research pipelines within durable budgets --- .../research-migration/acceptance/stage-3.md | 39 ++ .../issues/01-implementation.md | 4 +- backend/app/ai/contracts.py | 13 +- backend/app/main.py | 6 + backend/app/models.py | 35 ++ backend/app/research/experiments.py | 8 +- backend/app/research/routes.py | 51 ++ backend/app/research/runtime.py | 497 ++++++++++++++++++ backend/app/research/workflows.py | 284 ++++++++++ backend/app/research/workspace_contracts.py | 6 +- backend/app/research/workspace_tools.py | 44 ++ .../migrations/versions/0008_research_runs.py | 50 ++ backend/tests/ai_fake.py | 37 +- backend/tests/research_flows_postgres.py | 121 +++++ backend/tests/test_research_flows.py | 257 +++++++++ frontend/src/App.tsx | 20 +- frontend/src/ai/types.ts | 4 +- frontend/src/ai/workspace.ts | 2 + frontend/src/components/AppSidebar.tsx | 3 +- frontend/src/research/FlowLaunchForm.tsx | 306 +++++++++++ frontend/src/research/FlowRunView.tsx | 195 +++++++ frontend/src/research/PipelinePage.tsx | 125 +++++ frontend/src/research/ResearchToolCard.tsx | 52 +- frontend/src/research/flowTypes.ts | 78 +++ frontend/tests/research-flows.spec.ts | 157 ++++++ 25 files changed, 2357 insertions(+), 37 deletions(-) create mode 100644 .scratch/research-migration/acceptance/stage-3.md create mode 100644 backend/app/research/runtime.py create mode 100644 backend/app/research/workflows.py create mode 100644 backend/migrations/versions/0008_research_runs.py create mode 100644 backend/tests/research_flows_postgres.py create mode 100644 backend/tests/test_research_flows.py create mode 100644 frontend/src/research/FlowLaunchForm.tsx create mode 100644 frontend/src/research/FlowRunView.tsx create mode 100644 frontend/src/research/PipelinePage.tsx create mode 100644 frontend/src/research/flowTypes.ts create mode 100644 frontend/tests/research-flows.spec.ts diff --git a/.scratch/research-migration/acceptance/stage-3.md b/.scratch/research-migration/acceptance/stage-3.md new file mode 100644 index 0000000..26eb6ed --- /dev/null +++ b/.scratch/research-migration/acceptance/stage-3.md @@ -0,0 +1,39 @@ +# 第三阶段验收:固定自动研究 + +日期:2026-09-08。范围:固定研究配方、有限授权、运行和步骤持久化、候选池、暂停停止与恢复。未调用真实模型或 WorldQuant。 + +## 交付 + +- 新增研究流水线菜单和运行页面。六阶段配方按生成、校验设参、回测、评估、增强、重新展开执行,并保存固定输入与轮次节点。 +- 启动弹窗确认输入范围、初始模板版本、种子、假设、设置和三项预算。服务端保存不可变授权、账户身份、模型配置版本、配方、规则及元数据快照。 +- 模型调用、模拟条目先预留预算;预算均为有限正整数。没有足够预算执行下一步时停止,保留已产生的模板、候选和预览。 +- 回测预览先单独提交,再核验授权后使用原 Backtests.start 和调度 lane。幂等键固定到步骤,运行和回测关联在同一事务保存。 +- 暂停停止阻止后续步骤,已提交模拟继续收集。服务重启保留预览和已知回测;未完成模型请求标为中断,恢复需要新的调用预算,历史调用尝试不被覆盖。 +- 空候选校验阻塞不能通过恢复跳过。未知模拟提交保留 needs_review,不能自动重提。扩大预算或更换范围、账户身份、模型配置需重新确认新研究。 +- 普通 Chatbox 仅新增读取研究运行和预算的工具;原 start_backtest 仍需固定集合用户确认,来源字段不是自动执行授权。 + +## 验收结果 + +| 步骤 | 结果 | +|---|---| +| 两轮固定研究闭环 | 通过;16 个阶段步骤、4 条模拟、3 次模型调用,保留评估证据及增强候选 | +| 重复启动与冲突请求 | 通过;相同请求返回同一研究;请求标识复用到不同内容返回 409 | +| 预览先提交、模拟预算不足 | 通过;预览保留,模拟使用数为 0,没有平台 POST | +| 同进程并行推进、PostgreSQL 两执行器并发推进 | 通过;模型预算只预留一次,模拟预算不超限 | +| 暂停与停止已发模拟 | 通过;阻止后续评估/增强,已发条目仍保存结果 | +| 模型步骤重启恢复 | 通过;标记中断、不退回已预留预算,恢复保留中断和完成两次尝试记录 | +| 预览后重启、已知回测后重启 | 通过;复用原 preview_id、backtest_run_id,模拟只发送一次 | +| 未知模拟提交 | 通过;保留 needs_review 和已用预算,重复推进及恢复不重发 | +| 无效/布尔/小数预算、账户身份变化 | 通过;非法预算 422,身份变化中断执行,不发模型或模拟请求 | +| 最后一次展开全部无效后恢复 | 通过;409 拒绝跳过校验阻塞,不误报研究完成 | +| 后端全量回归 | 200 passed;后续新增及调整恢复、严格预算、AI 读取等专项最终 36 passed。日志 `/tmp/wq-stage-three-full-backend.log`、`/tmp/wq-stage-three-final-tests.log` | +| 前端构建与类型检查 | 通过;保留原有 lottie-web 构建警告 | +| 浏览器授权→执行→刷新恢复 | 通过;一轮执行 2 条模拟和 2 次模型调用,8 个步骤均完成。截图 `/tmp/wq-stage-three-run.png` 已检查 | +| 侧栏与 AI 浏览器回归 | 5 项通过;流水线用例的下拉动画层造成弹窗定位歧义,限定授权弹窗后复验 1 passed。日志 `/tmp/wq-stage-three-browser-regression.log`、`/tmp/wq-stage-three-browser-final.log` | +| PostgreSQL 17 增量迁移和恢复 | 0007→0008、Alembic check、两轮闭环、并发预算及 pg_dump/pg_restore 通过。旧研究备注保留。脚本 `backend/tests/research_flows_postgres.py`,日志 `/tmp/wq-stage-three-postgres.log` | + +独立只读核验没有发现预算超支或重启重复模拟问题;其指出的空候选恢复问题已修正并增加回归测试。自定义工作流入口尚未展示,留待第四阶段开放。 + +## 仍需联调 + +真实平台协议和真实模拟仍待此前授权问题得到当前对话确认;此处所有模型与平台均为合成响应,不作为真实平台验收记录。模型配置必须已测试启用,固定流水线使用现有 REGULAR / FASTEXPR / EQUITY 和单账户、单后端执行边界。 diff --git a/.scratch/research-migration/issues/01-implementation.md b/.scratch/research-migration/issues/01-implementation.md index 1e365c5..dbb157b 100644 --- a/.scratch/research-migration/issues/01-implementation.md +++ b/.scratch/research-migration/issues/01-implementation.md @@ -6,7 +6,7 @@ Status: ready-for-agent - [x] 一:算子、模板、表达式模块、两类变体与回测来源闭环。 - [x] 二:特征方案、保存视图、关系、比较和评估。 -- [ ] 三:固定自动研究、预算、持久步骤、恢复。 +- [x] 三:固定自动研究、预算、持久步骤、恢复。 - [ ] 四:原生 QuantFlow 画布与共用执行。 - [ ] 验证:后端、前端、浏览器、PostgreSQL 迁移与恢复。 @@ -17,3 +17,5 @@ Status: ready-for-agent 第一阶段已通过本地后端、浏览器和 PostgreSQL 验收;报告见 `../acceptance/stage-1.md`。真实 WorldQuant 联调因自动审批未认可已有授权而等待当前对话确认,不将合成测试记为真实联调。 第二阶段已通过后端、前端、浏览器和 PostgreSQL 迁移/恢复验证;报告见 `../acceptance/stage-2.md`。独立 Alpha 列表布局改动继续留在工作区,本次仅提交保存视图接入。 + +第三阶段固定研究及有限授权已验收,报告见 `../acceptance/stage-3.md`。 diff --git a/backend/app/ai/contracts.py b/backend/app/ai/contracts.py index 6e16a97..1a01090 100644 --- a/backend/app/ai/contracts.py +++ b/backend/app/ai/contracts.py @@ -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 diff --git a/backend/app/main.py b/backend/app/main.py index 84fa50c..600bc29 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -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) diff --git a/backend/app/models.py b/backend/app/models.py index 09b70e4..f53c322 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -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"),) diff --git a/backend/app/research/experiments.py b/backend/app/research/experiments.py index 38ce18b..d8c96b3 100644 --- a/backend/app/research/experiments.py +++ b/backend/app/research/experiments.py @@ -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": diff --git a/backend/app/research/routes.py b/backend/app/research/routes.py index e165a76..9741b98 100644 --- a/backend/app/research/routes.py +++ b/backend/app/research/routes.py @@ -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 diff --git a/backend/app/research/runtime.py b/backend/app/research/runtime.py new file mode 100644 index 0000000..4d315ac --- /dev/null +++ b/backend/app/research/runtime.py @@ -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" diff --git a/backend/app/research/workflows.py b/backend/app/research/workflows.py new file mode 100644 index 0000000..9c749ff --- /dev/null +++ b/backend/app/research/workflows.py @@ -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, + } diff --git a/backend/app/research/workspace_contracts.py b/backend/app/research/workspace_contracts.py index 826d352..4ec2530 100644 --- a/backend/app/research/workspace_contracts.py +++ b/backend/app/research/workspace_contracts.py @@ -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): diff --git a/backend/app/research/workspace_tools.py b/backend/app/research/workspace_tools.py index 3931a87..39008fe 100644 --- a/backend/app/research/workspace_tools.py +++ b/backend/app/research/workspace_tools.py @@ -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 不授予自动研究执行权限。" diff --git a/backend/migrations/versions/0008_research_runs.py b/backend/migrations/versions/0008_research_runs.py new file mode 100644 index 0000000..6d11196 --- /dev/null +++ b/backend/migrations/versions/0008_research_runs.py @@ -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") diff --git a/backend/tests/ai_fake.py b/backend/tests/ai_fake.py index 12fa969..3f7518a 100644 --- a/backend/tests/ai_fake.py +++ b/backend/tests/ai_fake.py @@ -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") diff --git a/backend/tests/research_flows_postgres.py b/backend/tests/research_flows_postgres.py new file mode 100644 index 0000000..1935ccb --- /dev/null +++ b/backend/tests/research_flows_postgres.py @@ -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" + ) diff --git a/backend/tests/test_research_flows.py b/backend/tests/test_research_flows.py new file mode 100644 index 0000000..655fe79 --- /dev/null +++ b/backend/tests/test_research_flows.py @@ -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 diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index 3309397..f6d912e 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -12,6 +12,7 @@ import { } from "@douyinfe/semi-ui-19"; import zhCN from "@douyinfe/semi-ui-19/lib/es/locale/source/zh_CN"; import { IconSetting, IconComment, IconHistory } from "@douyinfe/semi-icons"; +import { PipelinePage } from "./research/PipelinePage"; import { FeaturesPage } from "./research/FeaturesPage"; import { ResearchWorkspace } from "./research/ResearchWorkspace"; import { OperatorsPage } from "./research/OperatorsPage"; @@ -50,7 +51,11 @@ export default function App() { page: "templates", }); useEffect(() => { - if (["operators", "templates", "variants", "features"].includes(page)) + if ( + ["operators", "templates", "variants", "features", "pipeline"].includes( + page, + ) + ) setVisitedResearch((old) => (old.includes(page) ? old : [...old, page])); }, [page]); const [catalogModal, setCatalogModal] = useState(false); @@ -371,6 +376,15 @@ export default function App() { /> )} + {visitedResearch.includes("pipeline") && ( + + )} {visitedResearch.includes("features") && ( )}