Files
worldquant-alpha-system/backend/app/alphas.py
T

226 lines
8.3 KiB
Python
Raw Normal View History

"""Normalize upstream business data, without making platform or research decisions."""
import math
import re
from datetime import datetime
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)
)
SENSITIVE_KEYS = {
"password",
"token",
"cookies",
"cookie",
"authorization",
"credentials",
"secret",
"accesstoken",
"refreshtoken",
"sessiontoken",
"clientsecret",
"authorizationheader",
"setcookie",
"apikey",
"csrftoken",
"xsrftoken",
"authentication",
}
def sanitize(value):
if isinstance(value, dict):
return {
k: sanitize(v) for k, v in value.items() if re.sub(r"[^a-z]", "", k.lower()) not in SENSITIVE_KEYS
}
if isinstance(value, list):
return [sanitize(v) for v in value]
if isinstance(value, float) and not math.isfinite(value):
return None
return value
def number(value):
if value is None or isinstance(value, bool):
return None
try:
result = float(value)
return result if math.isfinite(result) else None
except (ValueError, TypeError):
return None
def date(value):
try:
return datetime.fromisoformat(value.replace("Z", "+00:00")) if value else None
except (ValueError, TypeError, AttributeError):
return None
def code(value):
return value.get("code") if isinstance(value, dict) else value if isinstance(value, str) else None
async def upsert_alpha(db, raw: dict):
alpha_id = raw.get("id")
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)
settings = raw.get("settings") or {}
metrics = raw.get("is") if isinstance(raw.get("is"), dict) else {}
item.name = raw.get("name")
item.expression = code(raw.get("regular"))
item.selection, item.combo = code(raw.get("selection")), code(raw.get("combo"))
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 [])
for key in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
setattr(item, key, number(metrics.get(key)))
item.date_created, item.date_submitted = date(raw.get("dateCreated")), date(raw.get("dateSubmitted"))
item.synced_at, item.raw = now(), sanitize(raw)
await db.flush()
if await db.get(Research, alpha_id) is None:
db.add(Research(alpha_id=alpha_id))
return item
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("_", "\\_") + "%"
query = query.where(
or_(
*(
getattr(Alpha, f).ilike(pattern, escape="\\")
for f in ("id", "name", "expression", "selection", "combo")
)
)
)
for name in ("region", "universe", "alpha_type", "language", "status", "stage", "hidden"):
value = getattr(filters, name)
if value is not None:
query = query.where(getattr(Alpha, name) == value)
if filters.research_state:
query = query.where(Research.state == filters.research_state)
if filters.favorite is not None:
query = query.where(Research.favorite == filters.favorite)
if filters.tag:
query = query.where(Alpha.id.in_(select(ResearchTag.alpha_id).where(ResearchTag.tag == filters.tag)))
if filters.created_from:
query = query.where(Alpha.date_created >= filters.created_from)
if filters.created_to:
query = query.where(Alpha.date_created <= filters.created_to)
for name in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
for suffix, compare in (("min", "ge"), ("max", "le")):
value = getattr(filters, f"{name}_{suffix}")
if value is not None:
column = getattr(Alpha, name)
query = query.where(column >= value if compare == "ge" else column <= value)
return query
def sorted_statement(query, sort, direction):
column = getattr(Alpha, sort)
order = column.desc() if direction == "desc" else column.asc()
return query.order_by(order.nullslast(), Alpha.id.asc())
def summary(item: Alpha, research: Research):
keys = (
"id",
"name",
"alpha_type",
"language",
"stage",
"status",
"hidden",
"region",
"universe",
"sharpe",
"fitness",
"returns",
"turnover",
"margin",
"drawdown",
"date_created",
"date_submitted",
"synced_at",
)
result = {k: getattr(item, k) for k in keys}
result["expression_preview"] = (item.expression or item.selection or "")[:240]
result["research"] = {
k: getattr(research, k) for k in ("note", "tags", "favorite", "state", "updated_at", "version")
}
return result
def pnl_points(raw):
"""Use schema column names, preserving missing values rather than creating zero PnL."""
records = raw.get("records")
schema = raw.get("schema") or {}
properties = schema.get("properties", []) if isinstance(schema, dict) else schema
if isinstance(properties, dict):
names = list(properties)
else:
names = [p.get("name", "") if isinstance(p, dict) else str(p) for p in properties]
normalized = [name.lower() for name in names]
if not isinstance(records, list):
raise ValueError("PnL 缺少 records")
points = []
for row in records:
if isinstance(row, dict):
row = {str(k).lower(): v for k, v in row.items()}
timestamp = next((row[k] for k in ("date", "datetime", "timestamp") if k in row), None)
value = next((row[k] for k in ("pnl", "value") if k in row), None)
else:
date_i = next(
(i for i, n in enumerate(normalized) if n in ("date", "datetime", "timestamp")), None
)
pnl_i = next((i for i, n in enumerate(normalized) if n in ("pnl", "value")), None)
if date_i is None or pnl_i is None or not isinstance(row, list) or len(row) <= max(date_i, pnl_i):
raise ValueError("PnL schema 无法识别日期或数值列")
timestamp, value = row[date_i], row[pnl_i]
if timestamp is not None:
if isinstance(timestamp, (int, float)):
from datetime import timezone
timestamp = datetime.fromtimestamp(
timestamp / 1000 if timestamp > 1e11 else timestamp, tz=timezone.utc
).isoformat()
points.append({"date": str(timestamp), "value": number(value)})
return sorted(points, key=lambda p: p["date"])