Files
worldquant-alpha-system/backend/app/research_access/queries.py
T

135 lines
7.2 KiB
Python
Raw Normal View History

"""Historical evidence reads, independent of transport and current Alpha refreshes."""
from collections import Counter
from sqlalchemy import func, select
from ..alphas import number, sanitize
from ..backtests.contracts import fingerprint
from ..models import BacktestItem, BacktestResult, BacktestRun, Pnl
from ..research.serialization import encode_snapshot
def page(items, total, limit, offset):
return {"items": items, "total": total, "limit": limit, "offset": offset,
"has_more": offset + len(items) < total}
def checks_summary(snapshot):
"""Preserve unknown check values; missing checks can never mean passed."""
checks = []
for section in ("is", "os"):
metrics = snapshot.get(section)
if isinstance(metrics, dict) and "checks" in metrics:
raw = metrics["checks"]
checks.extend({"section": section, "raw": c} for c in (raw if isinstance(raw, list) else [raw]))
if "checks" in snapshot:
raw = snapshot["checks"]
checks.extend({"section": "root", "raw": c} for c in (raw if isinstance(raw, list) else [raw]))
counts = Counter({key: 0 for key in ("PASS", "FAIL", "PENDING", "WARNING", "UNKNOWN")})
non_pass = []
for check in checks:
raw = check["raw"]
value = raw.get("result", raw.get("status")) if isinstance(raw, dict) else None
state = value if isinstance(value, str) and value in counts else "UNKNOWN"
counts[state] += 1
if state != "PASS":
non_pass.append({**check, "status": state})
return {"status": "unknown" if not checks else "reported", "counts": dict(counts),
"total": len(checks), "non_pass": non_pass}
def item_summary(item, result):
snapshot = sanitize(result.snapshot) if result else {}
metrics = {}
for section in ("is", "os"):
raw = snapshot.get(section)
raw = raw if isinstance(raw, dict) else {}
metrics[section] = {key: number(raw.get(key)) for key in
("sharpe", "fitness", "returns", "turnover", "margin", "drawdown")}
return encode_snapshot({
**{k: getattr(item, k) for k in (
"id", "run_id", "client_item_id", "expression", "settings", "attempt_id",
"platform_status", "collection_status", "persistence_status", "simulation_id", "alpha_id",
)},
"error": sanitize(item.error), "metrics": metrics,
"missing_metrics_reason": "来源未提供或非有限数字;null 不等于零",
"checks": checks_summary(snapshot),
"result": {"observed_at": result.observed_at, "complete": result.complete} if result else None,
"artifact_reference": {"item_id": item.id},
})
class EvidenceQueries:
def __init__(self, db):
self.db = db
async def history(self, args):
query = select(BacktestItem, BacktestResult, BacktestRun).join(
BacktestRun, BacktestRun.id == BacktestItem.run_id
).outerjoin(BacktestResult, BacktestResult.item_id == BacktestItem.id)
for key in ("source", "reference"):
value = getattr(args, key)
if value is not None:
query = query.where(BacktestRun.source["kind" if key == "source" else key].as_string() == value)
if args.status:
query = query.where(BacktestRun.status == args.status)
if args.created_from:
query = query.where(BacktestRun.created_at >= args.created_from)
if args.created_to:
query = query.where(BacktestRun.created_at <= args.created_to)
if args.scope:
for source, target in (("instrument_type", "instrumentType"), ("region", "region"), ("universe", "universe")):
query = query.where(BacktestItem.settings[target].as_string() == getattr(args.scope, source))
query = query.where(BacktestItem.settings["delay"].as_integer() == args.scope.delay)
if args.q:
escaped = args.q.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
query = query.where(BacktestItem.expression.ilike(f"%{escaped}%", escape="\\"))
matches = {}
if args.candidates:
for c in args.candidates:
matches.setdefault(fingerprint(c.platform_input()), []).append(c.client_item_id)
query = query.where(BacktestItem.fingerprint.in_(matches))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = (await self.db.execute(query.order_by(BacktestRun.created_at.desc(), BacktestRun.id, BacktestItem.ordinal)
.limit(args.limit).offset(args.offset))).all()
items = [{**item_summary(i, r), "source": run.source, "run_status": run.status,
"created_at": run.created_at, "matched_candidates": matches.get(i.fingerprint, []),
"match_type": "exact_input" if args.candidates else "filter"} for i, r, run in rows]
return encode_snapshot(page(items, total, args.limit, args.offset))
async def results(self, args):
query = select(BacktestItem, BacktestResult).outerjoin(
BacktestResult, BacktestResult.item_id == BacktestItem.id
).where(BacktestItem.run_id == args.run_id)
if args.item_ids:
query = query.where(BacktestItem.id.in_(args.item_ids))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = (await self.db.execute(query.order_by(BacktestItem.ordinal).limit(args.limit).offset(args.offset))).all()
return {"backtest_run_id": args.run_id, **page([item_summary(i, r) for i, r in rows], total, args.limit, args.offset)}
async def artifact(self, args):
from .service import ResearchError
item = await self.db.get(BacktestItem, args.item_id)
if not item:
raise ResearchError("NOT_FOUND", "候选不存在")
result = await self.db.get(BacktestResult, item.id)
if args.kind == "snapshot":
# Top-level entries retain complete nested values; no hidden string/list truncation.
entries = [{"key": k, "value": v} for k, v in sanitize(result.snapshot).items()] if result else []
return encode_snapshot({"item_id": item.id, "kind": args.kind,
"status": "available" if result else "not_available",
"observed_at": result.observed_at if result else None,
"complete": result.complete if result else False,
**page(entries[args.offset:args.offset + args.limit], len(entries), args.limit, args.offset)})
pnl = await self.db.get(Pnl, item.alpha_id) if item.alpha_id else None
points = pnl.points if pnl else []
points = [p for p in points if
(not args.date_from or p["date"][:10] >= args.date_from.isoformat()) and
(not args.date_to or p["date"][:10] <= args.date_to.isoformat())]
return encode_snapshot({"item_id": item.id, "alpha_id": item.alpha_id, "kind": args.kind,
"status": "available" if pnl else "not_cached", "fetched_at": pnl.fetched_at if pnl else None,
"units": "供应商原始累计值;未提供货币或规模单位",
**page(points[args.offset:args.offset + args.limit], len(points), args.limit, args.offset)})