Merge branch 'main' into feature/issue-579

This commit is contained in:
Mads Hvelplund 2026-06-24 07:25:12 +02:00 • committed by GitHub
commit eda171d91e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 288 additions and 21 deletions

View file

@ -57,6 +57,24 @@ strix --target <target> [options]
Path to a custom config file (JSON) to use instead of `~/.strix/cli-config.json`.
</ParamField>
<ParamField path="--max-budget-usd" type="number">
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.
</ParamField>
## Examples
```bash

View file

@ -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()

View file

@ -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)

View file

@ -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})"
)

View file

@ -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:

View file

@ -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()

View file

@ -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,

View file

@ -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

View file

@ -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,

View file

@ -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)

44
tests/test_execution.py Normal file
View file

@ -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)

108
tests/test_hooks.py Normal file
View file

@ -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)