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

319 lines
12 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
from .research.provenance import source_alpha_ids
METRIC_FIELDS = (
"sharpe", "fitness", "returns", "turnover", "margin", "drawdown",
"sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl",
)
def failed_checks(checks):
"""Return failed platform check names; local correlation never changes this list."""
return [
check.get("name") if isinstance(check.get("name"), str) else "未命名检查"
for check in checks if isinstance(check, dict) and check.get("result") == "FAIL"
] if isinstance(checks, list) else []
def snapshot_columns(settings, metrics, checks):
"""Derive list fields from a platform snapshot, preserving missing metrics as null.
Only explicit FAIL results count. Empty, malformed and unfinished checks are
pending; all known checks passing without PROD_CORRELATION is only a pre-check.
No submission eligibility or activity eligibility is inferred here.
"""
settings = settings if isinstance(settings, dict) else {}
metrics = metrics if isinstance(metrics, dict) else {}
checks = checks if isinstance(checks, list) else []
valid = [check for check in checks if isinstance(check, dict)]
failures = len(failed_checks(checks))
by_name = {check["name"]: check for check in valid if isinstance(check.get("name"), str)}
if failures:
check_type = "FAIL_1" if failures == 1 else "FAIL_2"
elif not checks or len(valid) != len(checks) or any(check.get("result") != "PASS" for check in valid):
check_type = "PENDING"
else:
check_type = "PASS" if "PROD_CORRELATION" in by_name else "PRE_CHECK"
neutralization = settings.get("neutralization")
return {
"check_type": check_type,
"neutralization": neutralization if isinstance(neutralization, str) else None,
"pnl": number(metrics.get("pnl")),
**{
field: number(by_name.get(name, {}).get("value"))
for field, name in (
("sub_universe_sharpe", "LOW_SUB_UNIVERSE_SHARPE"),
("robust_universe_sharpe", "LOW_ROBUST_UNIVERSE_SHARPE"),
("two_year_sharpe", "LOW_2Y_SHARPE"),
("prod_correlation", "PROD_CORRELATION"),
)
},
}
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, value in snapshot_columns(item.settings, item.is_metrics, item.checks).items():
setattr(item, key, value)
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))
source_filters = {k: getattr(filters, k) for k in ("source", "source_reference", "research_id", "backtest_run_id")}
if any(source_filters.values()):
query = query.where(Alpha.id.in_(source_alpha_ids(**source_filters)))
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", "check_type", "neutralization"):
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 METRIC_FIELDS:
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",
"check_type",
"neutralization",
"sub_universe_sharpe",
"robust_universe_sharpe",
"two_year_sharpe",
"prod_correlation",
"pnl",
)
result = {k: getattr(item, k) for k in keys}
result["failed_checks"] = failed_checks(item.checks)
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, column=None):
"""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]
value_names = (column,) if column else ("pnl", "value")
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 value_names 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 value_names), None)
if date_i is None or pnl_i is None or not isinstance(row, list) or len(row) <= date_i:
raise ValueError("PnL schema 无法识别日期或数值列")
if len(row) <= pnl_i and column is None:
raise ValueError("PnL schema 无法识别日期或数值列")
timestamp, value = row[date_i], row[pnl_i] if len(row) > pnl_i else None
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"])
def glb_pnl_series(raw, points):
"""Read GLB display series from cached raw data; keep the correlation baseline intact.
Missing columns are omitted, while missing values remain gaps. Legacy caches
containing only normalized points still return their overall PnL.
"""
series = [{"id": "pnl", "label": "总体 PnL", "points": points}]
schema = raw.get("schema") or {}
properties = schema.get("properties", []) if isinstance(schema, dict) else schema
names = properties if isinstance(properties, dict) else [
p.get("name", "") if isinstance(p, dict) else str(p) for p in properties
]
available = {name.lower() for name in names}
for row in raw.get("records", []):
if isinstance(row, dict):
available.update(str(key).lower() for key in row)
for column, label in (
("investability-constrained-pnl", "可投资性约束 PnL"),
("amer-pnl", "AMER PnL"),
("apac-pnl", "APAC PnL"),
("emea-pnl", "EMEA PnL"),
):
if column in available:
series.append({"id": column, "label": label, "points": pnl_points(raw, column)})
return series