diff --git a/strix/core/agents.py b/strix/core/agents.py index 4d3d65cc6..095d5fba8 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -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: diff --git a/strix/core/execution.py b/strix/core/execution.py index d5d94eb57..69a5cf878 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -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) diff --git a/strix/core/hooks.py b/strix/core/hooks.py index 21400c0b4..eae4d7c7f 100644 --- a/strix/core/hooks.py +++ b/strix/core/hooks.py @@ -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: diff --git a/strix/core/runner.py b/strix/core/runner.py index 9420b76ec..7fdf2662f 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -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) diff --git a/strix/tools/agents_graph/tools.py b/strix/tools/agents_graph/tools.py index da05bb1e7..2807b1c92 100644 --- a/strix/tools/agents_graph/tools.py +++ b/strix/tools/agents_graph/tools.py @@ -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" diff --git a/tests/test_budget_pause_policy.py b/tests/test_budget_pause_policy.py new file mode 100644 index 000000000..302b23baf --- /dev/null +++ b/tests/test_budget_pause_policy.py @@ -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]