Files
worldquant-alpha-system/backend/app/research/runtime.py
T

498 lines
22 KiB
Python
Raw Normal View History

"""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"