Files
worldquant-alpha-system/backend/app/research/experiments.py
T

433 lines
19 KiB
Python

"""Research producers share snapshot binding, candidate persistence and backtest previews."""
import difflib
import json
from collections import defaultdict
from fastapi import HTTPException
from sqlalchemy import func, select, update
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, SimulationSettings, Source
from ..backtests.service import Backtests, uid
from ..catalog.research_metadata import ResearchMetadata
from ..catalog.service import Catalog
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment, TemplateInput
from .assets import Assets
from .expressions import GROUPS, analyze, expand
from .serialization import encode_snapshot as jsonable_encoder
from .workspace_contracts import TemplateSpec
def scope_of(settings):
return {
"instrument_type": settings.instrumentType,
"region": settings.region,
"universe": settings.universe,
"delay": settings.delay,
}
def seed_settings(snapshot):
"""Decode executable settings, retaining returned historical dates in the parent snapshot.
startDate/endDate are result window metadata absent from POST settings. Unknown
execution parameters still fail strict validation rather than being discarded.
"""
return SimulationSettings.model_validate(
{k: v for k, v in snapshot.items() if k not in ("startDate", "endDate")}
)
class Experiments:
def __init__(self, db):
self.db = db
self.catalog = Catalog(db)
self.assets = Assets(db)
async def inputs(self, ids, scope=None):
if len(set(ids)) != len(ids):
raise HTTPException(422, "输入快照重复")
snapshots = [await self.catalog.input(input_id) for input_id in ids]
if scope and any(item["scope"] != scope for item in snapshots):
raise HTTPException(422, "输入快照与研究范围不一致,跨市场需要各自固定输入")
fields = {}
for item in snapshots:
for name, kind in item["field_types"].items():
if name not in item["field_ids"]:
continue
if name in fields and fields[name] != kind:
raise HTTPException(422, f"字段 {name} 在不同快照中类型不一致")
fields[name] = kind
return snapshots, fields
async def parents(self, alpha_ids, experiment_ids):
parents = []
for alpha_id in dict.fromkeys(alpha_ids):
alpha = await self.db.get(Alpha, alpha_id)
if not alpha:
raise HTTPException(404, f"种子 Alpha {alpha_id} 尚未同步")
if alpha.alpha_type != "REGULAR" or alpha.language != "FASTEXPR":
raise HTTPException(422, "变体生成仅支持 REGULAR + FASTEXPR")
parents.append(
{
"kind": "alpha",
"id": alpha.id,
"expression": alpha.expression,
"settings": alpha.settings,
"synced_at": jsonable_encoder(alpha.synced_at),
}
)
for experiment_id in dict.fromkeys(experiment_ids):
experiment = await self.get(experiment_id)
parents.append(
{
"kind": "experiment",
"id": experiment_id,
"candidates": experiment["candidates"],
"hypothesis": experiment["hypothesis"],
"input_references": [
{k: entry[k] for k in ("id", "collection_version", "scope", "dataset_id")}
for entry in experiment["inputs"]
],
"template_reference": {
k: experiment["evidence"].get("template", {}).get(k) for k in ("id", "version")
},
}
)
return parents
async def settings_check(self, settings):
snapshot = await ResearchMetadata(self.db).get("settings")
rows = snapshot["content"].get("items", [])
matches = [
row for row in rows if all(row.get(key) == value for key, value in scope_of(settings).items())
]
errors = []
if not matches:
errors.append("此市场设置尚未在平台设置快照中核实,请同步合法设置")
elif not any(settings.neutralization in row.get("neutralizations", []) for row in matches):
errors.append("中性化设置尚未在平台设置快照中核实")
return errors, snapshot
async def field_evidence(self, scope, fields):
rows = await self.db.scalars(select(CatalogResource).where(CatalogResource.kind == "availability"))
return {
row.content["field_id"]: ResearchMetadata.output(row)
for row in rows
if row.content.get("scope") == scope and row.content.get("field_id") in fields
}
@staticmethod
def validate(expression, fields, operators, scope, availability):
validation = analyze(expression, fields, operators)
for field in validation["fields"]:
if field not in availability:
continue # Published, scoped catalog membership is direct positive evidence.
content = availability[field]["content"]
if content.get("status") != "available" or scope not in content.get("items", []):
validation["availability"].append(
f"字段 {field} 的字段级可用性证据未确认目标范围,请重新核实"
)
if validation["availability"] and validation["status"] == "valid":
validation["status"] = "needs_review"
return validation
async def create(self, body, kind="template", extra_evidence=None, *, parent_snapshots=None):
asset = await self.assets.get(body.asset_id, body.version, "template") if body.asset_id else None
template = TemplateSpec.model_validate(asset["content"]) if asset else body.template
scope = scope_of(body.settings)
if template.scope and template.scope.model_dump() != scope:
raise HTTPException(422, "模板适用范围与候选设置不同")
snapshots, fields = await self.inputs(body.input_ids, scope)
parents = (
parent_snapshots
if parent_snapshots is not None
else await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
)
variables = {}
for name, variable in template.variables.items():
if variable.kind == "field":
for value in variable.values:
if fields.get(str(value)) != variable.field_type:
raise HTTPException(422, f"变量 {name} 的字段 {value} 不在固定输入中或类型不符")
if variable.kind == "group" and any(
str(v) not in GROUPS and fields.get(str(v)) != "GROUP" for v in variable.values
):
raise HTTPException(422, f"分组变量 {name} 未在固定输入中核实")
variables[name] = [
json.dumps(v, ensure_ascii=False) if variable.kind == "string" else v for v in variable.values
]
try:
expanded = expand(template.expression, variables, body.mode, body.limit, body.seed)
except ValueError as exc:
raise HTTPException(422, str(exc)) from None
operators_snapshot = await ResearchMetadata(self.db).get("operators")
operators = {item["name"] for item in operators_snapshot["content"].get("items", [])}
setting_errors, settings_snapshot = await self.settings_check(body.settings)
availability = await self.field_evidence(scope, fields)
candidates = []
for index, item in enumerate(expanded["items"]):
validation = self.validate(item["expression"], fields, operators, scope, availability)
validation["availability"].extend(setting_errors)
if setting_errors and validation["status"] == "valid":
validation["status"] = "needs_review"
candidates.append(
{
**Candidate(
client_item_id=f"c{index + 1}", expression=item["expression"], settings=body.settings
).model_dump(mode="json"),
"bindings": item["bindings"],
"input_ids": list(body.input_ids),
"validation": validation,
"changes": [
self.diff(parent.get("expression", ""), item["expression"])
for parent in parents
if parent["kind"] == "alpha"
],
}
)
evidence = {
"template": asset or {"content": template.model_dump(mode="json")},
"field_availability": availability,
"availability_basis": "各输入的已发布范围目录;若另有字段级证据,须同时满足",
"combination_count": expanded["combination_count"],
"seed": expanded["seed"],
"operators_snapshot": operators_snapshot,
"settings_snapshot": settings_snapshot,
**(extra_evidence or {}),
}
return await self.save(template.name, kind, body.hypothesis, snapshots, parents, candidates, evidence)
@staticmethod
def diff(before, after):
return [
{"operation": op, "before": before[i:j], "after": after[k:end], "start": i, "end": j}
for op, i, j, k, end in difflib.SequenceMatcher(a=before or "", b=after).get_opcodes()
if op != "equal"
]
async def save(self, name, kind, hypothesis, snapshots, parents, candidates, evidence):
row = ResearchExperiment(
id=uid(),
name=name,
kind=kind,
hypothesis=hypothesis,
inputs=jsonable_encoder(snapshots),
parents=jsonable_encoder(parents),
candidates=jsonable_encoder(candidates),
evidence=jsonable_encoder(evidence),
)
self.db.add(row)
await self.db.flush()
from ..models import ResearchParent
for parent_kind, parent_id in {(p["kind"], p["id"]) for p in parents}:
self.db.add(ResearchParent(child_id=row.id, parent_kind=parent_kind, parent_id=parent_id))
await self.db.flush()
return await self.get(row.id)
async def get(self, experiment_id):
row = await self.db.get(ResearchExperiment, experiment_id)
if not row:
raise HTTPException(404, "研究实验不存在")
runs = list(
await self.db.scalars(
select(BacktestRun.id).where(BacktestRun.source["research_id"].as_string() == row.id)
)
)
return jsonable_encoder(
{
**{
key: getattr(row, key)
for key in (
"id",
"name",
"kind",
"archived",
"hypothesis",
"inputs",
"parents",
"candidates",
"evidence",
"created_at",
)
},
"backtest_run_ids": runs,
}
)
async def list(self, kind=None, limit=25, offset=0):
query = select(ResearchExperiment).where(ResearchExperiment.archived.is_(False))
if kind:
query = query.where(ResearchExperiment.kind == kind)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(ResearchExperiment.created_at.desc()).limit(limit).offset(offset)
)
return jsonable_encoder(
{
"items": [
{
"id": row.id,
"name": row.name,
"kind": row.kind,
"total": len(row.candidates),
"created_at": row.created_at,
}
for row in rows
],
"total": total,
"limit": limit,
"offset": offset,
}
)
async def archive(self, experiment_id):
"""Hide an immutable experiment; backtests and lineage must still resolve it.
Returns an acknowledgement, or raises HTTP 404 for an unknown ID. Repeated
deletion is idempotent because candidate contents cannot change.
"""
result = await self.db.execute(
update(ResearchExperiment).where(ResearchExperiment.id == experiment_id).values(archived=True)
)
if result.rowcount != 1:
raise HTTPException(404, "研究实验不存在")
return {"ok": True}
async def preview(self, experiment_id, candidate_ids=None, source_kind=None, reference=None):
experiment = await self.get(experiment_id)
candidates = experiment["candidates"]
if candidate_ids is not None:
chosen = set(candidate_ids)
if len(chosen) != len(candidate_ids):
raise HTTPException(422, "候选选择包含重复项")
candidates = [item for item in candidates if item["client_item_id"] in chosen]
if len(candidates) != len(chosen):
raise HTTPException(422, "选择包含未知候选")
else:
candidates = [item for item in candidates if item["validation"]["status"] == "valid"]
if not candidates or any(item["validation"]["status"] != "valid" for item in candidates):
raise HTTPException(422, "候选存在语法、类型或可用性问题,请先解决;至少保留一条已核实候选")
inputs = experiment["inputs"]
return await Backtests(self.db).preview(
PreviewInput(
inline=DraftInput(
name=experiment["name"],
source=Source(
kind=source_kind or experiment["kind"],
reference=reference or experiment_id,
research_id=experiment_id,
template_input_id=inputs[0]["id"] if len(inputs) == 1 else None,
hypothesis=experiment["hypothesis"][:2000],
),
candidates=[
Candidate.model_validate(
{
key: item[key]
for key in ("client_item_id", "expression", "settings", "alpha_type")
}
)
for item in candidates
],
)
),
preserve_source=True,
)
async def setting_variants(self, body, *, parent_snapshot=None, extra_evidence=None, kind="variant"):
parents = (
[parent_snapshot] if parent_snapshot is not None else await self.parents([body.alpha_id], [])
)
original = parents[0]
base = seed_settings(original["settings"])
expression = original["expression"]
snapshots, _ = await self.inputs(body.input_ids)
groups = defaultdict(list)
for snapshot in snapshots:
groups[json.dumps(snapshot["scope"], sort_keys=True)].append(snapshot)
operators_snapshot = await ResearchMetadata(self.db).get("operators")
operators = {item["name"] for item in operators_snapshot["content"].get("items", [])}
candidates, rejected = [], []
for subset in groups.values():
scope = subset[0]["scope"]
try:
settings = SimulationSettings.model_validate(
{
**base.model_dump(),
"instrumentType": scope["instrument_type"],
**{key: scope[key] for key in ("region", "universe", "delay")},
}
)
except ValueError:
rejected.append({"scope": scope, "reason": "目标不属于当前支持的回测范围"})
continue
if scope == scope_of(base):
continue
_, fields = await self.inputs([s["id"] for s in subset], scope)
availability = await self.field_evidence(scope, fields)
validation = self.validate(expression, fields, operators, scope, availability)
errors, _ = await self.settings_check(settings)
validation["availability"].extend(errors)
if errors and validation["status"] == "valid":
validation["status"] = "needs_review"
candidates.append(
{
**Candidate(
client_item_id=f"v{len(candidates) + 1}", expression=expression, settings=settings
).model_dump(mode="json"),
"validation": validation,
"bindings": {},
"input_ids": [s["id"] for s in subset],
"field_availability": availability,
"changes": {
key: {"before": getattr(base, key), "after": getattr(settings, key)}
for key in ("region", "universe", "delay", "instrumentType")
if getattr(base, key) != getattr(settings, key)
},
}
)
return await self.save(
f"{body.alpha_id} · 设置变体",
kind,
body.hypothesis,
snapshots,
parents,
candidates,
{
**(extra_evidence or {}),
"method": "settings",
"rejected": rejected,
"operators_snapshot": operators_snapshot,
"settings_snapshot": await ResearchMetadata(self.db).get("settings"),
"availability_evidence": "各目标范围已发布的完整字段集合及固定输入;所有表达式字段必须存在",
},
)
async def generation_context(self, body):
snapshots, fields = await self.inputs(body.input_ids)
parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
metadata = await ResearchMetadata(self.db).operators(limit=100)
# This is a declared bounded context, not an assertion that a search page is the full input.
return {
"name": body.name,
"hypothesis": body.hypothesis,
"method": body.method,
"inputs": [
{"id": item["id"], "scope": item["scope"], "dataset_id": item["dataset_id"]}
for item in snapshots
],
"fields": dict(list(fields.items())[:300]),
"fields_total": len(fields),
"operators": [
{k: item.get(k) for k in ("name", "description", "definition")} for item in metadata["items"]
],
"parents": [{**p, "candidates": p.get("candidates", [])[:10]} for p in parents],
}
async def available_inputs(self, limit=100):
rows = await self.db.scalars(
select(TemplateInput).order_by(TemplateInput.created_at.desc()).limit(limit)
)
return {"items": [await self.catalog.input(row.id) for row in rows]}