"""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 import random from collections import defaultdict from fastapi import HTTPException from sqlalchemy import select from sqlalchemy.exc import SQLAlchemyError from ..backtests.contracts import ControlInput, SimulationSettings, 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, scope_of from .features import Features from .model import request_model from .workflows import validate_graph from .workspace_contracts import ( AssetWrite, EvaluateInput, Expansion, FeatureSpec, Generation, SettingVariants, 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 output_type = FeatureSpec if context.get("method") == "feature" else TemplateSpec result, evidence = await request_model(self.ai, context, output_type, 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 {"template": run.authorization.get("template")}), "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 == "feature": reference = run.authorization.get("node_assets", {}).get(key) if step.output.get("feature"): ref = step.output["feature"] reference = await Assets(db).get(ref["id"], ref["version"], "feature") if reference: template = await Features(db).to_template(reference["id"], reference["version"]) step.output = { **data, **step.output, "type": "context", "feature": {"id": reference["id"], "version": reference["version"]}, "input_ids": reference["content"]["input_ids"], "template": {"id": template["id"], "version": template["version"]}, } else: return await self.reserve_model(db, run, step, data, node) elif node.type == "expand": reference = run.authorization.get("node_assets", {}).get(key) or data.get("template") await self.expand(db, run, step, {**data, "template": reference}) return elif node.type == "variant": if node.config.get("method", "structure") == "structure": if step.output.get("template"): await self.expand(db, run, step, step.output) return return await self.reserve_model(db, run, step, data, node) experiment = await Experiments(db).setting_variants( SettingVariants( alpha_id=run.authorization["parent_alpha_ids"][0], input_ids=run.authorization["input_ids"], hypothesis=run.authorization["hypothesis"], ), parent_snapshot=run.authorization["parents"][0], extra_evidence={"flow_run_id": run.id, "node_id": key, "round": run.round}, kind=run.authorization["kind"], ) self.candidates(run, step, experiment, {}) 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 == "condition": step.output = {**data, "type": "evaluation"} elif node.type == "filter": report = (await Evaluations(db).get(data["evaluation_id"]))["report"] ids = [ r["client_item_id"] for r in report["records"] if r["verdict"] in node.config.get("verdicts", ["pass"]) ] experiment = await Experiments(db).get(data["experiment_id"]) template = experiment["evidence"].get("template") step.output = { **data, "type": "candidates", "candidate_ids": ids, "template": {"id": template["id"], "version": template["version"]} if template else None, } if not ids: step.status, step.updated_at = "skipped", now() return elif node.type == "summarize": # References keep joins bounded, without recursively copying the upstream graph. step.output = { "type": "summary", "artifacts": [ { "step_id": steps[e.source].id, "node_id": e.source, **{ k: steps[e.source].output[k] for k in ( "type", "template", "feature", "experiment_id", "evaluation_id", "backtest_run_id", "verdict", ) if k in 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) ], } 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=data.get("input_ids", run.authorization["input_ids"]), method="feature" if node.type == "feature" else "structure" if node.type == "variant" else "template", parent_experiment_ids=[data["experiment_id"]] if data.get("experiment_id") else [], ) context = await Experiments(db).generation_context(generation) context["method"] = generation.method 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 if isinstance(result, FeatureSpec) and (result.preparation_refs or set(result.input_ids) != set( [i["id"] for i in step.output["context"]["inputs"]] )): raise HTTPException(422, "模型不能改变已固定的输入范围") # A paused/stopped run may collect this already-issued model output, but cannot advance. asset = await Assets(db).save( AssetWrite( kind="feature" if isinstance(result, FeatureSpec) else "template", content=result.model_dump(mode="json"), ), provenance={ "flow_run_id": run_id, "step_id": step_id, "generation": evidence, "context": step.output["context"], }, ) feature = ( {"id": asset["id"], "version": asset["version"]} if isinstance(result, FeatureSpec) else None ) mark_model_attempt(step, "completed") step.output = { "model_attempts": step.output.get("model_attempts", []), "type": "context" if feature else "template", "feature": feature, "input_ids": result.input_ids if feature else [i["id"] for i in step.output["context"]["inputs"]], "template": None if feature else {"id": asset["id"], "version": asset["version"]}, "parent_experiment_ids": step.output.get("parent_experiment_ids", []), "generation": evidence, "reserved_call": step.output["reserved_call"], } node = next(n for n in run.definition["nodes"] if n["id"] == step.node_id) # Keep post-processing durable and separate: pause/account changes are checked again # before converting features or expanding variants on the next active tick. step.status = "generated" if node["type"] in ("variant", "feature") else "completed" step.updated_at = now() changed(run) def candidates(self, run, step, experiment, data): step.output = { **data, "type": "candidates", "experiment_id": experiment["id"], "candidate_ids": [ c["client_item_id"] for c in experiment["candidates"] if experiment["kind"] == "template" or c["validation"]["status"] == "valid" ], } ids = step.output["candidate_ids"] maximum = run.authorization["batch_candidates"] if len(ids) > maximum: step.output = { **step.output, "candidate_ids": random.Random(run.authorization["seed"] + run.round).sample(ids, maximum), } step.status, step.updated_at = "completed", now() if not step.output["candidate_ids"]: step.status, step.error = "blocked", "候选均未通过本地校验,请核实字段、算子和设置" if run.status in ACTIVE: halt(run, "needs_review", step.error) async def expand(self, db, run, step, data): reference = data.get("template") if not reference: raise HTTPException(422, "展开节点没有固定模板版本") scope = scope_of(SimulationSettings.model_validate(run.authorization["settings"])) ids = [ i["id"] for i in run.authorization["inputs"] if i["scope"] == scope and i["id"] in data.get("input_ids", run.authorization["input_ids"]) ] if not ids: raise HTTPException(422, "模板展开需要与基础设置匹配的固定输入") body = Expansion( asset_id=reference["id"], version=reference["version"], input_ids=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, run.authorization["kind"], { "flow_run_id": run.id, "node_id": step.node_id, "round": run.round, "method": run.authorization["kind"], "authorized_seed_snapshots": run.authorization["parents"], }, parent_snapshots=frozen_parents, ) self.candidates(run, step, experiment, data) 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"]}.issubset(set(run.authorization["input_ids"])) is False ): raise HTTPException(403, "候选不属于此研究运行的固定输入范围") if "backtest" not in run.authorization["methods"] or any( SimulationSettings.model_validate(c["settings"]).model_dump(mode="json") not in [ SimulationSettings.model_validate(value).model_dump(mode="json") for value in ( run.authorization.get("allowed_settings", [run.authorization["settings"]]) if experiment["evidence"].get("method") == "settings" else [run.authorization["settings"]] ) ] or not c.get("input_ids") or any( not any( i["id"] == input_id and i["scope"] == scope_of(SimulationSettings.model_validate(c["settings"])) for i in run.authorization["inputs"] ) for input_id in c["input_ids"] ) 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"