Simplify template candidate confirmation and direct batch backtesting
Deploy production / deploy (push) Successful in 57s

This commit is contained in:
yuxuanhui
2026-09-20 11:01:07 +08:00
parent 07dd767c52
commit 13a2168ca5
14 changed files with 510 additions and 245 deletions
+81 -23
View File
@@ -7,14 +7,14 @@ 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.contracts import Candidate, DraftInput, PreviewInput, SimulationSettings, Source, StartInput
from ..backtests.service import Backtests, uid
from ..catalog.research_metadata import ResearchMetadata
from ..catalog.service import Catalog
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment
from ..models import Account, Alpha, BacktestPreview, BacktestRun, CatalogResource, ResearchExperiment
from ..preparations.service import Preparations
from .assets import Assets
from .expressions import GROUPS, analyze, expand
from .expressions import GROUPS, ExpressionError, Parser, analyze, expand
from .serialization import encode_snapshot as jsonable_encoder
from .workspace_contracts import TemplateSpec
@@ -45,7 +45,7 @@ class Experiments:
self.catalog = Catalog(db)
self.assets = Assets(db)
async def inputs(self, ids, scope=None):
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]
@@ -56,7 +56,7 @@ class Experiments:
for name, kind in item["field_types"].items():
if name not in item["field_ids"]:
continue
if name in fields and fields[name] != kind:
if check_types and name in fields and fields[name] != kind:
raise HTTPException(422, f"字段 {name} 在不同快照中类型不一致")
fields[name] = kind
return snapshots, fields
@@ -138,7 +138,9 @@ class Experiments:
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)
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
@@ -155,9 +157,9 @@ class Experiments:
if not values:
raise HTTPException(422, f"变量 {name} 没有匹配的 {variable.field_type} 字段,请调整数据准备")
for value in values:
if fields.get(str(value)) != variable.field_type:
if kind != "template" and fields.get(str(value)) != variable.field_type:
raise HTTPException(422, f"变量 {name} 的字段 {value} 不在固定输入中或类型不符")
if variable.kind == "group" and any(
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} 未在固定输入中核实")
@@ -170,16 +172,29 @@ class Experiments:
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)
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"]):
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 = {}
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(
@@ -187,7 +202,7 @@ class Experiments:
).model_dump(mode="json"),
"bindings": item["bindings"],
"input_ids": list(body.input_ids),
"validation": validation,
**findings,
"changes": [
self.diff(parent.get("expression", ""), item["expression"])
for parent in parents
@@ -197,16 +212,21 @@ class Experiments:
)
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,
**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 [
@@ -314,11 +334,20 @@ class Experiments:
candidates = [item for item in candidates if item["client_item_id"] in chosen]
if len(candidates) != len(chosen):
raise HTTPException(422, "选择包含未知候选")
else:
elif experiment["kind"] != "template":
candidates = [item for item in candidates if item["validation"]["status"] == "valid"]
if not candidates or any(item["validation"]["status"] != "valid" for item in candidates):
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 await Backtests(self.db).preview(
PreviewInput(
inline=DraftInput(
@@ -345,6 +374,35 @@ class Experiments:
preserve_source=True,
)
async def start_template_backtest(self, experiment_id, body):
"""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, "候选集合已删除")
preview = await self.preview(experiment_id, body.candidate_ids)
return await Backtests(self.db).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], [])