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], [])
+10
View File
@@ -2,6 +2,7 @@
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from ..backtests.contracts import RunOutput
from ..security import require_auth
from .assets import Assets
from .comparisons import compare
@@ -22,6 +23,7 @@ from .workspace_contracts import (
ImportCommit,
ImportPreview,
SettingVariants,
TemplateBacktest,
WorkflowSpec,
)
@@ -149,6 +151,14 @@ async def preview(experiment_id: str, body: ExperimentPreview, request: Request)
return await Experiments(db).preview(experiment_id, body.candidate_ids)
@router.post("/experiments/{experiment_id}/backtest", status_code=202, response_model=RunOutput)
async def template_backtest(experiment_id: str, body: TemplateBacktest, request: Request):
async with request.app.state.sessions.begin() as db:
result = await Experiments(db).start_template_backtest(experiment_id, body)
request.app.state.runner.backtests.wake.set()
return result
@router.post("/variants/settings", status_code=201)
async def settings_variants(body: SettingVariants, request: Request):
async with request.app.state.sessions.begin() as db:
+2 -1
View File
@@ -522,7 +522,8 @@ class ResearchRuntime:
"type": "candidates",
"experiment_id": experiment["id"],
"candidate_ids": [
c["client_item_id"] for c in experiment["candidates"] if c["validation"]["status"] == "valid"
c["client_item_id"] for c in experiment["candidates"]
if experiment["kind"] == "template" or c["validation"]["status"] == "valid"
],
}
ids = step.output["candidate_ids"]
@@ -144,6 +144,11 @@ class ExperimentPreview(Contract):
candidate_ids: list[str] | None = Field(default=None, min_length=1, max_length=10000)
class TemplateBacktest(Contract):
candidate_ids: list[str] = Field(min_length=1, max_length=10000)
idempotency_key: str = Field(min_length=1, max_length=100)
class EvaluationRules(Contract):
version: Literal["research-v1"] = "research-v1"
sharpe_min: float = Field(default=1.0, allow_inf_nan=False)
+1 -1
View File
@@ -138,7 +138,7 @@ async def test_sdk_template_creation_frozen_evidence_and_web_expansion(app, logg
assert {c["expression"] for c in experiment["candidates"]} == {
f"rank({field}) + {offset}" for field in ["TEST_FIN_001", "TEST_FIN_002"] for offset in [0, 1, 5]
}
assert all(c["validation"]["status"] == "valid" for c in experiment["candidates"])
assert all("validation" not in c for c in experiment["candidates"])
assert experiment["evidence"]["template"]["provenance"] == stored["provenance"]
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 1
+91 -19
View File
@@ -8,7 +8,7 @@ from app.catalog.research_metadata import ResearchMetadata
from app.models import BacktestRun, CatalogResource, ResearchExperiment
from app.research.expressions import analyze, expand
from tests.conftest import alpha
from tests.test_backtests import execute, setup, start
from tests.test_backtests import execute, setup
from tests.test_catalog import SCOPE, prepare, sync
from tests.test_catalog import catalog as catalog_fixture
@@ -120,7 +120,7 @@ def test_bounded_sampling_and_repeated_placeholders():
expand(expression, values, "all", 100)
async def test_template_version_expansion_preview_and_backtest(app, logged_in, research_input):
async def test_template_version_expansion_direct_backtest_and_idempotency(app, logged_in, research_input):
saved = await logged_in.post("/api/v1/research/assets", json={"kind": "template", "content": template()})
assert saved.status_code == 201, saved.text
asset = saved.json()
@@ -130,7 +130,7 @@ async def test_template_version_expansion_preview_and_backtest(app, logged_in, r
assert generated.status_code == 201, generated.text
experiment = generated.json()
assert len(experiment["candidates"]) == 2
assert all(c["validation"]["status"] == "valid" for c in experiment["candidates"])
assert all("validation" not in c for c in experiment["candidates"])
modified = template()
modified["expression"] = "-rank({field})"
response = await logged_in.put(
@@ -147,10 +147,24 @@ async def test_template_version_expansion_preview_and_backtest(app, logged_in, r
)
).status_code == 409
platform, lane = await setup(app)
preview = await logged_in.post(f"/api/v1/research/experiments/{experiment['id']}/preview", json={})
assert preview.status_code == 201, preview.text
assert not platform.posts
run = await start(logged_in, preview.json(), "research-stage-one")
from app.models import BacktestPreview
url = f"/api/v1/research/experiments/{experiment['id']}/backtest"
request = {"candidate_ids": ["c2"], "idempotency_key": "template-confirmation"}
result = await logged_in.post(url, json=request)
assert result.status_code == 202, result.text
run = result.json()
assert run["total"] == 1
assert run["source"]["kind"] == "template"
assert run["source"]["input_snapshot_ids"] == [research_input["id"]]
retry = await logged_in.post(url, json=request)
assert retry.status_code == 202 and retry.json()["backtest_run_id"] == run["backtest_run_id"]
conflict = await logged_in.post(url, json={**request, "candidate_ids": ["c1"]})
assert conflict.status_code == 409
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 1
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 1
assert lane.wake.is_set()
await execute(app, lane, run["backtest_run_id"])
results = (await logged_in.get(f"/api/v1/backtests/runs/{run['backtest_run_id']}/results")).json()
aid = results["items"][0]["alpha_id"]
@@ -161,22 +175,34 @@ async def test_template_version_expansion_preview_and_backtest(app, logged_in, r
assert old["backtest_run_ids"] == [run["backtest_run_id"]]
async def test_invalid_fields_and_unknown_operators_never_start(app, logged_in, research_input):
async def test_template_generation_does_not_require_field_operator_or_settings_evidence(app, logged_in, research_input):
from sqlalchemy import delete
async with app.state.sessions.begin() as db:
await db.execute(delete(CatalogResource))
body = expansion(research_input["id"])
body["template"]["variables"]["field"]["values"] = ["other_field"]
assert (await logged_in.post("/api/v1/research/experiments", json=body)).status_code == 422
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(ResearchExperiment)) == 0
body = expansion(research_input["id"])
body["template"]["expression"] = "made_up({field})"
body["template"]["expression"] = "made_up({field}) + vec_avg(TEST_FIN_001)"
response = await logged_in.post("/api/v1/research/experiments", json=body)
assert response.status_code == 201, response.text
eid = response.json()["id"]
assert (await logged_in.post(f"/api/v1/research/experiments/{eid}/preview", json={})).status_code == 422
experiment = response.json()
assert all("validation" not in c for c in experiment["candidates"])
assert not {"operators_snapshot", "settings_snapshot", "field_availability"} & experiment["evidence"].keys()
assert (await logged_in.post(f"/api/v1/research/experiments/{experiment['id']}/preview", json={})).status_code == 201
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
@pytest.mark.parametrize("expression", ["rank({field}", "rank({field},,)", "x = {field}"])
async def test_template_syntax_errors_reject_whole_collection(app, logged_in, research_input, expression):
body = expansion(research_input["id"])
body["template"]["expression"] = expression
response = await logged_in.post("/api/v1/research/experiments", json=body)
assert response.status_code == 422 and "语法错误" in response.text
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(ResearchExperiment)) == 0
async def test_import_preview_conflict_and_explicit_commit(logged_in):
legacy = {
"name": "legacy",
@@ -367,7 +393,7 @@ def test_actual_cnhk_setting_choice_nesting_is_supported():
)
async def test_published_input_does_not_override_conflicting_field_evidence(app, logged_in, research_input):
async def test_template_candidates_ignore_conflicting_field_evidence(app, logged_in, research_input):
async with app.state.sessions.begin() as db:
await ResearchMetadata(db).publish(
"availability-fixture",
@@ -382,12 +408,11 @@ async def test_published_input_does_not_override_conflicting_field_evidence(app,
experiment = (
await logged_in.post("/api/v1/research/experiments", json=expansion(research_input["id"]))
).json()
assert experiment["candidates"][0]["validation"]["status"] == "needs_review"
assert experiment["candidates"][1]["validation"]["status"] == "valid"
assert all("validation" not in c for c in experiment["candidates"])
denied = await logged_in.post(
f"/api/v1/research/experiments/{experiment['id']}/preview", json={"candidate_ids": ["c1"]}
)
assert denied.status_code == 422
assert denied.status_code == 201
duplicate = await logged_in.post(
f"/api/v1/research/experiments/{experiment['id']}/preview", json={"candidate_ids": ["c2", "c2"]}
)
@@ -607,3 +632,50 @@ async def test_template_contract_rejects_removed_scope(logged_in):
assert response.status_code == 422, response.text
assert "scope" in response.text
assert "scope" not in TemplateSpec.model_json_schema()["properties"]
async def test_template_start_checks_selection_and_ignores_old_row_validation(app, logged_in, research_input):
from app.models import BacktestPreview
experiment = (await logged_in.post("/api/v1/research/experiments", json=expansion(research_input["id"]))).json()
async with app.state.sessions.begin() as db:
row = await db.get(ResearchExperiment, experiment["id"])
row.candidates = [{**c, "validation": {"status": "needs_review", "syntax": [], "types": [],
"availability": ["历史字段未核实"]}} for c in row.candidates]
await setup(app)
url = f"/api/v1/research/experiments/{experiment['id']}/backtest"
for ids in [[], ["unknown"], ["c1", "c1"]]:
response = await logged_in.post(url, json={"candidate_ids": ids, "idempotency_key": "confirm-old"})
assert response.status_code == 422, response.text
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 0
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
response = await logged_in.post(url, json={"candidate_ids": ["c1", "c2"], "idempotency_key": "confirm-old"})
assert response.status_code == 202 and response.json()["total"] == 2
retry = await logged_in.post(url, json={"candidate_ids": ["c2", "c1"], "idempotency_key": "confirm-old"})
assert retry.status_code == 202 and retry.json()["backtest_run_id"] == response.json()["backtest_run_id"]
@pytest.mark.parametrize("change", ["syntax", "scope", "archived", "disconnected"])
async def test_template_start_rejects_unusable_collection_without_partial_writes(app, logged_in, research_input, change):
from app.models import Account, BacktestPreview
experiment = (await logged_in.post("/api/v1/research/experiments", json=expansion(research_input["id"]))).json()
async with app.state.sessions.begin() as db:
row = await db.get(ResearchExperiment, experiment["id"])
if change == "syntax":
row.candidates = [{**c, "expression": "rank("} for c in row.candidates]
elif change == "scope":
row.candidates = [{**c, "settings": {**c["settings"], "region": "EUR"}} for c in row.candidates]
elif change == "archived":
row.archived = True
else:
(await db.get(Account, 1)).connection_status = "disconnected"
app.state.runner.backtests.wake.clear()
response = await logged_in.post(f"/api/v1/research/experiments/{experiment['id']}/backtest",
json={"candidate_ids": ["c1"], "idempotency_key": "invalid"})
assert response.status_code == (422 if change in ("syntax", "scope") else 409), response.text
assert not app.state.runner.backtests.wake.is_set()
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 0
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0