"""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 == "self_correlation": # 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 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 == "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 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. Missing references are reported as incomplete coverage. 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([], []) result["reason"] = "没有同地区已提交 Alpha 可供比较,请先同步已提交 Alpha" 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()