228 lines
15 KiB
Python
228 lines
15 KiB
Python
|
|
"""Versioned SUPER plans and deterministic candidate construction; never executes simulations."""
|
||
|
|
|
||
|
|
import math
|
||
|
|
import random
|
||
|
|
from uuid import uuid4
|
||
|
|
|
||
|
|
from fastapi import HTTPException
|
||
|
|
from sqlalchemy import func, select
|
||
|
|
|
||
|
|
from ..alphas import sanitize
|
||
|
|
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source, fingerprint
|
||
|
|
from ..backtests.service import Backtests
|
||
|
|
from ..models import Account, Alpha, Job, ResearchExperiment, ResearchRequest, SuperSelectionSnapshot, now
|
||
|
|
from ..research.assets import Assets
|
||
|
|
from ..research.expressions import PLACEHOLDER
|
||
|
|
from ..research.serialization import encode_snapshot
|
||
|
|
from ..research.workspace_contracts import AssetWrite
|
||
|
|
from .contracts import PlanSpec, SelectionPreview
|
||
|
|
from .evidence import actual_components, parse_components
|
||
|
|
from .settings import validate_settings
|
||
|
|
|
||
|
|
|
||
|
|
async def validate_source(db, source, candidates):
|
||
|
|
"""Verify server-owned provenance references without making assets mandatory for direct execution."""
|
||
|
|
if (source.get("superalpha_plan_id") or source.get("selection_snapshot_ids") or source.get("kind") == "superalpha") and any(c.get("alpha_type", "REGULAR") != "SUPER" for c in candidates):
|
||
|
|
raise HTTPException(422, "Super Alpha 方案或组件来源只能关联 SUPER 候选")
|
||
|
|
await validate_settings(db, [c["settings"] for c in candidates if c.get("alpha_type") == "SUPER"])
|
||
|
|
experiment = None
|
||
|
|
if bool(source.get("superalpha_plan_id")) != bool(source.get("superalpha_plan_version")):
|
||
|
|
raise HTTPException(422, "方案引用须同时指定 ID 和版本")
|
||
|
|
if source.get("superalpha_plan_id"):
|
||
|
|
if not source.get("superalpha_plan_version"):
|
||
|
|
raise HTTPException(422, "方案引用须指定版本")
|
||
|
|
await Assets(db).get(source["superalpha_plan_id"], source["superalpha_plan_version"], "superalpha_plan")
|
||
|
|
if source.get("research_id") and (source.get("kind") == "superalpha" or source.get("superalpha_plan_id") or any(c.get("alpha_type") == "SUPER" for c in candidates)):
|
||
|
|
experiment = await db.get(ResearchExperiment, source["research_id"])
|
||
|
|
if not experiment or experiment.kind != "superalpha":
|
||
|
|
raise HTTPException(404, "SUPER 候选构造记录不存在")
|
||
|
|
source["research_kind"] = "superalpha"
|
||
|
|
expected = {c["client_item_id"]: fingerprint(Candidate.model_validate(c).platform_input()) for c in experiment.candidates}
|
||
|
|
for value in candidates:
|
||
|
|
c = Candidate.model_validate(value)
|
||
|
|
if expected.get(c.client_item_id) != fingerprint(c.platform_input()):
|
||
|
|
raise HTTPException(409, "候选与引用的固定构造记录不一致")
|
||
|
|
ref = experiment.evidence.get("plan_reference", {})
|
||
|
|
if source.get("superalpha_plan_id") and ref != {
|
||
|
|
"id": source["superalpha_plan_id"], "version": source["superalpha_plan_version"]}:
|
||
|
|
raise HTTPException(409, "方案版本与构造来源不一致")
|
||
|
|
for snapshot_id in source.get("selection_snapshot_ids", []):
|
||
|
|
row = await db.get(SuperSelectionSnapshot, snapshot_id)
|
||
|
|
if not row or row.source != "preview":
|
||
|
|
raise HTTPException(404, "Selection 预览快照不存在")
|
||
|
|
query = SelectionPreview.model_validate(row.request).platform_query()
|
||
|
|
if not any(Candidate.model_validate(c).alpha_type == "SUPER" and SelectionPreview(
|
||
|
|
selection=c["selection"], settings=c["settings"]).platform_query() == query for c in (experiment.candidates if experiment else candidates)):
|
||
|
|
raise HTTPException(409, "组件预览与候选 Selection/范围不匹配")
|
||
|
|
|
||
|
|
|
||
|
|
class SuperResearch:
|
||
|
|
def __init__(self, db):
|
||
|
|
self.db = db
|
||
|
|
self.assets = Assets(db)
|
||
|
|
|
||
|
|
async def previous(self, operation, args):
|
||
|
|
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||
|
|
if not account:
|
||
|
|
raise HTTPException(409, "工作空间未初始化")
|
||
|
|
digest = fingerprint(args.model_dump(mode="json", exclude={"idempotency_key"}))
|
||
|
|
previous = await self.db.scalar(select(ResearchRequest).where(ResearchRequest.account_id == 1,
|
||
|
|
ResearchRequest.operation == operation, ResearchRequest.idempotency_key == args.idempotency_key))
|
||
|
|
if previous and previous.digest != digest:
|
||
|
|
raise HTTPException(409, "幂等键已用于不同内容")
|
||
|
|
return previous, digest
|
||
|
|
|
||
|
|
async def remember(self, operation, args, digest, result, business_id):
|
||
|
|
result = encode_snapshot(result)
|
||
|
|
result["_meta"] = {"schema_version": 1, "observed_at": now().isoformat(), "source": "system"}
|
||
|
|
self.db.add(ResearchRequest(id=str(uuid4()), account_id=1, operation=operation,
|
||
|
|
idempotency_key=args.idempotency_key, digest=digest, business_id=business_id, response=result))
|
||
|
|
await self.db.flush()
|
||
|
|
return result
|
||
|
|
|
||
|
|
async def provenance(self, plan):
|
||
|
|
result = {"reference": plan.reference}
|
||
|
|
if plan.parent_plan_id:
|
||
|
|
parent = await self.assets.get(plan.parent_plan_id, plan.parent_plan_version, "superalpha_plan")
|
||
|
|
result["parent_plan"] = {k: parent[k] for k in ("id", "version", "name")}
|
||
|
|
if plan.parent_alpha_id:
|
||
|
|
alpha = await self.db.get(Alpha, plan.parent_alpha_id)
|
||
|
|
if not alpha or alpha.alpha_type != "SUPER":
|
||
|
|
raise HTTPException(404, "父 SUPER Alpha 尚未导入")
|
||
|
|
result["parent_alpha"] = {"id": alpha.id, "snapshot": sanitize(alpha.raw), "observed_at": alpha.synced_at}
|
||
|
|
if plan.parent_experiment_id:
|
||
|
|
parent = await self.experiment(plan.parent_experiment_id)
|
||
|
|
result["parent_experiment"] = {"id": parent["id"], "created_at": parent["created_at"]}
|
||
|
|
return encode_snapshot(result)
|
||
|
|
|
||
|
|
async def save(self, args):
|
||
|
|
previous, digest = await self.previous("save_superalpha_plan", args)
|
||
|
|
if previous:
|
||
|
|
return previous.response
|
||
|
|
await validate_settings(self.db, [args.plan.settings])
|
||
|
|
result = await self.assets.save(AssetWrite(kind="superalpha_plan", content=args.plan.model_dump(mode="json"),
|
||
|
|
version=args.version), args.plan_id, await self.provenance(args.plan))
|
||
|
|
return await self.remember("save_superalpha_plan", args, digest, result, result["id"])
|
||
|
|
|
||
|
|
async def build(self, args):
|
||
|
|
previous, digest = await self.previous("build_superalpha_candidates", args)
|
||
|
|
if previous:
|
||
|
|
return previous.response
|
||
|
|
asset = await self.assets.get(args.plan_id, args.version, "superalpha_plan") if args.plan_id else None
|
||
|
|
plan = PlanSpec.model_validate(asset["content"]) if asset else args.plan
|
||
|
|
provenance = await self.provenance(plan)
|
||
|
|
names = list(plan.variables)
|
||
|
|
settings_names = list(plan.setting_variants)
|
||
|
|
axes = [plan.variables[k].values for k in names] + [plan.setting_variants[k] for k in settings_names]
|
||
|
|
count = math.prod(len(a) for a in axes)
|
||
|
|
if count > 10**12 or (args.mode == "all" and count > args.limit):
|
||
|
|
raise HTTPException(422, f"理论组合数 {count} 超出展开上限;缩小参数或采用随机采样")
|
||
|
|
indices = range(count) if args.mode == "all" else sorted(random.Random(args.seed).sample(range(count), min(count, args.limit)))
|
||
|
|
candidates, annotations, seen = [], {}, {}
|
||
|
|
for index in indices:
|
||
|
|
remaining, values = index, []
|
||
|
|
for axis in reversed(axes):
|
||
|
|
remaining, position = divmod(remaining, len(axis))
|
||
|
|
values.insert(0, axis[position])
|
||
|
|
bindings = dict(zip(names, values[:len(names)]))
|
||
|
|
def substitute(text):
|
||
|
|
return PLACEHOLDER.sub(lambda m: str(bindings[m.group(1)]), text)
|
||
|
|
selection, combo = substitute(plan.selection), substitute(plan.combo)
|
||
|
|
settings = {**plan.settings.model_dump(), **dict(zip(settings_names, values[len(names):]))}
|
||
|
|
variants = [(combo, combo == "1")] + ([("1", True)] if plan.include_baseline and combo != "1" else [])
|
||
|
|
for combo_value, baseline in variants:
|
||
|
|
candidate_id = f"super-{index + 1}{'-baseline' if baseline else ''}"
|
||
|
|
c = Candidate(client_item_id=candidate_id, alpha_type="SUPER", selection=selection,
|
||
|
|
combo=combo_value, settings=settings)
|
||
|
|
h = fingerprint(c.platform_input())
|
||
|
|
annotations[candidate_id] = {"baseline": baseline, "parameters": bindings,
|
||
|
|
"duplicate_of": seen.get(h), "request_hash": h}
|
||
|
|
seen.setdefault(h, candidate_id)
|
||
|
|
candidates.append(c.model_dump(mode="json"))
|
||
|
|
if len(candidates) > 10000:
|
||
|
|
raise HTTPException(422, "包含基线后超过 10000 项,请缩小候选数")
|
||
|
|
plan_reference = {"id": asset["id"], "version": asset["version"]} if asset else {}
|
||
|
|
source = Source(kind="superalpha", research_kind="superalpha", reference=plan.reference, hypothesis=plan.hypothesis,
|
||
|
|
superalpha_plan_id=args.plan_id, superalpha_plan_version=args.version,
|
||
|
|
selection_snapshot_ids=args.selection_snapshot_ids).model_dump(mode="json")
|
||
|
|
await validate_source(self.db, source, candidates)
|
||
|
|
experiment = ResearchExperiment(id=str(uuid4()), name=plan.name, kind="superalpha", hypothesis=plan.hypothesis,
|
||
|
|
inputs=[], parents=[], candidates=candidates, evidence={"plan": plan.model_dump(mode="json"),
|
||
|
|
"plan_reference": plan_reference, "provenance": provenance, "source": source,
|
||
|
|
"selection_snapshot_ids": args.selection_snapshot_ids, "annotations": annotations,
|
||
|
|
"combination_count": str(count), "mode": args.mode, "seed": args.seed})
|
||
|
|
self.db.add(experiment)
|
||
|
|
await self.db.flush()
|
||
|
|
result = await self.experiment(experiment.id)
|
||
|
|
return await self.remember("build_superalpha_candidates", args, digest, result, experiment.id)
|
||
|
|
|
||
|
|
async def experiment(self, experiment_id, limit=100, offset=0):
|
||
|
|
row = await self.db.get(ResearchExperiment, experiment_id)
|
||
|
|
if not row or row.kind != "superalpha":
|
||
|
|
raise HTTPException(404, "SUPER 研究记录不存在")
|
||
|
|
source = {**row.evidence["source"], "research_id": row.id}
|
||
|
|
visible = row.candidates[offset:offset + limit]
|
||
|
|
evidence = {**row.evidence, "annotations": {c["client_item_id"]: row.evidence["annotations"].get(c["client_item_id"], {}) for c in visible}}
|
||
|
|
return encode_snapshot({"id": row.id, "name": row.name, "kind": row.kind, "hypothesis": row.hypothesis,
|
||
|
|
"created_at": row.created_at, "evidence": evidence, "source": source,
|
||
|
|
"candidates": row.candidates[offset:offset + limit], "total": len(row.candidates),
|
||
|
|
"limit": limit, "offset": offset, "has_more": offset + limit < len(row.candidates)})
|
||
|
|
|
||
|
|
async def experiments(self, plan_id=None, limit=25, offset=0):
|
||
|
|
query = select(ResearchExperiment).where(ResearchExperiment.kind == "superalpha")
|
||
|
|
if plan_id:
|
||
|
|
query = query.where(ResearchExperiment.evidence["plan_reference"]["id"].as_string() == plan_id)
|
||
|
|
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
|
||
|
|
rows = await self.db.scalars(query.order_by(ResearchExperiment.created_at.desc(), ResearchExperiment.id).limit(limit).offset(offset))
|
||
|
|
return encode_snapshot({"items": [{"id": r.id, "name": r.name, "created_at": r.created_at,
|
||
|
|
"total": len(r.candidates)} for r in rows], "total": total, "limit": limit, "offset": offset})
|
||
|
|
|
||
|
|
async def preview(self, experiment_id, candidate_ids):
|
||
|
|
row = await self.db.get(ResearchExperiment, experiment_id)
|
||
|
|
await self.experiment(experiment_id)
|
||
|
|
selected = [c for c in row.candidates if c["client_item_id"] in set(candidate_ids)]
|
||
|
|
if len(selected) != len(set(candidate_ids)):
|
||
|
|
raise HTTPException(422, "候选不属于当前研究记录")
|
||
|
|
return await Backtests(self.db).preview(PreviewInput(inline=DraftInput(name=row.name,
|
||
|
|
candidates=selected, source={**row.evidence["source"], "research_id": row.id})), preserve_source=True)
|
||
|
|
|
||
|
|
async def selection_job(self, args):
|
||
|
|
if bool(args.plan_id) != bool(args.version):
|
||
|
|
raise HTTPException(422, "预览的方案来源需同时指定 ID 和版本")
|
||
|
|
if args.plan_id:
|
||
|
|
await self.assets.get(args.plan_id, args.version, "superalpha_plan")
|
||
|
|
await validate_settings(self.db, [args.settings])
|
||
|
|
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||
|
|
if not account or account.connection_status not in ("connected", "expired"):
|
||
|
|
raise HTTPException(409, "请先连接 WorldQuant")
|
||
|
|
payload = args.model_dump(mode="json")
|
||
|
|
jobs = await self.db.scalars(select(Job).where(Job.kind == "super_selection_preview",
|
||
|
|
Job.status.in_(("queued", "running", "waiting_auth", "waiting_connection"))))
|
||
|
|
job = next((j for j in jobs if j.payload == payload and not j.cancel_requested), None)
|
||
|
|
if not job:
|
||
|
|
job = Job(id=str(uuid4()), kind="super_selection_preview", payload=payload, total=1)
|
||
|
|
self.db.add(job)
|
||
|
|
await self.db.flush()
|
||
|
|
return {"job_id": job.id, "status": job.status, "read_with": "get_superalpha_selection"}
|
||
|
|
|
||
|
|
async def alpha(self, alpha_id, limit=25, offset=0):
|
||
|
|
from ..business import Business
|
||
|
|
from ..models import BacktestItem, BacktestResult
|
||
|
|
alpha = await self.db.get(Alpha, alpha_id)
|
||
|
|
if not alpha or alpha.alpha_type != "SUPER":
|
||
|
|
raise HTTPException(404, "SUPER Alpha 尚未导入")
|
||
|
|
item = await self.db.scalar(select(BacktestItem).join(BacktestResult, BacktestResult.item_id == BacktestItem.id)
|
||
|
|
.where(BacktestItem.alpha_id == alpha_id).order_by(BacktestResult.observed_at.desc()).limit(1))
|
||
|
|
components = await actual_components(self.db, item.id, limit, offset) if item else {
|
||
|
|
"status": "unknown", "complete": False, "source": "actual", "items": [], "total": 0}
|
||
|
|
if not item:
|
||
|
|
parsed = parse_components(alpha.raw.get("components", alpha.raw.get("selectedAlphas")))
|
||
|
|
components = {"source": "actual", "status": "available" if parsed["complete"] else "unknown",
|
||
|
|
"complete": parsed["complete"], "component_hash": parsed["component_hash"],
|
||
|
|
"reported_total": parsed["total"], "total": len(parsed["components"]), "warnings": parsed["warnings"],
|
||
|
|
"observed_at": alpha.synced_at, "items": parsed["components"][offset:offset + limit], "limit": limit, "offset": offset}
|
||
|
|
return {**await Business(self.db).get_alpha(alpha_id), "components": components,
|
||
|
|
"descriptions": {k: (alpha.raw.get(k) or {}).get("description", "")
|
||
|
|
if isinstance(alpha.raw.get(k), dict) else "" for k in ("selection", "combo")},
|
||
|
|
"sources": await Business(self.db).get_alpha_sources(alpha_id)}
|