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

258 lines
10 KiB
Python

"""Fixed research: durable previews, budget grants, pause/stop and recovery."""
import asyncio
import pytest
from sqlalchemy import func, select
from app.models import Account, BacktestRun
from app.research.runtime import ResearchRuntime
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
from tests.test_backtests import execute, setup
from tests.test_research_workspace import catalog, research_input, template
__all__ = ["catalog", "research_input"]
@pytest.fixture
async def flow_setup(app, logged_in, research_input, monkeypatch):
await configure(app, logged_in)
platform, lane = await setup(app)
calls = []
async def model(ai, context, output_type, revision):
calls.append(context)
value = template()
value["expression"] = f"rank({{field}}) + {len(calls)}"
return TemplateSpec.model_validate(value), {
"model": "fixture",
"revision": revision,
"usage": {"requests": 1},
}
monkeypatch.setattr("app.research.runtime.request_model", model)
body = {
"request_id": "finite-run",
"name": "有限研究",
"input_ids": [research_input["id"]],
"hypothesis": "排名稳定性",
"settings": {"region": "USA", "universe": "TOP3000", "delay": 1},
"budget": {"max_rounds": 2, "max_simulations": 4, "max_model_calls": 3},
"batch_candidates": 2,
}
return body, platform, lane, calls
async def begin(client, body):
result = await client.post("/api/v1/research/flows/runs", json=body)
assert result.status_code == 201, result.text
return result.json()
async def get(client, run_id):
return (await client.get(f"/api/v1/research/flows/runs/{run_id}")).json()
async def drive(app, client, run_id, lane, ticks=45):
for _ in range(ticks):
await app.state.research.advance(run_id)
run = await get(client, run_id)
for step in run["steps"]:
if step["status"] == "waiting" and step["backtest_run_id"]:
await execute(app, lane, step["backtest_run_id"])
if run["status"] not in ("queued", "running"):
return run
raise AssertionError(await get(client, run_id))
async def test_fixed_two_rounds_and_idempotent_start(app, logged_in, flow_setup):
body, platform, lane, calls = flow_setup
first = await begin(logged_in, body)
assert (await begin(logged_in, body))["id"] == first["id"]
assert (
await logged_in.post("/api/v1/research/flows/runs", json={**body, "name": "different"})
).status_code == 409
result = await drive(app, logged_in, first["id"], lane)
assert result["status"] == "completed", result
assert result["round"] == 2 and result["simulations_used"] == 4 and result["model_calls_used"] == 3
assert len(result["steps"]) == 16 and len(calls) == 3 and len(platform.posts) == 2
assert all(s["output"].get("digest") for s in result["steps"] if s["node_id"] == "simulate")
assert len(calls[1]["evaluation"]["records"]) == 2
async def test_preview_commits_before_budget_gate_and_parallel_ticks(app, logged_in, flow_setup):
body, platform, _, _ = flow_setup
body["budget"]["max_simulations"] = 1
run = await begin(logged_in, body)
for _ in range(6):
await asyncio.gather(app.state.research.advance(run["id"]), app.state.research.advance(run["id"]))
result = await get(logged_in, run["id"])
assert result["status"] == "budget_exhausted", result
step = next(s for s in result["steps"] if s["node_id"] == "simulate")
assert step["status"] == "previewed" and step["output"]["preview_id"]
assert result["simulations_used"] == 0 and not platform.posts
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
async def test_pause_stop_keep_known_simulation_and_block_next_steps(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
run = await begin(logged_in, body)
for _ in range(5):
await app.state.research.advance(run["id"])
current = await get(logged_in, run["id"])
simulation = next(s for s in current["steps"] if s["node_id"] == "simulate")
assert simulation["backtest_run_id"]
# Issue remote simulation first. Stop must continue collecting it.
async with app.state.sessions() as db:
from app.models import SimulationAttempt
aid = await db.scalar(
select(SimulationAttempt.id).where(SimulationAttempt.run_id == simulation["backtest_run_id"])
)
await lane.step(aid)
for action in ("pause", "stop"):
current = await get(logged_in, run["id"])
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": action, "version": current["version"]},
)
assert response.status_code == 200, response.text
await app.state.research.advance(run["id"])
await lane.step(aid)
result = (await logged_in.get(f"/api/v1/backtests/runs/{simulation['backtest_run_id']}/results")).json()
assert all(i["persistence_status"] == "saved" for i in result["items"])
final = await get(logged_in, run["id"])
assert final["status"] == "stopped" and len(final["steps"]) == 4 and len(platform.posts) == 1
async def test_recovery_marks_model_interrupted_without_refund(app, logged_in, flow_setup):
body, _, _, calls = flow_setup
run = await begin(logged_in, body)
await app.state.research.advance(run["id"])
work = await app.state.research.prepare(run["id"])
assert work and not calls
recovered = ResearchRuntime(app.state.sessions, app.state.ai, app.state.runner)
await recovered.recover()
current = await get(logged_in, run["id"])
assert current["status"] == "interrupted" and current["model_calls_used"] == 1
await recovered.advance(run["id"])
assert not calls
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": "resume", "version": current["version"]},
)
assert response.status_code == 200
await recovered.advance(run["id"])
current = await get(logged_in, run["id"])
assert current["model_calls_used"] == 2 and len(calls) == 1
assert [a["status"] for a in current["steps"][-1]["output"]["model_attempts"]] == [
"interrupted",
"completed",
]
@pytest.mark.parametrize(
"key,value",
[
("max_rounds", 0),
("max_simulations", -1),
("max_model_calls", 0),
("max_rounds", None),
("max_rounds", True),
("max_model_calls", 1.5),
],
)
async def test_finite_positive_budgets(logged_in, flow_setup, key, value):
body, _, _, _ = flow_setup
body["budget"][key] = value
assert (await logged_in.post("/api/v1/research/flows/runs", json=body)).status_code == 422
async def test_changed_account_stops_authorized_execution(app, logged_in, flow_setup):
body, platform, _, calls = flow_setup
run = await begin(logged_in, body)
async with app.state.sessions.begin() as db:
account = await db.get(Account, 1)
account.wq_user_id = "different"
await app.state.research.advance(run["id"])
assert (await get(logged_in, run["id"]))["status"] == "interrupted"
assert not platform.posts and not calls
async def test_unknown_submission_retains_budget_and_is_never_reposted(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
platform.reject = "unknown"
run = await begin(logged_in, body)
result = await drive(app, logged_in, run["id"], lane)
assert result["status"] == "needs_review" and result["simulations_used"] == 2
assert len(platform.posts) == 1
for _ in range(3):
await app.state.research.advance(run["id"])
await app.state.research.recover()
assert len(platform.posts) == 1
async def test_source_fields_do_not_create_automatic_grant(app, logged_in, research_input):
from app.ai.tools import CAPABILITIES
from app.models import ResearchFlowRun
from tests.test_backtests import candidate, preview
await preview(logged_in, [candidate()])
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(ResearchFlowRun)) == 0
assert CAPABILITIES["start_backtest"].requires_confirmation
assert not any("start_flow" in key for key in CAPABILITIES)
async def test_invalid_final_candidates_cannot_resume_into_completion(
app, logged_in, flow_setup, monkeypatch
):
body, _, lane, _ = flow_setup
body["budget"]["max_rounds"] = 1
run = await begin(logged_in, body)
for _ in range(6):
await app.state.research.advance(run["id"])
current = await get(logged_in, run["id"])
for step in current["steps"]:
if step["status"] == "waiting":
await execute(app, lane, step["backtest_run_id"])
async def invalid(ai, context, output_type, revision):
value = template()
value["expression"] = "unknown_operator({field})"
return TemplateSpec.model_validate(value), {"model": "fixture", "revision": revision}
monkeypatch.setattr("app.research.runtime.request_model", invalid)
current = await drive(app, logged_in, run["id"], lane)
assert current["status"] == "needs_review", current
response = await logged_in.post(
f"/api/v1/research/flows/runs/{run['id']}/control",
json={"action": "resume", "version": current["version"]},
)
assert response.status_code == 409
await app.state.research.advance(run["id"])
assert (await get(logged_in, run["id"]))["status"] == "needs_review"
async def test_restart_reuses_preview_and_known_backtest(app, logged_in, flow_setup):
body, platform, lane, _ = flow_setup
run = await begin(logged_in, body)
for _ in range(4):
await app.state.research.advance(run["id"])
before = await get(logged_in, run["id"])
preview_id = before["steps"][-1]["output"]["preview_id"]
runtime = ResearchRuntime(app.state.sessions, app.state.ai, app.state.runner)
await runtime.recover()
await runtime.advance(run["id"])
current = await get(logged_in, run["id"])
backtest_id = current["steps"][-1]["backtest_run_id"]
assert current["steps"][-1]["output"]["preview_id"] == preview_id
await runtime.recover()
await execute(app, lane, backtest_id)
await runtime.advance(run["id"])
after = await get(logged_in, run["id"])
assert after["simulations_used"] == 2 and len(platform.posts) == 1
assert next(s for s in after["steps"] if s["node_id"] == "simulate")["backtest_run_id"] == backtest_id