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}")
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Persist local self-correlation independently of platform snapshots."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0003"
|
||||
down_revision = "0002"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"self_correlations",
|
||||
sa.Column("alpha_id", sa.String(100), sa.ForeignKey("alphas.id"), primary_key=True),
|
||||
sa.Column("region", sa.String(50), nullable=True),
|
||||
sa.Column("result", sa.JSON(), nullable=False),
|
||||
sa.Column("stale", sa.Boolean(), nullable=False),
|
||||
sa.Column("calculated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
)
|
||||
op.create_index("ix_self_correlations_region", "self_correlations", ["region"])
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("self_correlations")
|
||||
@@ -45,7 +45,7 @@ def sample(index):
|
||||
"selection": {"code": "self_correlation < 0.5"} if super_alpha else None,
|
||||
"combo": {"code": "alpha"} if super_alpha else None,
|
||||
"settings": {
|
||||
"region": ["USA", "CHN", "EUR"][index % 3],
|
||||
"region": ["USA", "CHN", "EUR"][(index // 3) % 3],
|
||||
"universe": "TOP3000",
|
||||
"language": language,
|
||||
"delay": 1,
|
||||
@@ -64,7 +64,10 @@ def sample(index):
|
||||
"checks": [{"name": "LOW_SHARPE", "result": "PASS", "value": 2.1, "limit": 1.58}],
|
||||
},
|
||||
"os": {"sharpe": 1.1} if index % 3 == 0 else None,
|
||||
"dateCreated": (datetime(2025, 1, 1, tzinfo=timezone.utc) + timedelta(days=index)).isoformat(),
|
||||
"dateCreated": (datetime(2025, 1, 1, tzinfo=timezone.utc) + timedelta(seconds=index)).isoformat(),
|
||||
"dateSubmitted": (datetime(2025, 2, 1, tzinfo=timezone.utc) + timedelta(days=index % 2)).isoformat()
|
||||
if index % 3 == 0
|
||||
else None,
|
||||
}
|
||||
|
||||
|
||||
@@ -141,6 +144,21 @@ def create_test_app():
|
||||
matched = [
|
||||
r for r in records if (r["status"] == "UNSUBMITTED") == unsubmitted and r["hidden"] == hidden
|
||||
]
|
||||
for field in ("dateCreated", "dateSubmitted"):
|
||||
for suffix in (">=", "<"):
|
||||
boundary = request.url.params.get(field + suffix)
|
||||
if boundary:
|
||||
bound = datetime.fromisoformat(boundary).replace(tzinfo=timezone.utc)
|
||||
matched = [
|
||||
r
|
||||
for r in matched
|
||||
if r.get(field)
|
||||
and (
|
||||
datetime.fromisoformat(r[field]) >= bound
|
||||
if suffix == ">="
|
||||
else datetime.fromisoformat(r[field]) < bound
|
||||
)
|
||||
]
|
||||
offset, limit = (
|
||||
int(request.url.params.get("offset", 0)),
|
||||
int(request.url.params.get("limit", 100)),
|
||||
|
||||
@@ -0,0 +1,341 @@
|
||||
"""Behavioral coverage for submission scopes, daily recovery and local correlation."""
|
||||
|
||||
import csv
|
||||
import io
|
||||
from datetime import date, datetime, timedelta
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.alphas import upsert_alpha
|
||||
from app.correlation import calculate_correlation, daily_changes
|
||||
from app.jobs import Runner
|
||||
from app.models import Alpha, JobItem, Pnl, Research, SelfCorrelation
|
||||
from app.worldquant import WqClient, WqError
|
||||
from tests.conftest import alpha
|
||||
from tests.test_jobs import FakePlatform, ready_runner, result
|
||||
|
||||
PREFIX = "/api/v1"
|
||||
|
||||
|
||||
def points(changes, start=date(2025, 1, 1), multiplier=1):
|
||||
value, data = 100, [{"date": start.isoformat(), "value": 100}]
|
||||
for index, change in enumerate(changes, 1):
|
||||
value += change * multiplier
|
||||
data.append({"date": (start + timedelta(days=index)).isoformat(), "value": value})
|
||||
return data
|
||||
|
||||
|
||||
def reference(alpha_id, data):
|
||||
return {"alpha_id": alpha_id, "points": data, "fetched_at": "2025-03-01T00:00:00Z"}
|
||||
|
||||
|
||||
def test_signed_pearson_and_incomplete_coverage():
|
||||
changes = [((i * 7) % 19) - 8 for i in range(60)]
|
||||
target = points(changes)
|
||||
positive = reference("positive", points(changes, multiplier=2))
|
||||
negative = reference("negative", points(changes, multiplier=-1))
|
||||
computed = calculate_correlation(target, [positive, negative])
|
||||
assert computed["status"] == "high" and computed["compared_count"] == 2
|
||||
assert computed["most_correlated_alpha_id"] == "positive"
|
||||
assert computed["matches"][0]["sample_count"] == 60
|
||||
assert computed["max_correlation"] == pytest.approx(1)
|
||||
computed = calculate_correlation(target, [negative])
|
||||
assert computed["status"] == "low" and computed["max_correlation"] == pytest.approx(-1)
|
||||
incomplete = calculate_correlation(target, [negative, {"alpha_id": "missing", "error": "无权访问"}])
|
||||
assert incomplete["status"] == "partial" and incomplete["skipped_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("candidate", [[], points([1] * 60), points([1, 3, 2] * 9)])
|
||||
def test_insufficient_and_constant_series_are_never_passed(candidate):
|
||||
report = calculate_correlation(points([1, 3, 2] * 20), [reference("invalid", candidate)])
|
||||
assert report["status"] == "insufficient_data"
|
||||
assert report["max_correlation"] is None and report["compared_count"] == 0
|
||||
assert report["skipped_count"] == 1
|
||||
|
||||
|
||||
def test_missing_points_break_intervals_and_dates_align_before_comparison():
|
||||
original = points([1, 3, 2] * 20)
|
||||
missing = [dict(p) for p in original]
|
||||
missing[4]["value"] = None
|
||||
changes, _ = daily_changes(missing)
|
||||
assert date(2025, 1, 5) not in changes and date(2025, 1, 6) not in changes
|
||||
report = calculate_correlation(original, [reference("gap", list(reversed(missing)))])
|
||||
assert report["matches"][0]["sample_count"] == 58
|
||||
# A missing row must not pair a two-day increment with a one-day increment.
|
||||
removed = original[:4] + original[5:]
|
||||
report = calculate_correlation(original, [reference("gap", removed)])
|
||||
assert report["matches"][0]["sample_count"] == 58
|
||||
with pytest.raises(ValueError, match="同一天"):
|
||||
daily_changes(original + [original[0]])
|
||||
|
||||
|
||||
def test_common_four_year_window_is_anchored_to_target():
|
||||
changes = [1, 3, 2] * 20
|
||||
target = points(changes, date(2025, 1, 1))
|
||||
historical = points(changes, date(2019, 1, 1))
|
||||
report = calculate_correlation(target, [reference("old", historical)])
|
||||
assert report["status"] == "insufficient_data" and report["max_correlation"] is None
|
||||
assert report["window_from"] == "2021-03-02"
|
||||
|
||||
|
||||
async def test_submission_tabs_and_export_share_the_same_scope(app, logged_in):
|
||||
async with app.state.sessions() as db:
|
||||
for raw in (
|
||||
alpha("pending"),
|
||||
alpha("active", status="ACTIVE", stage="IS"),
|
||||
alpha("retired", status="DECOMMISSIONED"),
|
||||
alpha("unknown", status=None),
|
||||
):
|
||||
await upsert_alpha(db, raw)
|
||||
await db.commit()
|
||||
pending = (await logged_in.get(f"{PREFIX}/alphas?submission=UNSUBMITTED")).json()
|
||||
submitted = (await logged_in.get(f"{PREFIX}/alphas?submission=SUBMITTED")).json()
|
||||
assert [a["id"] for a in pending["items"]] == ["pending"]
|
||||
assert {a["id"] for a in submitted["items"]} == {"active", "retired"}
|
||||
exported = await logged_in.get(f"{PREFIX}/alphas/export?submission=SUBMITTED")
|
||||
assert {r["id"] for r in csv.DictReader(io.StringIO(exported.text.lstrip("\ufeff")))} == {
|
||||
"active",
|
||||
"retired",
|
||||
}
|
||||
assert (await logged_in.get(f"{PREFIX}/alphas?submission=BAD")).status_code == 422
|
||||
|
||||
|
||||
async def test_daily_job_validation_scope_deduplication_and_history(app, logged_in):
|
||||
await ready_runner(app)
|
||||
for payload in (
|
||||
{"kind": "daily_sync"},
|
||||
{"kind": "full_sync", "submission": "UNSUBMITTED"},
|
||||
{"kind": "full_sync", "date_from": "2025-01-01"},
|
||||
{"kind": "daily_sync", "submission": "SUBMITTED", "date_from": "2025-01-02", "date_to": "2025-01-01"},
|
||||
{"kind": "daily_sync", "submission": "SUBMITTED", "date_from": "2025-01-01", "date_to": "9999-01-01"},
|
||||
{"kind": "alpha_refresh", "alpha_ids": ["a"], "submission": "SUBMITTED"},
|
||||
):
|
||||
assert (await logged_in.post(f"{PREFIX}/sync-jobs", json=payload)).status_code == 422
|
||||
payload = {
|
||||
"kind": "daily_sync",
|
||||
"submission": "UNSUBMITTED",
|
||||
"date_from": "2025-01-01",
|
||||
"date_to": "2025-01-02",
|
||||
}
|
||||
first = (await logged_in.post(f"{PREFIX}/sync-jobs", json=payload)).json()
|
||||
duplicate = (await logged_in.post(f"{PREFIX}/sync-jobs", json=payload)).json()
|
||||
other = (await logged_in.post(f"{PREFIX}/sync-jobs", json={**payload, "submission": "SUBMITTED"})).json()
|
||||
assert first["id"] == duplicate["id"] != other["id"]
|
||||
assert first["payload"]["date_from"] == "2025-01-01"
|
||||
full = (await logged_in.post(f"{PREFIX}/sync-jobs", json={"kind": "full_sync"})).json()
|
||||
assert full["payload"]["submission"] == "SUBMITTED"
|
||||
|
||||
|
||||
class DailyPlatform(FakePlatform):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.daily_calls = []
|
||||
self.fail_daily = True
|
||||
self.records = [
|
||||
alpha("midnight", dateCreated="2025-01-01T00:00:00Z"),
|
||||
alpha("end", dateCreated="2025-01-01T23:59:59.999999Z"),
|
||||
alpha("next", dateCreated="2025-01-02T00:00:00Z"),
|
||||
alpha("hidden1", hidden=True, dateCreated="2025-01-02T01:00:00Z"),
|
||||
alpha("hidden2", hidden=True, dateCreated="2025-01-02T02:00:00Z"),
|
||||
alpha(
|
||||
"submitted",
|
||||
status="ACTIVE",
|
||||
dateCreated="2024-01-01T00:00:00Z",
|
||||
dateSubmitted="2025-01-02T00:00:00Z",
|
||||
),
|
||||
]
|
||||
|
||||
async def alphas(self, submission, hidden, offset, before, *, date_from=None, date_to=None):
|
||||
self.daily_calls.append((submission, hidden, offset, date_from, date_to))
|
||||
if self.fail_daily and hidden and offset == 1:
|
||||
raise WqError("模拟日内第二页失败", "network_error")
|
||||
field = "dateCreated" if submission == "UNSUBMITTED" else "dateSubmitted"
|
||||
matched = [
|
||||
r
|
||||
for r in self.records
|
||||
if (r["status"] == "UNSUBMITTED") == (submission == "UNSUBMITTED") and r["hidden"] == hidden
|
||||
]
|
||||
if date_from:
|
||||
matched = [
|
||||
r
|
||||
for r in matched
|
||||
if datetime.fromisoformat(date_from)
|
||||
<= datetime.fromisoformat(r[field].replace("Z", "+00:00"))
|
||||
< datetime.fromisoformat(date_to)
|
||||
]
|
||||
return {"results": matched[offset : offset + 1], "count": len(matched)}
|
||||
|
||||
|
||||
async def test_daily_sync_midnight_hidden_pages_restart_and_research_preservation(app, logged_in):
|
||||
runner = await ready_runner(app)
|
||||
runner.client = DailyPlatform()
|
||||
payload = {
|
||||
"kind": "daily_sync",
|
||||
"submission": "UNSUBMITTED",
|
||||
"date_from": "2025-01-01",
|
||||
"date_to": "2025-01-02",
|
||||
}
|
||||
job_id = (await logged_in.post(f"{PREFIX}/sync-jobs", json=payload)).json()["id"]
|
||||
await runner.execute(job_id)
|
||||
failed = await result(runner, job_id)
|
||||
assert failed.status == "failed" and failed.processed == 4
|
||||
assert failed.checkpoint["date"] == "2025-01-02" and failed.checkpoint["offset"] == 1
|
||||
async with runner.sessions() as db:
|
||||
research = await db.get(Research, "midnight")
|
||||
research.note, research.state = "preserved", "candidate"
|
||||
await db.commit()
|
||||
resumed = Runner(runner.sessions, runner.settings, DailyPlatform())
|
||||
resumed.client.fail_daily = False
|
||||
await resumed.execute(job_id)
|
||||
complete = await result(resumed, job_id)
|
||||
assert complete.status == "completed" and complete.processed == complete.total == 5
|
||||
assert complete.checkpoint["dates_completed"] == complete.checkpoint["dates_total"] == 2
|
||||
assert resumed.client.daily_calls[0][1:3] == (True, 1)
|
||||
async with runner.sessions() as db:
|
||||
assert len((await db.scalars(select(JobItem).where(JobItem.job_id == job_id))).all()) == 5
|
||||
assert (await db.get(Research, "midnight")).note == "preserved"
|
||||
assert await db.get(Alpha, "submitted") is None
|
||||
# Daily submitted sync uses submission date, although its creation was a year earlier.
|
||||
submitted_id = (
|
||||
await logged_in.post(f"{PREFIX}/sync-jobs", json={**payload, "submission": "SUBMITTED"})
|
||||
).json()["id"]
|
||||
await resumed.execute(submitted_id)
|
||||
assert (await result(resumed, submitted_id)).processed == 1
|
||||
full_id = (await logged_in.post(f"{PREFIX}/sync-jobs", json={"kind": "full_sync"})).json()["id"]
|
||||
resumed.client.daily_calls.clear()
|
||||
await resumed.execute(full_id)
|
||||
assert (await result(resumed, full_id)).processed == 1
|
||||
assert {call[0] for call in resumed.client.daily_calls} == {"SUBMITTED"}
|
||||
|
||||
|
||||
async def test_platform_ignoring_daily_filter_does_not_import_other_days(app, logged_in):
|
||||
runner = await ready_runner(app)
|
||||
|
||||
async def wrong_day(*args, **kwargs):
|
||||
return {"results": [alpha("wrong", dateCreated="2024-01-01T00:00:00Z")], "next": None}
|
||||
|
||||
runner.client.alphas = wrong_day
|
||||
job = (
|
||||
await logged_in.post(
|
||||
f"{PREFIX}/sync-jobs",
|
||||
json={
|
||||
"kind": "daily_sync",
|
||||
"submission": "UNSUBMITTED",
|
||||
"date_from": "2025-01-01",
|
||||
"date_to": "2025-01-01",
|
||||
},
|
||||
)
|
||||
).json()
|
||||
await runner.execute(job["id"])
|
||||
assert (await result(runner, job["id"])).status == "failed"
|
||||
async with runner.sessions() as db:
|
||||
assert await db.get(Alpha, "wrong") is None
|
||||
|
||||
|
||||
async def test_local_detection_uses_cache_excludes_self_and_keeps_research(app, logged_in):
|
||||
runner = app.state.runner # No credentials: any upstream call fails this test.
|
||||
data = points([1, 4, 2, -2] * 20)
|
||||
async with runner.sessions() as db:
|
||||
for raw in [
|
||||
alpha("target", status="ACTIVE"),
|
||||
alpha("peer", status="ACTIVE"),
|
||||
alpha("pending"),
|
||||
alpha("other-region", status="ACTIVE", settings={"region": "CHN"}),
|
||||
alpha("unknown", status=None),
|
||||
]:
|
||||
await upsert_alpha(db, raw)
|
||||
db.add(Pnl(alpha_id=raw["id"], raw={}, points=data))
|
||||
await db.flush()
|
||||
research = await db.get(Research, "target")
|
||||
research.note, research.state = "hypothesis", "candidate"
|
||||
await db.commit()
|
||||
job = (
|
||||
await logged_in.post(
|
||||
f"{PREFIX}/sync-jobs", json={"kind": "self_correlation", "alpha_ids": ["target"]}
|
||||
)
|
||||
).json()
|
||||
await runner.execute(job["id"])
|
||||
assert (await result(runner, job["id"])).status == "completed"
|
||||
report = (await logged_in.get(f"{PREFIX}/alphas/target/self-correlation")).json()["result"]
|
||||
assert report["candidate_count"] == report["compared_count"] == 1
|
||||
assert report["matches"][0]["alpha_id"] == "peer" and report["status"] == "high"
|
||||
for timestamp in (
|
||||
report["calculated_at"],
|
||||
report["target_pnl_fetched_at"],
|
||||
report["matches"][0]["pnl_fetched_at"],
|
||||
):
|
||||
assert datetime.fromisoformat(timestamp).utcoffset() == timedelta(0)
|
||||
async with runner.sessions() as db:
|
||||
assert (await db.get(Research, "target")).state == "candidate"
|
||||
assert (await db.get(Research, "target")).note == "hypothesis"
|
||||
assert (await db.get(Alpha, "target")).checks[0]["result"] == "FAIL"
|
||||
# A new submitted reference invalidates the old result without deleting it.
|
||||
await upsert_alpha(db, alpha("new-peer", status="ACTIVE"))
|
||||
await db.commit()
|
||||
assert (await db.get(SelfCorrelation, "target")).stale
|
||||
summary = (await logged_in.get(f"{PREFIX}/alphas?q=target")).json()["items"][0]["local_correlation"]
|
||||
assert summary["stale"] and summary["max_correlation"] == pytest.approx(1)
|
||||
assert datetime.fromisoformat(summary["calculated_at"]).utcoffset() == timedelta(0)
|
||||
|
||||
|
||||
async def test_missing_pnl_filled_once_and_partial_data_reported(app, logged_in):
|
||||
runner = await ready_runner(app)
|
||||
calls = []
|
||||
data = points([1, 3, -2] * 20)
|
||||
|
||||
async def pnl(alpha_id):
|
||||
calls.append(alpha_id)
|
||||
if alpha_id == "unavailable":
|
||||
raise WqError("无权访问", "access_denied")
|
||||
multiplier = 1 if alpha_id == "target" else -1
|
||||
return {"records": [{"date": p["date"], "pnl": p["value"] * multiplier} for p in data]}
|
||||
|
||||
runner.client.pnl = pnl
|
||||
async with runner.sessions() as db:
|
||||
for raw in [alpha("target"), alpha("peer", status="ACTIVE"), alpha("unavailable", status="ACTIVE")]:
|
||||
await upsert_alpha(db, raw)
|
||||
await db.commit()
|
||||
|
||||
async def run():
|
||||
job = (
|
||||
await logged_in.post(
|
||||
f"{PREFIX}/sync-jobs", json={"kind": "self_correlation", "alpha_ids": ["target"]}
|
||||
)
|
||||
).json()
|
||||
await runner.execute(job["id"])
|
||||
return (await logged_in.get(f"{PREFIX}/alphas/target/self-correlation")).json()["result"]
|
||||
|
||||
report = await run()
|
||||
assert report["status"] == "partial" and report["skipped_count"] == 1
|
||||
await run()
|
||||
assert calls.count("target") == calls.count("peer") == 1
|
||||
|
||||
|
||||
async def test_scoped_date_query_parameters_and_no_platform_check(settings):
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return httpx.Response(200, json={"results": []})
|
||||
|
||||
client = WqClient(settings, transport=httpx.MockTransport(handler))
|
||||
client.credentials, client.authenticated = ("test@example.com", "test"), True
|
||||
for submission in ("UNSUBMITTED", "SUBMITTED"):
|
||||
await client.alphas(
|
||||
submission,
|
||||
True,
|
||||
100,
|
||||
"2025-03-01T00:00:00+00:00",
|
||||
date_from="2025-01-01T00:00:00+00:00",
|
||||
date_to="2025-01-02T00:00:00+00:00",
|
||||
)
|
||||
first, second = [dict(r.url.params) for r in requests]
|
||||
assert first["dateCreated>="] == "2025-01-01T00:00:00+00:00"
|
||||
assert first["dateCreated<"] == "2025-01-02T00:00:00+00:00"
|
||||
assert second["dateSubmitted>="] == first["dateCreated>="]
|
||||
assert second["dateSubmitted<"] == first["dateCreated<"]
|
||||
assert "status!" in second and second["hidden"] == "true"
|
||||
assert all(r.method == "GET" and r.url.path == "/users/self/alphas" for r in requests)
|
||||
await client.close()
|
||||
Reference in New Issue
Block a user