feat: run fixed research pipelines within durable budgets
This commit is contained in:
@@ -0,0 +1,257 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user