strix/tests/test_budget_pause_policy.py
Ahmed Allam 463b149bdb fix(budget): parked agents count as active; park never overwrites a stop
active_agents_except (finish_scan, wait_for_message) treats budget_paused as
active, so a root cannot finish the scan over a parked child. park_for_budget
only transitions a running agent, and the wake back to running happens under
the coordinator lock.
2026-09-30 03:13:40 +03:00

540 lines
20 KiB
Python

"""``budget_policy="pause"``: agents park before a paid call and never hear about budgets."""
from __future__ import annotations
import asyncio
from typing import TYPE_CHECKING, Any
from unittest.mock import MagicMock, patch
import pytest
from strix.core import execution
from strix.core.agents import AgentCoordinator
from strix.core.execution import _start_child_runner, run_agent_loop
from strix.core.hooks import (
BudgetExceededError,
BudgetPausedError,
ReportUsageHooks,
recomputed_budget_flags,
)
from strix.core.sessions import open_agent_session
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Callable
from pathlib import Path
COST_PER_CALL = 1.0
_CALL_LATENCY_S = 0.005
class _FakeLedger:
def __init__(self) -> None:
self.cost = 0.0
self.calls: list[str] = []
self.remaining: dict[str, int] = {}
self.warned_inputs: list[list[Any]] = []
self.gate: asyncio.Event | None = None
self.in_flight = 0
def record_sdk_usage(self, **_kwargs: Any) -> None:
return
def get_total_llm_cost(self) -> float:
return self.cost
class _FakeStream:
"""One ``Runner.run_streamed`` call: several LLM turns, each guarded by the hooks.
Mirrors the SDK's ordering: ``on_llm_start`` runs before the paid request,
``on_llm_end`` after it. A ``BudgetPausedError`` from ``on_llm_start`` ends
the stream without spending, exactly like the SDK surfacing a hook error.
"""
def __init__(
self,
*,
ledger: _FakeLedger,
hooks: ReportUsageHooks,
context: dict[str, Any],
agent: Any,
coordinator: AgentCoordinator,
) -> None:
self._ledger = ledger
self._hooks = hooks
self._context = context
self._agent = agent
self._coordinator = coordinator
self.run_loop_exception: BaseException | None = None
self.final_output = None
async def stream_events(self) -> AsyncIterator[Any]:
agent_id = str(self._context.get("agent_id"))
ctx_wrapper = MagicMock()
ctx_wrapper.context = self._context
while self._ledger.remaining.get(agent_id, 0) > 0:
input_items: list[Any] = []
try:
await self._hooks.on_llm_start(ctx_wrapper, self._agent, None, input_items)
except BudgetPausedError as exc:
self.run_loop_exception = exc
return
self._ledger.warned_inputs.append(input_items)
if self._ledger.gate is not None:
self._ledger.in_flight += 1
await self._ledger.gate.wait()
self._ledger.in_flight -= 1
self._ledger.cost += COST_PER_CALL
self._ledger.calls.append(agent_id)
self._ledger.remaining[agent_id] -= 1
await self._hooks.on_llm_end(ctx_wrapper, self._agent, MagicMock())
await asyncio.sleep(_CALL_LATENCY_S)
if self._coordinator.statuses.get(agent_id) == "running":
await self._coordinator.set_status(agent_id, "completed")
items: tuple[Any, ...] = ()
for item in items:
yield item
def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002
return
def _fake_runner(ledger: _FakeLedger, coordinator: AgentCoordinator) -> Any:
class _FakeRunner:
@staticmethod
def run_streamed(
agent: Any,
input: Any, # noqa: A002, ARG004
*,
run_config: Any, # noqa: ARG004
context: dict[str, Any],
max_turns: int, # noqa: ARG004
session: Any, # noqa: ARG004
hooks: ReportUsageHooks,
) -> _FakeStream:
return _FakeStream(
ledger=ledger,
hooks=hooks,
context=context,
agent=agent,
coordinator=coordinator,
)
return _FakeRunner
async def _noop_compact(*_args: Any, **_kwargs: Any) -> bool:
return False
async def _wait_until(predicate: Callable[[], bool], *, timeout: float = 5.0) -> None:
async def _poll() -> None:
while not predicate():
await asyncio.sleep(0.001)
await asyncio.wait_for(_poll(), timeout=timeout)
def _all_parked(coordinator: AgentCoordinator, *agent_ids: str) -> bool:
return all(coordinator.statuses.get(aid) == "budget_paused" for aid in agent_ids)
class _Scan:
"""Root + children driven through the real non-interactive loops."""
def __init__(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
*,
max_budget_usd: float | None,
) -> None:
self.ledger = _FakeLedger()
self.hooks = ReportUsageHooks(
model="test-model", max_budget_usd=max_budget_usd, budget_policy="pause"
)
self.coordinator = AgentCoordinator()
self.coordinator.set_budget_policy("pause")
self.coordinator.set_budget_limit_setter(self.hooks.set_max_budget_usd)
monkeypatch.setattr(execution, "Runner", _fake_runner(self.ledger, self.coordinator))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
self.db_path = tmp_path / "agents.sqlite"
self.sessions: list[Any] = []
self.run_config = MagicMock()
self.root_ctx: dict[str, Any] = {
"agent_id": "root",
"parent_id": None,
"coordinator": self.coordinator,
}
self.root_task: asyncio.Task[Any] | None = None
async def start_root(self, *, calls: int) -> None:
self.ledger.remaining["root"] = calls
await self.coordinator.register("root", "strix", parent_id=None)
session = open_agent_session("root", self.db_path)
self.sessions.append(session)
self.root_task = asyncio.create_task(
run_agent_loop(
agent=MagicMock(),
initial_input=[],
run_config=self.run_config,
context=self.root_ctx,
max_turns=500,
coordinator=self.coordinator,
agent_id="root",
interactive=False,
session=session,
hooks=self.hooks,
)
)
async def start_child(self, child_id: str, *, calls: int) -> None:
self.ledger.remaining[child_id] = calls
await self.coordinator.register(child_id, "recon", parent_id="root")
await _start_child_runner(
parent_ctx=self.root_ctx,
coordinator=self.coordinator,
agents_db_path=self.db_path,
sessions_to_close=self.sessions,
run_config=self.run_config,
max_turns=500,
interactive=False,
child_agent=MagicMock(),
child_id=child_id,
name=f"recon-{child_id}",
parent_id="root",
task="probe things",
initial_input=[],
hooks=self.hooks,
)
def tasks(self) -> list[asyncio.Task[Any]]:
tasks = [self.root_task] if self.root_task is not None else []
tasks.extend(rt.task for rt in self.coordinator.runtimes.values() if rt.task is not None)
return tasks
async def teardown(self) -> None:
for task in self.tasks():
task.cancel()
await asyncio.gather(*self.tasks(), return_exceptions=True)
for session in self.sessions:
session.close()
@pytest.mark.asyncio
async def test_pause_policy_parks_every_agent_at_the_limit_and_resumes_in_place(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=5.0)
agents = ("root", "child-a", "child-b")
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
await scan.start_root(calls=100)
await scan.start_child("child-a", calls=100)
await scan.start_child("child-b", calls=100)
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
assert scan.ledger.cost == pytest.approx(5.0)
assert len(scan.ledger.calls) == 5
assert scan.coordinator.budget_paused is False
assert scan.coordinator.budget_stopped is False
assert scan.coordinator.reserve_stopped is False
assert all(not task.done() for task in scan.tasks())
woken = await scan.coordinator.resume_budget(max_budget_usd=8.0)
assert sorted(woken) == sorted(agents)
assert scan.hooks.max_budget_usd == 8.0
await _wait_until(lambda: scan.ledger.cost >= 8.0)
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
assert scan.ledger.cost == pytest.approx(8.0)
await scan.coordinator.resume_budget(max_budget_usd=9.0)
await _wait_until(lambda: scan.ledger.cost >= 9.0)
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
assert scan.ledger.cost == pytest.approx(9.0)
assert len(scan.ledger.calls) == 9
assert all(not task.done() for task in scan.tasks())
# The model never saw a budget message: no warning band, no resume note.
assert scan.ledger.warned_inputs
assert all(items == [] for items in scan.ledger.warned_inputs)
for session in scan.sessions:
assert await session.get_items() == []
# Stop while parked: an individual stop wakes that loop and it exits.
child_a_task = scan.coordinator.runtimes["child-a"].task
assert child_a_task is not None
await scan.coordinator.request_stop("child-a")
await asyncio.wait_for(child_a_task, timeout=5.0)
assert scan.coordinator.statuses["child-a"] == "stopped"
assert scan.coordinator.statuses["root"] == "budget_paused"
assert scan.coordinator.statuses["child-b"] == "budget_paused"
# A scan-wide cancel while parked tears the rest down cleanly.
await scan.teardown()
assert all(task.done() for task in scan.tasks())
assert scan.ledger.cost == pytest.approx(9.0)
@pytest.mark.asyncio
async def test_pause_policy_operator_pause_and_resume_without_a_limit(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=None)
agents = ("root", "child-a", "child-b")
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
await scan.start_root(calls=6)
await scan.start_child("child-a", calls=6)
await scan.start_child("child-b", calls=6)
await _wait_until(lambda: scan.ledger.cost >= 3.0)
await scan.coordinator.pause_budget()
spent_at_pause = scan.ledger.cost
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
await _wait_until(lambda: scan.coordinator.budget_paused, timeout=0.1)
assert scan.ledger.cost == pytest.approx(spent_at_pause)
assert all(not task.done() for task in scan.tasks())
await asyncio.sleep(0.05)
assert scan.ledger.cost == pytest.approx(spent_at_pause)
woken = await scan.coordinator.resume_budget()
assert sorted(woken) == sorted(agents)
await _wait_until(lambda: not scan.coordinator.budget_paused, timeout=0.1)
await asyncio.wait_for(asyncio.gather(*scan.tasks(), return_exceptions=True), timeout=5.0)
assert scan.ledger.cost == pytest.approx(18.0)
assert {aid: str(s) for aid, s in scan.coordinator.statuses.items()} == {
"root": "completed",
"child-a": "completed",
"child-b": "completed",
}
assert all(items == [] for items in scan.ledger.warned_inputs)
for session in scan.sessions:
session.close()
@pytest.mark.asyncio
async def test_pause_policy_keeps_in_flight_overshoot(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=1.0)
scan.ledger.gate = asyncio.Event()
await scan.coordinator.register("root", "strix", parent_id=None)
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
await scan.start_child("child-a", calls=100)
await scan.start_child("child-b", calls=100)
# Both calls were dispatched under the limit; neither is cancelled.
await _wait_until(lambda: scan.ledger.in_flight == 2)
assert scan.ledger.cost == 0.0
scan.ledger.gate.set()
await _wait_until(lambda: _all_parked(scan.coordinator, "child-a", "child-b"))
assert scan.ledger.cost == pytest.approx(2.0)
assert scan.hooks.max_budget_usd is not None
assert scan.ledger.cost > scan.hooks.max_budget_usd
assert scan.coordinator.budget_stopped is False
assert all(not task.done() for task in scan.tasks())
# Resuming at a limit that is already spent parks again without a call.
await scan.coordinator.resume_budget(max_budget_usd=1.5)
await asyncio.sleep(0.05)
await _wait_until(lambda: _all_parked(scan.coordinator, "child-a", "child-b"))
assert scan.ledger.cost == pytest.approx(2.0)
await scan.teardown()
@pytest.mark.asyncio
async def test_resume_between_park_decision_and_wait_is_not_missed() -> None:
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
await coordinator.register("a", "strix", parent_id=None)
parked_epoch = coordinator.resume_epoch
await coordinator.park_for_budget("a")
assert _all_parked(coordinator, "a")
await coordinator.resume_budget()
await asyncio.wait_for(
coordinator.wait_for_budget_resume("a", parked_epoch=parked_epoch), timeout=1.0
)
assert coordinator.statuses["a"] == "running"
@pytest.mark.asyncio
async def test_parked_wait_returns_on_stop_signals() -> None:
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
await coordinator.register("a", "strix", parent_id=None)
await coordinator.register("b", "recon", parent_id="a")
await coordinator.park_for_budget("a")
await coordinator.park_for_budget("b")
wait_a = asyncio.create_task(
coordinator.wait_for_budget_resume("a", parked_epoch=coordinator.resume_epoch)
)
wait_b = asyncio.create_task(
coordinator.wait_for_budget_resume("b", parked_epoch=coordinator.resume_epoch)
)
await asyncio.sleep(0.02)
assert not wait_a.done()
assert not wait_b.done()
await coordinator.request_stop("b")
await asyncio.wait_for(wait_b, timeout=1.0)
assert coordinator.statuses["b"] == "stopped"
assert not wait_a.done()
await coordinator.trigger_budget_stop()
await asyncio.wait_for(wait_a, timeout=1.0)
assert coordinator.statuses["a"] == "budget_paused"
@pytest.mark.asyncio
async def test_park_never_overwrites_a_stop_that_landed_first() -> None:
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
await coordinator.register("a", "strix", parent_id=None)
parked_epoch = coordinator.resume_epoch
await coordinator.request_stop("a")
assert await coordinator.park_for_budget("a") is False
assert coordinator.statuses["a"] == "stopped"
await asyncio.wait_for(
coordinator.wait_for_budget_resume("a", parked_epoch=parked_epoch), timeout=1.0
)
assert coordinator.statuses["a"] == "stopped"
@pytest.mark.asyncio
async def test_parked_children_keep_the_scan_open() -> None:
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
assert await coordinator.park_for_budget("child") is True
active = await coordinator.active_agents_except("root")
assert [a["agent_id"] for a in active] == ["child"]
assert active[0]["status"] == "budget_paused"
@pytest.mark.asyncio
async def test_resume_budget_replaces_the_limit_and_validates_it() -> None:
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
coordinator.set_budget_limit_setter(hooks.set_max_budget_usd)
await coordinator.resume_budget(max_budget_usd=25.0)
assert hooks.max_budget_usd == 25.0
await coordinator.resume_budget()
assert hooks.max_budget_usd == 25.0
with pytest.raises(ValueError, match="greater than 0"):
await coordinator.resume_budget(max_budget_usd=0.0)
with pytest.raises(ValueError, match="finite"):
await coordinator.resume_budget(max_budget_usd=float("inf"))
assert hooks.max_budget_usd == 25.0
def _ctx(coordinator: AgentCoordinator | None, *, parent_id: str | None = None) -> MagicMock:
wrapper = MagicMock()
wrapper.context = {"agent_id": "x", "parent_id": parent_id}
if coordinator is not None:
wrapper.context["coordinator"] = coordinator
return wrapper
@pytest.mark.asyncio
async def test_pause_hooks_never_warn_and_park_only_at_the_limit() -> None:
ledger = _FakeLedger()
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
coordinator = AgentCoordinator()
coordinator.set_budget_policy("pause")
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
for cost in (7.0, 8.5, 9.5, 9.99):
ledger.cost = cost
for parent_id in (None, "root"):
items: list[Any] = []
await hooks.on_llm_start(
_ctx(coordinator, parent_id=parent_id), MagicMock(), None, items
)
assert items == []
await hooks.on_llm_end(
_ctx(coordinator, parent_id=parent_id), MagicMock(), MagicMock()
)
ledger.cost = 10.0
await hooks.on_llm_end(_ctx(coordinator, parent_id="root"), MagicMock(), MagicMock())
with pytest.raises(BudgetPausedError) as at_limit:
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, [])
assert at_limit.value.resume_epoch == coordinator.resume_epoch
ledger.cost = 13.7
with pytest.raises(BudgetPausedError):
await hooks.on_llm_start(_ctx(coordinator, parent_id="root"), MagicMock(), None, [])
hooks.set_max_budget_usd(20.0)
items = []
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, items)
assert items == []
await coordinator.pause_budget()
with pytest.raises(BudgetPausedError, match="paused"):
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, [])
@pytest.mark.asyncio
async def test_pause_hooks_do_not_count_a_parked_turn() -> None:
ledger = _FakeLedger()
ledger.cost = 10.0
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
coordinator = AgentCoordinator()
ctx = _ctx(coordinator)
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
with pytest.raises(BudgetPausedError):
await hooks.on_llm_start(ctx, MagicMock(), None, [])
assert "llm_turn" not in ctx.context
hooks.set_max_budget_usd(11.0)
await hooks.on_llm_start(ctx, MagicMock(), None, [])
assert ctx.context["llm_turn"] == 1
@pytest.mark.asyncio
async def test_stop_policy_is_unchanged() -> None:
ledger = _FakeLedger()
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0)
assert hooks.budget_policy == "stop"
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
ledger.cost = 7.0
items: list[Any] = []
await hooks.on_llm_start(_ctx(None), MagicMock(), None, items)
assert len(items) == 1
assert "Scan cost budget" in str(items[0])
ledger.cost = 10.0
with pytest.raises(BudgetExceededError):
await hooks.on_llm_end(_ctx(None), MagicMock(), MagicMock())
assert recomputed_budget_flags(10.0, 10.0, interactive=False, budget_policy="stop") == (
True,
True,
)
assert recomputed_budget_flags(10.0, 10.0, interactive=False, budget_policy="pause") == (
False,
False,
)
def test_budget_policy_is_validated() -> None:
with pytest.raises(ValueError, match="budget_policy"):
ReportUsageHooks(model="m", budget_policy="later") # type: ignore[arg-type]