"""Transactional research interface. Callers own authorization and commit boundaries.""" from collections import Counter, defaultdict from uuid import uuid4 from fastapi import HTTPException from fastapi.encoders import jsonable_encoder from sqlalchemy import func, select, update from ..models import ( Account, BacktestConfig, BacktestDraft, BacktestEvent, BacktestItem, BacktestPreview, BacktestResult, BacktestRun, SimulationAttempt, now, ) from .contracts import Candidate, DraftInput, PreviewInput, Source, fingerprint, group_key def uid(): return str(uuid4()) async def event(db, run, kind, payload): """Append a run-local cursor under the run row lock, in the result's transaction.""" run.event_seq += 1 run.updated_at = now() db.add(BacktestEvent(run_id=run.id, seq=run.event_seq, kind=kind, payload=payload)) async def locked_run(db, run_id): run = await db.scalar(select(BacktestRun).where(BacktestRun.id == run_id).with_for_update()) if not run: raise HTTPException(404, "回测运行不存在") return run async def refresh_status(db, run): await db.flush() states = list(await db.scalars(select(SimulationAttempt.state).where(SimulationAttempt.run_id == run.id))) if any(s in ("needs_review", "collection_failed") for s in states): run.status = "needs_review" elif all(s in ("completed", "failed", "skipped") for s in states): run.status = ( "stopped" if run.control == "stopped" else "completed_with_errors" if "failed" in states else "completed" ) elif run.control == "paused": run.status = "paused" elif run.control == "stopped": run.status = "stopping" elif any(s in ("submitting", "submitted", "collecting") for s in states): run.status = "running" else: run.status = "queued" class Backtests: def __init__(self, db, ai_context=None): self.db = db self.ai_context = ai_context or {} async def config(self): row = await self.db.get(BacktestConfig, 1) return jsonable_encoder( { k: getattr(row, k) for k in ("concurrency", "batch_size", "version", "blocked_reason", "blocked_until") } ) async def configure(self, body): result = await self.db.execute( update(BacktestConfig) .where(BacktestConfig.id == 1, BacktestConfig.version == body.version) .values( concurrency=body.concurrency, batch_size=body.batch_size, version=BacktestConfig.version + 1, ) ) if result.rowcount != 1: raise HTTPException(409, "调度配置已变化,请刷新后重试") return await self.config() async def capabilities(self): return { "alpha_types": ["REGULAR", "SUPER"], "languages": ["FASTEXPR"], "instrument_types": ["EQUITY"], "settings_schema": Candidate.model_json_schema(), "scheduler": await self.config(), "max_candidates": 10000, "remote_cancel": False, "automatic_history_reuse": False, "confirmation": "每个固定运行确认一次;启动后返回 ID,不循环等待", "mapping": "完整输入匹配;证据不足待核对,不按 children 顺序匹配", "limits": "并发和批大小为本地配置,并非平台可用额度;账户外提交不在预算内", } async def bind_preparations(self, body): from ..preparations.service import Preparations from ..research.expressions import analyze if not body.preparation_refs and not body.input_ids: return if any(c.alpha_type == "SUPER" for c in body.candidates): raise HTTPException(422, "SUPER 组件快照不能使用字段数据准备集合") await Preparations(self.db).bind(body) snapshots = [await Preparations(self.db).snapshot(i) for i in body.input_ids] for candidate in body.candidates: scope = dict(instrument_type=candidate.settings.instrumentType, region=candidate.settings.region, universe=candidate.settings.universe, delay=candidate.settings.delay) if any(s["scope"] != scope for s in snapshots): raise HTTPException(422, "数据准备集合与回测范围不一致") fields = {} for snapshot in snapshots: for field, kind in snapshot["field_types"].items(): if field in fields and fields[field] != kind: raise HTTPException(422, "输入字段类型冲突") fields[field] = kind validation = analyze(candidate.expression, fields) if validation["syntax"] or validation["types"]: raise HTTPException(422, ";".join(validation["syntax"] + validation["types"])) body.source.input_snapshot_ids = body.input_ids body.source.input_snapshot_id = body.input_ids[0] if len(body.input_ids) == 1 else None async def save_draft(self, body, draft_id=None): await self.bind_preparations(body) data = body.model_dump(mode="json", exclude={"version", "preparation_refs", "input_ids"}) if draft_id: changed = await self.db.execute( update(BacktestDraft) .where(BacktestDraft.id == draft_id, BacktestDraft.version == body.version) .values( **data, version=BacktestDraft.version + 1, updated_at=now(), ) ) if changed.rowcount != 1: raise HTTPException(409, "草稿已变化或不存在;保留当前编辑并重新载入") else: draft_id = uid() self.db.add(BacktestDraft(id=draft_id, **data)) await self.db.flush() return await self.draft(draft_id) async def drafts(self, limit=25, offset=0, q="", sort="updated_at", direction="desc"): query = select(BacktestDraft) if q: query = query.where(BacktestDraft.name.contains(q, autoescape=True)) column = {"name": BacktestDraft.name, "updated_at": BacktestDraft.updated_at}[sort] order = column.asc() if direction == "asc" else column.desc() rows = ( await self.db.scalars( query .order_by(order, BacktestDraft.id) .limit(limit) .offset(offset) ) ).all() return { "items": [ jsonable_encoder( { "id": r.id, "version": r.version, "name": r.name, "total": len(r.candidates), "updated_at": r.updated_at, } ) for r in rows ], "total": await self.db.scalar(select(func.count()).select_from(query.subquery())), "limit": limit, "offset": offset, } async def draft(self, draft_id): row = await self.db.get(BacktestDraft, draft_id) if not row: raise HTTPException(404, "候选草稿不存在") return jsonable_encoder( {k: getattr(row, k) for k in ("id", "version", "name", "source", "candidates", "updated_at")} ) async def preview(self, body, *, preserve_source=False): """Fix inputs; new chatbox candidates inherit trusted generating-run provenance. Existing draft references and server-side subsets/reruns retain their producer. ai_context separately identifies whoever starts the execution. """ if body.inline: await self.bind_preparations(body.inline) data = body.inline.model_dump(mode="json") if self.ai_context and not preserve_source: data["source"] = { **data["source"], "kind": "chatbox", "reference": self.ai_context["conversation_id"], "research_id": self.ai_context["ai_run_id"], "parent_run_id": None, } else: draft = await self.db.scalar( select(BacktestDraft).where(BacktestDraft.id == body.draft_id).with_for_update() ) if not draft or draft.version != body.draft_version: raise HTTPException(409, "候选草稿已变化,请重新准备预览") candidates = draft.candidates if body.selection is not None: selection = set(body.selection) candidates = [c for c in candidates if c["client_item_id"] in selection] if len(candidates) != len(selection): raise HTTPException(422, "选择包含不属于当前草稿的候选") data = {"name": draft.name, "source": draft.source, "candidates": candidates} from ..superalpha.service import validate_source await validate_source(self.db, data["source"], data["candidates"]) candidates = DraftInput.model_validate(data).model_dump(mode="json")["candidates"] config = await self.db.get(BacktestConfig, 1) groups = defaultdict(list) hashes = [] for i, c in enumerate(candidates): groups[group_key(c)].append(i) hashes.append(fingerprint(Candidate.model_validate(c).platform_input())) # Query hashes in bounded chunks, including SQLite's bind-parameter limit. existing = set() for index in range(0, len(hashes), 400): existing.update( await self.db.scalars( select(BacktestItem.fingerprint) .where(BacktestItem.fingerprint.in_(hashes[index : index + 400])) .distinct() ) ) seen, duplicates = set(), [] for c, h in zip(candidates, hashes): if h in seen or h in existing: duplicates.append( { "client_item_id": c["client_item_id"], "historical": h in existing, "within_preview": h in seen, } ) seen.add(h) batches = [] for indices in groups.values(): local_batches = [] for index in indices: if candidates[index]["alpha_type"] == "SUPER": batches.append([index]) continue batch = next( ( b for b in local_batches if len(b) < config.batch_size and all(hashes[i] != hashes[index] for i in b) ), None, ) if batch is None: batch = [] local_batches.append(batch) batch.append(index) batches.extend(local_batches) row = BacktestPreview( id=uid(), name=data["name"], source=data["source"], candidates=candidates, batches=batches, batch_size=config.batch_size, digest=fingerprint({"candidates": candidates, "source": data["source"]}), duplicates=duplicates, ai_context=self.ai_context, ) self.db.add(row) await self.db.flush() return await self.get_preview(row.id) async def get_preview(self, preview_id, limit=25, offset=0): row = await self.db.get(BacktestPreview, preview_id) if not row: raise HTTPException(404, "回测预览不存在") return jsonable_encoder( { "preview_id": row.id, "version": row.version, "name": row.name, "source": row.source, "digest": row.digest, "total": len(row.candidates), "batch_count": len(row.batches), "batch_size": row.batch_size, "duplicate_count": len(row.duplicates), "duplicates": row.duplicates[offset : offset + limit], "items": row.candidates[offset : offset + limit], "limit": limit, "offset": offset, "has_more": offset + limit < len(row.candidates), "created_at": row.created_at, } ) async def start(self, body): # One account row serializes all starts; unique keys remain the final DB invariant. account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update()) previous = await self.db.scalar( select(BacktestRun).where(BacktestRun.idempotency_key == body.idempotency_key) ) if previous: if previous.preview_id != body.preview_id or body.version != 1: raise HTTPException(409, "幂等键已用于另一份预览") return await self.run(previous.id) preview = await self.db.get(BacktestPreview, body.preview_id) if not preview or preview.version != body.version: raise HTTPException(409, "预览不存在或版本不匹配") previous = await self.db.scalar(select(BacktestRun).where(BacktestRun.preview_id == preview.id)) if previous: return await self.run(previous.id) if not account or not account.wq_user_id or account.connection_status != "connected": raise HTTPException(409, "请先连接并确认 WorldQuant 账户身份") run = BacktestRun( id=uid(), preview_id=preview.id, idempotency_key=body.idempotency_key, name=preview.name, source=preview.source, total=len(preview.candidates), batch_size=preview.batch_size, ai_context=self.ai_context or preview.ai_context, event_seq=0, ) self.db.add(run) await self.db.flush() for n, indices in enumerate(preview.batches): candidates = [Candidate.model_validate(preview.candidates[i]) for i in indices] attempt = SimulationAttempt( id=uid(), run_id=run.id, ordinal=n, payload=[c.platform_input() for c in candidates] ) self.db.add(attempt) await self.db.flush() for i, c in zip(indices, candidates): self.db.add( BacktestItem( id=uid(), run_id=run.id, attempt_id=attempt.id, ordinal=i, client_item_id=c.client_item_id, expression=c.expression, alpha_type=c.alpha_type, selection=c.selection, combo=c.combo, settings=c.settings.model_dump(), fingerprint=fingerprint(c.platform_input()), ) ) await event(self.db, run, "created", {"total": run.total, "batch_count": len(preview.batches)}) await self.db.flush() return await self.run(run.id) async def runs(self, limit=25, offset=0, source=None, reference=None, research_id=None, q="", sort="created_at", direction="desc", alpha_type=None): query = select(BacktestRun) if alpha_type: query = query.where(BacktestRun.id.in_(select(BacktestItem.run_id).where(BacktestItem.alpha_type == alpha_type))) if q: query = query.where(BacktestRun.name.contains(q, autoescape=True)) column = {"name": BacktestRun.name, "created_at": BacktestRun.created_at}[sort] order = column.asc() if direction == "asc" else column.desc() for key, value in (("kind", source), ("reference", reference), ("research_id", research_id)): if value: query = query.where(BacktestRun.source[key].as_string() == value) total = await self.db.scalar(select(func.count()).select_from(query.subquery())) rows = ( await self.db.scalars( query.order_by(order, BacktestRun.id).limit(limit).offset(offset) ) ).all() return { "items": [await self.run(r.id) for r in rows], "total": total, "limit": limit, "offset": offset, } async def sources(self): kinds = await self.db.scalars(select(BacktestRun.source["kind"].as_string()).distinct()) return sorted({kind for kind in kinds if kind} | {"chatbox", "manual"}) async def run(self, run_id): row = await self.db.get(BacktestRun, run_id) if not row: raise HTTPException(404, "回测运行不存在") groups = ( await self.db.execute( select( BacktestItem.platform_status, BacktestItem.collection_status, BacktestItem.persistence_status, func.count(), ) .where(BacktestItem.run_id == run_id) .group_by( BacktestItem.platform_status, BacktestItem.collection_status, BacktestItem.persistence_status, ) ) ).all() counts = {"platform": Counter(), "collection": Counter(), "persistence": Counter()} for p, c, s, n in groups: for key, value in (("platform", p), ("collection", c), ("persistence", s)): counts[key][value] += n return jsonable_encoder( { "backtest_run_id": row.id, **{ k: getattr(row, k) for k in ( "preview_id", "name", "source", "ai_context", "control", "status", "version", "total", "batch_size", "created_at", "updated_at", ) }, "counts": counts, "cursor": row.event_seq, "scheduler": await self.config(), } ) async def results(self, run_id, limit=25, offset=0): run = await self.run(run_id) rows = ( await self.db.execute( select(BacktestItem, BacktestResult) .outerjoin(BacktestResult, BacktestResult.item_id == BacktestItem.id) .where(BacktestItem.run_id == run_id) .order_by(BacktestItem.ordinal) .limit(limit) .offset(offset) ) ).all() return jsonable_encoder( { "backtest_run_id": run_id, "total": run["total"], "limit": limit, "offset": offset, "items": [ { **{ k: getattr(i, k) for k in ( "id", "client_item_id", "expression", "alpha_type", "selection", "combo", "settings", "attempt_id", "platform_status", "collection_status", "persistence_status", "simulation_id", "alpha_id", "error", ) }, "result": { "snapshot": r.snapshot, "observed_at": r.observed_at, "complete": r.complete, } if r else None, } for i, r in rows ], } ) async def events(self, run_id, after=0, limit=100): await self.run(run_id) rows = ( await self.db.scalars( select(BacktestEvent) .where(BacktestEvent.run_id == run_id, BacktestEvent.seq > after) .order_by(BacktestEvent.seq) .limit(limit + 1) ) ).all() return jsonable_encoder( { "items": [ {"seq": r.seq, "kind": r.kind, "payload": r.payload, "created_at": r.created_at} for r in rows[:limit] ], "next_cursor": rows[min(len(rows), limit) - 1].seq if rows else after, "has_more": len(rows) > limit, } ) async def attempts(self, run_id): await self.run(run_id) rows = ( await self.db.scalars( select(SimulationAttempt) .where(SimulationAttempt.run_id == run_id) .order_by(SimulationAttempt.ordinal) ) ).all() return jsonable_encoder( [ { k: getattr(a, k) for k in ( "id", "state", "ordinal", "progress_url", "remote_complete", "children", "error", "error_code", "poll_count", "submit_count", "next_poll_at", ) } for a in rows ] ) async def control(self, run_id, body): run = await locked_run(self.db, run_id) if run.version != body.version: raise HTTPException(409, "运行控制已变化,请重新确认") attempts = ( await self.db.scalars(select(SimulationAttempt).where(SimulationAttempt.run_id == run_id)) ).all() if body.action == "recover": for a in attempts: if a.state in ("needs_review", "collection_failed") and a.progress_url: if len(a.children) != len(a.payload): # Re-enumerate missing children while retaining collected receipts/results. a.children = [] a.state, a.poll_count, a.next_poll_at, a.error, a.error_code = ( "submitted", 0, None, None, None, ) # Recovery never clears uncertain submissions or creates a new POST. elif body.action == "resume": if run.control == "stopped": raise HTTPException(409, "已停止的剩余项不能恢复,请生成重跑预览") run.control = "active" config = await self.db.scalar( select(BacktestConfig).where(BacktestConfig.id == 1).with_for_update() ) # An explicit resume may clear an indefinite quota block, never a Retry-After deadline. if config.blocked_until is None: config.blocked_reason = None elif body.action == "pause": if run.control == "stopped": raise HTTPException(409, "该运行已经停止") run.control = "paused" else: run.control = "stopped" for a in attempts: if a.state == "queued": a.state = "skipped" await self.db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values( platform_status="skipped", collection_status="not_required", persistence_status="not_required", ) ) run.version += 1 await refresh_status(self.db, run) await event(self.db, run, "control", {"action": body.action, "control": run.control}) await self.db.flush() return await self.run(run_id) async def rerun(self, run_id, body): run = await locked_run(self.db, run_id) rows = ( await self.db.scalars( select(BacktestItem).where(BacktestItem.run_id == run_id).order_by(BacktestItem.ordinal) ) ).all() selected = [r for r in rows if r.id in set(body.item_ids)] if len(selected) != len(set(body.item_ids)): raise HTTPException(422, "重跑项不属于指定运行") if any(r.platform_status not in ("completed", "failed", "skipped") for r in selected): raise HTTPException(409, "仍在执行或结果未知的项须先核对,不能直接重跑") return await self.preview( PreviewInput( inline=DraftInput( name=f"{run.name[:190]} · 重跑", source=Source.model_validate({**run.source, "parent_run_id": run.id}), candidates=[ Candidate( client_item_id=r.client_item_id, expression=r.expression, settings=r.settings, alpha_type=r.alpha_type, selection=r.selection, combo=r.combo ) for r in selected ], ) ), preserve_source=True, ) async def attach_reference(self, attempt_id, body): """Record a human-supplied original simulation; collection still verifies its input.""" a = await self.db.get(SimulationAttempt, attempt_id) if not a: raise HTTPException(404, "执行尝试不存在") run = await locked_run(self.db, a.run_id) if run.version != body.version or a.state != "needs_review" or a.progress_url: raise HTTPException(409, "执行状态已变化或已有平台引用,请重新读取") duplicate = await self.db.scalar( select(SimulationAttempt.id).where(SimulationAttempt.progress_url == body.progress_url) ) if duplicate: raise HTTPException(409, "此模拟引用已经关联其他执行尝试") a.progress_url, a.state, a.error, a.error_code = body.progress_url, "submitted", None, None a.next_poll_at, a.poll_count = None, 0 run.version += 1 await self.db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values(platform_status="submitted", error=None) ) await refresh_status(self.db, run) await event( self.db, run, "reference_attached", {"attempt_id": a.id, "progress_url": body.progress_url} ) return await self.run(run.id) async def subset(self, preview_id, body): parent = await self.db.get(BacktestPreview, preview_id) if not parent: raise HTTPException(404, "预览不存在") excluded = set(body.exclude_ids) if not excluded.issubset({c["client_item_id"] for c in parent.candidates}): raise HTTPException(422, "排除集合包含未知候选") candidates = [c for c in parent.candidates if c["client_item_id"] not in excluded] if not candidates: raise HTTPException(422, "至少保留一条候选") return await self.preview( PreviewInput(inline=DraftInput(name=parent.name, source=parent.source, candidates=candidates)), preserve_source=True, )