feat(budget): budget_policy=pause parks every agent at the limit until the operator resumes

Adds budget_policy: stop | pause to run_strix_scan / ReportUsageHooks /
AgentCoordinator, independent of interactive mode. Under pause the agents
get no budget warnings and no sub-agent reserve; each agent parks before
its next LLM call once spent >= limit or the scan is paused, sessions and
sandbox stay alive, and coordinator.resume_budget(max_budget_usd=...)
replaces the limit and wakes every parked agent without adding anything
to any session. coordinator.pause_budget() parks a running scan the same
way. In-flight calls are never cancelled, so spent may end above the
limit. Parked agents count as active for stop_agent.
This commit is contained in:
Ahmed Allam 2026-09-29 23:33:40 +00:00 • committed by Ahmed Allam
parent 0ff9f8c324
commit d355838ea0
6 changed files with 751 additions and 41 deletions

View file

@ -24,6 +24,8 @@ logger = logging.getLogger(__name__)
Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed", "budget_paused"]
BudgetPolicy = Literal["stop", "pause"]
TERMINAL_STATUSES: frozenset[str] = frozenset({"completed", "stopped", "crashed", "failed"})
# Why an agent parked. The user can message any agent, so this - not the agent's
@ -68,7 +70,10 @@ class AgentCoordinator:
self._budget_stopped = False
self._reserve_stopped = False
self._budget_paused = False
self._resume_epoch = 0
self._budget_policy: BudgetPolicy = "stop"
self._extend_budget: Callable[[], None] | None = None
self._set_budget_limit: Callable[[float | None], None] | None = None
def set_snapshot_path(self, path: Path) -> None:
self._snapshot_path = path
@ -95,15 +100,89 @@ class AgentCoordinator:
def budget_paused(self) -> bool:
return self._budget_paused
@property
def resume_epoch(self) -> int:
"""Bumped by every ``resume_budget``; an agent parks against the value it read."""
return self._resume_epoch
@property
def budget_policy(self) -> BudgetPolicy:
return self._budget_policy
def set_budget_policy(self, policy: BudgetPolicy) -> None:
self._budget_policy = policy
def set_budget_extender(self, extend: Callable[[], None]) -> None:
self._extend_budget = extend
def set_budget_limit_setter(self, setter: Callable[[float | None], None]) -> None:
self._set_budget_limit = setter
async def pause_for_budget(self, agent_id: str) -> None:
async with self._lock:
self._budget_paused = True
await self.set_status(agent_id, "budget_paused")
async def park_for_budget(self, agent_id: str) -> None:
"""Record that ``agent_id`` parked before an LLM call (pause policy)."""
await self.set_status(agent_id, "budget_paused")
async def pause_budget(self) -> None:
"""Operator pause: every agent parks before its next LLM call.
Agents mid-call or mid-tool finish that step first, so their spend still
lands; nothing is cancelled.
"""
async with self._lock:
self._budget_paused = True
logger.info("scan paused by the operator")
await self._maybe_snapshot()
async def resume_budget(self, *, max_budget_usd: float | None = None) -> list[str]:
"""Lift the pause and wake every parked agent; returns the woken agent ids.
With ``max_budget_usd`` the scan's limit is replaced first (``None`` keeps
the current one). Agents continue with the LLM call they parked on; no
message is added to any session. An agent that parks again on its next
call (the new limit is already spent) is not an error.
"""
if max_budget_usd is not None and self._set_budget_limit is not None:
self._set_budget_limit(max_budget_usd)
async with self._lock:
self._budget_paused = False
self._resume_epoch += 1
woken = [aid for aid, status in self.statuses.items() if status == "budget_paused"]
for aid in woken:
self.runtimes.setdefault(aid, AgentRuntime()).wake.set()
logger.info("scan resumed; woke %d parked agent(s)", len(woken))
await self._maybe_snapshot()
return woken
async def wait_for_budget_resume(self, agent_id: str, *, parked_epoch: int) -> None:
"""Block until a resume newer than ``parked_epoch``, a scan-wide stop, or
the agent itself being stopped while parked.
``parked_epoch`` is the ``resume_epoch`` the agent read when it decided to
park, so a resume that lands between that decision and this wait is not
missed. ``agent_id`` is ``running`` again on return unless it was stopped.
"""
while True:
async with self._lock:
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
if (
self._budget_stopped
or self._resume_epoch != parked_epoch
or self.statuses.get(agent_id) != "budget_paused"
):
break
wake = runtime.wake
wake.clear()
await wake.wait()
if not self._budget_stopped and self.statuses.get(agent_id) == "budget_paused":
await self.set_status(agent_id, "running")
async def resume_from_budget_pause(self, *, exclude: str | None = None) -> None:
"""Legacy interactive resume: extend by the original budget and nudge agents."""
async with self._lock:
if not self._budget_paused:
return
@ -308,7 +387,7 @@ class AgentCoordinator:
unknown, or it is terminal and its loop does not park for wake-ups.
"""
from_user = message.get("from") == "user"
if from_user and self._budget_paused:
if from_user and self._budget_paused and self._budget_policy != "pause":
await self.resume_from_budget_pause(exclude=target_agent_id)
async with self._lock:
if target_agent_id not in self.statuses:

View file

@ -522,33 +522,51 @@ async def _run_until_lifecycle(
await coordinator.set_status(agent_id, "stopped")
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
if interactive:
result = await _run_cycle_parked(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
event_sink=event_sink,
hooks=hooks,
)
else:
result = await _run_cycle(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=False,
event_sink=event_sink,
hooks=hooks,
)
try:
if interactive:
result = await _run_cycle_parked(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
event_sink=event_sink,
hooks=hooks,
)
else:
result = await _run_cycle(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=False,
event_sink=event_sink,
hooks=hooks,
)
except BudgetPausedError as exc:
if coordinator.budget_policy != "pause":
raise
# The agent parked right before an LLM call; everything up to that
# point is already in its session. Once resumed, the same call goes
# out with nothing added to the conversation.
await coordinator.wait_for_budget_resume(agent_id, parked_epoch=exc.resume_epoch)
if (
not coordinator.budget_stopped
and await _agent_status(coordinator, agent_id) != "running"
):
# Stopped while parked (operator stop or a parent's stop_agent).
await coordinator.reset_recovery(agent_id)
return result
if session is not None:
input_data = []
continue
status = await _agent_status(coordinator, agent_id)
if status != "running":
@ -759,7 +777,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
await coordinator.detach_stream(agent_id, stream)
except BudgetPausedError as exc:
logger.info("agent %s paused at the scan budget limit: %s", agent_id, exc)
await coordinator.pause_for_budget(agent_id)
if coordinator.budget_policy == "pause":
await coordinator.park_for_budget(agent_id)
else:
await coordinator.pause_for_budget(agent_id)
raise
except SubagentBudgetReservedError as exc:
logger.info("sub-agent %s stopped at the budget reserve: %s", agent_id, exc)

View file

@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any
from agents.lifecycle import RunHooks
from strix.core.agents import BudgetPolicy, coordinator_from_context
from strix.report.state import get_global_report_state
@ -22,6 +23,26 @@ logger = logging.getLogger(__name__)
LLM_TURN_KEY = "llm_turn"
# ``BudgetPolicy`` decides what happens when the accumulated LLM cost reaches
# ``max_budget_usd``.
#
# ``stop``: the agents are warned as the limit approaches, sub-agents are cut at
# a reserve so the root can write its report, and the scan ends at the limit.
#
# ``pause``: the agents are never told a limit exists. Every agent parks right
# before its next LLM call once the limit is reached (or an operator pauses the
# scan), keeping its session, context and sandbox alive, and continues with that
# same call when the operator raises the limit or resumes.
__all__ = [
"LLM_TURN_KEY",
"BudgetExceededError",
"BudgetPausedError",
"BudgetPolicy",
"ReportUsageHooks",
"SubagentBudgetReservedError",
"recomputed_budget_flags",
]
_STAGE_LABELS: tuple[str, ...] = ("NOTICE", "URGENT", "CRITICAL")
_TURN_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
_ROOT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
@ -38,7 +59,20 @@ class SubagentBudgetReservedError(RuntimeError):
class BudgetPausedError(RuntimeError):
"""Raised to park one agent when an interactive scan reaches its budget."""
"""Raised to park one agent until the scan budget is raised or the pause lifted.
``resume_epoch`` is the coordinator's ``resume_epoch`` at the moment the agent
decided to park; the agent waits for a resume newer than that.
"""
def __init__(self, message: str, *, resume_epoch: int = 0) -> None:
super().__init__(message)
self.resume_epoch = resume_epoch
def _validate_budget(max_budget_usd: float | None) -> None:
if max_budget_usd is not None and (not math.isfinite(max_budget_usd) or max_budget_usd <= 0):
raise ValueError("max_budget_usd must be a finite number greater than 0")
def recomputed_budget_flags(
@ -46,11 +80,12 @@ def recomputed_budget_flags(
max_budget_usd: float | None,
*,
interactive: bool,
budget_policy: BudgetPolicy = "stop",
) -> tuple[bool, bool]:
"""Return the (budget_stopped, reserve_stopped) flags a resumed scan should carry."""
if max_budget_usd is None:
return False, False
if interactive:
if interactive or budget_policy == "pause":
return False, False
budget_stopped = cost >= max_budget_usd
reserve_stopped = cost >= max_budget_usd * _SUBAGENT_BUDGET_RESERVE
@ -121,18 +156,32 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
max_budget_usd: float | None = None,
max_turns: int | None = None,
interactive: bool = False,
budget_policy: BudgetPolicy = "stop",
) -> None:
if max_budget_usd is not None and (
not math.isfinite(max_budget_usd) or max_budget_usd <= 0
):
raise ValueError("max_budget_usd must be a finite number greater than 0")
_validate_budget(max_budget_usd)
if max_turns is not None and max_turns <= 0:
raise ValueError("max_turns must be a positive integer")
if budget_policy not in ("stop", "pause"):
raise ValueError(f"unknown budget_policy: {budget_policy!r}")
self._model = model
self._max_budget_usd = max_budget_usd
self._budget_increment = max_budget_usd
self._max_turns = max_turns
self._interactive = interactive
self._budget_policy: BudgetPolicy = budget_policy
@property
def max_budget_usd(self) -> float | None:
return self._max_budget_usd
@property
def budget_policy(self) -> BudgetPolicy:
return self._budget_policy
def set_max_budget_usd(self, max_budget_usd: float | None) -> None:
"""Replace the scan's cost limit; ``None`` removes it."""
_validate_budget(max_budget_usd)
self._max_budget_usd = max_budget_usd
def extend_budget(self) -> None:
if self._max_budget_usd is None or self._budget_increment is None:
@ -146,6 +195,8 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
system_prompt: str | None, # noqa: ARG002
input_items: list[TResponseInputItem],
) -> None:
if self._budget_policy == "pause":
self._pause_if_limited(context)
context.context[LLM_TURN_KEY] = int(context.context.get(LLM_TURN_KEY, 0)) + 1
try:
self._maybe_warn_turns(context, input_items)
@ -153,6 +204,32 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
except Exception:
logger.exception("budget/turn warning injection failed")
def _pause_if_limited(self, context: RunContextWrapper[dict[str, Any]]) -> None:
"""Park the agent before a paid call when the scan is at its limit or paused.
Only calls that have already returned are counted, so calls in flight on
other agents still land and are paid for: ``spent`` may end up above the
limit, which is expected and never an error under this policy.
"""
coordinator = coordinator_from_context(context.context)
epoch = coordinator.resume_epoch if coordinator is not None else 0
if coordinator is not None and coordinator.budget_paused:
raise BudgetPausedError(
"scan paused; waiting for the operator to resume", resume_epoch=epoch
)
if self._max_budget_usd is None:
return
report_state = get_global_report_state()
if report_state is None:
return
cost = report_state.get_total_llm_cost()
if cost >= self._max_budget_usd:
raise BudgetPausedError(
f"Scan budget of ${self._max_budget_usd:.2f} reached (spent ${cost:.4f}); "
"pausing until the operator raises the limit",
resume_epoch=epoch,
)
def _maybe_warn_turns(
self,
context: RunContextWrapper[dict[str, Any]],
@ -182,7 +259,7 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
context: RunContextWrapper[dict[str, Any]],
input_items: list[TResponseInputItem],
) -> None:
if self._max_budget_usd is None:
if self._max_budget_usd is None or self._budget_policy == "pause":
return
report_state = get_global_report_state()
if report_state is None:
@ -250,6 +327,11 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
except Exception:
logger.exception("failed to record SDK usage for agent %s", agent_id)
if self._budget_policy == "pause":
# The finished call is paid for and its tool calls still run for free;
# the agent parks before its next call, in ``on_llm_start``.
return
if self._max_budget_usd is not None:
cost = report_state.get_total_llm_cost()
if cost >= self._max_budget_usd:

View file

@ -26,7 +26,7 @@ from strix.config.models import (
uses_chat_completions_tool_schema,
)
from strix.config.settings import DEFAULT_MAX_TURNS
from strix.core.agents import AgentCoordinator
from strix.core.agents import AgentCoordinator, BudgetPolicy
from strix.core.execution import (
respawn_subagents,
run_agent_loop,
@ -189,6 +189,7 @@ async def run_strix_scan(
interactive: bool = False,
max_turns: int = DEFAULT_MAX_TURNS,
max_budget_usd: float | None = None,
budget_policy: BudgetPolicy = "stop",
model: str | None = None,
cleanup_on_exit: bool = True,
event_sink: StreamEventSink | None = None,
@ -208,6 +209,12 @@ async def run_strix_scan(
``extra_system_prompt_context`` is merged into the root agent's scan
context before prompt rendering. Child agents keep the standard scan prompt
and context.
``budget_policy`` decides what happens when the LLM spend reaches
``max_budget_usd``: ``"stop"`` warns the agents as the limit approaches and
ends the scan at it; ``"pause"`` tells the agents nothing and parks every
agent before its next LLM call until the caller resumes the scan through
``coordinator.resume_budget()`` (optionally with a higher limit) or cancels
it. ``coordinator.pause_budget()`` parks a running scan the same way.
``mcp_connection_requests`` supplies the run's MCP connections from any
source: when given, the engine connects those requests; when ``None`` (the
command-line default) it reads ``~/.strix/mcp-servers.json`` itself. Either
@ -256,9 +263,12 @@ async def run_strix_scan(
if not strict_tool_schemas:
logger.info("Sending non-strict tool schemas: %s caps strict tools", resolved_model)
if budget_policy not in ("stop", "pause"):
raise ValueError(f"unknown budget_policy: {budget_policy!r}")
if coordinator is None:
coordinator = AgentCoordinator()
coordinator.set_snapshot_path(agents_path)
coordinator.set_budget_policy(budget_policy)
from strix.tools.coverage.tools import hydrate_coverage_from_disk
from strix.tools.notes.tools import hydrate_notes_from_disk
@ -289,11 +299,17 @@ async def run_strix_scan(
report_state.get_total_llm_cost(),
max_budget_usd,
interactive=interactive,
budget_policy=budget_policy,
)
# Under the pause policy the hooks re-park at the first call if the
# spend is still at the limit, so a restored pause flag would only
# hold agents back after the limit was raised.
await coordinator.reset_budget_stops(
budget_stopped=budget_stopped,
reserve_stopped=reserve_stopped,
budget_paused=interactive and coordinator.budget_paused,
budget_paused=(
interactive and budget_policy != "pause" and coordinator.budget_paused
),
)
for aid, parent in coordinator.parent_of.items():
if parent is None:
@ -372,8 +388,10 @@ async def run_strix_scan(
max_budget_usd=max_budget_usd,
max_turns=max_turns,
interactive=interactive,
budget_policy=budget_policy,
)
if interactive:
coordinator.set_budget_limit_setter(hooks.set_max_budget_usd)
if interactive and budget_policy != "pause":
coordinator.set_budget_extender(hooks.extend_budget)
scope_context = build_scope_context(scan_config)

View file

@ -19,7 +19,7 @@ from strix.report.state import get_global_report_state
from strix.skills import validate_requested_skills
_ACTIVE_STATUSES: frozenset[str] = frozenset({"running", "waiting"})
_ACTIVE_STATUSES: frozenset[str] = frozenset({"running", "waiting", "budget_paused"})
logger = logging.getLogger(__name__)
@ -816,7 +816,7 @@ async def stop_agent(
"success": False,
"error": (
f"Agent {target_agent_id} is already '{current_status}'; "
"stop_agent only acts on running/waiting agents — use "
"stop_agent only acts on running/waiting/paused agents — use "
"view_agent_graph to find still-active descendants and "
"stop them individually, or send_message_to_agent if you "
"want to wake this one with new instructions"

View file

@ -0,0 +1,510 @@
"""``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_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]