Files
yuxuanhui 69c19ed25f
Deploy production / deploy (push) Successful in 51s
Refactor project components and workflows
2026-09-20 11:20:51 +08:00

541 lines
26 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,
StartInput,
fingerprint,
)
from ..backtests.service import Backtests, uid
from ..catalog.research_metadata import ResearchMetadata
from ..catalog.service import Catalog
from ..models import Account, Alpha, BacktestPreview, BacktestRun, CatalogResource, ResearchExperiment
from ..preparations.service import Preparations
from .assets import Assets
from .expressions import GROUPS, ExpressionError, Parser, 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, *, check_types=True):
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 check_types and 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", "preparation_id", "preparation_version", "scope", "dataset_ids")}
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):
await Preparations(self.db).bind(body)
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)
snapshots, fields = await self.inputs(body.input_ids, scope, check_types=kind != "template")
if kind == "template" and not snapshots:
raise HTTPException(422, "请先选择数据准备")
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():
values = variable.values
# Empty field definitions bind only to the selected immutable input scope.
# Existing explicit domains remain restrictions and are never silently widened.
if variable.kind == "field":
if not values:
values = sorted(field for field, kind in fields.items() if kind == variable.field_type)
if not values:
raise HTTPException(422, f"变量 {name} 没有匹配的 {variable.field_type} 字段,请调整数据准备")
for value in values:
if kind != "template" and fields.get(str(value)) != variable.field_type:
raise HTTPException(422, f"变量 {name} 的字段 {value} 不在固定输入中或类型不符")
if kind != "template" and 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} 未在固定输入中核实")
if not values:
raise HTTPException(422, f"变量 {name} 缺少候选取值,请通过模板接口补充,或将固定参数直接写入表达式")
variables[name] = [
json.dumps(v, ensure_ascii=False) if variable.kind == "string" else v for v in values
]
try:
expanded = expand(template.expression, variables, body.mode, body.limit, body.seed)
except ValueError as exc:
raise HTTPException(422, str(exc)) from None
validation_evidence = {}
if kind != "template":
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)
validation_evidence = {
"field_availability": availability,
"availability_basis": "各输入的已发布范围目录;若另有字段级证据,须同时满足",
"operators_snapshot": operators_snapshot,
"settings_snapshot": settings_snapshot,
}
candidates = []
for index, item in enumerate(expanded["items"]):
findings = {}
if kind == "template":
self.check_syntax(item["expression"], f"候选 {index + 1}")
else:
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"
findings["validation"] = validation
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),
**findings,
"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")},
"combination_count": expanded["combination_count"],
"seed": expanded["seed"],
**validation_evidence,
**(extra_evidence or {}),
}
return await self.save(template.name, kind, body.hypothesis, snapshots, parents, candidates, evidence)
@staticmethod
def check_syntax(expression, label="表达式"):
"""Reject unsupported syntax before persistence; platform semantics are not inferred."""
try:
Parser(expression).parse()
except (ExpressionError, RecursionError) as exc:
raise HTTPException(422, f"{label}语法错误:{exc}") from None
@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 template_candidates(self, experiment_id, limit=25, offset=0):
"""Read a bounded page of stored template candidates and their immutable references."""
experiment = await self.get(experiment_id)
if experiment["kind"] != "template":
raise HTTPException(422, "此入口仅用于模板候选集合")
candidates = experiment["candidates"]
template = experiment["evidence"].get("template", {})
return {
"id": experiment_id, "experiment_id": experiment_id,
"name": experiment["name"], "archived": experiment["archived"],
"template": {k: template.get(k) for k in ("id", "version", "name")},
"inputs": [{k: item.get(k) for k in ("id", "preparation_id", "preparation_version", "scope")}
for item in experiment["inputs"]],
"items": [{k: c[k] for k in ("client_item_id", "expression", "settings", "alpha_type", "bindings") if k in c}
for c in candidates[offset:offset + limit]],
"total": len(candidates), "limit": limit, "offset": offset,
"has_more": offset + limit < len(candidates),
"backtest_run_ids": experiment["backtest_run_ids"],
}
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 backtest_input(self, experiment_id, candidate_ids=None, source_kind=None, reference=None):
"""Build the complete fixed selection and enforce its domain checks."""
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, "选择包含未知候选")
elif experiment["kind"] != "template":
candidates = [item for item in candidates if item["validation"]["status"] == "valid"]
if not candidates:
raise HTTPException(422, "请至少选择一条候选")
if experiment["kind"] != "template" and any(item["validation"]["status"] != "valid" for item in candidates):
raise HTTPException(422, "候选存在语法、类型或可用性问题,请先解决;至少保留一条已核实候选")
inputs = experiment["inputs"]
if experiment["kind"] == "template":
# Historical collections follow the same syntax/scope contract; old row findings are irrelevant.
for candidate in candidates:
self.check_syntax(candidate["expression"], candidate["client_item_id"])
scope = scope_of(SimulationSettings.model_validate(candidate["settings"]))
if not inputs or any(item["scope"] != scope for item in inputs):
raise HTTPException(422, "数据准备与回测参数组合不一致,请重新生成候选集合")
return DraftInput(
name=experiment["name"],
source=Source(
kind=source_kind or experiment["kind"],
reference=reference or experiment_id,
research_id=experiment_id,
input_snapshot_ids=[i["id"] for i in inputs],
input_snapshot_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
],
)
async def preview(self, experiment_id, candidate_ids=None, source_kind=None, reference=None, *, backtests=None):
draft = await self.backtest_input(experiment_id, candidate_ids, source_kind, reference)
return await (backtests or Backtests(self.db)).preview(PreviewInput(inline=draft), preserve_source=True)
async def start_template_backtest(self, experiment_id, body, *, backtests=None, confirmed_preview=None):
"""Start the explicitly selected immutable collection in the caller's transaction.
Account locking covers preview creation as well as run creation, so concurrent
retries share one run. Reusing a key for another collection/selection raises 409.
The caller must wake the runner only after committing this transaction.
"""
chosen = set(body.candidate_ids)
if len(chosen) != len(body.candidate_ids):
raise HTTPException(422, "候选选择包含重复项")
await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
previous = await self.db.scalar(select(BacktestRun).where(BacktestRun.idempotency_key == body.idempotency_key))
if previous:
saved = await self.db.get(BacktestPreview, previous.preview_id)
if previous.source.get("research_id") != experiment_id or chosen != {
c["client_item_id"] for c in saved.candidates
}:
raise HTTPException(409, "幂等键已用于另一候选集合或选择")
return await Backtests(self.db).run(previous.id)
experiment = await self.get(experiment_id)
if experiment["kind"] != "template":
raise HTTPException(422, "此入口仅用于模板候选集合")
if experiment["archived"]:
raise HTTPException(409, "候选集合已删除")
service = backtests or Backtests(self.db)
if confirmed_preview is None:
preview = await self.preview(experiment_id, body.candidate_ids, backtests=service)
else:
# Approval authorizes all persisted candidates, not the first display page.
draft = (await self.backtest_input(experiment_id, body.candidate_ids)).model_dump(mode="json")
expected = fingerprint({"candidates": draft["candidates"], "source": draft["source"]})
saved = await self.db.get(BacktestPreview, confirmed_preview["preview_id"])
if (
saved is None
or saved.version != confirmed_preview["version"]
or saved.digest != confirmed_preview["digest"]
or saved.digest != expected
or fingerprint({"candidates": saved.candidates, "source": saved.source}) != expected
):
raise HTTPException(409, "回测候选与确认内容不匹配,请重新确认")
preview = confirmed_preview
return await service.start(StartInput(
preview_id=preview["preview_id"], version=preview["version"], idempotency_key=body.idempotency_key,
))
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"]
await Preparations(self.db).bind(body)
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):
await Preparations(self.db).bind(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_ids": item["dataset_ids"], "name": item["name"], "fields": item["fields"][:100]}
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],
}