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