feat: add scoped alpha sync and local self-correlation

This commit is contained in:
yuxuanhui
2026-09-08 10:40:49 +08:00
parent 404a4d8a04
commit d2ccd94721
28 changed files with 1656 additions and 69 deletions
+218 -19
View File
@@ -8,14 +8,15 @@ a distributed lease and session coordinator.
import asyncio
import logging
from contextlib import suppress
from datetime import timedelta
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 pnl_points, sanitize, upsert_alpha
from .models import Account, Alpha, Job, JobItem, Pnl, now
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
@@ -235,7 +236,10 @@ class Runner:
async with self.sessions() as db:
job = await db.get(Job, job_id)
kind, payload = job.kind, job.payload
if kind == "verify":
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:
@@ -245,7 +249,7 @@ class Runner:
await self.ensure_connected(force=kind == "connect")
if kind in ("connect", "profile"):
await self.refresh_profile()
elif kind == "full_sync":
elif kind in ("full_sync", "daily_sync"):
await self.sync_all(job_id)
else:
await self.sync_ids(job_id, kind, payload["alpha_ids"])
@@ -296,25 +300,68 @@ class Runner:
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()
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 = partitions[partition]
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:
await self.checkpoint(job_id, {"next_retry_at": None})
raw = await self.client.alphas(submission, hidden, offset, before)
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:
@@ -335,9 +382,12 @@ class Runner:
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:
@@ -389,11 +439,7 @@ class Runner:
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()
await self.save_pnl(db, alpha_id, raw, points)
else:
await upsert_alpha(db, raw)
previous.error = error
@@ -403,3 +449,156 @@ class Runner:
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()