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
+1
View File
@@ -0,0 +1 @@
"""WorldQuant research execution; callers never manage platform batches or polling."""
+218
View File
@@ -0,0 +1,218 @@
"""Fixed, typed inputs shared by HTTP, AI and research producers."""
import hashlib
import json
from typing import Literal
from pydantic import Field, field_validator, model_validator
from ..schemas import Contract
class SimulationSettings(Contract):
instrumentType: Literal["EQUITY"] = "EQUITY"
region: str = Field(min_length=1, max_length=50, pattern=r"^[A-Z0-9_]+$")
universe: str = Field(min_length=1, max_length=100, pattern=r"^[A-Z0-9_]+$")
delay: Literal[0, 1]
decay: int = Field(default=0, ge=0, le=10000)
neutralization: str = Field(default="INDUSTRY", min_length=1, max_length=50, pattern=r"^[A-Z_]+$")
truncation: float = Field(default=0.08, ge=0, le=1)
pasteurization: Literal["ON", "OFF"] = "ON"
unitHandling: Literal["VERIFY"] = "VERIFY"
nanHandling: Literal["ON", "OFF"] = "OFF"
language: Literal["FASTEXPR"] = "FASTEXPR"
visualization: bool = False
maxTrade: Literal["ON", "OFF"] = "OFF"
class Candidate(Contract):
client_item_id: str = Field(min_length=1, max_length=100)
expression: str = Field(min_length=1, max_length=20000)
settings: SimulationSettings
alpha_type: Literal["REGULAR"] = "REGULAR"
@field_validator("expression")
@classmethod
def nonempty(cls, value):
value = value.strip()
if not value:
raise ValueError("表达式不能为空")
return value
def platform_input(self):
return {"type": self.alpha_type, "regular": self.expression, "settings": self.settings.model_dump()}
class Source(Contract):
kind: str = Field(default="manual", min_length=1, max_length=100)
reference: str | None = Field(default=None, max_length=200)
batch_id: str | None = Field(default=None, max_length=200)
template_input_id: str | None = Field(default=None, max_length=200)
research_id: str | None = Field(default=None, max_length=200)
parent_run_id: str | None = Field(default=None, max_length=36)
class DraftInput(Contract):
name: str = Field(min_length=1, max_length=200)
source: Source = Field(default_factory=Source)
candidates: list[Candidate] = Field(min_length=1, max_length=10000)
@model_validator(mode="after")
def unique_ids(self):
if len({c.client_item_id for c in self.candidates}) != len(self.candidates):
raise ValueError("client_item_id 在候选集合内必须唯一")
return self
class DraftUpdate(DraftInput):
version: int = Field(ge=1)
class PreviewInput(Contract):
draft_id: str | None = Field(default=None, max_length=36)
draft_version: int | None = Field(default=None, ge=1)
selection: list[str] | None = Field(default=None, min_length=1, max_length=10000)
inline: DraftInput | None = None
@model_validator(mode="after")
def one_input(self):
if (self.inline is None) == (self.draft_id is None):
raise ValueError("必须提供 inline 或 draft_id 之一")
if self.draft_id and self.draft_version is None:
raise ValueError("引用草稿时必须提供 draft_version")
if self.inline and (self.draft_version is not None or self.selection is not None):
raise ValueError("inline 已经是完整固定集合")
return self
class StartInput(Contract):
preview_id: str = Field(min_length=1, max_length=36)
version: int = Field(default=1, ge=1)
idempotency_key: str = Field(min_length=1, max_length=100)
class ControlInput(Contract):
action: Literal["pause", "resume", "stop", "recover"]
version: int = Field(ge=1)
class RerunInput(Contract):
item_ids: list[str] = Field(min_length=1, max_length=10000)
class SchedulerInput(Contract):
concurrency: int = Field(default=3, ge=1, le=8)
batch_size: int = Field(default=8, ge=1, le=10)
version: int = Field(ge=1)
def fingerprint(payload: dict) -> str:
return hashlib.sha256(json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
def group_key(candidate: dict):
settings = candidate["settings"]
return tuple(settings[k] for k in ("region", "delay", "language", "instrumentType"))
class ReferenceInput(Contract):
progress_url: str = Field(min_length=1, max_length=2000)
version: int = Field(ge=1)
class SubsetInput(Contract):
exclude_ids: list[str] = Field(min_length=1, max_length=10000)
# OpenAPI outputs deliberately keep platform snapshots as extensible objects.
class SchedulerOutput(Contract):
concurrency: int
batch_size: int
version: int
blocked_reason: str | None
blocked_until: str | None
class PreviewOutput(Contract):
preview_id: str
version: int
name: str
source: Source
digest: str
total: int
batch_count: int
batch_size: int
duplicate_count: int
duplicates: list[dict]
items: list[Candidate]
limit: int
offset: int
has_more: bool
created_at: str
class RunOutput(Contract):
backtest_run_id: str
preview_id: str
name: str
source: Source
ai_context: dict
control: Literal["active", "paused", "stopped"]
status: str
version: int
total: int
batch_size: int
created_at: str
updated_at: str
counts: dict[str, dict[str, int]]
cursor: int
scheduler: SchedulerOutput
class RunPage(Contract):
items: list[RunOutput]
total: int
limit: int
offset: int
class ResultSnapshot(Contract):
snapshot: dict
observed_at: str
complete: bool
class ItemOutput(Contract):
id: str
client_item_id: str
expression: str
settings: SimulationSettings
attempt_id: str
platform_status: str
collection_status: str
persistence_status: str
simulation_id: str | None
alpha_id: str | None
error: str | None
result: ResultSnapshot | None
class ResultPage(Contract):
backtest_run_id: str
total: int
limit: int
offset: int
items: list[ItemOutput]
class EventOutput(Contract):
seq: int
kind: str
payload: dict
created_at: str
class EventPage(Contract):
items: list[EventOutput]
next_cursor: int
has_more: bool
+166
View File
@@ -0,0 +1,166 @@
"""Authenticated adapters; every mutation is committed before the execution lane wakes."""
from fastapi import APIRouter, Depends, Query, Request
from ..business import Business
from ..security import require_auth
from .contracts import (
ControlInput,
DraftInput,
DraftUpdate,
EventPage,
PreviewInput,
PreviewOutput,
ReferenceInput,
RerunInput,
ResultPage,
RunOutput,
RunPage,
SchedulerInput,
SchedulerOutput,
StartInput,
SubsetInput,
)
router = APIRouter(prefix="/api/v1/backtests", tags=["backtests"], dependencies=[Depends(require_auth)])
@router.get("/capabilities")
async def capabilities(request: Request):
async with request.app.state.sessions() as db:
return await Business(db).backtests.capabilities()
@router.get("/config", response_model=SchedulerOutput)
async def config(request: Request):
async with request.app.state.sessions() as db:
return await Business(db).backtests.config()
@router.put("/config", response_model=SchedulerOutput)
async def configure(body: SchedulerInput, request: Request):
async with request.app.state.sessions.begin() as db:
result = await Business(db).backtests.configure(body)
request.app.state.runner.backtests.wake.set()
return result
@router.get("/drafts")
async def drafts(request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)):
async with request.app.state.sessions() as db:
return await Business(db).backtests.drafts(limit, offset)
@router.post("/drafts", status_code=201)
async def save_draft(body: DraftInput, request: Request):
async with request.app.state.sessions.begin() as db:
return await Business(db).backtests.save_draft(body)
@router.get("/drafts/{draft_id}")
async def draft(draft_id: str, request: Request):
async with request.app.state.sessions() as db:
return await Business(db).backtests.draft(draft_id)
@router.put("/drafts/{draft_id}")
async def update_draft(draft_id: str, body: DraftUpdate, request: Request):
async with request.app.state.sessions.begin() as db:
return await Business(db).backtests.save_draft(body, draft_id)
@router.post("/previews", status_code=201, response_model=PreviewOutput)
async def preview(body: PreviewInput, request: Request):
async with request.app.state.sessions.begin() as db:
return await Business(db).backtests.preview(body)
@router.get("/previews/{preview_id}", response_model=PreviewOutput)
async def get_preview(
preview_id: str, request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)
):
async with request.app.state.sessions() as db:
return await Business(db).backtests.get_preview(preview_id, limit, offset)
@router.post("/runs", status_code=202, response_model=RunOutput)
async def start(body: StartInput, request: Request):
async with request.app.state.sessions.begin() as db:
result = await Business(db).backtests.start(body)
request.app.state.runner.backtests.wake.set()
return result
@router.get("/runs", response_model=RunPage)
async def runs(
request: Request,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
source: str | None = Query(None, max_length=100),
):
async with request.app.state.sessions() as db:
return await Business(db).backtests.runs(limit, offset, source)
@router.get("/runs/{run_id}", response_model=RunOutput)
async def run(run_id: str, request: Request):
async with request.app.state.sessions() as db:
return await Business(db).backtests.run(run_id)
@router.get("/runs/{run_id}/results", response_model=ResultPage)
async def results(
run_id: str, request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)
):
async with request.app.state.sessions() as db:
return await Business(db).backtests.results(run_id, limit, offset)
@router.get("/runs/{run_id}/events", response_model=EventPage)
async def events(
run_id: str, request: Request, after: int = Query(0, ge=0), limit: int = Query(100, ge=1, le=100)
):
async with request.app.state.sessions() as db:
return await Business(db).backtests.events(run_id, after, limit)
@router.get("/runs/{run_id}/attempts")
async def attempts(run_id: str, request: Request):
async with request.app.state.sessions() as db:
return await Business(db).backtests.attempts(run_id)
@router.post("/runs/{run_id}/control", response_model=RunOutput)
async def control(run_id: str, body: ControlInput, request: Request):
async with request.app.state.sessions.begin() as db:
result = await Business(db).backtests.control(run_id, body)
request.app.state.runner.backtests.wake.set()
return result
@router.post("/runs/{run_id}/rerun-preview", status_code=201, response_model=PreviewOutput)
async def rerun(run_id: str, body: RerunInput, request: Request):
async with request.app.state.sessions.begin() as db:
return await Business(db).backtests.rerun(run_id, body)
@router.post("/attempts/{attempt_id}/reference", response_model=RunOutput)
async def attach_reference(attempt_id: str, body: ReferenceInput, request: Request):
from fastapi import HTTPException
from ..worldquant import WqError
try:
body.progress_url = request.app.state.runner.client.simulation_url(body.progress_url)
except WqError as exc:
raise HTTPException(422, str(exc)) from None
async with request.app.state.sessions.begin() as db:
result = await Business(db).backtests.attach_reference(attempt_id, body)
request.app.state.runner.backtests.wake.set()
return result
@router.post("/previews/{preview_id}/subset", status_code=201, response_model=PreviewOutput)
async def subset(preview_id: str, body: SubsetInput, request: Request):
async with request.app.state.sessions.begin() as db:
return await Business(db).backtests.subset(preview_id, body)
+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,
},
)
+607
View File
@@ -0,0 +1,607 @@
"""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"],
"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 save_draft(self, body, draft_id=None):
data = body.model_dump(mode="json", exclude={"version"})
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):
rows = (
await self.db.scalars(
select(BacktestDraft)
.order_by(BacktestDraft.updated_at.desc(), 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(BacktestDraft)),
"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):
if body.inline:
data = body.inline.model_dump(mode="json")
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}
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:
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,
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):
query = select(BacktestRun)
if source:
query = query.where(BacktestRun.source["kind"].as_string() == source)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = (
await self.db.scalars(
query.order_by(BacktestRun.created_at.desc(), 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 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",
"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
)
for r in selected
],
)
)
)
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))
)