From 962d4459d936a136e403500021c1198f0381d7a9 Mon Sep 17 00:00:00 2001 From: Mads Hvelplund Date: Mon, 22 Jun 2026 17:17:08 +0200 Subject: [PATCH 1/2] Add configurable token / cost usage limits (#576) * fix: resolve pre-commit check failures - Change RuntimeError to TypeError for type validation in report/writer.py - Update pyupgrade to v3.21.2 for Python 3.14 compatibility * feat(cli): add --max-budget-usd flag Raises BudgetExceededError in ReportUsageHooks after each LLM call when accumulated cost reaches the limit, with clean "stopped" status and child-agent cancellation in non-interactive mode. * test: add budget enforcement unit tests 7 tests covering no-budget, under-budget, at-limit, over-limit, error message content, None report state, and exception hierarchy. Also adds pytest/pytest-asyncio to dev deps and a mypy override for tests. * fix(budget): validate positive budget and check the live cost ledger Two hardening fixes for --max-budget-usd enforcement: - Reject non-positive budgets. ReportUsageHooks now raises ValueError for max_budget_usd <= 0, and the CLI validates the flag via a custom argparse type so '--max-budget-usd 0' fails fast with a friendly message instead of silently killing the scan on the first model response. - Read the live cost. The budget check now reads ReportState.get_total_llm_cost() (the live ledger) instead of the persisted run-record snapshot, so it stays accurate even when a usage save fails after a model call. * fix(budget): stop the entire scan deterministically when the limit is hit Previously a BudgetExceededError was handled per-agent: it was swallowed in interactive mode (the loop kept waiting), a child's error escaped its detached task as an unretrieved-exception warning, the parent was never released from wait_for_message, and the stop was logged at ERROR with a traceback as if the agent had failed. Replace that with a single scan-wide signal on the coordinator: - AgentCoordinator.trigger_budget_stop() sets a flag and wakes every parked agent; wait_for_message returns as soon as the flag is set. - The run loops check coordinator.budget_stopped and raise to exit cleanly, marking themselves 'stopped'. The root's exception reaches run_strix_scan's handler, which cancels descendants and tears the scan down once; child exceptions are swallowed in their detached task. - The budget stop is logged at INFO, not as a failure. This is deterministic regardless of tree depth or which agent first sees the limit, fixing the interactive/TUI hang where a deep agent's stop never reached a parked root. Also re-raises BudgetExceededError explicitly in the stream handler so it can't be mistaken for the LiteLLM 'after shutdown' race. * fix(budget): treat a budget stop as a clean stop in the TUI Add an explicit BudgetExceededError handler in the TUI scan thread so that, if the error ever reaches it, the budget stop is logged as a graceful stop rather than surfaced as a red scan error by the broad 'except Exception'. The runner normally absorbs the error and returns cleanly, so this is defensive depth for a money-spending feature. * docs(cli): document --max-budget-usd behavior and limitations Clarify that the budget is cumulative across all agents, checked after each model response, that the scan stops cleanly (not as a failure), that the value must be > 0, and that spend can slightly overshoot due to in-flight calls and best-effort cost estimation. * Apply suggestions from code review Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- .pre-commit-config.yaml | 3 +- docs/usage/cli.mdx | 18 +++++++ pyproject.toml | 9 ++++ strix/core/agents.py | 14 ++++- strix/core/execution.py | 62 +++++++++++++++------ strix/core/hooks.py | 17 +++++- strix/core/runner.py | 12 ++++- strix/interface/cli.py | 1 + strix/interface/main.py | 18 +++++++ strix/interface/tui/app.py | 7 +++ strix/report/state.py | 4 ++ strix/report/usage.py | 4 ++ strix/report/writer.py | 2 +- tests/__init__.py | 0 tests/test_execution.py | 44 +++++++++++++++ tests/test_hooks.py | 108 +++++++++++++++++++++++++++++++++++++ uv.lock | 51 ++++++++++++++++++ 17 files changed, 351 insertions(+), 23 deletions(-) create mode 100644 tests/__init__.py create mode 100644 tests/test_execution.py create mode 100644 tests/test_hooks.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index d0a95523..ab2b5563 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -19,6 +19,7 @@ repos: types-python-dateutil, pydantic, fastapi, + pytest, "openai-agents[litellm]==0.14.6", ] args: [--install-types, --non-interactive] @@ -46,7 +47,7 @@ repos: # Additional Python code quality checks - repo: https://github.com/asottile/pyupgrade - rev: v3.20.0 + rev: v3.21.2 hooks: - id: pyupgrade args: [--py312-plus] diff --git a/docs/usage/cli.mdx b/docs/usage/cli.mdx index bb320096..88ec8d05 100644 --- a/docs/usage/cli.mdx +++ b/docs/usage/cli.mdx @@ -43,6 +43,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/pyproject.toml b/pyproject.toml index 35f3c016..4485d779 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -55,8 +55,13 @@ dev = [ "bandit>=1.8.3", "pre-commit>=4.2.0", "pyinstaller>=6.17.0; python_version >= '3.12' and python_version < '3.15'", + "pytest>=8.3", + "pytest-asyncio>=0.24", ] +[tool.pytest.ini_options] +asyncio_mode = "auto" + [build-system] requires = ["hatchling"] build-backend = "hatchling.build" @@ -104,6 +109,10 @@ module = [ ignore_missing_imports = true disable_error_code = ["import-untyped"] +[[tool.mypy.overrides]] +module = ["tests.*"] +disallow_untyped_decorators = false + # ============================================================================ # Ruff Configuration (Fast Python Linter & Formatter) # ============================================================================ 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 0e9f9140..06dc3ddf 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -15,6 +15,7 @@ from agents.sandbox.errors import ExecTransportError from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore] 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 @@ -97,6 +98,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, @@ -278,6 +283,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, @@ -360,6 +369,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 @@ -377,6 +390,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: if ( image_strips < 3 @@ -527,21 +547,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 a7371e9f..457056e1 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 b5c56ad4..aac3fae7 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -301,6 +301,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", @@ -424,6 +435,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 e8fac0e4..4e43d979 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/strix/report/writer.py b/strix/report/writer.py index 8118fe9f..a7c2146d 100644 --- a/strix/report/writer.py +++ b/strix/report/writer.py @@ -27,7 +27,7 @@ def read_run_record(run_dir: Path) -> dict[str, Any]: except (OSError, json.JSONDecodeError) as exc: raise RuntimeError(f"run.json at {path} is unreadable: {exc}") from exc if not isinstance(data, dict): - raise RuntimeError(f"run.json at {path} is not an object") + raise TypeError(f"run.json at {path} is not an object") return data diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..e69de29b 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) diff --git a/uv.lock b/uv.lock index 29df8e94..cc7cb0a1 100644 --- a/uv.lock +++ b/uv.lock @@ -775,6 +775,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a0/d9/a1e041c5e7caa9a05c925f4bdbdfb7f006d1f74996af53467bc394c97be7/importlib_metadata-8.5.0-py3-none-any.whl", hash = "sha256:45e54197d28b7a7f1559e60b95e7c567032b602131fbd588f1497f47880aa68b", size = 26514, upload-time = "2024-09-11T14:56:07.019Z" }, ] +[[package]] +name = "iniconfig" +version = "2.3.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/72/34/14ca021ce8e5dfedc35312d08ba8bf51fdd999c576889fc2c24cb97f4f10/iniconfig-2.3.0.tar.gz", hash = "sha256:c76315c77db068650d49c5b56314774a7804df16fee4402c1f19d6d15d8c4730", size = 20503, upload-time = "2025-10-18T21:55:43.219Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484, upload-time = "2025-10-18T21:55:41.639Z" }, +] + [[package]] name = "jinja2" version = "3.1.6" @@ -1347,6 +1356,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/63/d7/97f7e3a6abb67d8080dd406fd4df842c2be0efaf712d1c899c32a075027c/platformdirs-4.9.4-py3-none-any.whl", hash = "sha256:68a9a4619a666ea6439f2ff250c12a853cd1cbd5158d258bd824a7df6be2f868", size = 21216, upload-time = "2026-03-05T18:34:12.172Z" }, ] +[[package]] +name = "pluggy" +version = "1.6.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/f9/e2/3e91f31a7d2b083fe6ef3fa267035b518369d9511ffab804f839851d2779/pluggy-1.6.0.tar.gz", hash = "sha256:7dcc130b76258d33b90f61b658791dede3486c3e6bfb003ee5c9bfb396dd22f3", size = 69412, upload-time = "2025-05-15T12:30:07.975Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538, upload-time = "2025-05-15T12:30:06.134Z" }, +] + [[package]] name = "pre-commit" version = "4.5.1" @@ -1633,6 +1651,35 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/0c/82/a2c93e32800940d9573fb28c346772a14778b84ba7524e691b324620ab89/pyright-1.1.408-py3-none-any.whl", hash = "sha256:090b32865f4fdb1e0e6cd82bf5618480d48eecd2eb2e70f960982a3d9a4c17c1", size = 6399144, upload-time = "2026-01-08T08:07:37.082Z" }, ] +[[package]] +name = "pytest" +version = "9.1.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "iniconfig" }, + { name = "packaging" }, + { name = "pluggy" }, + { name = "pygments" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/84/0e/b5858858d74958632c49b72cb25a3976ff9f632397626715be71c89d3971/pytest-9.1.0.tar.gz", hash = "sha256:41dd9148c08072446394cefd3d79701701335a9f4cae69ba92e39f6c7f5c061c", size = 1634181, upload-time = "2026-06-13T18:52:45.983Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/8b/5a/ba30a81239b909821b3153e303e7def45178bf353da4f72380e6c5e8793b/pytest-9.1.0-py3-none-any.whl", hash = "sha256:8ebb0e7888bdf2bdfc602ec51f8f62d50200af37356c74e503c79a94f5c81f32", size = 386453, upload-time = "2026-06-13T18:52:44.045Z" }, +] + +[[package]] +name = "pytest-asyncio" +version = "1.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/43/7c/d36d04db312ecf4298932ef77e6e4a9e8ad017906e24e34f0b0c361a2473/pytest_asyncio-1.4.0.tar.gz", hash = "sha256:c6c0d2259945122819f171a32ecea2c349ead889ee28176caaf492143424be42", size = 58514, upload-time = "2026-05-26T09:56:04.083Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/e2/08a497ef684b88559c9cc5f4ad53a37e7b99e727094a86d6ea32536d5d3c/pytest_asyncio-1.4.0-py3-none-any.whl", hash = "sha256:933ca923a23075a87fb7070c0ec272a6848489824d887c85c812670932835aa1", size = 16930, upload-time = "2026-05-26T09:56:02.576Z" }, +] + [[package]] name = "python-discovery" version = "1.2.0" @@ -2056,6 +2103,8 @@ dev = [ { name = "pre-commit" }, { name = "pyinstaller", marker = "python_full_version < '3.15'" }, { name = "pyright" }, + { name = "pytest" }, + { name = "pytest-asyncio" }, { name = "ruff" }, ] @@ -2079,6 +2128,8 @@ dev = [ { name = "pre-commit", specifier = ">=4.2.0" }, { name = "pyinstaller", marker = "python_full_version >= '3.12' and python_full_version < '3.15'", specifier = ">=6.17.0" }, { name = "pyright", specifier = ">=1.1.401" }, + { name = "pytest", specifier = ">=8.3" }, + { name = "pytest-asyncio", specifier = ">=0.24" }, { name = "ruff", specifier = ">=0.11.13" }, ] From 7141ccff6204b36e9150770babf9e590797bc054 Mon Sep 17 00:00:00 2001 From: Mads Hvelplund Date: Mon, 22 Jun 2026 18:41:42 +0200 Subject: [PATCH 2/2] Support large target repos with with bind-mount option. (#577) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: resolve pre-commit check failures - Change RuntimeError to TypeError for type validation in report/writer.py - Update pyupgrade to v3.21.2 for Python 3.14 compatibility * chore: add pytest test infrastructure Mirror the layout introduced on feature/438-token_budget: pytest + pytest-asyncio dev deps, asyncio_mode auto, a tests.* mypy override, and pytest in the mypy pre-commit hook deps so the tests/ package type-checks. * feat: add --mount and large-target pre-flight for local repos (#492) Large local targets were copied into the sandbox file-by-file via the SDK LocalDir entry, which stalls on big repos and could leave /workspace empty. - --mount bind-mounts a host directory read-only at /workspace/ instead of copying it, bypassing the per-file stream. - A size pre-flight (STRIX_MAX_LOCAL_COPY_MB, default 1024) fails fast with a clear message suggesting --mount when a non-mounted local target is too big. * fix: reject empty --mount paths An empty or whitespace-only --mount value resolves to the current working directory and would silently bind-mount it into the sandbox. Reject it. * fix: dedupe local targets so a dir is never both copied and mounted If the same directory is passed via --target and --mount (or as duplicate values), it previously produced two targets — copied AND bind-mounted, and the copied one could trip the size pre-flight. Dedupe by resolved path, preferring the bind mount. * fix: treat non-positive STRIX_MAX_LOCAL_COPY_MB as disabled Previously a value of 0 (or negative) made every local target count as oversized, aborting all local scans. Now <= 0 disables the pre-flight. * fix: log unreadable subtrees during size pre-flight os.walk silently swallowed directory-listing errors, so a permission-denied subtree could make a large repo under-count and slip past the pre-flight. Surface such omissions via an onerror warning. * docs: document --mount and STRIX_MAX_LOCAL_COPY_MB Add CLI reference + example for --mount, document the size pre-flight env var, note the read-only-is-not-a-hard-boundary caveat and that remote repos are not size-checked, and clarify the backends docstring on when bind mounts apply. * Update strix/interface/main.py * Update strix/runtime/docker_client.py --------- --- docs/advanced/configuration.mdx | 4 + docs/usage/cli.mdx | 17 +++ strix/config/settings.py | 5 + strix/core/inputs.py | 3 +- strix/core/runner.py | 2 +- strix/interface/main.py | 46 +++++++- strix/interface/utils.py | 123 +++++++++++++++++++- strix/runtime/backends.py | 14 ++- strix/runtime/docker_client.py | 20 ++++ strix/runtime/session_manager.py | 50 ++++++-- tests/test_local_sources.py | 188 +++++++++++++++++++++++++++++++ tests/test_session_entries.py | 67 +++++++++++ 12 files changed, 516 insertions(+), 23 deletions(-) create mode 100644 tests/test_local_sources.py create mode 100644 tests/test_session_entries.py diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 9ab7f017..5d56c9a1 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -79,6 +79,10 @@ When remote vars are set, Strix dual-writes telemetry to both local JSONL and th Runtime backend for the sandbox environment. + + Maximum size (in MB) of a local directory target that Strix will copy into the sandbox file-by-file. Larger targets exit early with a suggestion to use `--mount` instead. Set to `0` to disable the check. + + ## Sandbox Configuration diff --git a/docs/usage/cli.mdx b/docs/usage/cli.mdx index 88ec8d05..58d3e99d 100644 --- a/docs/usage/cli.mdx +++ b/docs/usage/cli.mdx @@ -15,6 +15,20 @@ strix --target [options] Target to test. Accepts URLs, repositories, local directories, domains, or IP addresses. Can be specified multiple times. + + Bind-mount a local directory into the sandbox (read-only) instead of copying it in file-by-file. Use this for large repositories that are too big to stream into the container. Can be specified multiple times. + + Strix copies local `--target` directories into the sandbox one file at a time, which stalls on very large trees. When a local target exceeds the copy limit (see `STRIX_MAX_LOCAL_COPY_MB`, default 1024 MB) Strix exits early and asks you to re-run with `--mount`. + + + The mount is read-only to protect your source from accidental modification. This is not a hard security boundary: a root process inside the container can remount it writable, so treat `--mount` as "scan my own code", not as isolation from untrusted code. + + + + The size pre-flight only covers local directory targets. Remote repositories (cloned at scan time) are not size-checked. + + + Custom instructions for the scan. Use for credentials, focus areas, or specific testing approaches. @@ -81,6 +95,9 @@ strix -n --target ./ --scan-mode quick --scope-mode diff --diff-base origin/main # Multi-target white-box testing strix -t https://github.com/org/app -t https://staging.example.com + +# Large local repository — bind-mount instead of copying it in +strix --mount ./huge-monorepo ``` ## Exit Codes diff --git a/strix/config/settings.py b/strix/config/settings.py index 1458e1ff..91fbdef1 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -47,6 +47,11 @@ class RuntimeSettings(BaseSettings): alias="STRIX_IMAGE", ) backend: str = Field(default="docker", alias="STRIX_RUNTIME_BACKEND") + # Hard cap on a local target's size before we refuse to stream it into the + # sandbox file-by-file (the SDK copies every file individually, which stalls + # on large repos). Above this, the user must bind-mount via ``--mount``. + # Set to 0 (or less) to disable the pre-flight check entirely. + max_local_copy_mb: int = Field(default=1024, alias="STRIX_MAX_LOCAL_COPY_MB") class TelemetrySettings(BaseSettings): diff --git a/strix/core/inputs.py b/strix/core/inputs.py index b86daa4c..0505ef38 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -44,7 +44,8 @@ def build_root_task(scan_config: dict[str, Any]) -> str: ) elif ttype == "local_code": path = details.get("target_path", "unknown") - sections["Local Codebases"].append(f"- {path} (available at: {workspace_path})") + suffix = ", read-only mount" if details.get("mount") else "" + sections["Local Codebases"].append(f"- {path} (available at: {workspace_path}{suffix})") elif ttype == "web_application": sections["URLs"].append(f"- {details.get('target_url', '')}") elif ttype == "ip_address": diff --git a/strix/core/runner.py b/strix/core/runner.py index 457056e1..82e653d9 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -55,7 +55,7 @@ async def run_strix_scan( scan_config: dict[str, Any], scan_id: str | None = None, image: str, - local_sources: list[dict[str, str]] | None = None, + local_sources: list[dict[str, Any]] | None = None, coordinator: AgentCoordinator | None = None, interactive: bool = False, max_turns: int = DEFAULT_MAX_TURNS, diff --git a/strix/interface/main.py b/strix/interface/main.py index aac3fae7..4eae0527 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -33,9 +33,12 @@ from strix.interface.tui import run_tui from strix.interface.utils import ( assign_workspace_subdirs, build_final_stats_text, + build_mount_targets_info, check_docker_connection, clone_repository, collect_local_sources, + dedupe_local_targets, + find_oversized_local_targets, generate_run_name, image_exists, infer_target_type, @@ -328,6 +331,9 @@ Examples: # Local code analysis strix --target ./my-project + # Large local repository (bind-mounted read-only instead of copied) + strix --mount ./huge-monorepo + # Domain penetration test strix --target example.com @@ -363,6 +369,15 @@ Examples: "Can be specified multiple times for multi-target scans. " "Required for fresh runs; loaded from disk when ``--resume`` is set.", ) + parser.add_argument( + "--mount", + type=str, + action="append", + metavar="PATH", + help="Bind-mount a local directory into the sandbox (read-only) instead of " + "copying it file-by-file. Use this for large repositories that are too big to " + "stream into the container. Can be specified multiple times.", + ) parser.add_argument( "--instruction", type=str, @@ -473,9 +488,9 @@ Examples: args.user_explicit_instruction = args.instruction if args.resume else None if args.resume: - if args.target: + if args.target or args.mount: parser.error( - "Cannot combine --resume with --target. --resume picks up where " + "Cannot combine --resume with --target/--mount. --resume picks up where " "the prior run left off, including the original target list." ) _load_resume_state(args, parser) @@ -488,13 +503,13 @@ Examples: f"or remove --resume to start over with the same targets." ) else: - if not args.target: + if not args.target and not args.mount: parser.error( - "the following arguments are required: -t/--target " + "the following arguments are required: -t/--target or --mount " "(or use --resume to continue a prior scan)" ) args.targets_info = [] - for target in args.target: + for target in args.target or []: try: target_type, target_dict = infer_target_type(target) @@ -509,9 +524,30 @@ Examples: except ValueError: parser.error(f"Invalid target '{target}'") + try: + args.targets_info.extend(build_mount_targets_info(args.mount or [])) + except ValueError as e: + parser.error(str(e)) + + args.targets_info = dedupe_local_targets(args.targets_info) + assign_workspace_subdirs(args.targets_info) rewrite_localhost_targets(args.targets_info, HOST_GATEWAY_HOSTNAME) + max_local_copy_mb = load_settings().runtime.max_local_copy_mb + max_copy_bytes = max_local_copy_mb * 1024 * 1024 + oversized = find_oversized_local_targets(args.targets_info, max_copy_bytes) + if oversized: + details = "; ".join( + f"{path} ({size / (1024 * 1024):.0f} MB)" for path, size in oversized + ) + parser.error( + f"Local target too large to stream into the sandbox: {details}. " + f"The limit is {max_local_copy_mb} MB " + "(set STRIX_MAX_LOCAL_COPY_MB to change it). Re-run with " + "--mount to bind-mount the directory instead of copying it." + ) + return args diff --git a/strix/interface/utils.py b/strix/interface/utils.py index ff53a6cd..bffc0d47 100644 --- a/strix/interface/utils.py +++ b/strix/interface/utils.py @@ -1,5 +1,6 @@ import ipaddress import json +import logging import os import re import secrets @@ -23,6 +24,9 @@ from rich.text import Text from strix.config import load_settings +logger = logging.getLogger(__name__) + + def get_severity_color(severity: str) -> str: severity_colors = { "critical": "#dc2626", @@ -1185,8 +1189,8 @@ def is_whitebox_scan(targets_info: list[dict[str, Any]]) -> bool: return any(t.get("type") == "local_code" for t in targets_info or []) -def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str, str]]: - local_sources: list[dict[str, str]] = [] +def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str, Any]]: + local_sources: list[dict[str, Any]] = [] for target_info in targets_info: details = target_info["details"] @@ -1197,6 +1201,7 @@ def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str, { "source_path": details["target_path"], "workspace_subdir": workspace_subdir, + "mount": bool(details.get("mount", False)), } ) @@ -1205,12 +1210,126 @@ def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str, { "source_path": details["cloned_repo_path"], "workspace_subdir": workspace_subdir, + "mount": False, } ) return local_sources +def directory_size_bytes(path: Path) -> int: + """Total size in bytes of regular files under ``path`` (symlinks not followed). + + Best-effort: files that disappear or can't be stat'd mid-walk are skipped. + Used as a cheap (stat-only) pre-flight to estimate the cost of streaming a + local target into the sandbox before we actually try to copy it. + + Directories that can't be listed (e.g. permission denied) are logged and + skipped rather than silently dropped — so an under-count is at least + visible — but the returned total then excludes their contents. + """ + + def _on_walk_error(error: OSError) -> None: + logger.warning("Could not read %s while measuring size: %s", error.filename, error) + + total = 0 + for root, _dirs, files in os.walk(path, followlinks=False, onerror=_on_walk_error): + for name in files: + file_path = os.path.join(root, name) # noqa: PTH118 + try: + if os.path.islink(file_path): # noqa: PTH114 + continue + total += os.path.getsize(file_path) # noqa: PTH202 + except OSError: + continue + return total + + +def find_oversized_local_targets( + targets_info: list[dict[str, Any]], max_bytes: int +) -> list[tuple[str, int]]: + """Return ``(path, size_bytes)`` for non-mounted local targets over ``max_bytes``. + + Mounted targets are bind-mounted rather than copied, so their size is + irrelevant and they are excluded. A ``max_bytes`` of zero or less disables + the check entirely (returns no targets). + """ + if max_bytes <= 0: + return [] + oversized: list[tuple[str, int]] = [] + for target in targets_info: + if target.get("type") != "local_code": + continue + details = target.get("details") or {} + if details.get("mount"): + continue + target_path = details.get("target_path") + if not target_path: + continue + size = directory_size_bytes(Path(target_path)) + if size > max_bytes: + oversized.append((target_path, size)) + return oversized + + +def build_mount_targets_info(mount_paths: list[str]) -> list[dict[str, Any]]: + """Build ``targets_info`` entries for ``--mount`` directories. + + Each path must be an existing local directory; it is bind-mounted into the + sandbox (read-only) instead of being copied file-by-file. Raises + ``ValueError`` for an empty path, or one that does not exist or is not a + directory. + """ + targets_info: list[dict[str, Any]] = [] + for raw in mount_paths: + if not raw or not raw.strip(): + raise ValueError("--mount path must not be empty.") + path = Path(raw).expanduser() + try: + resolved = path.resolve() + is_dir = resolved.is_dir() + except (OSError, RuntimeError) as e: + raise ValueError(f"Invalid mount path '{raw}': {e!s}") from e + if not is_dir: + raise ValueError( + f"Mount path '{raw}' is not an existing directory. " + "--mount requires a path to a local directory." + ) + targets_info.append( + { + "type": "local_code", + "details": {"target_path": str(resolved), "mount": True}, + "original": str(resolved), + } + ) + return targets_info + + +def dedupe_local_targets(targets_info: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Collapse local_code targets that resolve to the same path. + + When a directory is supplied both as a copied ``--target`` and via + ``--mount`` (or as duplicate values of either), keep one entry and prefer + the bind-mounted one — so the same tree is never both streamed in and + mounted. Order is preserved; non-local targets pass through untouched. + """ + result: list[dict[str, Any]] = [] + index_by_path: dict[str, int] = {} + for target in targets_info: + details = target.get("details") or {} + path = details.get("target_path") + if target.get("type") != "local_code" or not path: + result.append(target) + continue + existing = index_by_path.get(path) + if existing is None: + index_by_path[path] = len(result) + result.append(target) + elif details.get("mount") and not (result[existing].get("details") or {}).get("mount"): + result[existing] = target # bind mount supersedes the copied entry + return result + + def _is_localhost_host(host: str) -> bool: host_lower = host.lower().strip("[]") diff --git a/strix/runtime/backends.py b/strix/runtime/backends.py index 9f241a3a..d7eba335 100644 --- a/strix/runtime/backends.py +++ b/strix/runtime/backends.py @@ -22,6 +22,7 @@ async def _docker_backend( image: str, manifest: Manifest, exposed_ports: tuple[int, ...], + bind_mounts: list[dict[str, Any]] | None = None, ) -> tuple[Any, Any]: """Bring up a session backed by the local Docker daemon. @@ -31,11 +32,15 @@ async def _docker_backend( backend don't need the docker-py library installed. ``session.start()`` is what materializes the manifest entries - (LocalDir copies, mount setup, etc.) into the running container — - the SDK's ``client.create()`` only builds the inner session object - without applying the manifest. ``async with session:`` would call it - too, but Strix manages session lifetime explicitly via + (LocalDir copies and manifest-declared volume/FUSE mounts) into the + running container — the SDK's ``client.create()`` only builds the inner + session object without applying the manifest. ``async with session:`` + would call it too, but Strix manages session lifetime explicitly via ``client.delete()`` so we trigger ``start()`` ourselves. + + ``bind_mounts`` are host directories (e.g. large repos passed via + ``--mount``) bind-mounted read-only; unlike manifest entries they are + applied by Docker at container-create time, not by ``start()``. """ import docker from agents.sandbox.sandboxes.docker import DockerSandboxClientOptions @@ -43,6 +48,7 @@ async def _docker_backend( from strix.runtime.docker_client import StrixDockerSandboxClient client = StrixDockerSandboxClient(docker.from_env()) + client.strix_bind_mounts = bind_mounts or [] options = DockerSandboxClientOptions(image=image, exposed_ports=exposed_ports) session = await client.create(options=options, manifest=manifest) await session.start() diff --git a/strix/runtime/docker_client.py b/strix/runtime/docker_client.py index 2a753834..497ae2f2 100644 --- a/strix/runtime/docker_client.py +++ b/strix/runtime/docker_client.py @@ -38,6 +38,7 @@ from agents.sandbox.sandboxes.docker import ( from agents.sandbox.session.sandbox_session import SandboxSession from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore] from docker.models.containers import Container # type: ignore[import-untyped, unused-ignore] +from docker.types import Mount as DockerSDKMount # type: ignore[import-untyped, unused-ignore] from docker.utils import parse_repository_tag # type: ignore[import-untyped, unused-ignore] @@ -45,6 +46,10 @@ logger = logging.getLogger(__name__) class StrixDockerSandboxClient(DockerSandboxClient): + # Host directories to bind-mount into the container, set by the docker + # backend before ``create()``. Each item is ``{source, target, read_only}``. + strix_bind_mounts: list[dict[str, Any]] = [] # overridden per-instance in backends.py + async def _create_container( self, image: str, @@ -111,6 +116,21 @@ class StrixDockerSandboxClient(DockerSandboxClient): extra_hosts = create_kwargs.setdefault("extra_hosts", {}) extra_hosts["host.docker.internal"] = "host-gateway" + # Strix injection: host bind mounts (e.g. large repos passed via --mount) + # that bypass the SDK's file-by-file LocalDir copy. + bind_mounts = getattr(self, "strix_bind_mounts", ()) + if bind_mounts: + mounts = create_kwargs.setdefault("mounts", []) + for spec in bind_mounts: + mounts.append( + DockerSDKMount( + target=spec["target"], + source=spec["source"], + type="bind", + read_only=spec.get("read_only", True), + ) + ) + logger.debug( "Creating sandbox container: image=%s caps=%s exposed_ports=%s", image, diff --git a/strix/runtime/session_manager.py b/strix/runtime/session_manager.py index 6f3e2733..a19495cc 100644 --- a/strix/runtime/session_manager.py +++ b/strix/runtime/session_manager.py @@ -23,30 +23,59 @@ _CONTAINER_CAIDO_PORT = 48080 _SESSION_CACHE: dict[str, dict[str, Any]] = {} +# Manifest root inside the container; entry keys hang off this path. +_WORKSPACE_ROOT = "/workspace" + + +def build_session_entries( + local_sources: list[dict[str, Any]], +) -> tuple[dict[str | Path, BaseEntry], list[dict[str, Any]]]: + """Split local sources into copied manifest entries and host bind mounts. + + Sources flagged ``mount`` are bind-mounted read-only at + ``/workspace/`` (not added to the manifest, so the SDK + does not stream them in file-by-file). Every other source becomes a + ``LocalDir`` entry copied into the container as before. + """ + entries: dict[str | Path, BaseEntry] = {} + bind_mounts: list[dict[str, Any]] = [] + for src in local_sources: + ws_subdir = src.get("workspace_subdir") or "" + host_path = src.get("source_path") or "" + if not ws_subdir or not host_path: + continue + resolved = Path(host_path).expanduser().resolve() + if src.get("mount"): + bind_mounts.append( + { + "source": str(resolved), + "target": f"{_WORKSPACE_ROOT}/{ws_subdir}", + "read_only": True, + } + ) + else: + entries[ws_subdir] = LocalDir(src=resolved) + return entries, bind_mounts + async def create_or_reuse( scan_id: str, *, image: str, - local_sources: list[dict[str, str]], + local_sources: list[dict[str, Any]], ) -> dict[str, Any]: """Return the existing session bundle for ``scan_id`` or create a new one. - Each ``local_sources`` entry mounts its host ``source_path`` at - ``/workspace/`` inside the container. + Each ``local_sources`` entry exposes its host ``source_path`` at + ``/workspace/`` inside the container — copied in, or + bind-mounted read-only when the entry is flagged ``mount``. """ cached = _SESSION_CACHE.get(scan_id) if cached is not None: logger.info("Reusing existing sandbox session for scan %s", scan_id) return cached - entries: dict[str | Path, BaseEntry] = {} - for src in local_sources: - ws_subdir = src.get("workspace_subdir") or "" - host_path = src.get("source_path") or "" - if not ws_subdir or not host_path: - continue - entries[ws_subdir] = LocalDir(src=Path(host_path).expanduser().resolve()) + entries, bind_mounts = build_session_entries(local_sources) # Caido runs as an in-container sidecar; HTTP(S) traffic from any # process started via ``session.exec`` (the SDK's Shell tool, etc.) @@ -81,6 +110,7 @@ async def create_or_reuse( image=image, manifest=manifest, exposed_ports=(_CONTAINER_CAIDO_PORT,), + bind_mounts=bind_mounts, ) caido_endpoint = await session.resolve_exposed_port(_CONTAINER_CAIDO_PORT) diff --git a/tests/test_local_sources.py b/tests/test_local_sources.py new file mode 100644 index 00000000..bd3448de --- /dev/null +++ b/tests/test_local_sources.py @@ -0,0 +1,188 @@ +"""Tests for local-source sizing and ``--mount`` target helpers in interface.utils.""" + +from __future__ import annotations + +import logging +import os +import sys +from typing import TYPE_CHECKING, Any + +import pytest + + +if TYPE_CHECKING: + from pathlib import Path + +from strix.interface.utils import ( + build_mount_targets_info, + collect_local_sources, + dedupe_local_targets, + directory_size_bytes, + find_oversized_local_targets, +) + + +def _write_file(path: Path, size: int) -> None: + path.write_bytes(b"x" * size) + + +def _local_target(target_path: str, *, mount: bool = False) -> dict[str, Any]: + details: dict[str, Any] = {"target_path": target_path, "workspace_subdir": "repo"} + if mount: + details["mount"] = True + return {"type": "local_code", "details": details, "original": target_path} + + +def test_directory_size_empty_dir_is_zero(tmp_path: Path) -> None: + assert directory_size_bytes(tmp_path) == 0 + + +def test_directory_size_sums_flat_and_nested_files(tmp_path: Path) -> None: + _write_file(tmp_path / "a.txt", 100) + nested = tmp_path / "sub" / "deep" + nested.mkdir(parents=True) + _write_file(nested / "b.txt", 250) + assert directory_size_bytes(tmp_path) == 350 + + +def test_directory_size_skips_symlinks(tmp_path: Path) -> None: + _write_file(tmp_path / "real.txt", 100) + (tmp_path / "link.txt").symlink_to(tmp_path / "real.txt") + # The symlink target is counted once via the real file, not doubled. + assert directory_size_bytes(tmp_path) == 100 + + +@pytest.mark.skipif(sys.platform == "win32", reason="relies on POSIX permissions") +def test_directory_size_logs_and_skips_unreadable_subdir( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + if hasattr(os, "geteuid") and os.geteuid() == 0: + pytest.skip("root bypasses directory permissions") + _write_file(tmp_path / "top.txt", 100) + locked = tmp_path / "locked" + locked.mkdir() + _write_file(locked / "secret.bin", 9999) + locked.chmod(0o000) + try: + with caplog.at_level(logging.WARNING): + size = directory_size_bytes(tmp_path) + finally: + locked.chmod(0o755) + # The unreadable subtree is excluded (not silently treated as readable) and + # the omission is logged rather than vanishing without a trace. + assert size == 100 + assert any("Could not read" in record.message for record in caplog.records) + + +def test_find_oversized_returns_nothing_under_limit(tmp_path: Path) -> None: + _write_file(tmp_path / "a.txt", 100) + targets = [_local_target(str(tmp_path))] + assert find_oversized_local_targets(targets, max_bytes=1000) == [] + + +def test_find_oversized_returns_target_over_limit(tmp_path: Path) -> None: + _write_file(tmp_path / "big.bin", 500) + targets = [_local_target(str(tmp_path))] + result = find_oversized_local_targets(targets, max_bytes=100) + assert result == [(str(tmp_path), 500)] + + +def test_find_oversized_ignores_mounted_targets(tmp_path: Path) -> None: + _write_file(tmp_path / "big.bin", 500) + targets = [_local_target(str(tmp_path), mount=True)] + assert find_oversized_local_targets(targets, max_bytes=100) == [] + + +def test_find_oversized_ignores_non_local_targets() -> None: + targets = [{"type": "web_application", "details": {"target_url": "https://x"}}] + assert find_oversized_local_targets(targets, max_bytes=1) == [] + + +@pytest.mark.parametrize("disabled", [0, -1]) +def test_find_oversized_disabled_for_non_positive_limit(tmp_path: Path, disabled: int) -> None: + _write_file(tmp_path / "big.bin", 500) + targets = [_local_target(str(tmp_path))] + assert find_oversized_local_targets(targets, max_bytes=disabled) == [] + + +def test_collect_local_sources_propagates_mount_flag() -> None: + copied = _local_target("/copied") + copied["details"]["workspace_subdir"] = "copied" + mounted = _local_target("/mounted", mount=True) + mounted["details"]["workspace_subdir"] = "mounted" + + sources = collect_local_sources([copied, mounted]) + + by_path = {s["source_path"]: s for s in sources} + assert by_path["/copied"]["mount"] is False + assert by_path["/mounted"]["mount"] is True + + +def test_collect_local_sources_repository_is_never_mounted() -> None: + repo = { + "type": "repository", + "details": {"cloned_repo_path": "/clone", "workspace_subdir": "clone"}, + } + sources = collect_local_sources([repo]) + assert sources == [{"source_path": "/clone", "workspace_subdir": "clone", "mount": False}] + + +def test_build_mount_targets_info_for_valid_dir(tmp_path: Path) -> None: + result = build_mount_targets_info([str(tmp_path)]) + assert len(result) == 1 + entry = result[0] + assert entry["type"] == "local_code" + assert entry["details"]["mount"] is True + assert entry["details"]["target_path"] == str(tmp_path.resolve()) + + +def test_build_mount_targets_info_rejects_missing_path(tmp_path: Path) -> None: + missing = tmp_path / "does-not-exist" + with pytest.raises(ValueError, match="not an existing directory"): + build_mount_targets_info([str(missing)]) + + +def test_build_mount_targets_info_rejects_file(tmp_path: Path) -> None: + file_path = tmp_path / "a-file.txt" + _write_file(file_path, 10) + with pytest.raises(ValueError, match="not an existing directory"): + build_mount_targets_info([str(file_path)]) + + +@pytest.mark.parametrize("empty", ["", " "]) +def test_build_mount_targets_info_rejects_empty_path(empty: str) -> None: + # An empty path would otherwise resolve to the current working directory + # and silently bind-mount it into the sandbox. + with pytest.raises(ValueError, match="must not be empty"): + build_mount_targets_info([empty]) + + +def test_dedupe_keeps_distinct_targets_in_order() -> None: + targets = [ + _local_target("/a"), + {"type": "web_application", "details": {"target_url": "https://x"}}, + _local_target("/b", mount=True), + ] + assert dedupe_local_targets(targets) == targets + + +def test_dedupe_mount_supersedes_copied_same_path() -> None: + copied = _local_target("/repo") + mounted = _local_target("/repo", mount=True) + + # Copied first, then mounted: the single surviving entry is the mount. + result = dedupe_local_targets([copied, mounted]) + assert len(result) == 1 + assert result[0]["details"]["mount"] is True + + # Order-independent: mounted first, copied second also yields the mount. + result_rev = dedupe_local_targets([mounted, copied]) + assert len(result_rev) == 1 + assert result_rev[0]["details"]["mount"] is True + + +def test_dedupe_collapses_duplicate_mounts() -> None: + result = dedupe_local_targets( + [_local_target("/repo", mount=True), _local_target("/repo", mount=True)] + ) + assert len(result) == 1 diff --git a/tests/test_session_entries.py b/tests/test_session_entries.py new file mode 100644 index 00000000..8288c1d0 --- /dev/null +++ b/tests/test_session_entries.py @@ -0,0 +1,67 @@ +"""Tests for build_session_entries: splitting copied vs bind-mounted sources.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from agents.sandbox.entries import LocalDir + +from strix.runtime.session_manager import build_session_entries + + +if TYPE_CHECKING: + from pathlib import Path + + +def _source(subdir: str, path: str, *, mount: bool = False) -> dict[str, Any]: + return {"source_path": path, "workspace_subdir": subdir, "mount": mount} + + +def test_copied_source_becomes_localdir_entry(tmp_path: Path) -> None: + entries, bind_mounts = build_session_entries([_source("repo", str(tmp_path))]) + + assert bind_mounts == [] + assert isinstance(entries["repo"], LocalDir) + assert entries["repo"].src == tmp_path.resolve() + + +def test_mounted_source_becomes_bind_mount(tmp_path: Path) -> None: + entries, bind_mounts = build_session_entries([_source("repo", str(tmp_path), mount=True)]) + + assert entries == {} + assert bind_mounts == [ + { + "source": str(tmp_path.resolve()), + "target": "/workspace/repo", + "read_only": True, + } + ] + + +def test_mixed_sources_split_correctly(tmp_path: Path) -> None: + copied = tmp_path / "copied" + mounted = tmp_path / "mounted" + copied.mkdir() + mounted.mkdir() + + entries, bind_mounts = build_session_entries( + [ + _source("copied", str(copied)), + _source("mounted", str(mounted), mount=True), + ] + ) + + assert list(entries) == ["copied"] + assert isinstance(entries["copied"], LocalDir) + assert [m["target"] for m in bind_mounts] == ["/workspace/mounted"] + + +def test_incomplete_sources_are_skipped() -> None: + entries, bind_mounts = build_session_entries( + [ + {"source_path": "", "workspace_subdir": "x"}, + {"source_path": "/p", "workspace_subdir": ""}, + ] + ) + assert entries == {} + assert bind_mounts == []