feat: add durable WorldQuant backtests with UI and AI confirmation

This commit is contained in:
yuxuanhui
2026-09-08 10:06:00 +08:00
parent 404a4d8a04
commit a4b93200c5
34 changed files with 4437 additions and 23 deletions
+547
View File
@@ -0,0 +1,547 @@
"""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", "alpha", "regular", "settings", "location") 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.expression == expression
and isinstance(settings, dict)
and all(k in settings and settings[k] == v for k, v in i.settings.items())
]
if count == 1:
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 []
)
# 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
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,
},
)