feat: add versioned research templates and alpha variants
This commit is contained in:
@@ -0,0 +1,446 @@
|
||||
"""Stage-one public API and durable provenance, with isolated platform HTTP."""
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from app.alphas import upsert_alpha
|
||||
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_catalog import SCOPE, prepare, sync
|
||||
from tests.test_catalog import catalog as catalog_fixture
|
||||
|
||||
catalog = catalog_fixture
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def research_input(catalog, app):
|
||||
client, _, _ = catalog
|
||||
await sync(catalog)
|
||||
version = (await sync(catalog, "TEST_FIN"))["id"]
|
||||
response = await prepare(client, version)
|
||||
assert response.status_code == 201
|
||||
async with app.state.sessions.begin() as db:
|
||||
db.add(
|
||||
CatalogResource(
|
||||
key="operators",
|
||||
kind="operators",
|
||||
content={
|
||||
"items": [
|
||||
{"name": "rank", "category": "Cross Sectional"},
|
||||
{"name": "ts_mean", "category": "Time Series"},
|
||||
{"name": "vec_avg", "category": "Vector"},
|
||||
{"name": "group_rank", "category": "Group"},
|
||||
]
|
||||
},
|
||||
)
|
||||
)
|
||||
db.add(
|
||||
CatalogResource(
|
||||
key="settings",
|
||||
kind="settings",
|
||||
content={
|
||||
"items": [
|
||||
{**SCOPE, "neutralizations": ["INDUSTRY", "NONE"]},
|
||||
]
|
||||
},
|
||||
)
|
||||
)
|
||||
return response.json()
|
||||
|
||||
|
||||
def template():
|
||||
return {
|
||||
"name": "测试字段排序",
|
||||
"expression": "rank({field})",
|
||||
"description": "测试经济假设",
|
||||
"variables": {
|
||||
"field": {"kind": "field", "field_type": "MATRIX", "values": ["TEST_FIN_001", "TEST_FIN_002"]}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def expansion(input_id, **kwargs):
|
||||
return {
|
||||
"template": template(),
|
||||
"input_ids": [input_id],
|
||||
"hypothesis": "排序比较",
|
||||
"settings": {key: SCOPE[key] for key in ("region", "universe", "delay")},
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expression,expected",
|
||||
[
|
||||
("x = close; rank(x)", "valid"),
|
||||
("rank(close)", "valid"),
|
||||
("rank(unknown)", "needs_review"),
|
||||
("rank(vec_avg(v))", "valid"),
|
||||
("rank(v)", "invalid"),
|
||||
("v", "invalid"),
|
||||
("abs(v)", "invalid"),
|
||||
("x=v; vec_avg(x)", "valid"),
|
||||
("rank(v + 1)", "invalid"),
|
||||
("group_rank(close, true)", "invalid"),
|
||||
("group_rank(close,industry)", "valid"),
|
||||
("rank()", "invalid"),
|
||||
("ts_mean(close)", "invalid"),
|
||||
("rank(close @)", "invalid"),
|
||||
("x=close", "invalid"),
|
||||
("rank(future)", "needs_review"),
|
||||
("rank(close,,)", "invalid"),
|
||||
],
|
||||
)
|
||||
def test_expression_provenance(expression, expected):
|
||||
result = analyze(
|
||||
expression,
|
||||
{"close": "MATRIX", "v": "VECTOR", "future": "FUTURE"},
|
||||
{"rank", "vec_avg", "group_rank", "ts_mean"},
|
||||
)
|
||||
assert result["status"] == expected, result
|
||||
if expression.startswith("x ="):
|
||||
assert result["fields"] == ["close"] and result["locals"] == ["x"]
|
||||
|
||||
|
||||
def test_bounded_sampling_and_repeated_placeholders():
|
||||
values = {f"p{i}": list(range(100)) for i in range(20)}
|
||||
expression = "+".join("{" + name + "}" for name in values)
|
||||
a = expand(expression, values, "random", 50, 7)
|
||||
assert a == expand(expression, values, "random", 50, 7)
|
||||
assert len(a["items"]) == 50 and len({r["expression"] for r in a["items"]}) == 50
|
||||
assert a["combination_count"] == str(100**20)
|
||||
assert expand("<x/> + {x}", {"x": [1, 2]})["items"] == [
|
||||
{"expression": "1 + 1", "bindings": {"x": 1}},
|
||||
{"expression": "2 + 2", "bindings": {"x": 2}},
|
||||
]
|
||||
with pytest.raises(ValueError):
|
||||
expand(expression, values, "all", 100)
|
||||
|
||||
|
||||
async def test_template_version_expansion_preview_and_backtest(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()
|
||||
body = expansion(research_input["id"], asset_id=asset["id"], version=1)
|
||||
body.pop("template")
|
||||
generated = await logged_in.post("/api/v1/research/experiments", json=body)
|
||||
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"])
|
||||
modified = template()
|
||||
modified["expression"] = "-rank({field})"
|
||||
response = await logged_in.put(
|
||||
f"/api/v1/research/assets/{asset['id']}", json={"kind": "template", "version": 1, "content": modified}
|
||||
)
|
||||
assert response.status_code == 200 and response.json()["version"] == 2
|
||||
assert (await logged_in.get(f"/api/v1/research/assets/{asset['id']}?version=1")).json()["content"][
|
||||
"expression"
|
||||
] == "rank({field})"
|
||||
assert (
|
||||
await logged_in.put(
|
||||
f"/api/v1/research/assets/{asset['id']}",
|
||||
json={"kind": "template", "version": 1, "content": modified},
|
||||
)
|
||||
).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")
|
||||
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"]
|
||||
origins = (await logged_in.get(f"/api/v1/alphas/{aid}/sources")).json()
|
||||
assert origins["items"][0]["source"]["research_id"] == experiment["id"]
|
||||
old = (await logged_in.get(f"/api/v1/research/experiments/{experiment['id']}")).json()
|
||||
assert old["evidence"]["template"]["version"] == 1
|
||||
assert old["backtest_run_ids"] == [run["backtest_run_id"]]
|
||||
|
||||
|
||||
async def test_invalid_fields_and_unknown_operators_never_start(app, logged_in, research_input):
|
||||
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})"
|
||||
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
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
||||
|
||||
|
||||
async def test_import_preview_conflict_and_explicit_commit(logged_in):
|
||||
legacy = {
|
||||
"name": "legacy",
|
||||
"expression": "rank(<field/>)",
|
||||
"templateConfigurations": {"field": {"variables": ["close"]}},
|
||||
}
|
||||
preview = (
|
||||
await logged_in.post("/api/v1/research/templates/import-preview", json={"templates": [legacy]})
|
||||
).json()
|
||||
assert preview["templates"][0]["expression"] == "rank({field})"
|
||||
body = {"templates": preview["templates"], "digest": preview["digest"]}
|
||||
assert (await logged_in.post("/api/v1/research/templates/import", json=body)).status_code == 201
|
||||
assert (await logged_in.post("/api/v1/research/templates/import", json=body)).status_code == 409
|
||||
second = (
|
||||
await logged_in.post("/api/v1/research/templates/import-preview", json={"templates": [legacy]})
|
||||
).json()
|
||||
assert second["conflicts"][0]["name"] == "legacy"
|
||||
|
||||
|
||||
async def test_operator_annotation_and_refresh_preserve_local(app, logged_in, research_input):
|
||||
response = await logged_in.patch(
|
||||
"/api/v1/catalog/operators/rank/research", json={"note": "排名", "favorite": True, "version": 0}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
async with app.state.sessions.begin() as db:
|
||||
await ResearchMetadata(db).publish(
|
||||
"operators", "operators", {"items": [{"name": "rank", "category": "updated"}]}
|
||||
)
|
||||
result = (await logged_in.get("/api/v1/catalog/operators?favorite=true")).json()
|
||||
assert result["items"][0]["local"]["note"] == "排名"
|
||||
assert (
|
||||
await logged_in.patch("/api/v1/catalog/operators/rank/research", json={"note": "wrong", "version": 0})
|
||||
).status_code == 409
|
||||
|
||||
|
||||
async def test_settings_variant_requires_all_fields_in_target(app, logged_in, research_input, catalog):
|
||||
from app.models import CatalogScope, TemplateInput
|
||||
from app.research.experiments import Experiments
|
||||
from app.research.workspace_contracts import SettingVariants
|
||||
|
||||
target_scope = {**SCOPE, "region": "EUR"}
|
||||
target_key = f"EQUITY|EUR|{SCOPE['universe']}|1"
|
||||
async with app.state.sessions.begin() as db:
|
||||
await upsert_alpha(
|
||||
db,
|
||||
alpha(
|
||||
"seed",
|
||||
regular={"code": "x = TEST_FIN_001; rank(x + TEST_FIN_002)"},
|
||||
settings={**{k: SCOPE[k] for k in ("region", "universe", "delay")}, "language": "FASTEXPR"},
|
||||
),
|
||||
)
|
||||
# Reuse an immutable field batch, with an explicit target-scope test snapshot.
|
||||
db.add(CatalogScope(key=target_key, scope=target_scope))
|
||||
await db.flush()
|
||||
db.add(
|
||||
TemplateInput(
|
||||
id="target",
|
||||
scope_key=target_key,
|
||||
dataset_id="TEST_FIN",
|
||||
collection_version=research_input["collection_version"],
|
||||
selection="explicit",
|
||||
field_ids=["TEST_FIN_001"],
|
||||
field_types={"TEST_FIN_001": "MATRIX"},
|
||||
)
|
||||
)
|
||||
await db.flush()
|
||||
metadata = await db.get(CatalogResource, "settings")
|
||||
metadata.content = {
|
||||
"items": metadata.content["items"] + [{**target_scope, "neutralizations": ["INDUSTRY"]}]
|
||||
}
|
||||
result = await Experiments(db).setting_variants(
|
||||
SettingVariants(alpha_id="seed", input_ids=["target"])
|
||||
)
|
||||
assert result["candidates"][0]["validation"]["status"] == "needs_review"
|
||||
assert "TEST_FIN_002" in str(result["candidates"][0]["validation"]["availability"])
|
||||
assert result["candidates"][0]["expression"] == "x = TEST_FIN_001; rank(x + TEST_FIN_002)"
|
||||
|
||||
|
||||
async def test_workspace_auth(client):
|
||||
assert (await client.get("/api/v1/research/assets")).status_code == 401
|
||||
assert (
|
||||
await client.post("/api/v1/catalog/operators/refresh", headers={"Origin": "https://other.test"})
|
||||
).status_code == 403
|
||||
|
||||
|
||||
async def test_model_generation_evidence_is_persisted(app, logged_in, research_input):
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from pydantic_ai.messages import ModelResponse, ToolCallPart
|
||||
from pydantic_ai.models.function import FunctionModel
|
||||
|
||||
from tests.test_ai import configure
|
||||
|
||||
await configure(app, logged_in)
|
||||
calls = []
|
||||
|
||||
def complete(messages, info):
|
||||
calls.append(messages)
|
||||
return ModelResponse(parts=[ToolCallPart(info.output_tools[0].name, template())])
|
||||
|
||||
@asynccontextmanager
|
||||
async def factory(config, settings):
|
||||
yield FunctionModel(function=complete, model_name="research-test")
|
||||
|
||||
app.state.ai.model_factory = factory
|
||||
response = await logged_in.post(
|
||||
"/api/v1/research/generate",
|
||||
json={"name": "生成测试", "hypothesis": "比较字段排序", "input_ids": [research_input["id"]]},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
saved = (await logged_in.get(f"/api/v1/research/assets/{response.json()['id']}?version=1")).json()
|
||||
assert saved["provenance"]["generation"]["model"] == "test-model"
|
||||
assert saved["provenance"]["context"]["inputs"][0]["id"] == research_input["id"]
|
||||
assert saved["provenance"]["generation"]["usage"]["requests"] == 1
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
async def test_native_ai_tools_share_experiment_and_confirmation_boundary(app, logged_in, research_input):
|
||||
from app.ai.capabilities import ToolContext
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.business import Business
|
||||
|
||||
async with app.state.sessions.begin() as db:
|
||||
ctx = ToolContext(Business(db))
|
||||
experiment = await CAPABILITIES["expand_research_template"].invoke(
|
||||
ctx, expansion(research_input["id"])
|
||||
)
|
||||
preview = await CAPABILITIES["prepare_experiment_backtest"].invoke(
|
||||
ctx, {"experiment_id": experiment["id"]}
|
||||
)
|
||||
assert preview["source"]["research_id"] == experiment["id"]
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
||||
assert CAPABILITIES["start_backtest"].requires_confirmation
|
||||
|
||||
|
||||
async def test_comparison_aligns_only_common_dates_and_preserves_nulls(app, logged_in):
|
||||
from app.models import Pnl
|
||||
|
||||
async with app.state.sessions.begin() as db:
|
||||
await upsert_alpha(db, alpha("baseline", **{"is": {"sharpe": None}}))
|
||||
await upsert_alpha(db, alpha("candidate", settings={"region": "EUR"}))
|
||||
await db.flush()
|
||||
db.add(
|
||||
Pnl(
|
||||
alpha_id="baseline",
|
||||
raw={},
|
||||
points=[
|
||||
{"date": "2025-01-01", "value": 1},
|
||||
{"date": "2025-01-02", "value": 3},
|
||||
{"date": "2025-01-03", "value": 5},
|
||||
],
|
||||
)
|
||||
)
|
||||
db.add(
|
||||
Pnl(
|
||||
alpha_id="candidate",
|
||||
raw={},
|
||||
points=[
|
||||
{"date": "2025-01-02", "value": 8},
|
||||
{"date": "2025-01-03", "value": 7},
|
||||
{"date": "2025-01-04", "value": 12},
|
||||
],
|
||||
)
|
||||
)
|
||||
response = await logged_in.post("/api/v1/research/compare", json={"alpha_ids": ["baseline", "candidate"]})
|
||||
assert response.status_code == 200, response.text
|
||||
result = response.json()
|
||||
assert result["common_dates"] == ["2025-01-02", "2025-01-03"]
|
||||
assert result["items"][0]["metrics"]["sharpe"] is None
|
||||
assert result["items"][1]["pnl"] == [
|
||||
{"date": "2025-01-02", "value": 0},
|
||||
{"date": "2025-01-03", "value": -1},
|
||||
]
|
||||
assert result["different_settings"]
|
||||
|
||||
|
||||
def test_actual_cnhk_setting_choice_nesting_is_supported():
|
||||
from app.catalog.research_metadata import setting_rows
|
||||
from tests.catalog_fake import platform_response
|
||||
|
||||
response = platform_response()
|
||||
children = response["actions"]["POST"]["settings"]["children"]
|
||||
for key in ("region", "delay", "universe"):
|
||||
children[key]["choices"] = children[key]["choices"]["instrumentType"]
|
||||
children["neutralization"] = {"choices": [{"value": "NONE"}]}
|
||||
rows = setting_rows(response)
|
||||
assert any(
|
||||
row["region"] == "USA" and row["delay"] == 0 and row["neutralizations"] == ["NONE"] for row in rows
|
||||
)
|
||||
|
||||
|
||||
async def test_published_input_does_not_override_conflicting_field_evidence(app, logged_in, research_input):
|
||||
async with app.state.sessions.begin() as db:
|
||||
await ResearchMetadata(db).publish(
|
||||
"availability-fixture",
|
||||
"availability",
|
||||
{
|
||||
"field_id": "TEST_FIN_001",
|
||||
"scope": SCOPE,
|
||||
"status": "available",
|
||||
"items": [{**SCOPE, "universe": "TOP1000"}],
|
||||
},
|
||||
)
|
||||
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"
|
||||
denied = await logged_in.post(
|
||||
f"/api/v1/research/experiments/{experiment['id']}/preview", json={"candidate_ids": ["c1"]}
|
||||
)
|
||||
assert denied.status_code == 422
|
||||
duplicate = await logged_in.post(
|
||||
f"/api/v1/research/experiments/{experiment['id']}/preview", json={"candidate_ids": ["c2", "c2"]}
|
||||
)
|
||||
assert duplicate.status_code == 422
|
||||
|
||||
|
||||
async def test_target_scope_full_input_and_parent_template_are_traceable(
|
||||
app, logged_in, research_input, catalog
|
||||
):
|
||||
target = {**SCOPE, "universe": "TOP1000"}
|
||||
await sync(catalog, scope=target)
|
||||
version = (await sync(catalog, "TEST_FIN", scope=target))["id"]
|
||||
target_input = (await prepare(logged_in, version, scope=target)).json()
|
||||
async with app.state.sessions.begin() as db:
|
||||
await upsert_alpha(db, alpha("seed", regular={"code": "x = TEST_FIN_001; rank(x)"}))
|
||||
metadata = await db.get(CatalogResource, "settings")
|
||||
metadata.content = {
|
||||
"items": metadata.content["items"] + [{**target, "neutralizations": ["INDUSTRY"]}]
|
||||
}
|
||||
response = await logged_in.post(
|
||||
"/api/v1/research/variants/settings",
|
||||
json={"alpha_id": "seed", "input_ids": [research_input["id"], target_input["id"]]},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
variant = response.json()
|
||||
candidate = variant["candidates"][0]
|
||||
assert candidate["validation"]["status"] == "valid"
|
||||
assert candidate["input_ids"] == [target_input["id"]]
|
||||
assert candidate["expression"] == "x = TEST_FIN_001; rank(x)"
|
||||
assert len(variant["inputs"]) == 2
|
||||
preview = (await logged_in.post(f"/api/v1/research/experiments/{variant['id']}/preview", json={})).json()
|
||||
assert preview["source"]["research_id"] == variant["id"]
|
||||
assert preview["items"][0]["client_item_id"] == candidate["client_item_id"]
|
||||
child = (
|
||||
await logged_in.post(
|
||||
"/api/v1/research/experiments",
|
||||
json=expansion(research_input["id"], parent_experiment_ids=[variant["id"]]),
|
||||
)
|
||||
).json()
|
||||
assert child["parents"][0]["input_references"][1]["id"] == target_input["id"]
|
||||
|
||||
|
||||
def test_partial_availability_and_deep_expression_fail_closed():
|
||||
from app.catalog.research_metadata import normalize_availability
|
||||
|
||||
result = normalize_availability(
|
||||
{
|
||||
"availability": [
|
||||
{"instrumentType": "EQUITY", "region": "USA", "universe": "TOP3000", "delay": 1},
|
||||
{"region": "USA"},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert result["status"] == "needs_review"
|
||||
assert analyze("+".join(["close"] * 2000), {"close": "MATRIX"}, set())["status"] == "invalid"
|
||||
Reference in New Issue
Block a user