Files

228 lines
15 KiB
Python
Raw Permalink Normal View History

"""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)}