Files
worldquant-alpha-system/backend/app/jobs.py
T
yuxuanhui e256d6fef1
Deploy production / deploy (push) Successful in 56s
feat: add Super Alpha research, management and MCP workflows
2026-09-13 12:32:16 +08:00

641 lines
29 KiB
Python

"""Durable, single-process task runner.
Each page and its checkpoint commit in one transaction. A single Uvicorn process
owns the runner and shared upstream session; do not scale replicas before adding
a distributed lease and session coordinator.
"""
import asyncio
import logging
from contextlib import suppress
from datetime import date, datetime, time, timedelta, timezone
from uuid import uuid4
from sqlalchemy import select, update
from sqlalchemy.exc import SQLAlchemyError
from .alphas import invalidate_correlations, pnl_points, sanitize, submission_condition, upsert_alpha
from .correlation import calculate_correlation
from .models import Account, Alpha, Job, JobItem, Pnl, SelfCorrelation, now
from .security import cipher
from .worldquant import VerificationRequired, WqClient, WqError
logger = logging.getLogger(__name__)
ACTIVE = ("queued", "running", "waiting_auth", "waiting_connection")
AUTH_KINDS = ("connect", "verify", "profile")
async def create_job(db, kind, payload=None):
job = Job(id=str(uuid4()), kind=kind, payload=payload or {}, checkpoint={})
db.add(job)
await db.commit()
return job
class Runner:
def __init__(self, sessions, settings, client=None):
self.sessions, self.settings = sessions, settings
self.client = client or WqClient(settings)
self.loop_task = None
self.active_task = None
self.active_id = None
self.stopping = False
self.disconnecting = False
self.control_lock = asyncio.Lock()
self.recover_database = False
self.wake = asyncio.Event()
from .backtests.runtime import BacktestLane
self.backtests = BacktestLane(self)
async def start(self):
async with self.sessions() as db:
await db.execute(update(Job).where(Job.status == "running").values(status="queued"))
account = await db.get(Account, 1)
if account.connection_status in ("connected", "connecting", "verification_required"):
account.connection_status = "expired"
account.verification_url = None
await db.commit()
self.loop_task = asyncio.create_task(self.run_loop())
await self.backtests.start()
async def stop(self):
self.stopping = True
self.wake.set()
if self.active_task:
self.active_task.cancel()
if self.loop_task:
await self.loop_task
await self.backtests.stop()
await self.client.close()
async def cancel(self, job_id):
if self.active_id == job_id and self.active_task:
self.active_task.cancel()
with suppress(asyncio.CancelledError):
await self.active_task
async def disconnect(self):
# Prevent the scheduler from starting another request while credentials are cleared.
async with self.control_lock:
self.disconnecting = True
try:
if self.active_task:
await self.cancel(self.active_id)
await self.backtests.interrupt()
self.client.disconnect()
async with self.sessions() as db:
account = await db.get(Account, 1)
account.connection_status = "disconnected"
account.verification_url = None
account.connection_error = None
await db.execute(
update(Job).where(Job.status.in_(ACTIVE)).values(status="waiting_connection")
)
await db.commit()
finally:
self.disconnecting = False
async def run_loop(self):
while not self.stopping:
try:
await self.run_next()
except (SQLAlchemyError, OSError):
# A temporary DB outage must not silently kill the scheduler.
self.recover_database = True
logger.warning("Task runner is waiting for database recovery")
await self.wait_for_work()
async def wait_for_work(self):
self.wake.clear()
if not self.stopping:
try:
await asyncio.wait_for(self.wake.wait(), timeout=1)
except TimeoutError:
pass
async def run_next(self):
async with self.control_lock:
async with self.sessions() as db:
if self.recover_database:
await db.execute(update(Job).where(Job.status == "running").values(status="queued"))
await db.commit()
self.recover_database = False
job = (
await db.scalars(
select(Job)
.where(Job.status == "queued")
.order_by(Job.kind.in_(AUTH_KINDS).desc(), Job.created_at)
.limit(1)
)
).first()
job_id = job.id if job else None
if job_id and not self.stopping:
self.active_id = job_id
self.active_task = asyncio.create_task(self.execute(job_id))
if job_id and self.active_task:
try:
with suppress(asyncio.CancelledError):
await self.active_task
finally:
self.active_id, self.active_task = None, None
else:
await self.wait_for_work()
async def set_account(self, status, error=None, url=None):
async with self.sessions() as db:
account = await db.get(Account, 1)
account.connection_status, account.connection_error, account.verification_url = status, error, url
await db.commit()
async def ensure_connected(self, force=False):
async with self.sessions() as db:
account = await db.get(Account, 1)
if (
not account.email
or not account.password_encrypted
or (account.connection_status == "disconnected" and not force)
):
raise WqError("请先配置并连接 WorldQuant", "disconnected")
password = cipher(self.settings).decrypt(account.password_encrypted.encode()).decode()
email = account.email
if self.client.verification_url and not force:
raise VerificationRequired(self.client.verification_url)
await self.client.authenticate(email, password, force=force)
await self.set_account("connected")
async def refresh_profile(self):
raw = await self.client.profile()
# Confirm identity before fetching or storing additional account data.
async with self.sessions() as db:
account = await db.get(Account, 1)
user_id = raw.get("id")
if not user_id or (account.wq_user_id and account.wq_user_id != str(user_id)):
self.client.disconnect()
raise WqError("平台账户身份不匹配,请核对凭据", "identity_mismatch")
usage = await self.client.account_usage()
async with self.sessions() as db:
account = await db.get(Account, 1)
account.wq_user_id = str(user_id)
# An allowlist avoids persisting unknown personal/security fields.
account.profile = sanitize(
{
k: raw[k]
for k in (
"id",
"username",
"name",
"email",
"firstName",
"lastName",
"fullName",
"level",
"geniusLevel",
"verified",
"approved",
"dateCreated",
"dateVerified",
"dateApproved",
"role",
"roles",
"permissions",
"type",
)
if k in raw
}
)
if self.client.permissions is not None:
account.profile = {**account.profile, "permissions": self.client.permissions}
account.profile = {**account.profile, "usage": usage}
account.connection_status, account.connection_error = "connected", None
account.verification_url, account.last_synced_at = None, now()
await db.execute(
update(Job)
.where(Job.status.in_(("waiting_auth", "waiting_connection")), ~Job.kind.in_(AUTH_KINDS))
.values(status="queued", error=None)
)
# Superseded connect/verify attempts are complete after identity confirmation.
await db.execute(
update(Job)
.where(Job.status.in_(("waiting_auth", "waiting_connection")), Job.kind.in_(AUTH_KINDS))
.values(status="completed", error=None)
)
await db.commit()
async def checkpoint(self, job_id, values):
async with self.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
for key, value in values.items():
setattr(job, key, value)
job.updated_at = now()
await db.commit()
async def execute(self, job_id):
async def retrying(delay):
await self.checkpoint(job_id, {"next_retry_at": now() + timedelta(seconds=delay)})
self.client.on_retry = retrying
try:
await self.checkpoint(job_id, {"status": "running", "error": None, "next_retry_at": None})
async with self.sessions() as db:
job = await db.get(Job, job_id)
kind, payload = job.kind, job.payload
if kind in ("self_correlation", "self_correlation_recheck"):
# Cached local comparisons also work while the platform is disconnected.
await self.check_correlations(job_id, payload["alpha_ids"])
elif kind == "verify":
if not self.client.verification_url:
await self.ensure_connected(force=True)
else:
await self.client.verify()
await self.refresh_profile()
else:
await self.ensure_connected(force=kind == "connect")
if kind in ("connect", "profile"):
await self.refresh_profile()
elif kind == "catalog_full_sync":
from .catalog.sync import sync_full_catalog
await sync_full_catalog(self, job_id, payload)
elif kind in ("catalog_sync", "field_sync"):
from .catalog.sync import sync_catalog
await sync_catalog(self, job_id, payload)
elif kind in ("full_sync", "daily_sync"):
await self.sync_all(job_id)
elif kind == "super_selection_preview":
from .superalpha.jobs import run_selection
await run_selection(self, job_id, payload)
elif kind == "submission_check":
from .submission import run_check
await run_check(self, job_id, payload)
else:
await self.sync_ids(job_id, kind, payload["alpha_ids"])
async with self.sessions() as db:
job = await db.get(Job, job_id)
job.status = "completed_with_errors" if job.failed else "completed"
job.next_retry_at, job.updated_at = None, now()
await db.commit()
except asyncio.CancelledError:
async with self.sessions() as db:
job = await db.get(Job, job_id)
job.status = (
"cancelled"
if job.cancel_requested
else "waiting_connection"
if self.disconnecting
else "queued"
if self.stopping
else "cancelled"
)
job.next_retry_at, job.updated_at = None, now()
await db.commit()
except VerificationRequired as exc:
await self.set_account("verification_required", str(exc), exc.url)
await self.checkpoint(
job_id, {"status": "waiting_auth", "error": str(exc), "next_retry_at": None}
)
except WqError as exc:
waiting = exc.code in ("disconnected", "authentication_failed", "identity_mismatch")
if waiting:
await self.set_account("disconnected" if exc.code == "disconnected" else "error", str(exc))
await self.checkpoint(
job_id,
{
"status": "waiting_connection" if waiting or (kind in ("catalog_full_sync", "super_selection_preview") and exc.code == "network_error") else "failed",
"error": str(exc),
"next_retry_at": None,
},
)
except Exception:
# Do not expose upstream bodies, decrypted credentials, or SQL bind parameters.
logger.error("Job %s failed with an internal error", job_id)
await self.checkpoint(
job_id,
{"status": "failed", "error": "任务内部错误,请检查服务日志后重试", "next_retry_at": None},
)
finally:
self.client.on_retry = None
async def sync_all(self, job_id):
async with self.sessions() as db:
job = await db.get(Job, job_id)
checkpoint = job.checkpoint
before = job.created_at.isoformat()
payload, daily = job.payload, job.kind == "daily_sync"
days = [None]
if daily:
first, last = date.fromisoformat(payload["date_from"]), date.fromisoformat(payload["date_to"])
days = [first + timedelta(days=i) for i in range((last - first).days + 1)]
# Preserve the scope and partition indexes of pre-upgrade queued jobs.
# New JobInput always fixes a full sync to SUBMITTED.
submissions = [payload["submission"]] if payload.get("submission") else ["UNSUBMITTED", "SUBMITTED"]
partitions = [
(submission, hidden, day)
for day in days
for submission in submissions
for hidden in (False, True)
]
start_partition, offset = checkpoint.get("partition", 0), checkpoint.get("offset", 0)
for partition in range(start_partition, len(partitions)):
submission, hidden, day = partitions[partition]
day_params = {}
if day:
day_params = {
"date_from": datetime.combine(day, time.min, timezone.utc).isoformat(),
"date_to": datetime.combine(day + timedelta(days=1), time.min, timezone.utc).isoformat(),
}
while True:
progress = {"partition": partition, "offset": offset}
if daily:
progress.update(
date=day.isoformat(), dates_completed=partition // 2, dates_total=len(days)
)
await self.checkpoint(job_id, {"next_retry_at": None, "checkpoint": progress})
raw = await self.client.alphas(submission, hidden, offset, before, **day_params)
rows = raw.get("results")
if not isinstance(rows, list):
raise WqError("Alpha 列表缺少 results,已保留当前进度", "invalid_response")
if payload.get("submission"):
for raw_alpha in rows:
status = raw_alpha.get("status")
if not status or (status == "UNSUBMITTED") != (submission == "UNSUBMITTED"):
raise WqError(
"平台返回的 Alpha 不属于请求的提交分组,已保留进度", "invalid_response"
)
if day:
field = "dateCreated" if submission == "UNSUBMITTED" else "dateSubmitted"
try:
timestamp = datetime.fromisoformat(raw_alpha[field].replace("Z", "+00:00"))
timestamp = (
timestamp.replace(tzinfo=timezone.utc)
if timestamp.tzinfo is None
else timestamp
)
matches_day = timestamp.astimezone(timezone.utc).date() == day
except (ValueError, KeyError, AttributeError, TypeError):
matches_day = False
if not matches_day:
raise WqError(
"平台未按请求的日期返回 Alpha,已保留进度,请核对平台日期筛选支持",
"invalid_response",
)
async with self.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
await db.scalar(select(Account).where(Account.id == 1).with_for_update())
for raw_alpha in rows:
await upsert_alpha(db, raw_alpha)
if not await db.get(JobItem, (job_id, raw_alpha["id"])):
db.add(JobItem(job_id=job_id, alpha_id=raw_alpha["id"]))
job.processed += 1
# Trust explicit next/count when available; otherwise an empty page ends iteration.
count = raw.get("count")
more = bool(rows)
if "next" in raw:
more = bool(raw["next"])
elif isinstance(count, int):
more = offset + len(rows) < count and bool(rows)
if more and not rows:
raise WqError("平台分页未前进", "invalid_response")
offset += len(rows)
job.checkpoint = {
**progress,
"partition": partition if more else partition + 1,
"offset": offset if more else 0,
}
if daily:
job.checkpoint["dates_completed"] = job.checkpoint["partition"] // 2
job.updated_at = now()
await db.commit()
if not more:
break
offset = 0
async with self.sessions() as db:
job = await db.get(Job, job_id)
job.total = job.processed
account = await db.get(Account, 1)
account.last_synced_at = now()
await db.commit()
async def sync_ids(self, job_id, kind, alpha_ids):
"""Process fixed IDs, preserving existing PnL during a missing-only backfill.
Successful items commit individually so cancellation and retries keep
completed work. Cache presence is rechecked when a queued item runs.
"""
pnl_job = kind in ("pnl_refresh", "pnl_backfill")
await self.checkpoint(job_id, {"total": len(alpha_ids)})
for alpha_id in alpha_ids:
async with self.sessions() as db:
previous = await db.get(JobItem, (job_id, alpha_id))
if previous and not previous.error:
continue
cached = await db.get(Pnl, alpha_id) if kind == "pnl_backfill" else None
await self.checkpoint(job_id, {"next_retry_at": None, "checkpoint": {"alpha_id": alpha_id}})
error = None
try:
raw = cached.raw if cached is not None else await (
self.client.pnl(alpha_id) if pnl_job else self.client.alpha(alpha_id)
)
if not pnl_job and raw.get("id") != alpha_id:
raise ValueError("平台返回的 Alpha ID 与请求不一致")
points = cached.points if cached is not None else pnl_points(raw) if pnl_job else None
except VerificationRequired:
raise
except WqError as exc:
if exc.code not in ("not_found", "access_denied"):
raise
error = str(exc)
except ValueError as exc:
error = str(exc)
async with self.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
previous = await db.get(JobItem, (job_id, alpha_id))
if previous and previous.error:
job.failed -= 1
if not previous:
previous = JobItem(job_id=job_id, alpha_id=alpha_id)
db.add(previous)
if not error:
if pnl_job:
if not await db.get(Alpha, alpha_id):
error = "请先导入此 Alpha"
elif kind != "pnl_backfill" or await db.get(Pnl, alpha_id) is None:
await self.save_pnl(db, alpha_id, raw, points)
else:
await db.scalar(select(Account).where(Account.id == 1).with_for_update())
await upsert_alpha(db, raw)
previous.error = error
if error:
job.failed += 1
else:
job.processed += 1
job.updated_at = now()
await db.commit()
async def save_pnl(self, db, alpha_id, raw, points):
"""Persist a cache and invalidate conclusions that depend on the changed series."""
alpha = await db.get(Alpha, alpha_id)
pnl = await db.get(Pnl, alpha_id)
if pnl is None or pnl.points != points:
regions = (
[alpha.region] if alpha.region and alpha.status and alpha.status != "UNSUBMITTED" else []
)
await invalidate_correlations(db, alpha_id, regions)
if pnl is None:
pnl = Pnl(alpha_id=alpha_id)
db.add(pnl)
pnl.raw, pnl.points, pnl.fetched_at = sanitize(raw), points, now()
return pnl
async def correlation_pnl(self, job_id, alpha_id):
"""Use a local cache, filling a missing one through the read-only adapter."""
await self.checkpoint(job_id, {"next_retry_at": None})
async with self.sessions() as db:
pnl = await db.get(Pnl, alpha_id)
if pnl is not None:
return {
"alpha_id": alpha_id,
"points": pnl.points,
# SQLite drops the offset; stored datetimes are still UTC.
"fetched_at": pnl.fetched_at.replace(
tzinfo=pnl.fetched_at.tzinfo or timezone.utc
).isoformat(),
}
await self.ensure_connected()
raw = await self.client.pnl(alpha_id)
points = pnl_points(raw)
async with self.sessions() as db:
if (await db.get(Job, job_id)).cancel_requested:
raise asyncio.CancelledError()
pnl = await self.save_pnl(db, alpha_id, raw, points)
await db.commit()
return {"alpha_id": alpha_id, "points": points, "fetched_at": pnl.fetched_at.isoformat()}
async def check_correlations(self, job_id, alpha_ids):
"""Check fixed targets against same-region submitted caches, with per-target recovery.
An empty same-region reference set passes locally with correlation zero. Authentication,
transient upstream errors and cancellation retain the task for retry.
"""
await self.checkpoint(job_id, {"total": len(alpha_ids)})
reference_caches = {}
for alpha_id in alpha_ids:
async with self.sessions() as db:
previous = await db.get(JobItem, (job_id, alpha_id))
if previous and not previous.error:
continue
alpha = await db.get(Alpha, alpha_id)
region = alpha.region if alpha else None
reference_ids = (
list(
(
await db.scalars(
select(Alpha.id)
.where(
submission_condition("SUBMITTED"),
Alpha.region == region,
Alpha.id != alpha_id,
)
.order_by(Alpha.id)
)
).all()
)
if region
else []
)
await self.checkpoint(
job_id,
{
"checkpoint": {
"alpha_id": alpha_id,
"phase": "pnl",
"references_total": len(reference_ids),
"references_loaded": 0,
}
},
)
error, result = None, None
try:
if not alpha:
raise ValueError("Alpha 尚未同步")
if not region:
result = calculate_correlation([], [])
result["reason"] = "平台未提供目标 Alpha 的地区"
elif not reference_ids:
result = calculate_correlation([], [])
# A local policy default, not a measured Pearson coefficient.
result.update(
status="low",
max_correlation=0.0,
reason="没有同地区已提交 Alpha 可供比较,按本地规则视为通过,自相关值记为 0",
)
else:
target = await self.correlation_pnl(job_id, alpha_id)
references = []
for index, reference_id in enumerate(reference_ids):
if reference_id not in reference_caches:
try:
reference_caches[reference_id] = await self.correlation_pnl(
job_id, reference_id
)
except WqError as exc:
if exc.code not in ("not_found", "access_denied"):
raise
reference_caches[reference_id] = {"alpha_id": reference_id, "error": str(exc)}
except ValueError as exc:
reference_caches[reference_id] = {"alpha_id": reference_id, "error": str(exc)}
references.append(reference_caches[reference_id])
await self.checkpoint(
job_id,
{
"checkpoint": {
"alpha_id": alpha_id,
"phase": "pnl",
"references_total": len(reference_ids),
"references_loaded": index + 1,
}
},
)
await self.checkpoint(
job_id, {"checkpoint": {"alpha_id": alpha_id, "phase": "calculating"}}
)
result = await asyncio.to_thread(calculate_correlation, target["points"], references)
result["target_pnl_fetched_at"] = target["fetched_at"]
except WqError as exc:
if exc.code not in ("not_found", "access_denied"):
raise
error = str(exc)
except ValueError as exc:
error = str(exc)
async with self.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
previous = await db.get(JobItem, (job_id, alpha_id))
if previous and previous.error:
job.failed -= 1
if not previous:
previous = JobItem(job_id=job_id, alpha_id=alpha_id)
db.add(previous)
previous.error = error
if error:
job.failed += 1
else:
row = await db.get(SelfCorrelation, alpha_id)
if row is None:
row = SelfCorrelation(alpha_id=alpha_id)
db.add(row)
row.region, row.result, row.calculated_at, row.stale = region, result, now(), False
job.processed += 1
job.updated_at = now()
await db.commit()