diff --git a/docs/usage/cli.mdx b/docs/usage/cli.mdx index ab2cc47d..58d3e99d 100644 --- a/docs/usage/cli.mdx +++ b/docs/usage/cli.mdx @@ -57,6 +57,24 @@ strix --target [options] Path to a custom config file (JSON) to use instead of `~/.strix/cli-config.json`. + + Maximum LLM spend in USD for the whole scan, counted cumulatively across the + root agent and every child agent. The budget is checked after each model + response; once the running cost reaches the threshold, the scan stops cleanly + with a `stopped` status (not a failure) and the sandbox is torn down. + + Must be greater than `0`. Omit the flag for no limit. + + **Limitations** + + - The check fires *after* a response is returned, so the final spend can + slightly overshoot the limit by any calls already in flight when the + threshold is crossed (most relevant with several child agents running + concurrently). + - Cost is a best-effort estimate derived from token usage and model pricing; + providers that do not expose priced usage may under-count. + + ## Examples ```bash diff --git a/strix/core/agents.py b/strix/core/agents.py index 298aa2be..24a7185f 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -43,6 +43,7 @@ class AgentCoordinator: self._lock = asyncio.Lock() self._snapshot_path: Path | None = None self.is_shutting_down = False + self._budget_stopped = False def set_snapshot_path(self, path: Path) -> None: self._snapshot_path = path @@ -50,6 +51,17 @@ class AgentCoordinator: def mark_shutting_down(self) -> None: self.is_shutting_down = True + @property + def budget_stopped(self) -> bool: + return self._budget_stopped + + async def trigger_budget_stop(self) -> None: + """Signal a scan-wide budget stop and wake every parked agent so it exits.""" + async with self._lock: + self._budget_stopped = True + for runtime in self.runtimes.values(): + runtime.wake.set() + async def register( self, agent_id: str, @@ -143,7 +155,7 @@ class AgentCoordinator: async def wait_for_message(self, agent_id: str) -> None: while True: async with self._lock: - if self.pending_counts.get(agent_id, 0) > 0: + if self._budget_stopped or self.pending_counts.get(agent_id, 0) > 0: return wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake wake.clear() diff --git a/strix/core/execution.py b/strix/core/execution.py index b3d03f78..0e4aabd4 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -16,6 +16,7 @@ from docker import errors as docker_errors # type: ignore[import-untyped, unuse from litellm.exceptions import ContextWindowExceededError from openai import APIError +from strix.core.hooks import BudgetExceededError from strix.core.inputs import child_initial_input from strix.core.sessions import open_agent_session, strip_all_images_from_session @@ -98,6 +99,10 @@ async def run_agent_loop( except asyncio.CancelledError: return result + if coordinator.budget_stopped: + await coordinator.set_status(agent_id, "stopped") + raise BudgetExceededError("scan budget reached") + await coordinator.consume_pending(agent_id) result = await _run_cycle( agent, @@ -279,6 +284,10 @@ async def _run_noninteractive_until_lifecycle( invalid_final_output_limit = max(1, max_turns) while True: + if coordinator.budget_stopped: + await coordinator.set_status(agent_id, "stopped") + raise BudgetExceededError("scan budget reached") + result = await _run_cycle( agent, coordinator, @@ -361,6 +370,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 logger.exception("stream event sink failed for %s", agent_id) if stream.run_loop_exception is not None: raise stream.run_loop_exception + except BudgetExceededError: + # A RuntimeError subclass: re-raise explicitly so it is never + # mistaken for the LiteLLM "after shutdown" race below. + raise except RuntimeError as stream_exc: if "after shutdown" not in str(stream_exc): raise @@ -378,6 +391,13 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 ) finally: await coordinator.detach_stream(agent_id, stream) + except BudgetExceededError as exc: + logger.info( + "agent %s reached the scan budget limit; stopping the scan: %s", agent_id, exc + ) + await coordinator.set_status(agent_id, "stopped") + await coordinator.trigger_budget_stop() + raise except Exception as exc: # ContextWindowExceededError carries status_code=400, which would otherwise # match _INPUT_REJECTION_CODES and trigger image-strip recovery — a path @@ -541,21 +561,29 @@ async def _start_child_runner( child_ctx["parent_id"] = parent_id child_ctx["task"] = task - task_handle = asyncio.create_task( - run_agent_loop( - agent=child_agent, - initial_input=initial_input, - run_config=run_config, - context=child_ctx, - max_turns=max_turns, - coordinator=coordinator, - agent_id=child_id, - interactive=interactive, - session=session, - start_parked=start_parked, - event_sink=event_sink, - hooks=hooks, - ), - name=f"agent-{name}-{child_id}", - ) + async def _child_loop() -> None: + # A budget stop is a clean scan-wide shutdown, not a child failure: the + # child's status and parent notification are already settled in + # ``_run_cycle``. Swallow it here so the detached task does not surface a + # spurious "Task exception was never retrieved" warning. The root agent + # hits the same limit on its next call and tears the scan down. + try: + await run_agent_loop( + agent=child_agent, + initial_input=initial_input, + run_config=run_config, + context=child_ctx, + max_turns=max_turns, + coordinator=coordinator, + agent_id=child_id, + interactive=interactive, + session=session, + start_parked=start_parked, + event_sink=event_sink, + hooks=hooks, + ) + except BudgetExceededError: + logger.info("child %s stopped after reaching the scan budget limit", child_id) + + task_handle = asyncio.create_task(_child_loop(), name=f"agent-{name}-{child_id}") await coordinator.attach_runtime(child_id, task=task_handle) diff --git a/strix/core/hooks.py b/strix/core/hooks.py index f7f888fa..6b0d5924 100644 --- a/strix/core/hooks.py +++ b/strix/core/hooks.py @@ -19,11 +19,19 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +class BudgetExceededError(RuntimeError): + """Raised when the accumulated LLM cost reaches the configured budget.""" + + class ReportUsageHooks(RunHooks[dict[str, Any]]): """Persist SDK-native usage after every model response.""" - def __init__(self, *, model: str) -> None: + def __init__(self, *, model: str, max_budget_usd: float | None = None) -> None: + import math + 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") self._model = model + self._max_budget_usd = max_budget_usd async def on_llm_end( self, @@ -52,3 +60,10 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]): ) except Exception: logger.exception("failed to record SDK usage for agent %s", agent_id) + + if self._max_budget_usd is not None: + cost = report_state.get_total_llm_cost() + if cost >= self._max_budget_usd: + raise BudgetExceededError( + f"Token budget of ${self._max_budget_usd:.2f} exceeded (spent ${cost:.4f})" + ) diff --git a/strix/core/runner.py b/strix/core/runner.py index 9ab120bc..82e653d9 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -27,7 +27,7 @@ from strix.core.execution import ( from strix.core.execution import ( spawn_child_agent as start_child_agent, ) -from strix.core.hooks import ReportUsageHooks +from strix.core.hooks import BudgetExceededError, ReportUsageHooks from strix.core.inputs import ( DEFAULT_MAX_TURNS, build_root_task, @@ -59,6 +59,7 @@ async def run_strix_scan( coordinator: AgentCoordinator | None = None, interactive: bool = False, max_turns: int = DEFAULT_MAX_TURNS, + max_budget_usd: float | None = None, model: str | None = None, cleanup_on_exit: bool = True, event_sink: StreamEventSink | None = None, @@ -164,7 +165,7 @@ async def run_strix_scan( sandbox=SandboxRunConfig(client=bundle["client"], session=bundle["session"]), trace_include_sensitive_data=False, ) - hooks = ReportUsageHooks(model=resolved_model) + hooks = ReportUsageHooks(model=resolved_model, max_budget_usd=max_budget_usd) scope_context = build_scope_context(scan_config) @@ -300,6 +301,13 @@ async def run_strix_scan( str(final)[:300], ) return result # noqa: TRY300 + except BudgetExceededError as exc: + logger.info("Scan %s stopped: %s", scan_id, exc) + if root_id is not None: + await coordinator.cancel_descendants(root_id) + with contextlib.suppress(Exception): + await coordinator.set_status(root_id, "stopped") + return None except BaseException: logger.exception("Strix scan %s failed", scan_id) if root_id is not None: diff --git a/strix/interface/cli.py b/strix/interface/cli.py index 1f7cbce0..f5079120 100644 --- a/strix/interface/cli.py +++ b/strix/interface/cli.py @@ -183,6 +183,7 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915 image=_resolve_sandbox_image(), local_sources=getattr(args, "local_sources", None) or [], interactive=bool(getattr(args, "interactive", False)), + max_budget_usd=getattr(args, "max_budget_usd", None), ) finally: stop_updates.set() diff --git a/strix/interface/main.py b/strix/interface/main.py index a928ee30..4eae0527 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -304,6 +304,17 @@ def get_version() -> str: return "unknown" +def _positive_budget(value: str) -> float: + try: + budget = float(value) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"invalid float value: {value!r}") from exc + import math + if not math.isfinite(budget) or budget <= 0: + raise argparse.ArgumentTypeError("must be a finite number greater than 0") + return budget + + def parse_arguments() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Strix Multi-Agent Cybersecurity Penetration Testing Tool", @@ -439,6 +450,13 @@ Examples: help="Path to a custom config file (JSON) to use instead of ~/.strix/cli-config.json", ) + parser.add_argument( + "--max-budget-usd", + type=_positive_budget, + default=None, + help="Maximum LLM cost in USD (> 0). The scan stops cleanly when this limit is reached.", + ) + parser.add_argument( "--resume", type=str, diff --git a/strix/interface/tui/app.py b/strix/interface/tui/app.py index 58b20925..76dded4b 100644 --- a/strix/interface/tui/app.py +++ b/strix/interface/tui/app.py @@ -31,6 +31,7 @@ from textual.widgets import Button, Label, Static, TextArea, Tree from textual.widgets.tree import TreeNode from strix.config import load_settings +from strix.core.hooks import BudgetExceededError from strix.core.runner import run_strix_scan from strix.interface.tui.live_view import TuiLiveView from strix.interface.tui.messages import send_user_message_to_agent @@ -1369,12 +1370,18 @@ class StrixTUIApp(App): # type: ignore[misc] local_sources=getattr(self.args, "local_sources", None) or [], coordinator=self.coordinator, interactive=True, + max_budget_usd=getattr(self.args, "max_budget_usd", None), event_sink=self._capture_sdk_event, ), ) except (KeyboardInterrupt, asyncio.CancelledError): logger.info("Scan interrupted by user") + except BudgetExceededError: + # Defensive: the runner stops the scan cleanly on budget and + # returns, so this normally never propagates. Treat it as a + # graceful stop, not a scan error, if it ever does. + logger.info("Scan stopped: --max-budget-usd limit reached") except (ConnectionError, TimeoutError) as e: logging.exception("Network error during scan") self._scan_error = e diff --git a/strix/report/state.py b/strix/report/state.py index 1ee72c51..6f626c13 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -236,6 +236,10 @@ class ReportState: def get_total_llm_usage(self) -> dict[str, Any]: return dict(self.run_record.get("llm_usage") or self._build_llm_usage_record()) + def get_total_llm_cost(self) -> float: + """Live accumulated LLM cost, independent of the persisted run-record snapshot.""" + return self._llm_usage.total_cost + def update_scan_final_fields( self, executive_summary: str, diff --git a/strix/report/usage.py b/strix/report/usage.py index 58977c66..b4f2b786 100644 --- a/strix/report/usage.py +++ b/strix/report/usage.py @@ -52,6 +52,10 @@ class LLMUsageLedger: if isinstance(cost, int | float) and cost > 0: self._total_cost += float(cost) + @property + def total_cost(self) -> float: + return _round_cost(self._total_cost) + def to_record(self) -> dict[str, Any]: record = serialize_usage(self._total_usage) record["cost"] = _round_cost(self._total_cost) diff --git a/tests/test_execution.py b/tests/test_execution.py new file mode 100644 index 00000000..59a37e65 --- /dev/null +++ b/tests/test_execution.py @@ -0,0 +1,44 @@ +"""Tests for the scan-wide budget-stop signal on the agent coordinator.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from strix.core.agents import AgentCoordinator + + +@pytest.mark.asyncio +async def test_budget_stop_sets_flag() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + assert coordinator.budget_stopped is False + await coordinator.trigger_budget_stop() + assert coordinator.budget_stopped is True + + +@pytest.mark.asyncio +async def test_budget_stop_unblocks_parked_agent() -> None: + # A parent parked in wait_for_message (awaiting a child) must be released so + # it can exit, no matter where in the tree the budget limit was hit. + coordinator = AgentCoordinator() + await coordinator.register("parent", "strix", parent_id=None) + + waiter = asyncio.create_task(coordinator.wait_for_message("parent")) + await asyncio.sleep(0) # let the waiter park + assert not waiter.done() + + await coordinator.trigger_budget_stop() + await asyncio.wait_for(waiter, timeout=1.0) + + +@pytest.mark.asyncio +async def test_wait_for_message_returns_immediately_after_budget_stop() -> None: + coordinator = AgentCoordinator() + await coordinator.register("agent", "recon", parent_id="parent") + await coordinator.trigger_budget_stop() + + # No pending messages, but the stop flag short-circuits the wait. + await asyncio.wait_for(coordinator.wait_for_message("agent"), timeout=1.0) diff --git a/tests/test_hooks.py b/tests/test_hooks.py new file mode 100644 index 00000000..5dcbed4f --- /dev/null +++ b/tests/test_hooks.py @@ -0,0 +1,108 @@ +"""Tests for budget enforcement in ReportUsageHooks.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from strix.core.hooks import BudgetExceededError, ReportUsageHooks + + +def _make_hooks(max_budget: float | None) -> ReportUsageHooks: + return ReportUsageHooks(model="test-model", max_budget_usd=max_budget) + + +def _make_report_state(cost: float) -> MagicMock: + state = MagicMock() + state.get_total_llm_cost.return_value = cost + state.record_sdk_usage = MagicMock() + return state + + +def _make_context(agent_id: str = "test-agent") -> MagicMock: + ctx: MagicMock = MagicMock() + ctx.context = {"agent_id": agent_id} + return ctx + + +@pytest.mark.asyncio +async def test_no_budget_never_raises() -> None: + hooks = _make_hooks(None) + state = _make_report_state(9999.0) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_under_budget_does_not_raise() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(9.99) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_at_budget_raises() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(10.0) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetExceededError), + ): + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_over_budget_raises() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(10.01) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetExceededError), + ): + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_budget_check_uses_live_cost_accessor() -> None: + # The check must read the live ledger, not the persisted run-record snapshot, + # so it stays accurate even when a save fails after a usage record. + hooks = _make_hooks(5.0) + state = _make_report_state(6.0) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetExceededError), + ): + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + state.get_total_llm_cost.assert_called_once() + state.get_total_llm_usage.assert_not_called() + + +@pytest.mark.asyncio +async def test_error_message_includes_amounts() -> None: + hooks = _make_hooks(5.0) + state = _make_report_state(7.1234) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + with pytest.raises(BudgetExceededError, match=r"\$5\.00") as exc_info: + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + assert "7.1234" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_no_raise_when_report_state_none() -> None: + hooks = _make_hooks(1.0) + with patch("strix.core.hooks.get_global_report_state", return_value=None): + # Should return early without raising, even with budget set + await hooks.on_llm_end(_make_context(), MagicMock(), MagicMock()) + + +@pytest.mark.parametrize("bad_budget", [0.0, -0.01, -5.0]) +def test_non_positive_budget_rejected(bad_budget: float) -> None: + with pytest.raises(ValueError, match="greater than 0"): + ReportUsageHooks(model="test-model", max_budget_usd=bad_budget) + + +def test_budget_exceeded_error_is_runtime_error() -> None: + err = BudgetExceededError("test") + assert isinstance(err, RuntimeError)