Files
worldquant-alpha-system/backend/app/jobs.py
T

406 lines
17 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 timedelta
from uuid import uuid4
from sqlalchemy import select, update
from sqlalchemy.exc import SQLAlchemyError
from .alphas import pnl_points, sanitize, upsert_alpha
from .models import Account, Alpha, Job, JobItem, Pnl, 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()
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())
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.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)
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 == "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 == "full_sync":
await self.sync_all(job_id)
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 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):
partitions = [
("UNSUBMITTED", False),
("UNSUBMITTED", True),
("SUBMITTED", False),
("SUBMITTED", True),
]
async with self.sessions() as db:
job = await db.get(Job, job_id)
checkpoint = job.checkpoint
before = job.created_at.isoformat()
start_partition, offset = checkpoint.get("partition", 0), checkpoint.get("offset", 0)
for partition in range(start_partition, len(partitions)):
submission, hidden = partitions[partition]
while True:
await self.checkpoint(job_id, {"next_retry_at": None})
raw = await self.client.alphas(submission, hidden, offset, before)
rows = raw.get("results")
if not isinstance(rows, list):
raise WqError("Alpha 列表缺少 results,已保留当前进度", "invalid_response")
async with self.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
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 = {
"partition": partition if more else partition + 1,
"offset": offset if more else 0,
}
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):
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
await self.checkpoint(job_id, {"next_retry_at": None})
error = None
try:
raw = await (
self.client.pnl(alpha_id) if kind == "pnl_refresh" else self.client.alpha(alpha_id)
)
if kind != "pnl_refresh" and raw.get("id") != alpha_id:
raise ValueError("平台返回的 Alpha ID 与请求不一致")
points = pnl_points(raw) if kind == "pnl_refresh" 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 kind == "pnl_refresh":
if not await db.get(Alpha, alpha_id):
error = "请先导入此 Alpha"
else:
pnl = await db.get(Pnl, alpha_id)
if pnl is None:
pnl = Pnl(alpha_id=alpha_id)
db.add(pnl)
pnl.raw, pnl.points, pnl.fetched_at = sanitize(raw), points, now()
else:
await upsert_alpha(db, raw)
previous.error = error
if error:
job.failed += 1
else:
job.processed += 1
job.updated_at = now()
await db.commit()