feat: run fixed research pipelines within durable budgets
This commit is contained in:
@@ -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 和单账户、单后端执行边界。
|
||||
@@ -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`。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"),)
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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):
|
||||
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
+19
-1
@@ -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() {
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
{visitedResearch.includes("pipeline") && (
|
||||
<div className="alpha-page-view" hidden={page !== "pipeline"}>
|
||||
<PipelinePage
|
||||
active={page === "pipeline"}
|
||||
onAction={handleAction}
|
||||
onContext={setResearchContext}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
{visitedResearch.includes("features") && (
|
||||
<div className="alpha-page-view" hidden={page !== "features"}>
|
||||
<FeaturesPage
|
||||
@@ -448,6 +462,10 @@ export default function App() {
|
||||
onClose={() => setChatOpen(false)}
|
||||
context={
|
||||
{
|
||||
pipeline:
|
||||
researchContext.page === "pipeline"
|
||||
? researchContext
|
||||
: { page: "pipeline" as const },
|
||||
features:
|
||||
researchContext.page === "features"
|
||||
? researchContext
|
||||
|
||||
@@ -19,7 +19,9 @@ export type PageContext = {
|
||||
| "operators"
|
||||
| "templates"
|
||||
| "variants"
|
||||
| "features";
|
||||
| "features"
|
||||
| "pipeline";
|
||||
research_run_id?: string | null;
|
||||
research_asset_id?: string;
|
||||
research_experiment_id?: string;
|
||||
catalog_scope?: {
|
||||
|
||||
@@ -7,6 +7,8 @@ const contextLabels: Record<WorkspacePage, (context: PageContext) => string> = {
|
||||
operators: () => "上下文:算子库",
|
||||
templates: (context) =>
|
||||
`上下文:模板工坊${context.research_asset_id ? ` · ${context.research_asset_id}` : ""}${context.research_experiment_id ? ` · 实验 ${context.research_experiment_id}` : ""}`,
|
||||
pipeline: (context) =>
|
||||
`上下文:研究流水线${context.research_run_id ? ` · ${context.research_run_id}` : ""}`,
|
||||
features: () => "上下文:特征工程",
|
||||
variants: () => "上下文:Alpha 变体",
|
||||
account: () => "上下文:个人信息页",
|
||||
|
||||
@@ -27,6 +27,7 @@ const navigation = [
|
||||
{ id: "features", label: "特征工程", icon: IconBeaker, group: "研究实验" },
|
||||
{ id: "variants", label: "Alpha 变体", icon: IconBeaker, group: "研究实验" },
|
||||
{ id: "backtests", label: "回测研究", icon: IconBeaker, group: "研究实验" },
|
||||
{ id: "pipeline", label: "研究流水线", icon: IconBeaker, group: "研究编排" },
|
||||
{ id: "alphas", label: "Alpha 管理", icon: IconGridView, group: "研究成果" },
|
||||
{ id: "account", label: "个人信息", icon: IconUser, group: "" },
|
||||
] as const;
|
||||
@@ -170,7 +171,7 @@ export function AppSidebar({
|
||||
</div>
|
||||
)}
|
||||
<nav className="sidebar-nav" aria-label="主导航">
|
||||
{["数据与素材", "研究实验", "研究成果"].map((group) => (
|
||||
{["数据与素材", "研究实验", "研究编排", "研究成果"].map((group) => (
|
||||
<div key={group}>
|
||||
{!collapsed && (
|
||||
<div className="sidebar-group-heading">{group}</div>
|
||||
|
||||
@@ -0,0 +1,306 @@
|
||||
import { useEffect, useState } from "react";
|
||||
import {
|
||||
Banner,
|
||||
Button,
|
||||
Input,
|
||||
InputNumber,
|
||||
Modal,
|
||||
TextArea,
|
||||
Toast,
|
||||
} from "@douyinfe/semi-ui-19";
|
||||
import { api, post } from "../api";
|
||||
import type { InputSnapshot, Asset } from "./workspaceTypes";
|
||||
import type { FlowLaunch, FlowRun, Budget } from "./flowTypes";
|
||||
import type { SimulationSettings } from "../backtests/types";
|
||||
import { ResearchSelect } from "./ResearchSelect";
|
||||
export function FlowLaunchForm({
|
||||
onStarted,
|
||||
workflow,
|
||||
}: {
|
||||
onStarted: (r: FlowRun) => void;
|
||||
workflow?: { id: string; version: number; name: string };
|
||||
}) {
|
||||
const [inputs, setInputs] = useState<InputSnapshot[]>([]);
|
||||
const [templates, setTemplates] = useState<Asset[]>([]);
|
||||
const [ids, setIds] = useState<string[]>([]);
|
||||
const [template, setTemplate] = useState<string>();
|
||||
const [name, setName] = useState("固定研究流水线");
|
||||
const [hypothesis, setHypothesis] = useState("");
|
||||
const [parent, setParent] = useState("");
|
||||
const [budget, setBudget] = useState<Budget>({
|
||||
max_rounds: 3,
|
||||
max_simulations: 24,
|
||||
max_model_calls: 5,
|
||||
});
|
||||
const [batch, setBatch] = useState(8);
|
||||
const [seed, setSeed] = useState(0);
|
||||
const [confirmation, setConfirmation] = useState<FlowLaunch | null>(null);
|
||||
const [busy, setBusy] = useState(false);
|
||||
const [settings, setSettings] = useState<SimulationSettings>({
|
||||
instrumentType: "EQUITY",
|
||||
region: "USA",
|
||||
universe: "TOP3000",
|
||||
delay: 1,
|
||||
decay: 0,
|
||||
neutralization: "INDUSTRY",
|
||||
truncation: 0.08,
|
||||
pasteurization: "ON",
|
||||
unitHandling: "VERIFY",
|
||||
nanHandling: "OFF",
|
||||
language: "FASTEXPR",
|
||||
visualization: false,
|
||||
maxTrade: "OFF",
|
||||
});
|
||||
useEffect(() => {
|
||||
const c = new AbortController();
|
||||
Promise.all([
|
||||
api<{ items: InputSnapshot[] }>("/research/inputs", { signal: c.signal }),
|
||||
api<{ items: Asset[] }>("/research/assets?kind=template&limit=100", {
|
||||
signal: c.signal,
|
||||
}),
|
||||
])
|
||||
.then(([i, t]) => {
|
||||
setInputs(i.items);
|
||||
setTemplates(t.items);
|
||||
})
|
||||
.catch((e) => {
|
||||
if (!c.signal.aborted) Toast.error(e.message);
|
||||
});
|
||||
return () => c.abort();
|
||||
}, []);
|
||||
function preview() {
|
||||
const fixed = templates.find((t) => t.id === template);
|
||||
setConfirmation({
|
||||
request_id: crypto.randomUUID(),
|
||||
name,
|
||||
input_ids: ids,
|
||||
hypothesis,
|
||||
settings,
|
||||
budget,
|
||||
batch_candidates: batch,
|
||||
seed,
|
||||
parent_alpha_ids: parent.split(/[,,\s]+/).filter(Boolean),
|
||||
...(fixed
|
||||
? { template_id: fixed.id, template_version: fixed.version }
|
||||
: {}),
|
||||
...(workflow
|
||||
? { workflow_id: workflow.id, workflow_version: workflow.version }
|
||||
: {}),
|
||||
});
|
||||
}
|
||||
async function start() {
|
||||
if (!confirmation) return;
|
||||
setBusy(true);
|
||||
try {
|
||||
const run = await post<FlowRun>("/research/flows/runs", confirmation);
|
||||
setConfirmation(null);
|
||||
onStarted(run);
|
||||
} catch (e) {
|
||||
Toast.error((e as Error).message);
|
||||
} finally {
|
||||
setBusy(false);
|
||||
}
|
||||
}
|
||||
return (
|
||||
<section className="research-card">
|
||||
<h3>
|
||||
{workflow
|
||||
? `启动 ${workflow.name} v${workflow.version}`
|
||||
: "新建固定研究"}
|
||||
</h3>
|
||||
<div className="research-form-grid">
|
||||
<label>
|
||||
研究名称
|
||||
<Input aria-label="自动研究名称" value={name} onChange={setName} />
|
||||
</label>
|
||||
<label>
|
||||
种子 Alpha(可选)
|
||||
<Input
|
||||
aria-label="自动研究种子"
|
||||
value={parent}
|
||||
onChange={setParent}
|
||||
placeholder="多个 ID 用逗号分隔"
|
||||
/>
|
||||
</label>
|
||||
</div>
|
||||
<label>
|
||||
研究假设
|
||||
<TextArea
|
||||
aria-label="自动研究假设"
|
||||
value={hypothesis}
|
||||
onChange={setHypothesis}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
固定数据范围
|
||||
<ResearchSelect
|
||||
label="自动研究固定输入"
|
||||
multiple
|
||||
filter
|
||||
value={ids}
|
||||
optionList={inputs.map((i) => ({
|
||||
value: i.id,
|
||||
label: `${i.dataset_id} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.id.slice(0, 8)}`,
|
||||
}))}
|
||||
onChange={(v) => {
|
||||
const next = v as string[];
|
||||
setIds(next);
|
||||
const first = inputs.find((i) => i.id === next[0]);
|
||||
if (first)
|
||||
setSettings((s) => ({
|
||||
...s,
|
||||
region: first.scope.region,
|
||||
universe: first.scope.universe,
|
||||
delay: first.scope.delay,
|
||||
}));
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
初始模板版本(可选)
|
||||
<ResearchSelect
|
||||
label="自动研究初始模板"
|
||||
filter
|
||||
showClear
|
||||
value={template}
|
||||
optionList={templates.map((t) => ({
|
||||
value: t.id,
|
||||
label: `${t.name} · v${t.version}`,
|
||||
}))}
|
||||
onChange={(v) => setTemplate(v ? String(v) : undefined)}
|
||||
/>
|
||||
</label>
|
||||
<p>
|
||||
{settings.region} / {settings.universe} / D{settings.delay} · REGULAR /
|
||||
FASTEXPR / EQUITY
|
||||
</p>
|
||||
<div className="research-form-grid">
|
||||
<label>
|
||||
中性化
|
||||
<Input
|
||||
aria-label="自动研究中性化"
|
||||
value={settings.neutralization}
|
||||
onChange={(neutralization) =>
|
||||
setSettings({ ...settings, neutralization })
|
||||
}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
Decay
|
||||
<InputNumber
|
||||
aria-label="自动研究 Decay"
|
||||
min={0}
|
||||
value={settings.decay}
|
||||
onChange={(v) => {
|
||||
if (typeof v === "number") setSettings({ ...settings, decay: v });
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
Truncation
|
||||
<InputNumber
|
||||
aria-label="自动研究 Truncation"
|
||||
min={0}
|
||||
max={1}
|
||||
step={0.01}
|
||||
value={settings.truncation}
|
||||
onChange={(v) => {
|
||||
if (typeof v === "number")
|
||||
setSettings({ ...settings, truncation: v });
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
</div>
|
||||
<h4>本次授权预算</h4>
|
||||
<div className="research-form-grid">
|
||||
{(
|
||||
[
|
||||
["max_rounds", "最大轮数"],
|
||||
["max_simulations", "最大模拟条目数"],
|
||||
["max_model_calls", "最大模型调用数"],
|
||||
] as const
|
||||
).map(([key, label]) => (
|
||||
<label key={key}>
|
||||
{label}
|
||||
<InputNumber
|
||||
aria-label={label}
|
||||
min={1}
|
||||
precision={0}
|
||||
value={budget[key]}
|
||||
onChange={(v) => {
|
||||
if (typeof v === "number") setBudget({ ...budget, [key]: v });
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
))}
|
||||
<label>
|
||||
每轮候选上限
|
||||
<InputNumber
|
||||
aria-label="每轮候选上限"
|
||||
min={1}
|
||||
max={100}
|
||||
precision={0}
|
||||
value={batch}
|
||||
onChange={(v) => {
|
||||
if (typeof v === "number") setBatch(v);
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
采样种子
|
||||
<InputNumber
|
||||
aria-label="自动研究采样种子"
|
||||
precision={0}
|
||||
value={seed}
|
||||
onChange={(v) => {
|
||||
if (typeof v === "number") setSeed(v);
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
</div>
|
||||
<Banner
|
||||
type="info"
|
||||
description="启动后,在确认的有限预算内自动生成、校验、回测、评估和增强。超出预算即停止推进;扩大范围或预算需要重新确认。"
|
||||
/>
|
||||
<Button
|
||||
theme="solid"
|
||||
disabled={!ids.length || !hypothesis.trim() || !name.trim()}
|
||||
onClick={preview}
|
||||
>
|
||||
核对并启动研究
|
||||
</Button>
|
||||
<Modal
|
||||
title="确认自动研究授权"
|
||||
visible={Boolean(confirmation)}
|
||||
onCancel={() => {
|
||||
if (!busy) setConfirmation(null);
|
||||
}}
|
||||
onOk={() => void start()}
|
||||
confirmLoading={busy}
|
||||
okText="确认并开始自动研究"
|
||||
>
|
||||
{confirmation && (
|
||||
<>
|
||||
<p>
|
||||
{confirmation.name}:{confirmation.hypothesis}
|
||||
</p>
|
||||
<p>
|
||||
{confirmation.input_ids.length} 个固定输入 · {settings.region}/
|
||||
{settings.universe}/D{settings.delay}
|
||||
</p>
|
||||
<p>
|
||||
最多 {confirmation.budget.max_rounds} 轮、
|
||||
{confirmation.budget.max_simulations} 条模拟、
|
||||
{confirmation.budget.max_model_calls} 次模型调用;每轮最多{" "}
|
||||
{confirmation.batch_candidates} 个候选。
|
||||
</p>
|
||||
<p>
|
||||
评估规则
|
||||
research-v1。模型输出不能扩展固定范围;每轮候选预览会先保存,再由服务端核验授权后启动回测。
|
||||
</p>
|
||||
</>
|
||||
)}
|
||||
</Modal>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
import { Banner, Button, Tag, Toast } from "@douyinfe/semi-ui-19";
|
||||
import { useState } from "react";
|
||||
import { post } from "../api";
|
||||
import type { UIAction } from "../ai/types";
|
||||
import type { FlowRun } from "./flowTypes";
|
||||
import { flowStatus } from "./flowTypes";
|
||||
export function FlowRunView({
|
||||
run,
|
||||
onRefresh,
|
||||
onAction,
|
||||
}: {
|
||||
run: FlowRun;
|
||||
onRefresh: () => void;
|
||||
onAction: (a: UIAction) => void;
|
||||
}) {
|
||||
const [busy, setBusy] = useState(false);
|
||||
async function control(action: string) {
|
||||
setBusy(true);
|
||||
try {
|
||||
await post(`/research/flows/runs/${run.id}/control`, {
|
||||
action,
|
||||
version: run.version,
|
||||
});
|
||||
} catch (e) {
|
||||
Toast.error((e as Error).message);
|
||||
} finally {
|
||||
setBusy(false);
|
||||
onRefresh();
|
||||
}
|
||||
}
|
||||
const budget = run.authorization.budget;
|
||||
return (
|
||||
<section className="research-card" aria-label="研究运行详情">
|
||||
<div className="section-toolbar">
|
||||
<h3>{run.name}</h3>
|
||||
<Tag>{flowStatus[run.status] || run.status}</Tag>
|
||||
</div>
|
||||
<p>{run.authorization.hypothesis}</p>
|
||||
<div className="research-form-grid">
|
||||
<div>
|
||||
研究轮数{" "}
|
||||
<strong>
|
||||
{run.round} / {budget.max_rounds}
|
||||
</strong>
|
||||
</div>
|
||||
<div>
|
||||
模拟条目{" "}
|
||||
<strong>
|
||||
{run.simulations_used} / {budget.max_simulations}
|
||||
</strong>
|
||||
</div>
|
||||
<div>
|
||||
模型调用{" "}
|
||||
<strong>
|
||||
{run.model_calls_used} / {budget.max_model_calls}
|
||||
</strong>
|
||||
</div>
|
||||
</div>
|
||||
{run.error && <Banner type="warning" description={run.error} />}
|
||||
<div className="inline-actions">
|
||||
<Button
|
||||
disabled={busy || !["queued", "running"].includes(run.status)}
|
||||
onClick={() => void control("pause")}
|
||||
>
|
||||
暂停研究
|
||||
</Button>
|
||||
<Button
|
||||
disabled={
|
||||
busy ||
|
||||
!["paused", "interrupted", "needs_review"].includes(run.status) ||
|
||||
run.steps.some((step) => step.status === "blocked")
|
||||
}
|
||||
onClick={() => void control("resume")}
|
||||
>
|
||||
恢复研究
|
||||
</Button>
|
||||
<Button
|
||||
type="danger"
|
||||
disabled={busy || ["completed", "stopped"].includes(run.status)}
|
||||
onClick={() => void control("stop")}
|
||||
>
|
||||
停止研究
|
||||
</Button>
|
||||
</div>
|
||||
<p className="research-hint">
|
||||
暂停和停止阻止后续步骤,已发出的模拟继续收集。预算不足或范围、模型配置变化时,需要重新确认新研究;未知提交请在关联回测中核实。
|
||||
</p>
|
||||
<details>
|
||||
<summary>本次固定授权</summary>
|
||||
<p>
|
||||
{run.authorization.settings.region} /{" "}
|
||||
{run.authorization.settings.universe} / D
|
||||
{run.authorization.settings.delay} ·{" "}
|
||||
{run.authorization.methods.join(" → ")}
|
||||
</p>
|
||||
<pre className="code-block">
|
||||
{JSON.stringify(run.authorization, null, 2)}
|
||||
</pre>
|
||||
</details>
|
||||
<div className="research-table-scroll">
|
||||
<table className="research-table">
|
||||
<thead>
|
||||
<tr>
|
||||
<th>轮次 / 阶段</th>
|
||||
<th>状态</th>
|
||||
<th>产物与候选池</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{run.steps.map((step) => (
|
||||
<tr key={step.id}>
|
||||
<td>
|
||||
{step.round} ·{" "}
|
||||
{run.definition.nodes.find((n) => n.id === step.node_id)
|
||||
?.label || step.node_id}
|
||||
</td>
|
||||
<td>
|
||||
{flowStatus[step.status] || step.status}
|
||||
{step.error && <p>{step.error}</p>}
|
||||
</td>
|
||||
<td>
|
||||
<div className="inline-actions">
|
||||
{step.output.template && (
|
||||
<Button
|
||||
size="small"
|
||||
onClick={() =>
|
||||
onAction({
|
||||
type: "open_template",
|
||||
asset_id: step.output.template!.id,
|
||||
version: step.output.template!.version,
|
||||
nonce: Date.now(),
|
||||
})
|
||||
}
|
||||
>
|
||||
模板 v{step.output.template.version}
|
||||
</Button>
|
||||
)}
|
||||
{step.output.experiment_id && (
|
||||
<Button
|
||||
size="small"
|
||||
onClick={() =>
|
||||
onAction({
|
||||
type: "open_experiment",
|
||||
experiment_id: step.output.experiment_id!,
|
||||
nonce: Date.now(),
|
||||
})
|
||||
}
|
||||
>
|
||||
候选实验 {step.output.candidate_ids?.length ?? ""}
|
||||
</Button>
|
||||
)}
|
||||
{step.output.preview_id && (
|
||||
<Button
|
||||
size="small"
|
||||
onClick={() =>
|
||||
onAction({
|
||||
type: "open_backtest_preview",
|
||||
preview_id: step.output.preview_id!,
|
||||
nonce: Date.now(),
|
||||
})
|
||||
}
|
||||
>
|
||||
固定预览
|
||||
</Button>
|
||||
)}
|
||||
{step.backtest_run_id && (
|
||||
<Button
|
||||
size="small"
|
||||
onClick={() =>
|
||||
onAction({
|
||||
type: "open_backtest",
|
||||
run_id: step.backtest_run_id!,
|
||||
nonce: Date.now(),
|
||||
})
|
||||
}
|
||||
>
|
||||
关联回测
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
<details>
|
||||
<summary>步骤产物</summary>
|
||||
<pre className="code-block">
|
||||
{JSON.stringify(step.output, null, 2)}
|
||||
</pre>
|
||||
</details>
|
||||
</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
import { useEffect, useState } from "react";
|
||||
import { Banner, Button, Pagination } from "@douyinfe/semi-ui-19";
|
||||
import { api, formatTime } from "../api";
|
||||
import type { PageContext, UIAction } from "../ai/types";
|
||||
import { FlowLaunchForm } from "./FlowLaunchForm";
|
||||
import { FlowRunView } from "./FlowRunView";
|
||||
import type { FlowRun } from "./flowTypes";
|
||||
import { flowStatus } from "./flowTypes";
|
||||
import "./workspace.css";
|
||||
export function PipelinePage({
|
||||
active,
|
||||
onAction,
|
||||
onContext,
|
||||
}: {
|
||||
active: boolean;
|
||||
onAction: (a: UIAction) => void;
|
||||
onContext: (c: PageContext) => void;
|
||||
}) {
|
||||
const [items, setItems] = useState<FlowRun[]>([]);
|
||||
const [total, setTotal] = useState(0);
|
||||
const [page, setPage] = useState(1);
|
||||
const [selected, setSelected] = useState<string | null>(() =>
|
||||
localStorage.getItem("research-selected-flow"),
|
||||
);
|
||||
const [run, setRun] = useState<FlowRun | null>(null);
|
||||
const [creating, setCreating] = useState(!selected);
|
||||
const [revision, setRevision] = useState(0);
|
||||
const [error, setError] = useState("");
|
||||
useEffect(() => {
|
||||
if (active) onContext({ page: "pipeline", research_run_id: selected });
|
||||
}, [active, selected, onContext]);
|
||||
useEffect(() => {
|
||||
if (!active) return;
|
||||
const c = new AbortController();
|
||||
let timer: ReturnType<typeof setTimeout>;
|
||||
async function refresh() {
|
||||
try {
|
||||
const list = await api<{ items: FlowRun[]; total: number }>(
|
||||
`/research/flows/runs?offset=${(page - 1) * 25}`,
|
||||
{ signal: c.signal },
|
||||
);
|
||||
setItems(list.items);
|
||||
setTotal(list.total);
|
||||
if (selected)
|
||||
setRun(
|
||||
await api(`/research/flows/runs/${selected}`, { signal: c.signal }),
|
||||
);
|
||||
setError("");
|
||||
} catch (e) {
|
||||
if (!c.signal.aborted) setError((e as Error).message);
|
||||
} finally {
|
||||
if (!c.signal.aborted) timer = setTimeout(() => void refresh(), 1500);
|
||||
}
|
||||
}
|
||||
void refresh();
|
||||
return () => {
|
||||
c.abort();
|
||||
clearTimeout(timer);
|
||||
};
|
||||
}, [active, page, selected, revision]);
|
||||
function select(id: string) {
|
||||
setSelected(id);
|
||||
localStorage.setItem("research-selected-flow", id);
|
||||
setCreating(false);
|
||||
}
|
||||
return (
|
||||
<section className="research-workspace">
|
||||
<div className="section-toolbar">
|
||||
<div>
|
||||
<h2>研究流水线</h2>
|
||||
<p className="muted">
|
||||
生成 → 校验与设参 → 回测 → 评估决策 → 增强 → 重新展开
|
||||
</p>
|
||||
</div>
|
||||
<Button onClick={() => setCreating(true)}>新建自动研究</Button>
|
||||
</div>
|
||||
{error && <Banner type="danger" description={error} />}
|
||||
<div className="research-layout">
|
||||
<aside className="research-library">
|
||||
<h3>研究运行</h3>
|
||||
{items.map((item) => (
|
||||
<button
|
||||
className={`research-library-item ${selected === item.id ? "selected" : ""}`}
|
||||
key={item.id}
|
||||
onClick={() => select(item.id)}
|
||||
>
|
||||
<strong>{item.name}</strong>
|
||||
<span>
|
||||
{flowStatus[item.status]} · 第 {item.round} 轮
|
||||
</span>
|
||||
<span>{formatTime(item.created_at)}</span>
|
||||
</button>
|
||||
))}
|
||||
<Pagination
|
||||
size="small"
|
||||
total={total}
|
||||
currentPage={page}
|
||||
pageSize={25}
|
||||
onPageChange={setPage}
|
||||
/>
|
||||
</aside>
|
||||
<div className="research-main">
|
||||
{creating ? (
|
||||
<FlowLaunchForm
|
||||
key="new-flow"
|
||||
onStarted={(r) => {
|
||||
select(r.id);
|
||||
setRun(r);
|
||||
setRevision((v) => v + 1);
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
run && (
|
||||
<FlowRunView
|
||||
run={run}
|
||||
onRefresh={() => setRevision((v) => v + 1)}
|
||||
onAction={onAction}
|
||||
/>
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -8,7 +8,9 @@ export function ResearchToolCard({ call, onAction }: ToolCardProps) {
|
||||
{typeof result.name === "string" ? result.name : "研究资料"}
|
||||
{typeof result.version === "number" ? ` · v${result.version}` : ""}
|
||||
</p>
|
||||
{typeof result.id === "string" && !result.report && (
|
||||
{typeof result.id === "string" &&
|
||||
!result.report &&
|
||||
!result.authorization && (
|
||||
<Button
|
||||
onClick={() =>
|
||||
onAction(
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import type { SimulationSettings } from "../backtests/types";
|
||||
export type Budget = {
|
||||
max_rounds: number;
|
||||
max_simulations: number;
|
||||
max_model_calls: number;
|
||||
};
|
||||
export type FlowNode = {
|
||||
id: string;
|
||||
type: string;
|
||||
label: string;
|
||||
x: number;
|
||||
y: number;
|
||||
config: Record<string, unknown>;
|
||||
};
|
||||
export type Workflow = {
|
||||
name: string;
|
||||
nodes: FlowNode[];
|
||||
edges: { source: string; target: string; branch?: string | null }[];
|
||||
};
|
||||
export type FlowLaunch = {
|
||||
request_id: string;
|
||||
name: string;
|
||||
input_ids: string[];
|
||||
hypothesis: string;
|
||||
settings: SimulationSettings;
|
||||
budget: Budget;
|
||||
batch_candidates: number;
|
||||
seed: number;
|
||||
parent_alpha_ids: string[];
|
||||
template_id?: string;
|
||||
template_version?: number;
|
||||
workflow_id?: string;
|
||||
workflow_version?: number;
|
||||
};
|
||||
export type FlowRun = {
|
||||
id: string;
|
||||
name: string;
|
||||
status: string;
|
||||
version: number;
|
||||
round: number;
|
||||
simulations_used: number;
|
||||
model_calls_used: number;
|
||||
error: string | null;
|
||||
created_at: string;
|
||||
definition: Workflow;
|
||||
authorization: FlowLaunch & { kind: string; methods: string[] };
|
||||
steps: {
|
||||
id: string;
|
||||
node_id: string;
|
||||
round: number;
|
||||
status: string;
|
||||
output: {
|
||||
type?: string;
|
||||
template?: { id: string; version: number };
|
||||
experiment_id?: string;
|
||||
evaluation_id?: string;
|
||||
preview_id?: string;
|
||||
candidate_ids?: string[];
|
||||
[key: string]: unknown;
|
||||
};
|
||||
backtest_run_id: string | null;
|
||||
error: string | null;
|
||||
}[];
|
||||
};
|
||||
export const flowStatus: Record<string, string> = {
|
||||
queued: "等待推进",
|
||||
running: "执行中",
|
||||
completed: "已完成",
|
||||
skipped: "已跳过",
|
||||
blocked: "候选校验未通过",
|
||||
waiting: "收集回测",
|
||||
previewed: "候选预览已固定",
|
||||
paused: "已暂停",
|
||||
stopped: "已停止",
|
||||
interrupted: "步骤中断",
|
||||
needs_review: "需要处理",
|
||||
budget_exhausted: "预算不足",
|
||||
};
|
||||
@@ -0,0 +1,157 @@
|
||||
import { expect, test } from "@playwright/test";
|
||||
const headers = { "X-WQ-Request": "1" };
|
||||
const scope = {
|
||||
instrument_type: "EQUITY",
|
||||
region: "USA",
|
||||
universe: "TOP3000",
|
||||
delay: 1,
|
||||
};
|
||||
test("finite pipeline confirms scope and budget, executes native steps and restores run", async ({
|
||||
page,
|
||||
}) => {
|
||||
const errors: string[] = [];
|
||||
page.on("pageerror", (e) => errors.push(e.message));
|
||||
await page.goto("/");
|
||||
await page.getByLabel("密码", { exact: true }).fill("browser-test-password");
|
||||
await page.getByRole("button", { name: "进入工作空间" }).click();
|
||||
await page
|
||||
.getByRole("navigation", { name: "主导航" })
|
||||
.getByRole("button", { name: "研究流水线", exact: true })
|
||||
.click();
|
||||
await expect(
|
||||
page.getByRole("heading", { name: "研究流水线", exact: true }),
|
||||
).toBeVisible();
|
||||
await page.request.put("/api/v1/account/credentials", {
|
||||
headers,
|
||||
data: { email: "test@example.com", password: "synthetic-password" },
|
||||
});
|
||||
await page.request.post("/api/v1/account/connect", { headers });
|
||||
await expect
|
||||
.poll(
|
||||
async () =>
|
||||
(await (await page.request.get("/api/v1/account")).json())
|
||||
.connection_status,
|
||||
)
|
||||
.toBe("connected");
|
||||
const config = { base_url: "https://model.test/v1", model: "test-model" };
|
||||
await page.request.put("/api/v1/ai/settings", {
|
||||
headers,
|
||||
data: { ...config, api_key: "synthetic-key" },
|
||||
});
|
||||
expect(
|
||||
(
|
||||
await (
|
||||
await page.request.post("/api/v1/ai/settings/test", { headers })
|
||||
).json()
|
||||
).ready,
|
||||
).toBe(true);
|
||||
await page.request.put("/api/v1/ai/settings", {
|
||||
headers,
|
||||
data: { ...config, enabled: true },
|
||||
});
|
||||
for (const dataset_id of [null, "TEST_FIN"]) {
|
||||
const job = await (
|
||||
await page.request.post("/api/v1/catalog/sync-jobs", {
|
||||
headers,
|
||||
data: { scope, dataset_id },
|
||||
})
|
||||
).json();
|
||||
await expect
|
||||
.poll(
|
||||
async () =>
|
||||
(await (await page.request.get(`/api/v1/sync-jobs/${job.id}`)).json())
|
||||
.status,
|
||||
)
|
||||
.toBe("completed");
|
||||
}
|
||||
const params = new URLSearchParams(
|
||||
Object.entries(scope).map(([k, v]) => [k, String(v)]),
|
||||
);
|
||||
const fields = await (
|
||||
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
|
||||
).json();
|
||||
const fixed = await (
|
||||
await page.request.post("/api/v1/catalog/inputs", {
|
||||
headers,
|
||||
data: {
|
||||
scope,
|
||||
dataset_id: "TEST_FIN",
|
||||
collection_version: fields.collection_version,
|
||||
selection: "all",
|
||||
},
|
||||
})
|
||||
).json();
|
||||
expect(
|
||||
(
|
||||
await page.request.post("/api/v1/catalog/operators/refresh", { headers })
|
||||
).ok(),
|
||||
).toBe(true);
|
||||
expect(
|
||||
(
|
||||
await page.request.post("/api/v1/catalog/setting-options/refresh", {
|
||||
headers,
|
||||
})
|
||||
).ok(),
|
||||
).toBe(true);
|
||||
await page.reload();
|
||||
await page.getByLabel("自动研究名称").fill("浏览器有限研究");
|
||||
await page.getByLabel("自动研究假设").fill("研究排名稳定性");
|
||||
await page.getByRole("combobox", { name: "自动研究固定输入" }).click();
|
||||
await page
|
||||
.getByRole("option")
|
||||
.filter({ hasText: fixed.id.slice(0, 8) })
|
||||
.click();
|
||||
await page.getByRole("heading", { name: "研究流水线", exact: true }).click();
|
||||
for (const [label, value] of [
|
||||
["最大轮数", "1"],
|
||||
["最大模拟条目数", "2"],
|
||||
["最大模型调用数", "2"],
|
||||
["每轮候选上限", "2"],
|
||||
])
|
||||
await page.getByLabel(label, { exact: true }).fill(value);
|
||||
await page.getByRole("button", { name: "核对并启动研究" }).click();
|
||||
await expect(page.getByRole("dialog", {name: "确认自动研究授权"})).toContainText(
|
||||
"最多 1 轮、2 条模拟、2 次模型调用",
|
||||
);
|
||||
const started = page.waitForResponse(
|
||||
(r) =>
|
||||
r.url().endsWith("/research/flows/runs") &&
|
||||
r.request().method() === "POST",
|
||||
);
|
||||
await page
|
||||
.getByRole("dialog", {name: "确认自动研究授权"})
|
||||
.getByRole("button", { name: "confirm", exact: true })
|
||||
.click();
|
||||
const response = await started;
|
||||
expect(response.status()).toBe(201);
|
||||
const run = await response.json();
|
||||
await expect
|
||||
.poll(
|
||||
async () =>
|
||||
(
|
||||
await (
|
||||
await page.request.get(`/api/v1/research/flows/runs/${run.id}`)
|
||||
).json()
|
||||
).status,
|
||||
{ timeout: 40000 },
|
||||
)
|
||||
.toBe("completed");
|
||||
const completed = await (
|
||||
await page.request.get(`/api/v1/research/flows/runs/${run.id}`)
|
||||
).json();
|
||||
expect(completed.simulations_used).toBe(2);
|
||||
expect(completed.model_calls_used).toBe(2);
|
||||
expect(completed.steps).toHaveLength(8);
|
||||
await page.reload();
|
||||
await expect(
|
||||
page.getByRole("region", { name: "研究运行详情" }),
|
||||
).toContainText("浏览器有限研究");
|
||||
await expect(
|
||||
page.getByRole("button", { name: "关联回测", exact: true }),
|
||||
).toBeVisible();
|
||||
await page.screenshot({
|
||||
path: "/tmp/wq-stage-three-run.png",
|
||||
fullPage: true,
|
||||
});
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
Reference in New Issue
Block a user