"""One account execution lane owned by Runner; DB intent always precedes a POST. No HTTP retry can replay an uncertain submission. Each short worker owns its DB transactions; network waits never hold DB row locks or the sync execution lane. """ import asyncio import logging import re from datetime import timedelta from sqlalchemy import func, select, update from sqlalchemy.exc import SQLAlchemyError from ..alphas import code, sanitize, upsert_alpha from ..models import ( Account, BacktestConfig, BacktestItem, BacktestResult, BacktestRun, SimulationAttempt, now, ) from ..worldquant import SimulationDeferred, VerificationRequired, WqError from .service import event, locked_run, refresh_status logger = logging.getLogger(__name__) REMOTE = ("submitting", "submitted", "collecting", "needs_review", "collection_failed") TERMINAL = ("COMPLETE", "FAILED", "ERROR", "WARNING") class BacktestLane: def __init__(self, owner): self.owner, self.sessions, self.client = owner, owner.sessions, owner.client self.loop_task = None self.tasks = {} self.wake = asyncio.Event() self.last_run = None self.poll_interval = 5 self.poll_limit = 300 self.stopping = False self.receipt_cache = {} async def start(self): self.stopping = False async with self.sessions.begin() as db: attempts = ( await db.scalars(select(SimulationAttempt).where(SimulationAttempt.state == "submitting")) ).all() for a in attempts: run = await locked_run(db, a.run_id) a.state = "submitted" if a.progress_url else "needs_review" a.error = None if a.progress_url else "服务在提交期间中断,结果未知,禁止自动重提" a.error_code = None if a.progress_url else "submission_unknown" await db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values(platform_status="submitted" if a.progress_url else "unknown") ) await refresh_status(db, run) await event(db, run, "recovered_after_restart", {"attempt_id": a.id, "state": a.state}) 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 await self.interrupt() async def interrupt(self): tasks = list(self.tasks.values()) for task in tasks: task.cancel() await asyncio.gather(*tasks, return_exceptions=True) self.tasks.clear() async def loop(self): while not self.stopping: try: await self.tick() except (SQLAlchemyError, OSError): logger.warning("Backtest lane 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 in list(self.tasks): if self.tasks[key].done(): task = self.tasks.pop(key) try: task.result() except asyncio.CancelledError: pass except Exception: logger.warning("Backtest worker interrupted; will reconcile durable state") if self.stopping or self.owner.disconnecting: return async with self.sessions() as db: account = await db.get(Account, 1) if ( not account or account.connection_status not in ("connected", "expired") or not account.wq_user_id ): return config = await db.get(BacktestConfig, 1) attempts = ( await db.scalars( select(SimulationAttempt) .join(BacktestRun) .where(SimulationAttempt.state.in_(("queued", "submitted", "collecting", "submitting"))) .order_by(BacktestRun.created_at, SimulationAttempt.ordinal) ) ).all() active = await db.scalar( select(func.count()) .select_from(SimulationAttempt) .where(SimulationAttempt.state.in_(REMOTE), SimulationAttempt.remote_complete.is_(False)) ) controls = dict((await db.execute(select(BacktestRun.id, BacktestRun.control))).all()) blocked = config.blocked_reason is not None and ( config.blocked_until is None or config.blocked_until.replace(tzinfo=now().tzinfo) > now() ) capacity = max(0, config.concurrency - active) runnable = [] for a in attempts: if a.id in self.tasks or ( a.next_poll_at and a.next_poll_at.replace(tzinfo=now().tzinfo) > now() ): continue if a.state != "queued": runnable.append(a.id) run_ids = list( dict.fromkeys( a.run_id for a in attempts if a.state == "queued" and controls[a.run_id] == "active" ) ) if self.last_run in run_ids: p = run_ids.index(self.last_run) + 1 run_ids = run_ids[p:] + run_ids[:p] while capacity and run_ids and not blocked: next_ids = [] for run_id in run_ids: match = next( ( a for a in attempts if a.run_id == run_id and a.state == "queued" and a.id not in self.tasks and a.id not in runnable and ( a.next_poll_at is None or a.next_poll_at.replace(tzinfo=now().tzinfo) <= now() ) ), None, ) if match and capacity: runnable.append(match.id) self.last_run = run_id capacity -= 1 next_ids.append(run_id) run_ids = next_ids # DB claims happen in workers and recheck control, budget and account. for attempt_id in runnable: self.tasks[attempt_id] = asyncio.create_task(self.step(attempt_id)) async def step(self, attempt_id): try: async with self.sessions() as db: a = await db.get(SimulationAttempt, attempt_id) state = a.state if state not in ("queued", "submitting", "submitted", "collecting"): return await self.owner.ensure_connected() if state == "queued": await self.submit(attempt_id) elif state == "submitting": if attempt_id in self.receipt_cache: await self.accept(attempt_id, self.receipt_cache[attempt_id]) else: await self.mark( attempt_id, "needs_review", "提交状态未知,禁止自动重提", "submission_unknown" ) else: await self.collect(attempt_id) except asyncio.CancelledError: # A killed POST is ambiguous; its durable 'submitting' state remains for reconciliation. raise except VerificationRequired as exc: await self.owner.set_account("verification_required", str(exc), exc.url) except SimulationDeferred as exc: await self.defer(attempt_id, exc) except WqError as exc: if exc.code in ("disconnected", "authentication_failed", "identity_mismatch"): await self.owner.set_account( "disconnected" if exc.code == "disconnected" else "error", str(exc) ) else: await self.mark( attempt_id, "needs_review" if exc.code in ("submission_unknown", "mapping_unknown") else "failed" if exc.code == "submission_rejected" else "collection_failed", str(exc), exc.code, ) except (SQLAlchemyError, OSError): # Receipt/raw data already persisted are retried without POST. Volatile Location is a cache only. logger.warning("Backtest persistence interrupted; durable attempt retained") except Exception: logger.error("Backtest internal failure: %s", attempt_id) await self.mark( attempt_id, "needs_review", "执行内部异常;已保留提交阶段,请核对后恢复", "internal_error" ) finally: self.wake.set() async def submit(self, attempt_id): async with self.owner.control_lock: if self.owner.disconnecting or self.stopping: return async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) config = await db.scalar( select(BacktestConfig).where(BacktestConfig.id == 1).with_for_update() ) account = await db.get(Account, 1) active = await db.scalar( select(func.count()) .select_from(SimulationAttempt) .where(SimulationAttempt.state.in_(REMOTE), SimulationAttempt.remote_complete.is_(False)) ) blocked = config.blocked_reason and ( not config.blocked_until or config.blocked_until.replace(tzinfo=now().tzinfo) > now() ) if ( a.state != "queued" or run.control != "active" or active >= config.concurrency or blocked or account.connection_status != "connected" ): return a.state, a.submit_count = "submitting", a.submit_count + 1 payload = a.payload await db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values(platform_status="submitting") ) await refresh_status(db, run) await event(db, run, "submitting", {"attempt_id": a.id}) url = await self.client.submit_simulations(payload) self.receipt_cache[attempt_id] = url await self.accept(attempt_id, url) async def accept(self, attempt_id, url): async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) a.progress_url, a.state, a.error, a.next_poll_at = url, "submitted", None, None await db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values(platform_status="submitted") ) await refresh_status(db, run) await event(db, run, "accepted", {"attempt_id": a.id}) self.receipt_cache.pop(attempt_id, None) async def defer(self, attempt_id, exc): async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) if a.state == "submitting": a.state = ( "skipped" if run.control == "stopped" else "queued" if a.submit_count < self.owner.settings.retry_attempts else "failed" ) await db.execute( update(BacktestItem) .where(BacktestItem.attempt_id == a.id) .values( platform_status="pending" if a.state == "queued" else a.state, collection_status="pending" if a.state == "queued" else "not_required", persistence_status="pending" if a.state == "queued" else "not_required", ) ) a.error, a.error_code = str(exc), exc.code a.next_poll_at = now() + timedelta(seconds=exc.delay) if exc.code == "rate_limited": config = await db.scalar( select(BacktestConfig).where(BacktestConfig.id == 1).with_for_update() ) if not config.blocked_reason or ( config.blocked_until and config.blocked_until.replace(tzinfo=now().tzinfo) < a.next_poll_at ): config.blocked_reason, config.blocked_until = str(exc), a.next_poll_at await refresh_status(db, run) await event(db, run, "deferred", {"attempt_id": a.id, "code": exc.code}) async def mark(self, attempt_id, state, message, code_value): async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) a.state, a.error, a.error_code = state, message, code_value items = (await db.scalars(select(BacktestItem).where(BacktestItem.attempt_id == a.id))).all() for i in items: if i.persistence_status == "saved" or i.platform_status == "failed": continue i.error = message if state == "failed": i.platform_status, i.collection_status, i.persistence_status = ( "failed", "not_required", "not_required", ) elif state == "needs_review": i.platform_status = "unknown" else: i.collection_status = "failed" await refresh_status(db, run) await event( db, run, "attention", {"attempt_id": a.id, "state": state, "code": code_value, "error": message}, ) async def checkpoint_receipt(self, attempt_id, simulation_id, receipt): async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) a.receipts = {**a.receipts, simulation_id: sanitize(receipt)} a.state = "collecting" await event(db, run, "received", {"attempt_id": a.id, "simulation_id": simulation_id}) async def collect(self, attempt_id): async with self.sessions() as db: a = await db.get(SimulationAttempt, attempt_id) url, children, receipts, count = a.progress_url, a.children, dict(a.receipts), len(a.payload) if a.poll_count >= self.poll_limit: raise WqError("轮询预算已用完,可找回原模拟,不会重新提交", "poll_timeout") delay = self.poll_interval if not children: parent, retry = await self.client.poll_simulation(url) delay = max(delay, retry) status = parent.get("status") if count == 1 and status in TERMINAL: children = [url.rsplit("/", 1)[-1]] receipts[children[0]] = {"progress": self.safe_progress(parent)} elif count > 1 and isinstance(parent.get("children"), list) and parent["children"]: children = parent["children"] if any( not isinstance(c, str) or not re.fullmatch(r"[A-Za-z0-9_-]+", c) for c in children ) or len(set(children)) != len(children): raise WqError("子模拟引用不合法或重复", "mapping_unknown") elif status in ("FAILED", "ERROR", "WARNING"): await self.quota(parent) await self.mark( attempt_id, "failed", "平台父模拟失败,请检查输入后创建重跑预览", "platform_failed" ) return async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) a.children = children a.receipts = sanitize(receipts) collection_errors = [] for child in children: try: receipt = receipts.get(child, {}) progress = receipt.get("progress", {}) if progress.get("status") not in TERMINAL: progress, retry = await self.client.poll_simulation(f"/simulations/{child}") progress = self.safe_progress(progress) delay = max(delay, retry) if progress.get("status") not in TERMINAL: continue receipt = {"progress": progress} receipts[child] = receipt await self.checkpoint_receipt(attempt_id, child, receipt) await self.quota(progress) await self.persist_receipt(attempt_id, child, receipt, count) alpha_id = progress.get("alpha") if alpha_id and progress.get("status") in ("COMPLETE", "WARNING") and "detail" not in receipt: if not isinstance(alpha_id, str) or not re.fullmatch(r"[A-Za-z0-9_-]+", alpha_id): raise WqError("平台 Alpha 标识无法确认", "mapping_unknown") detail = await self.client.alpha(alpha_id) if detail.get("id") != alpha_id: raise WqError("平台结果标识与请求不一致", "mapping_unknown") receipt = {**receipt, "detail": sanitize(detail), "observed_at": now().isoformat()} receipts[child] = receipt await self.checkpoint_receipt(attempt_id, child, receipt) await self.persist_receipt(attempt_id, child, receipt, count) except (VerificationRequired, SimulationDeferred): raise except WqError as exc: if exc.code in ("authentication_failed", "disconnected"): raise collection_errors.append((child, str(exc))) async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) items = (await db.scalars(select(BacktestItem).where(BacktestItem.attempt_id == a.id))).all() terminal = all(i.persistence_status == "saved" or i.platform_status == "failed" for i in items) all_children_done = bool(children) and all( receipts.get(c, {}).get("progress", {}).get("status") in TERMINAL for c in children ) a.remote_complete = all_children_done and len(children) == count if collection_errors: a.state, a.error, a.error_code = ( "collection_failed", collection_errors[0][1], "collection_failed", ) for i in items: if i.persistence_status != "saved" and i.platform_status != "failed": i.collection_status, i.error = "failed", a.error elif terminal and len(children) == count: a.state = "failed" if any(i.platform_status == "failed" for i in items) else "completed" a.error, a.error_code = None, None elif all_children_done: a.state, a.error, a.error_code = ( "needs_review", "部分子结果缺失或不能唯一匹配输入,请核对", "mapping_unknown", ) for i in items: if i.persistence_status != "saved" and i.platform_status != "failed": i.platform_status, i.error = "unknown", a.error a.poll_count += 1 a.next_poll_at = now() + timedelta(seconds=delay) await refresh_status(db, run) await event(db, run, "progress", {"attempt_id": a.id, "state": a.state}) def safe_progress(self, value): # Store useful protocol evidence, never arbitrary upstream diagnostics or credentials. result = {k: value[k] for k in ("status", "type", "alpha", "regular", "selection", "combo", "settings", "location", "warnings") if k in value} message = value.get("error") or value.get("message") if isinstance(message, str): for secret in list(self.client.credentials or ()) + list(self.client.client.cookies.values()): if secret: message = message.replace(secret, "[redacted]") result["message"] = message[:1000] return sanitize(result) async def quota(self, progress): location = progress.get("location") if isinstance(location, dict) and location.get("type") == "DAILY_SIMULATION_LIMIT": async with self.sessions.begin() as db: config = await db.get(BacktestConfig, 1) config.blocked_reason, config.blocked_until = ( "平台反馈每日模拟限额;恢复额度后显式继续运行", None, ) async def persist_receipt(self, attempt_id, child, receipt, count): progress, detail = receipt["progress"], receipt.get("detail") async with self.sessions.begin() as db: a = await db.get(SimulationAttempt, attempt_id) run = await locked_run(db, a.run_id) items = list( await db.scalars( select(BacktestItem).where(BacktestItem.attempt_id == a.id).order_by(BacktestItem.ordinal) ) ) bound = next((i for i in items if i.simulation_id == child), None) if bound and bound.persistence_status == "saved": return evidence = detail or progress expression, settings = code(evidence.get("regular")), evidence.get("settings") matched = [ i for i in items if ((i.alpha_type == "REGULAR" and evidence.get("type", "REGULAR") == "REGULAR" and i.expression == expression) or (i.alpha_type == "SUPER" and evidence.get("type") == "SUPER" and i.selection == code(evidence.get("selection")) and i.combo == code(evidence.get("combo")))) and isinstance(settings, dict) and all(k in settings and settings[k] == v for k, v in i.settings.items()) ] if count == 1 and items[0].alpha_type == "REGULAR": matched = ( items if (expression == items[0].expression or (not expression and detail is None)) and ( not isinstance(settings, dict) or all(k not in settings or settings[k] == v for k, v in items[0].settings.items()) ) else [] ) if count == 1 and items[0].alpha_type == "SUPER" and detail is None: # A known receipt can record progress, but saving SUPER requires full type/input evidence. matched = items if not any(k in evidence for k in ("type", "selection", "combo", "settings")) else matched # Identical inputs within a multi-submit are intentionally not position-matched. if len(matched) != 1 or (matched[0].simulation_id not in (None, child)): return item = matched[0] item.simulation_id = child if detail is not None: item.platform_status, item.collection_status = "completed", "complete" item.alpha_id, item.error = detail["id"], None # Account lock also serializes Alpha upserts against the sync lane. await db.scalar(select(Account).where(Account.id == 1).with_for_update()) await upsert_alpha(db, detail) if not await db.get(BacktestResult, item.id): from datetime import datetime if item.alpha_type == "SUPER": from ..superalpha.evidence import save_actual_components await save_actual_components(db, item, detail, receipt["observed_at"]) db.add( BacktestResult( item_id=item.id, attempt_id=a.id, alpha_id=detail["id"], snapshot=sanitize(detail), observed_at=datetime.fromisoformat(receipt["observed_at"]), complete=True, ) ) item.persistence_status = "saved" elif progress.get("alpha") and progress.get("status") in ("COMPLETE", "WARNING"): item.platform_status, item.collection_status = "completed", "collecting" item.alpha_id, item.error = progress["alpha"], None elif progress.get("status") in TERMINAL: item.platform_status, item.collection_status, item.persistence_status = ( "failed", "not_required", "not_required", ) item.error = progress.get("message") or "平台模拟失败或未返回 Alpha 标识" await event( db, run, "item_result", { "item_id": item.id, "platform_status": item.platform_status, "persistence_status": item.persistence_status, "alpha_id": item.alpha_id, }, )