2026-09-08 12:43:00 +08:00
|
|
|
"""Public research workflow; only the model and WorldQuant HTTP are synthetic."""
|
|
|
|
|
|
|
|
|
|
import copy
|
|
|
|
|
|
|
|
|
|
import pytest
|
|
|
|
|
from fastapi import HTTPException
|
|
|
|
|
from sqlalchemy import func, select
|
|
|
|
|
|
2026-09-08 19:28:44 +08:00
|
|
|
from app.ai.capabilities import ToolContext
|
|
|
|
|
from app.ai.tools import CAPABILITIES
|
2026-09-08 12:43:00 +08:00
|
|
|
from app.alphas import upsert_alpha
|
|
|
|
|
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
|
|
|
|
|
from app.business import Business
|
2026-09-12 01:24:02 +08:00
|
|
|
from app.models import BacktestPreview, BacktestRun, Research, ResearchInputSnapshot
|
2026-09-08 12:43:00 +08:00
|
|
|
from tests.test_ai import configure, single_tool_factory
|
|
|
|
|
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 fixed_input(catalog):
|
|
|
|
|
client, _, _ = catalog
|
|
|
|
|
await sync(catalog)
|
|
|
|
|
version = (await sync(catalog, "TEST_FIN"))["id"]
|
|
|
|
|
response = await prepare(client, version)
|
|
|
|
|
assert response.status_code == 201
|
|
|
|
|
return response.json()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def ask(client, conversation, message, context=None, request_id="research"):
|
|
|
|
|
response = await client.post(
|
|
|
|
|
f"/api/v1/ai/conversations/{conversation}/runs",
|
|
|
|
|
json={
|
|
|
|
|
"request_id": request_id,
|
|
|
|
|
"message": message,
|
|
|
|
|
"context": context or {},
|
|
|
|
|
},
|
|
|
|
|
)
|
|
|
|
|
assert response.status_code == 200, response.text
|
|
|
|
|
return (await client.get(f"/api/v1/ai/runs/{response.headers['x-ai-run-id']}")).json()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def construction(input_id):
|
|
|
|
|
return {
|
|
|
|
|
"name": "字段研究",
|
|
|
|
|
"hypothesis": "显式字段排序",
|
2026-09-12 01:24:02 +08:00
|
|
|
"input_snapshot_id": input_id,
|
2026-09-08 12:43:00 +08:00
|
|
|
"candidates": [
|
|
|
|
|
{
|
|
|
|
|
"client_item_id": "one",
|
|
|
|
|
"expression_template": "rank({signal})",
|
|
|
|
|
"bindings": {"signal": {"field_id": "TEST_FIN_001", "field_type": "MATRIX"}},
|
|
|
|
|
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
|
|
|
|
|
}
|
|
|
|
|
],
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.parametrize("use_saved_input", [True, False])
|
|
|
|
|
async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixed_input, use_saved_input):
|
|
|
|
|
platform, lane = await setup(app)
|
|
|
|
|
await configure(app, logged_in)
|
|
|
|
|
conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
|
|
|
|
|
context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"}
|
|
|
|
|
if use_saved_input:
|
2026-09-12 01:24:02 +08:00
|
|
|
context["input_snapshot_id"] = fixed_input["id"]
|
2026-09-08 12:43:00 +08:00
|
|
|
run = await ask(logged_in, conversation, "研究此输入" if use_saved_input else "自行选字段研究", context)
|
|
|
|
|
assert run["status"] == "waiting_approval", run
|
|
|
|
|
assert not platform.posts
|
|
|
|
|
approval = next(c for c in run["tools"] if c["name"] == "start_backtest")
|
|
|
|
|
source = approval["preview"]["backtest"]["source"]
|
|
|
|
|
assert source["kind"] == "chatbox"
|
|
|
|
|
assert source["reference"] == conversation
|
|
|
|
|
assert source["research_id"] == run["id"]
|
2026-09-12 01:24:02 +08:00
|
|
|
assert source["input_snapshot_id"]
|
2026-09-08 12:43:00 +08:00
|
|
|
assert approval["preview"]["backtest"]["items"][0]["expression"] == "rank(TEST_FIN_001)"
|
|
|
|
|
async with app.state.sessions() as db:
|
|
|
|
|
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
|
|
|
|
for _ in range(2):
|
|
|
|
|
assert (
|
|
|
|
|
await logged_in.post(f"/api/v1/ai/approvals/{approval['id']}/decision", json={"approved": True})
|
|
|
|
|
).status_code == 200
|
|
|
|
|
completed_chat = (await logged_in.get(f"/api/v1/ai/runs/{run['id']}")).json()
|
|
|
|
|
assert completed_chat["status"] == "completed", completed_chat
|
|
|
|
|
runs = (
|
|
|
|
|
await logged_in.get("/api/v1/backtests/runs", params={"source": "chatbox", "reference": conversation})
|
|
|
|
|
).json()
|
|
|
|
|
assert runs["total"] == 1
|
|
|
|
|
rid = runs["items"][0]["backtest_run_id"]
|
|
|
|
|
assert runs["items"][0]["source"] == source
|
|
|
|
|
await execute(app, lane, rid)
|
|
|
|
|
followup = await ask(logged_in, conversation, "解读研究结果", request_id="results")
|
|
|
|
|
assert followup["status"] == "completed", followup
|
|
|
|
|
result = next(c["result"] for c in followup["tools"] if c["name"] == "get_backtest_results")
|
|
|
|
|
item = result["items"][0]
|
|
|
|
|
assert item["persistence_status"] == "saved"
|
|
|
|
|
assert item["result"]["is"]["sharpe"] is None
|
|
|
|
|
aid = item["alpha_id"]
|
|
|
|
|
origins = (await logged_in.get(f"/api/v1/alphas/{aid}/sources")).json()
|
|
|
|
|
assert origins["items"][0]["source"] == source
|
|
|
|
|
filtered = (
|
|
|
|
|
await logged_in.get("/api/v1/alphas", params={"source": "chatbox", "research_id": run["id"]})
|
|
|
|
|
).json()
|
|
|
|
|
assert filtered["total"] == 1 and filtered["items"][0]["id"] == aid
|
|
|
|
|
assert filtered["items"][0]["source_kinds"] == ["chatbox"]
|
|
|
|
|
assert len(platform.posts) == 1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
|
|
|
"invalid", ["type", "field", "scope", "placeholder", "duplicate", "unknown_type", "missing_input"]
|
|
|
|
|
)
|
|
|
|
|
async def test_invalid_construction_has_no_partial_preview(logged_in, app, fixed_input, invalid):
|
|
|
|
|
body = construction(fixed_input["id"])
|
|
|
|
|
item = body["candidates"][0]
|
|
|
|
|
if invalid == "type":
|
|
|
|
|
item["bindings"]["signal"]["field_type"] = "VECTOR"
|
|
|
|
|
elif invalid == "field":
|
|
|
|
|
item["bindings"]["signal"]["field_id"] = "OTHER_001"
|
|
|
|
|
elif invalid == "scope":
|
|
|
|
|
item["settings"]["delay"] = 0
|
|
|
|
|
elif invalid == "placeholder":
|
|
|
|
|
item["expression_template"] = "rank({missing})"
|
|
|
|
|
elif invalid == "duplicate":
|
|
|
|
|
body["candidates"].append(copy.deepcopy(item))
|
|
|
|
|
elif invalid == "unknown_type":
|
|
|
|
|
item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"}
|
|
|
|
|
else:
|
2026-09-12 01:24:02 +08:00
|
|
|
body["input_snapshot_id"] = "missing"
|
2026-09-08 12:43:00 +08:00
|
|
|
response = await logged_in.post("/api/v1/backtests/research-previews", json=body)
|
|
|
|
|
assert response.status_code == (404 if invalid == "missing_input" else 422), response.text
|
|
|
|
|
async with app.state.sessions() as db:
|
|
|
|
|
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in, catalog, fixed_input):
|
|
|
|
|
async def tool(name, args):
|
|
|
|
|
async with app.state.sessions.begin() as db:
|
2026-09-08 19:28:44 +08:00
|
|
|
return await CAPABILITIES[name].invoke(ToolContext(Business(db)), args)
|
2026-09-08 12:43:00 +08:00
|
|
|
|
|
|
|
|
page = await tool("get_research_input", {"input_id": fixed_input["id"], "offset": 100, "limit": 25})
|
|
|
|
|
assert page["field_count"] == page["total"] == 123 and len(page["items"]) == 23
|
|
|
|
|
assert page["items"][-1]["field_type"] == "FUTURE_TYPE" and not page["has_more"]
|
|
|
|
|
assert page["_meta"]["source"] == "local_database"
|
2026-09-12 01:24:02 +08:00
|
|
|
subset = (await logged_in.post("/api/v1/data-preparations", json={"name": "one field", "scope": SCOPE,
|
|
|
|
|
"fields": [{"scope": SCOPE, "dataset_id": "TEST_FIN", "field_id": "TEST_FIN_001", "source": "local",
|
|
|
|
|
"collection_version": fixed_input["fields"][0]["collection_version"]}]})).json()
|
|
|
|
|
selected = await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
2026-09-08 12:43:00 +08:00
|
|
|
assert selected["field_count"] == 1
|
|
|
|
|
bad = construction(selected["id"])
|
|
|
|
|
bad["candidates"][0]["bindings"]["signal"]["field_id"] = "TEST_FIN_002"
|
|
|
|
|
assert (await logged_in.post("/api/v1/backtests/research-previews", json=bad)).status_code == 422
|
|
|
|
|
state = catalog[2]
|
|
|
|
|
state["fields"][1]["type"] = "VECTOR"
|
|
|
|
|
await sync(catalog, "TEST_FIN")
|
|
|
|
|
old = await tool("get_research_input", {"input_id": fixed_input["id"], "q": "TEST_FIN_001"})
|
|
|
|
|
assert old["items"][0]["field_type"] == "MATRIX"
|
|
|
|
|
response = await logged_in.post(
|
|
|
|
|
"/api/v1/backtests/research-previews", json=construction(fixed_input["id"])
|
|
|
|
|
)
|
|
|
|
|
assert response.status_code == 201, response.text
|
2026-09-12 01:24:02 +08:00
|
|
|
await logged_in.patch(f"/api/v1/data-preparations/{subset['id']}", json={"name": "edited", "version": subset["version"]})
|
2026-09-08 12:43:00 +08:00
|
|
|
with pytest.raises(HTTPException) as exc:
|
2026-09-12 01:24:02 +08:00
|
|
|
await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
2026-09-08 12:43:00 +08:00
|
|
|
assert exc.value.status_code == 409
|
|
|
|
|
async with app.state.sessions() as db:
|
2026-09-12 01:24:02 +08:00
|
|
|
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 2
|
2026-09-08 12:43:00 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
async def test_multiple_origins_preserve_research_and_do_not_duplicate_alphas(app, logged_in):
|
|
|
|
|
platform, lane = await setup(app)
|
|
|
|
|
platform.existing_alpha_ids = ["shared"]
|
|
|
|
|
inputs = {
|
|
|
|
|
"name": "多来源研究",
|
|
|
|
|
"candidates": [
|
|
|
|
|
{
|
|
|
|
|
"client_item_id": "one",
|
|
|
|
|
"expression": "rank(close)",
|
|
|
|
|
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
|
|
|
|
|
}
|
|
|
|
|
],
|
|
|
|
|
}
|
|
|
|
|
ids = []
|
|
|
|
|
for kind in ("manual", "template"):
|
|
|
|
|
p = (
|
|
|
|
|
await logged_in.post(
|
|
|
|
|
"/api/v1/backtests/previews",
|
|
|
|
|
json={"inline": {**inputs, "source": {"kind": kind, "research_id": kind}}},
|
|
|
|
|
)
|
|
|
|
|
).json()
|
|
|
|
|
rid = (await start(logged_in, p, kind))["backtest_run_id"]
|
|
|
|
|
ids.append(rid)
|
|
|
|
|
await execute(app, lane, rid)
|
|
|
|
|
async with app.state.sessions.begin() as db:
|
|
|
|
|
research = await db.get(Research, "shared")
|
|
|
|
|
research.note = "保留人工结论"
|
|
|
|
|
await upsert_alpha(db, platform.alphas["shared"])
|
|
|
|
|
origins = (await logged_in.get("/api/v1/alphas/shared/sources", params={"limit": 1})).json()
|
|
|
|
|
assert origins["total"] == 2 and len(origins["items"]) == 1
|
|
|
|
|
assert (await logged_in.get("/api/v1/alphas/shared/sources", params={"limit": 1, "offset": 1})).json()[
|
|
|
|
|
"items"
|
|
|
|
|
][0]["source"]["kind"] == "manual"
|
|
|
|
|
alphas = (await logged_in.get("/api/v1/alphas")).json()
|
|
|
|
|
assert alphas["total"] == 1 and alphas["items"][0]["source_kinds"] == ["manual", "template"]
|
|
|
|
|
assert alphas["items"][0]["research"]["note"] == "保留人工结论"
|
|
|
|
|
assert (
|
|
|
|
|
await logged_in.get("/api/v1/alphas", params={"source": "manual", "research_id": "template"})
|
|
|
|
|
).json()["total"] == 0
|
|
|
|
|
assert (await logged_in.get("/api/v1/alphas", params={"backtest_run_id": ids[0]})).json()["total"] == 1
|
|
|
|
|
exported = await logged_in.get("/api/v1/alphas/export", params={"source": "manual"})
|
|
|
|
|
assert exported.text.count("shared") == 1
|
|
|
|
|
assert (await logged_in.get("/api/v1/alphas/facets")).json()["source"] == ["manual", "template"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def test_source_assignment_drafts_subsets_and_reruns(app, logged_in):
|
|
|
|
|
_, lane = await setup(app)
|
|
|
|
|
await configure(app, logged_in)
|
|
|
|
|
inline = {
|
|
|
|
|
"name": "直接聊天研究",
|
|
|
|
|
"source": {"kind": "forged", "reference": "wrong", "research_id": "wrong"},
|
|
|
|
|
"candidates": [
|
|
|
|
|
{
|
|
|
|
|
"client_item_id": "one",
|
|
|
|
|
"expression": "rank(close)",
|
|
|
|
|
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
|
|
|
|
|
}
|
|
|
|
|
],
|
|
|
|
|
}
|
|
|
|
|
app.state.ai.model_factory = single_tool_factory("prepare_backtest", {"inline": inline})
|
|
|
|
|
conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
|
|
|
|
|
ai = await ask(logged_in, conversation, "准备")
|
|
|
|
|
p = ai["tools"][0]["result"]
|
|
|
|
|
assert (
|
|
|
|
|
p["source"]["kind"] == "chatbox"
|
|
|
|
|
and p["source"]["reference"] == conversation
|
|
|
|
|
and p["source"]["research_id"] == ai["id"]
|
|
|
|
|
)
|
|
|
|
|
rid = (await start(logged_in, p))["backtest_run_id"]
|
|
|
|
|
await execute(app, lane, rid)
|
|
|
|
|
items = (await logged_in.get(f"/api/v1/backtests/runs/{rid}/results")).json()["items"]
|
|
|
|
|
draft = (
|
|
|
|
|
await logged_in.post("/api/v1/backtests/drafts", json={**inline, "source": {"kind": "template"}})
|
|
|
|
|
).json()
|
|
|
|
|
async with app.state.sessions.begin() as db:
|
|
|
|
|
business = Business(db, {"conversation_id": "later-conversation", "ai_run_id": "later-run"})
|
|
|
|
|
referenced = await business.backtests.preview(
|
|
|
|
|
PreviewInput(draft_id=draft["id"], draft_version=draft["version"])
|
|
|
|
|
)
|
|
|
|
|
assert referenced["source"]["kind"] == "template"
|
|
|
|
|
rerun = await business.backtests.rerun(rid, RerunInput(item_ids=[items[0]["id"]]))
|
|
|
|
|
assert rerun["source"] == {**p["source"], "parent_run_id": rid}
|
|
|
|
|
# Add a second candidate, then exclude it via the public fixed-snapshot contract.
|
|
|
|
|
two = {
|
|
|
|
|
**inline,
|
|
|
|
|
"candidates": inline["candidates"] + [{**inline["candidates"][0], "client_item_id": "two"}],
|
|
|
|
|
}
|
|
|
|
|
original = await business.backtests.preview(PreviewInput(inline=two))
|
|
|
|
|
subset = await business.backtests.subset(original["preview_id"], SubsetInput(exclude_ids=["two"]))
|
|
|
|
|
assert subset["source"] == original["source"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def test_new_interfaces_require_login_and_same_origin(client):
|
|
|
|
|
assert (await client.get("/api/v1/alphas/any/sources")).status_code == 401
|
|
|
|
|
assert (await client.get("/api/v1/backtests/sources")).status_code == 401
|
|
|
|
|
assert (
|
|
|
|
|
await client.post("/api/v1/backtests/research-previews", json=construction("none"))
|
|
|
|
|
).status_code == 401
|
|
|
|
|
assert (
|
|
|
|
|
await client.post(
|
|
|
|
|
"/api/v1/backtests/research-previews",
|
|
|
|
|
headers={"Origin": "https://evil.test"},
|
|
|
|
|
json=construction("none"),
|
|
|
|
|
)
|
|
|
|
|
).status_code == 403
|