feat: add scoped alpha sync and local self-correlation
This commit is contained in:
@@ -60,7 +60,7 @@ CATALOG = {
|
||||
"bulk_update_research": (BulkInput, "提出固定 1–100 个 Alpha 的批量标签或研究状态修改,等待用户确认。"),
|
||||
"create_sync_job": (
|
||||
JobInput,
|
||||
"提出全量同步、指定 Alpha 刷新或 PnL 刷新任务,等待确认;创建后立即返回任务 ID。",
|
||||
"提出同步或本地自相关任务,等待确认。full_sync 仅同步已提交;待提交必须用 daily_sync 并指定 submission、date_from/date_to(UTC),待提交按创建日、已提交按提交日逐天同步。alpha_refresh/pnl_refresh/self_correlation 使用固定 alpha_ids;自相关缺失 PnL 时自动补取,不触发平台检查。创建后立即返回任务 ID。",
|
||||
),
|
||||
"cancel_job": (JobArgs, "提出取消指定同步任务,等待用户确认。"),
|
||||
"retry_job": (JobArgs, "提出重试指定失败或暂停的同步任务,等待用户确认。"),
|
||||
|
||||
+28
-2
@@ -4,9 +4,24 @@ import math
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy import or_, select, update
|
||||
|
||||
from .models import Alpha, Research, ResearchTag, SelfCorrelation, now
|
||||
|
||||
|
||||
def submission_condition(submission):
|
||||
"""Match the platform list contract; a missing status is never assumed submitted."""
|
||||
return Alpha.status == "UNSUBMITTED" if submission == "UNSUBMITTED" else Alpha.status != "UNSUBMITTED"
|
||||
|
||||
|
||||
async def invalidate_correlations(db, alpha_id, regions=()):
|
||||
"""A changed baseline or PnL invalidates local conclusions without touching platform checks."""
|
||||
await db.execute(
|
||||
update(SelfCorrelation)
|
||||
.where(or_(SelfCorrelation.alpha_id == alpha_id, SelfCorrelation.region.in_(regions)))
|
||||
.values(stale=True)
|
||||
)
|
||||
|
||||
from .models import Alpha, Research, ResearchTag, now
|
||||
|
||||
SENSITIVE_KEYS = {
|
||||
"password",
|
||||
@@ -67,6 +82,8 @@ async def upsert_alpha(db, raw: dict):
|
||||
if not isinstance(alpha_id, str) or not alpha_id:
|
||||
raise ValueError("Alpha 数据缺少 ID")
|
||||
item = await db.get(Alpha, alpha_id)
|
||||
previous_region = item.region if item else None
|
||||
previous_status = item.status if item else None
|
||||
if item is None:
|
||||
item = Alpha(id=alpha_id)
|
||||
db.add(item)
|
||||
@@ -78,6 +95,13 @@ async def upsert_alpha(db, raw: dict):
|
||||
item.alpha_type, item.language = raw.get("type"), settings.get("language")
|
||||
item.stage, item.status, item.hidden = raw.get("stage"), raw.get("status"), raw.get("hidden") is True
|
||||
item.region, item.universe = settings.get("region"), settings.get("universe")
|
||||
if previous_region != item.region or previous_status != item.status:
|
||||
regions = {
|
||||
region
|
||||
for region, status in ((previous_region, previous_status), (item.region, item.status))
|
||||
if region and status and status != "UNSUBMITTED"
|
||||
}
|
||||
await invalidate_correlations(db, alpha_id, regions)
|
||||
item.settings, item.is_metrics = sanitize(settings), sanitize(metrics)
|
||||
item.os_metrics = sanitize(raw.get("os")) if isinstance(raw.get("os"), dict) else {}
|
||||
item.checks = sanitize(metrics.get("checks") or raw.get("checks") or [])
|
||||
@@ -93,6 +117,8 @@ async def upsert_alpha(db, raw: dict):
|
||||
|
||||
def list_statement(filters):
|
||||
query = select(Alpha, Research).join(Research, Research.alpha_id == Alpha.id)
|
||||
if filters.submission:
|
||||
query = query.where(submission_condition(filters.submission))
|
||||
q = filters.q
|
||||
if q:
|
||||
pattern = "%" + q.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
|
||||
|
||||
+52
-4
@@ -4,6 +4,7 @@ Mutations never commit here, so the AI executor can atomically save their audit
|
||||
Job runner notifications must happen after commit, using ``notify_job``.
|
||||
"""
|
||||
|
||||
from datetime import timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException
|
||||
@@ -11,7 +12,7 @@ from sqlalchemy import delete, func, select, update
|
||||
|
||||
from .alphas import list_statement, sorted_statement, summary
|
||||
from .jobs import ACTIVE
|
||||
from .models import Account, Alpha, Job, JobItem, Pnl, Research, ResearchTag, now
|
||||
from .models import Account, Alpha, Job, JobItem, Pnl, Research, ResearchTag, SelfCorrelation, now
|
||||
from .schemas import AlphaDetail, AlphaPage, BulkUpdate, JobInput, JobOutput, ResearchUpdate, normalize_tags
|
||||
|
||||
|
||||
@@ -29,10 +30,51 @@ class Business:
|
||||
.offset(filters.offset)
|
||||
)
|
||||
).all()
|
||||
correlations = {
|
||||
row.alpha_id: self.correlation_summary(row)
|
||||
for row in (
|
||||
await self.db.scalars(
|
||||
select(SelfCorrelation).where(SelfCorrelation.alpha_id.in_([a.id for a, _ in rows]))
|
||||
)
|
||||
).all()
|
||||
}
|
||||
return AlphaPage(
|
||||
items=[summary(a, r) for a, r in rows], total=total, limit=filters.limit, offset=filters.offset
|
||||
items=[{**summary(a, r), "local_correlation": correlations.get(a.id)} for a, r in rows],
|
||||
total=total,
|
||||
limit=filters.limit,
|
||||
offset=filters.offset,
|
||||
).model_dump(mode="json")
|
||||
|
||||
@staticmethod
|
||||
def correlation_summary(row):
|
||||
return {
|
||||
**{
|
||||
key: row.result.get(key)
|
||||
for key in ("status", "max_correlation", "compared_count", "skipped_count")
|
||||
},
|
||||
"stale": row.stale,
|
||||
"calculated_at": row.calculated_at.replace(
|
||||
tzinfo=row.calculated_at.tzinfo or timezone.utc
|
||||
).isoformat(),
|
||||
}
|
||||
|
||||
async def get_self_correlation(self, alpha_id):
|
||||
if not await self.db.get(Alpha, alpha_id):
|
||||
raise HTTPException(404, "Alpha 尚未同步")
|
||||
row = await self.db.get(SelfCorrelation, alpha_id)
|
||||
return {
|
||||
"cached": row is not None,
|
||||
"result": {
|
||||
**row.result,
|
||||
"stale": row.stale,
|
||||
"calculated_at": row.calculated_at.replace(
|
||||
tzinfo=row.calculated_at.tzinfo or timezone.utc
|
||||
).isoformat(),
|
||||
}
|
||||
if row
|
||||
else None,
|
||||
}
|
||||
|
||||
async def get_alpha_facets(self):
|
||||
result = {}
|
||||
for key in ("region", "universe", "alpha_type", "language", "status", "stage"):
|
||||
@@ -126,9 +168,15 @@ class Business:
|
||||
|
||||
async def create_sync_job(self, body: JobInput):
|
||||
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||||
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
|
||||
if body.kind != "self_correlation" and (
|
||||
not account.password_encrypted or account.connection_status in ("disconnected", "error")
|
||||
):
|
||||
raise HTTPException(409, "请先连接 WorldQuant")
|
||||
payload = {"alpha_ids": body.alpha_ids}
|
||||
if body.kind == "self_correlation":
|
||||
found = set((await self.db.scalars(select(Alpha.id).where(Alpha.id.in_(body.alpha_ids)))).all())
|
||||
if found != set(body.alpha_ids):
|
||||
raise HTTPException(404, "部分 Alpha 尚未同步")
|
||||
payload = body.model_dump(mode="json", exclude={"kind"}, exclude_none=True)
|
||||
for job in (
|
||||
await self.db.scalars(select(Job).where(Job.kind == body.kind, Job.status.in_(ACTIVE)))
|
||||
).all():
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
"""Local Pearson comparison of daily PnL changes; no platform eligibility decisions.
|
||||
|
||||
The four-year window and signed 0.7 threshold follow the legacy research tool.
|
||||
Thirty paired changes is a local minimum, not a claim about BRAIN's own checks.
|
||||
"""
|
||||
|
||||
import math
|
||||
from datetime import datetime, timezone
|
||||
from statistics import StatisticsError, correlation
|
||||
|
||||
THRESHOLD = 0.7
|
||||
MIN_SAMPLES = 30
|
||||
WINDOW_YEARS = 4
|
||||
|
||||
|
||||
def daily_changes(points):
|
||||
"""Return dated changes and the final date, rejecting ambiguous daily records.
|
||||
|
||||
Parameters are normalized PnL points. Missing values break a change interval;
|
||||
each change retains its start date so differently spaced samples never pair.
|
||||
Raises ValueError for malformed dates, duplicate days or non-finite values.
|
||||
"""
|
||||
daily = {}
|
||||
for point in points:
|
||||
try:
|
||||
timestamp = datetime.fromisoformat(point["date"].replace("Z", "+00:00"))
|
||||
day = (
|
||||
(timestamp.replace(tzinfo=timezone.utc) if timestamp.tzinfo is None else timestamp)
|
||||
.astimezone(timezone.utc)
|
||||
.date()
|
||||
)
|
||||
except (ValueError, TypeError, KeyError, AttributeError):
|
||||
raise ValueError("PnL 日期无法识别") from None
|
||||
if day in daily:
|
||||
raise ValueError("PnL 同一天存在多条记录")
|
||||
value = point.get("value")
|
||||
if value is not None and (
|
||||
isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value)
|
||||
):
|
||||
raise ValueError("PnL 包含无效数值")
|
||||
daily[day] = value
|
||||
changes, previous_day, previous = {}, None, None
|
||||
for day, value in sorted(daily.items()):
|
||||
if value is not None and previous is not None:
|
||||
delta = value - previous
|
||||
if not math.isfinite(delta):
|
||||
raise ValueError("PnL 变化超出有效数值范围")
|
||||
changes[day] = (previous_day, delta)
|
||||
previous_day, previous = day, value
|
||||
return changes, max(daily) if daily else None
|
||||
|
||||
|
||||
def calculate_correlation(target_points, references):
|
||||
"""Compare a target with supplied reference caches and return a JSON-safe report.
|
||||
|
||||
Each reference has alpha_id, points, fetched_at and optionally error. Callers
|
||||
select the same-region submitted set and exclude the target. Incomplete
|
||||
coverage can report high correlation, but can never report an all-clear.
|
||||
"""
|
||||
result = {
|
||||
"status": "insufficient_data",
|
||||
"threshold": THRESHOLD,
|
||||
"min_samples": MIN_SAMPLES,
|
||||
"window_years": WINDOW_YEARS,
|
||||
"window_from": None,
|
||||
"window_to": None,
|
||||
"max_correlation": None,
|
||||
"most_correlated_alpha_id": None,
|
||||
"candidate_count": len(references),
|
||||
"compared_count": 0,
|
||||
"skipped_count": 0,
|
||||
"matches": [],
|
||||
"skipped": [],
|
||||
"reason": None,
|
||||
}
|
||||
try:
|
||||
target, latest = daily_changes(target_points)
|
||||
except ValueError as exc:
|
||||
result["reason"] = str(exc)
|
||||
return result
|
||||
if latest is None or not target:
|
||||
result["reason"] = "目标 Alpha 没有可用的 PnL 日变化"
|
||||
return result
|
||||
try:
|
||||
cutoff = latest.replace(year=latest.year - WINDOW_YEARS)
|
||||
except ValueError:
|
||||
cutoff = latest.replace(year=latest.year - WINDOW_YEARS, day=28)
|
||||
result["window_from"], result["window_to"] = cutoff.isoformat(), latest.isoformat()
|
||||
matches, skipped = [], []
|
||||
for reference in references:
|
||||
alpha_id, reason = reference["alpha_id"], reference.get("error")
|
||||
if not reason:
|
||||
try:
|
||||
changes, _ = daily_changes(reference["points"])
|
||||
days = sorted(
|
||||
day
|
||||
for day in target.keys() & changes.keys()
|
||||
if cutoff < day <= latest and target[day][0] == changes[day][0]
|
||||
)
|
||||
if len(days) < MIN_SAMPLES:
|
||||
reason = f"共同有效样本不足 {MIN_SAMPLES} 个(实际 {len(days)})"
|
||||
else:
|
||||
coefficient = correlation([target[d][1] for d in days], [changes[d][1] for d in days])
|
||||
if not math.isfinite(coefficient):
|
||||
reason = "无法计算有效相关系数"
|
||||
else:
|
||||
matches.append(
|
||||
{
|
||||
"alpha_id": alpha_id,
|
||||
"correlation": coefficient,
|
||||
"sample_count": len(days),
|
||||
"date_from": days[0].isoformat(),
|
||||
"date_to": days[-1].isoformat(),
|
||||
"pnl_fetched_at": reference.get("fetched_at"),
|
||||
}
|
||||
)
|
||||
except StatisticsError:
|
||||
reason = "目标或基准 PnL 日变化为常量"
|
||||
except (ValueError, OverflowError) as exc:
|
||||
reason = str(exc)
|
||||
if reason:
|
||||
skipped.append({"alpha_id": alpha_id, "reason": reason})
|
||||
matches.sort(key=lambda row: (-row["correlation"], row["alpha_id"]))
|
||||
result.update(
|
||||
compared_count=len(matches), skipped_count=len(skipped), matches=matches[:10], skipped=skipped[:100]
|
||||
)
|
||||
if matches:
|
||||
maximum = matches[0]["correlation"]
|
||||
result.update(
|
||||
max_correlation=maximum,
|
||||
most_correlated_alpha_id=matches[0]["alpha_id"],
|
||||
status="high" if maximum >= THRESHOLD else "partial" if skipped else "low",
|
||||
)
|
||||
else:
|
||||
result["reason"] = (
|
||||
"没有同地区已提交 Alpha 可供比较" if not references else "所有基准均缺少足够的有效样本"
|
||||
)
|
||||
return result
|
||||
+218
-19
@@ -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()
|
||||
|
||||
@@ -40,6 +40,7 @@ from .schemas import (
|
||||
PnlOutput,
|
||||
PreferencesInput,
|
||||
ResearchUpdate,
|
||||
SelfCorrelationOutput,
|
||||
SessionOutput,
|
||||
)
|
||||
from .security import bootstrap, cipher, issue_session, require_auth, token_hash, valid_password
|
||||
@@ -341,6 +342,11 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
async with sessions() as db:
|
||||
return await Business(db).get_alpha_pnl(alpha_id)
|
||||
|
||||
@api.get("/alphas/{alpha_id}/self-correlation", response_model=SelfCorrelationOutput, tags=["alphas"])
|
||||
async def self_correlation(alpha_id: str):
|
||||
async with sessions() as db:
|
||||
return await Business(db).get_self_correlation(alpha_id)
|
||||
|
||||
@api.post("/sync-jobs", status_code=202, response_model=JobOutput, tags=["sync-jobs"])
|
||||
async def new_job(body: JobInput):
|
||||
async with sessions.begin() as db:
|
||||
|
||||
@@ -111,6 +111,15 @@ class ResearchTag(Base):
|
||||
tag: Mapped[str] = mapped_column(String(60), primary_key=True, index=True)
|
||||
|
||||
|
||||
class SelfCorrelation(Base):
|
||||
__tablename__ = "self_correlations"
|
||||
alpha_id: Mapped[str] = mapped_column(ForeignKey("alphas.id"), primary_key=True)
|
||||
region: Mapped[str | None] = mapped_column(String(50), index=True)
|
||||
result: Mapped[dict] = mapped_column(JSON)
|
||||
stale: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
calculated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class Job(Base):
|
||||
__tablename__ = "sync_jobs"
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
|
||||
+32
-4
@@ -1,13 +1,14 @@
|
||||
"""Validated public API contracts. Platform state is intentionally not a closed enum."""
|
||||
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from datetime import date, datetime, timezone
|
||||
from typing import Literal
|
||||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
ResearchState = Literal["inbox", "candidate", "optimizing", "archived"]
|
||||
Submission = Literal["UNSUBMITTED", "SUBMITTED"]
|
||||
SortField = Literal[
|
||||
"id",
|
||||
"name",
|
||||
@@ -62,6 +63,7 @@ class PreferencesInput(Contract):
|
||||
|
||||
|
||||
class AlphaFilters(Contract):
|
||||
submission: Submission | None = None
|
||||
q: str | None = Field(default=None, max_length=300)
|
||||
region: str | None = None
|
||||
universe: str | None = None
|
||||
@@ -160,16 +162,34 @@ class BulkUpdate(BulkInput):
|
||||
|
||||
|
||||
class JobInput(Contract):
|
||||
kind: Literal["full_sync", "alpha_refresh", "pnl_refresh"]
|
||||
kind: Literal["full_sync", "daily_sync", "alpha_refresh", "pnl_refresh", "self_correlation"]
|
||||
alpha_ids: list[str] = Field(default_factory=list)
|
||||
submission: Submission | None = None
|
||||
date_from: date | None = None
|
||||
date_to: date | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_ids(self):
|
||||
if self.kind == "full_sync":
|
||||
if self.kind in ("full_sync", "daily_sync"):
|
||||
if self.alpha_ids:
|
||||
raise ValueError("全量同步不接受 Alpha ID")
|
||||
raise ValueError("列表同步不接受 Alpha ID")
|
||||
if self.kind == "full_sync":
|
||||
if self.submission == "UNSUBMITTED":
|
||||
raise ValueError("待提交 Alpha 必须选择日期逐天同步")
|
||||
self.submission = "SUBMITTED"
|
||||
if self.date_from is not None or self.date_to is not None:
|
||||
raise ValueError("全量同步不接受日期范围")
|
||||
else:
|
||||
if not self.submission or not self.date_from or not self.date_to:
|
||||
raise ValueError("按天同步必须选择待提交/已提交和起止日期")
|
||||
if self.date_from > self.date_to:
|
||||
raise ValueError("开始日期不能晚于结束日期")
|
||||
if self.date_to > datetime.now(timezone.utc).date():
|
||||
raise ValueError("同步日期不能晚于今天(UTC)")
|
||||
else:
|
||||
self.alpha_ids = valid_ids(self.alpha_ids)
|
||||
if self.submission is not None or self.date_from is not None or self.date_to is not None:
|
||||
raise ValueError("按 ID 操作不接受分组或日期范围")
|
||||
return self
|
||||
|
||||
|
||||
@@ -199,6 +219,7 @@ class AlphaSummary(BaseModel):
|
||||
date_submitted: datetime | None
|
||||
synced_at: datetime
|
||||
research: ResearchOutput
|
||||
local_correlation: dict | None = None
|
||||
|
||||
|
||||
class AlphaDetail(AlphaSummary):
|
||||
@@ -230,6 +251,8 @@ class JobOutput(BaseModel):
|
||||
next_retry_at: datetime | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
payload: dict = Field(default_factory=dict)
|
||||
checkpoint: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class PlatformSessionOutput(BaseModel):
|
||||
@@ -286,6 +309,11 @@ class PnlOutput(BaseModel):
|
||||
fetched_at: datetime | None
|
||||
|
||||
|
||||
class SelfCorrelationOutput(BaseModel):
|
||||
cached: bool
|
||||
result: dict | None
|
||||
|
||||
|
||||
class FacetsOutput(BaseModel):
|
||||
region: list[str]
|
||||
universe: list[str]
|
||||
|
||||
+17
-12
@@ -257,20 +257,25 @@ class WqClient:
|
||||
errors[key] = str(exc)
|
||||
return usage_snapshot(data, errors)
|
||||
|
||||
async def alphas(self, submission, hidden, offset, before):
|
||||
async def alphas(self, submission, hidden, offset, before, *, date_from=None, date_to=None):
|
||||
# Cover every platform stage; submitted records are not assumed to be OS only.
|
||||
status_key = "status" if submission == "UNSUBMITTED" else "status!"
|
||||
return await self.get(
|
||||
"/users/self/alphas",
|
||||
{
|
||||
status_key: "UNSUBMITTED",
|
||||
"hidden": str(hidden).lower(),
|
||||
"limit": 100,
|
||||
"offset": offset,
|
||||
"order": "dateCreated",
|
||||
"dateCreated<": before,
|
||||
},
|
||||
)
|
||||
params = {
|
||||
status_key: "UNSUBMITTED",
|
||||
"hidden": str(hidden).lower(),
|
||||
"limit": 100,
|
||||
"offset": offset,
|
||||
"order": "dateCreated",
|
||||
"dateCreated<": before,
|
||||
}
|
||||
if date_from is not None and date_to is not None:
|
||||
# Daily intervals are [UTC midnight, next midnight). Submitted dates
|
||||
# use the actual submission time, independently of creation or stage.
|
||||
field = "dateCreated" if submission == "UNSUBMITTED" else "dateSubmitted"
|
||||
params.update({f"{field}>=": date_from, f"{field}<": date_to, "order": field})
|
||||
if field == "dateCreated":
|
||||
params["dateCreated<"] = min(before, date_to)
|
||||
return await self.get("/users/self/alphas", params)
|
||||
|
||||
async def alpha(self, alpha_id):
|
||||
return await self.get(f"/alphas/{alpha_id}")
|
||||
|
||||
Reference in New Issue
Block a user