Files
worldquant-alpha-system/backend/tests/test_research_integration.py
T

289 lines
13 KiB
Python

"""Public research workflow; only the model and WorldQuant HTTP are synthetic."""
import copy
import pytest
from fastapi import HTTPException
from sqlalchemy import func, select
from app.ai.capabilities import ToolContext
from app.ai.tools import CAPABILITIES
from app.alphas import upsert_alpha
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
from app.business import Business
from app.models import BacktestPreview, BacktestRun, Research, TemplateInput
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": "显式字段排序",
"template_input_id": input_id,
"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:
context["template_input_id"] = fixed_input["id"]
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"]
assert source["template_input_id"]
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:
body["template_input_id"] = "missing"
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:
return await CAPABILITIES[name].invoke(ToolContext(Business(db)), args)
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"
selected = await tool(
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
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
with pytest.raises(HTTPException) as exc:
await tool(
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
assert exc.value.status_code == 409
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 2
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