406 lines
17 KiB
Python
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()
|