diff --git a/docs/usage/cli.mdx b/docs/usage/cli.mdx index 3b10f9de..d5acb402 100644 --- a/docs/usage/cli.mdx +++ b/docs/usage/cli.mdx @@ -61,11 +61,28 @@ strix (--target | --target-list | --mount ) [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. + response. + + In non-interactive mode (`-n`), once the running cost reaches the threshold, + the scan stops cleanly with a `stopped` status (not a failure) and the sandbox + is torn down. Sub-agents are stopped early, at 90% of the budget, reserving + the final slice for the root agent to wind down and produce the final report. + + In interactive mode, reaching the budget pauses the scan instead of ending + it: every agent parks, and sending any message resumes the scan with the cap + extended by the original budget amount. There is no sub-agent reserve in + interactive mode. + + As the budget is approached, graduated wrap-up warnings are surfaced to + **every** agent so they can finish their work and call their lifecycle tool + before the hard stop. The bands sit just below each role's own stop point: the + root is warned at **70%, 85% and 95%** (it stops at 100%), while sub-agents are + warned at **75%, 80% and 85%** (they stop at the 90% reserve). In interactive + mode every agent uses the **70%, 85% and 95%** bands. Percentages shown in the + warnings are the real cumulative spend against the full budget. Must be greater than `0`. Omit the flag for no limit. @@ -84,6 +101,19 @@ strix (--target | --target-list | --mount ) [options] counts. + + Maximum number of turns (one model response plus its tool round) allotted to + **each** agent, applied per run. When an agent reaches this limit it is + force-stopped. + + As the limit is approached, graduated wrap-up warnings (at 70%, 85% and 95%) + are injected into that agent's next model turn so it can prioritise its + remaining work and call its lifecycle tool (`finish_scan` for the root agent, + `agent_finish` for sub-agents) before the hard stop. + + Must be greater than `0`. + + ## Examples ```bash @@ -99,6 +129,9 @@ strix --target api.example.com --instruction "Focus on IDOR and auth bypass" # CI/CD mode strix -n --target ./ --scan-mode quick +# Cap cost and per-agent turns +strix --target https://example.com --max-budget 25 --max-turns 300 + # Force diff-scope against a specific base ref strix -n --target ./ --scan-mode quick --scope-mode diff --diff-base origin/main diff --git a/strix/core/agents.py b/strix/core/agents.py index c08b24d6..93aa4975 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -14,13 +14,15 @@ from strix.core.sessions import session_write_lock if TYPE_CHECKING: + from collections.abc import Callable + from agents.items import TResponseInputItem from agents.memory import Session logger = logging.getLogger(__name__) -Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed"] +Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed", "budget_paused"] @dataclass(slots=True) @@ -47,6 +49,9 @@ class AgentCoordinator: self._snapshot_path: Path | None = None self.is_shutting_down = False self._budget_stopped = False + self._reserve_stopped = False + self._budget_paused = False + self._extend_budget: Callable[[], None] | None = None def set_snapshot_path(self, path: Path) -> None: self._snapshot_path = path @@ -65,6 +70,71 @@ class AgentCoordinator: for runtime in self.runtimes.values(): runtime.wake.set() + @property + def reserve_stopped(self) -> bool: + return self._reserve_stopped + + @property + def budget_paused(self) -> bool: + return self._budget_paused + + def set_budget_extender(self, extend: Callable[[], None]) -> None: + self._extend_budget = extend + + async def pause_for_budget(self, agent_id: str) -> None: + async with self._lock: + self._budget_paused = True + await self.set_status(agent_id, "budget_paused") + + async def resume_from_budget_pause(self, *, exclude: str | None = None) -> None: + async with self._lock: + if not self._budget_paused: + return + self._budget_paused = False + paused = [aid for aid, status in self.statuses.items() if status == "budget_paused"] + if self._extend_budget is not None: + self._extend_budget() + for aid in paused: + await self.set_status(aid, "waiting") + if aid != exclude: + await self.send( + aid, + { + "from": "system", + "type": "budget_extended", + "content": ( + "[Budget] The user extended the scan budget \u2014 continue your " + "current task." + ), + }, + ) + + async def reset_budget_stops( + self, + *, + budget_stopped: bool, + reserve_stopped: bool, + budget_paused: bool = False, + ) -> None: + async with self._lock: + self._budget_stopped = budget_stopped + self._reserve_stopped = reserve_stopped + if not budget_paused: + self._budget_paused = False + for aid, status in self.statuses.items(): + if status == "budget_paused": + self.statuses[aid] = "waiting" + await self._maybe_snapshot() + + async def claim_reserve_notification(self) -> str | None: + async with self._lock: + if self._reserve_stopped: + return None + self._reserve_stopped = True + for runtime in self.runtimes.values(): + runtime.wake.set() + return next((aid for aid, parent in self.parent_of.items() if parent is None), None) + async def register( self, agent_id: str, @@ -132,6 +202,8 @@ class AgentCoordinator: async def send(self, target_agent_id: str, message: dict[str, Any]) -> bool: """Deliver a user/peer message by appending it to the target SDK session.""" + if message.get("from") == "user" and self._budget_paused: + await self.resume_from_budget_pause(exclude=target_agent_id) async with self._lock: if target_agent_id not in self.statuses: logger.debug("agent.send dropped unknown target=%s", target_agent_id) @@ -166,7 +238,8 @@ class AgentCoordinator: async def wait_for_message(self, agent_id: str) -> None: while True: async with self._lock: - if self._budget_stopped or self.pending_counts.get(agent_id, 0) > 0: + reserve_exit = self._reserve_stopped and self.parent_of.get(agent_id) is not None + if self._budget_stopped or reserve_exit or self.pending_counts.get(agent_id, 0) > 0: return wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake wake.clear() @@ -300,6 +373,9 @@ class AgentCoordinator: "metadata": {aid: dict(md) for aid, md in self.metadata.items()}, "pending_counts": dict(self.pending_counts), "errors": dict(self.errors), + "budget_stopped": self._budget_stopped, + "reserve_stopped": self._reserve_stopped, + "budget_paused": self._budget_paused, } async def restore(self, snap: dict[str, Any]) -> None: @@ -310,6 +386,9 @@ class AgentCoordinator: self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()} self.pending_counts = dict(snap.get("pending_counts", {})) self.errors = dict(snap.get("errors", {})) + self._budget_stopped = bool(snap.get("budget_stopped", False)) + self._reserve_stopped = bool(snap.get("reserve_stopped", False)) + self._budget_paused = bool(snap.get("budget_paused", False)) for aid in self.statuses: self.runtimes.setdefault(aid, AgentRuntime()) diff --git a/strix/core/execution.py b/strix/core/execution.py index ea65612d..4efd5f9d 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -21,7 +21,11 @@ from openai import ( RateLimitError, ) -from strix.core.hooks import BudgetExceededError +from strix.core.hooks import ( + BudgetExceededError, + BudgetPausedError, + SubagentBudgetReservedError, +) from strix.core.inputs import child_initial_input from strix.core.sessions import ( enforce_image_budget, @@ -134,21 +138,34 @@ async def run_agent_loop( ) result: RunResultBase | None = None + budget_stopped = coordinator.budget_stopped + reserve_stopped = coordinator.reserve_stopped + if budget_stopped: + await coordinator.set_status(agent_id, "stopped") + raise BudgetExceededError("scan budget reached") + if reserve_stopped and context.get("parent_id") is not None: + await coordinator.set_status(agent_id, "stopped") + raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve") + + if reserve_stopped and start_parked and interactive and context.get("parent_id") is None: + await coordinator.send(agent_id, _reserve_notice()) + if not (start_parked and interactive): if interactive: - result = await _run_cycle( - agent, - coordinator, - agent_id, - input_data=initial_input, - run_config=run_config, - context=context, - max_turns=max_turns, - session=session, - interactive=interactive, - event_sink=event_sink, - hooks=hooks, - ) + with contextlib.suppress(BudgetPausedError): + result = await _run_cycle( + agent, + coordinator, + agent_id, + input_data=initial_input, + run_config=run_config, + context=context, + max_turns=max_turns, + session=session, + interactive=interactive, + event_sink=event_sink, + hooks=hooks, + ) else: result = await _run_noninteractive_until_lifecycle( agent, @@ -176,20 +193,25 @@ async def run_agent_loop( await coordinator.set_status(agent_id, "stopped") raise BudgetExceededError("scan budget reached") + if coordinator.reserve_stopped and context.get("parent_id") is not None: + await coordinator.set_status(agent_id, "stopped") + raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve") + await coordinator.consume_pending(agent_id) - result = await _run_cycle( - agent, - coordinator, - agent_id, - input_data=[], - run_config=run_config, - context=context, - max_turns=max_turns, - session=session, - interactive=interactive, - event_sink=event_sink, - hooks=hooks, - ) + with contextlib.suppress(BudgetPausedError): + result = await _run_cycle( + agent, + coordinator, + agent_id, + input_data=[], + run_config=run_config, + context=context, + max_turns=max_turns, + session=session, + interactive=interactive, + event_sink=event_sink, + hooks=hooks, + ) async def spawn_child_agent( @@ -361,6 +383,10 @@ async def _run_noninteractive_until_lifecycle( await coordinator.set_status(agent_id, "stopped") raise BudgetExceededError("scan budget reached") + if coordinator.reserve_stopped and context.get("parent_id") is not None: + await coordinator.set_status(agent_id, "stopped") + raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve") + result = await _run_cycle( agent, coordinator, @@ -456,9 +482,7 @@ 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. + except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError): raise except RuntimeError as stream_exc: if "after shutdown" not in str(stream_exc): @@ -477,6 +501,15 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 ) finally: await coordinator.detach_stream(agent_id, stream) + except BudgetPausedError as exc: + logger.info("agent %s paused at the scan budget limit: %s", agent_id, exc) + await coordinator.pause_for_budget(agent_id) + raise + except SubagentBudgetReservedError as exc: + logger.info("sub-agent %s stopped at the budget reserve: %s", agent_id, exc) + await coordinator.set_status(agent_id, "stopped") + await _notify_root_on_budget_reserve(coordinator) + raise except BudgetExceededError as exc: logger.info( "agent %s reached the scan budget limit; stopping the scan: %s", agent_id, exc @@ -639,6 +672,28 @@ async def _notify_parent_on_crash( ) +def _reserve_notice() -> dict[str, Any]: + return { + "from": "system", + "type": "budget_reserve_stop", + "priority": "high", + "content": ( + "[Budget reserve] The scan has reached the sub-agent budget reserve: every " + "sub-agent is being force-stopped as soon as its in-flight turn completes, and " + "none will send a completion report. Their confirmed vulnerabilities are " + "already filed as they were found. Do not wait on any sub-agents and do not " + "spawn new ones — wrap up now and call finish_scan." + ), + } + + +async def _notify_root_on_budget_reserve(coordinator: AgentCoordinator) -> None: + root = await coordinator.claim_reserve_notification() + if root is None: + return + await coordinator.send(root, _reserve_notice()) + + async def _start_child_runner( *, parent_ctx: dict[str, Any], @@ -690,6 +745,8 @@ async def _start_child_runner( ) except BudgetExceededError: logger.info("child %s stopped after reaching the scan budget limit", child_id) + except SubagentBudgetReservedError: + logger.info("child %s stopped at the sub-agent budget reserve", 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 133375ab..bfeafd45 100644 --- a/strix/core/hooks.py +++ b/strix/core/hooks.py @@ -14,26 +14,210 @@ from strix.report.state import get_global_report_state if TYPE_CHECKING: from agents import RunContextWrapper from agents.agent import Agent - from agents.items import ModelResponse + from agents.items import ModelResponse, TResponseInputItem logger = logging.getLogger(__name__) +_STAGE_LABELS: tuple[str, ...] = ("NOTICE", "URGENT", "CRITICAL") +_TURN_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95) +_ROOT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95) +_SUBAGENT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.75, 0.80, 0.85) +_SUBAGENT_BUDGET_RESERVE = 0.90 + + 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.""" +class SubagentBudgetReservedError(RuntimeError): + """Raised to stop a single sub-agent once the reserve threshold is crossed.""" - def __init__(self, *, model: str, max_budget_usd: float | None = None) -> None: + +class BudgetPausedError(RuntimeError): + """Raised to park one agent when an interactive scan reaches its budget.""" + + +def recomputed_budget_flags( + cost: float, + max_budget_usd: float | None, + *, + interactive: bool, +) -> tuple[bool, bool]: + """Return the (budget_stopped, reserve_stopped) flags a resumed scan should carry.""" + if max_budget_usd is None: + return False, False + if interactive: + return False, False + budget_stopped = cost >= max_budget_usd + reserve_stopped = cost >= max_budget_usd * _SUBAGENT_BUDGET_RESERVE + return budget_stopped, reserve_stopped + + +def _crossed_stage(fraction: float, bands: tuple[float, ...]) -> int | None: + crossed: int | None = None + for index, band in enumerate(bands): + if fraction >= band: + crossed = index + return crossed + + +_ROOT_DIRECTIVES: tuple[str, ...] = ( + ( + "As the root agent, begin planning your wind-down of the whole scan: avoid " + "starting large new lines of investigation, and keep your required objectives on " + "track so you can call finish_scan comfortably before the limit." + ), + ( + "As the root agent, prioritize wrapping up the whole scan now: stop opening new " + "lines of investigation, close out only what is essential, and move toward calling " + "finish_scan to compile and deliver the final report." + ), + ( + "As the root agent, STOP all other work on the whole scan and finish immediately: " + "secure your findings and call finish_scan now — anything left unfinished when the " + "limit is hit is discarded." + ), +) +_SUBAGENT_DIRECTIVES: tuple[str, ...] = ( + ( + "As a sub-agent, begin planning your wind-down: avoid starting large new subtasks, " + "and if you are close to a confirmed, validated vulnerability, drive it to a result " + "you can report." + ), + ( + "As a sub-agent, prioritize wrapping up your task now: report any confirmed, " + "validated vulnerability, finish work that is nearly done rather than starting " + "anything new, and prepare to call agent_finish." + ), + ( + "As a sub-agent, STOP all other work and finish immediately: report any confirmed " + "vulnerability right now and call agent_finish to hand your results back to your " + "parent before you are cut off." + ), +) + + +def _wrapup_directive(context: RunContextWrapper[dict[str, Any]], stage: int) -> str: + is_root = context.context.get("parent_id") is None + directives = _ROOT_DIRECTIVES if is_root else _SUBAGENT_DIRECTIVES + return directives[stage] + + +def _urgency(stage: int) -> str: + return _STAGE_LABELS[stage] + + +class ReportUsageHooks(RunHooks[dict[str, Any]]): + """Persist SDK-native usage and warn/stop as turn and cost budgets are consumed.""" + + def __init__( + self, + *, + model: str, + max_budget_usd: float | None = None, + max_turns: int | None = None, + interactive: bool = False, + ) -> None: if max_budget_usd is not None and ( not math.isfinite(max_budget_usd) or max_budget_usd <= 0 ): raise ValueError("max_budget_usd must be a finite number greater than 0") + if max_turns is not None and max_turns <= 0: + raise ValueError("max_turns must be a positive integer") self._model = model self._max_budget_usd = max_budget_usd + self._budget_increment = max_budget_usd + self._max_turns = max_turns + self._interactive = interactive + + def extend_budget(self) -> None: + if self._max_budget_usd is None or self._budget_increment is None: + return + self._max_budget_usd += self._budget_increment + + async def on_llm_start( + self, + context: RunContextWrapper[dict[str, Any]], + agent: Agent[dict[str, Any]], # noqa: ARG002 + system_prompt: str | None, # noqa: ARG002 + input_items: list[TResponseInputItem], + ) -> None: + try: + self._maybe_warn_turns(context, input_items) + self._maybe_warn_budget(context, input_items) + except Exception: + logger.exception("budget/turn warning injection failed") + + def _maybe_warn_turns( + self, + context: RunContextWrapper[dict[str, Any]], + input_items: list[TResponseInputItem], + ) -> None: + if not self._max_turns: + return + usage = getattr(context, "usage", None) + requests = getattr(usage, "requests", None) + if not isinstance(requests, int): + return + turns_used = requests + 1 + stage = _crossed_stage(turns_used / self._max_turns, _TURN_WARN_BANDS) + if stage is None: + return + remaining = max(self._max_turns - turns_used, 0) + pct = round(100 * turns_used / self._max_turns) + content = ( + f"[{_urgency(stage)}] Turn budget: {turns_used}/{self._max_turns} used ({pct}%). " + f"About {remaining} turn(s) remain before this agent is force-stopped and any " + f"in-progress work is discarded. {_wrapup_directive(context, stage)}" + ) + input_items.append({"role": "user", "content": content}) + + def _maybe_warn_budget( + self, + context: RunContextWrapper[dict[str, Any]], + input_items: list[TResponseInputItem], + ) -> None: + if self._max_budget_usd is None: + return + report_state = get_global_report_state() + if report_state is None: + return + cost = report_state.get_total_llm_cost() + is_root = context.context.get("parent_id") is None + if self._interactive: + bands = _ROOT_BUDGET_WARN_BANDS + else: + bands = _ROOT_BUDGET_WARN_BANDS if is_root else _SUBAGENT_BUDGET_WARN_BANDS + stage = _crossed_stage(cost / self._max_budget_usd, bands) + if stage is None: + return + pct = round(100 * cost / self._max_budget_usd) + reserve_pct = round(_SUBAGENT_BUDGET_RESERVE * 100) + if self._interactive: + content = ( + f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} " + f"spent ({pct}%). This budget is shared across every agent in the scan; when it " + "is reached all agents are paused until the user chooses to continue. " + f"{_wrapup_directive(context, stage)}" + ) + elif is_root: + content = ( + f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} " + f"spent ({pct}%). This budget is shared across every agent in the scan; when it " + "is reached the whole scan is stopped immediately, and sub-agents are stopped at " + f"{reserve_pct}% to reserve the remainder for your final report. " + f"{_wrapup_directive(context, stage)}" + ) + else: + content = ( + f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} " + f"spent ({pct}%). This budget is shared across every agent in the scan; " + f"sub-agents are stopped at {reserve_pct}% to leave the remainder for the root " + f"agent's final report. {_wrapup_directive(context, stage)}" + ) + input_items.append({"role": "user", "content": content}) async def on_llm_end( self, @@ -66,6 +250,21 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]): if self._max_budget_usd is not None: cost = report_state.get_total_llm_cost() if cost >= self._max_budget_usd: + if self._interactive: + raise BudgetPausedError( + f"Scan budget of ${self._max_budget_usd:.2f} reached " + f"(spent ${cost:.4f}); pausing until the user continues" + ) raise BudgetExceededError( f"Token budget of ${self._max_budget_usd:.2f} exceeded (spent ${cost:.4f})" ) + is_root = ctx.get("parent_id") is None + if not self._interactive and not is_root: + reserve_limit = self._max_budget_usd * _SUBAGENT_BUDGET_RESERVE + if cost >= reserve_limit: + raise SubagentBudgetReservedError( + f"Sub-agent budget reserve reached: spent ${cost:.4f} of " + f"${self._max_budget_usd:.2f} " + f"(>= {round(_SUBAGENT_BUDGET_RESERVE * 100)}% reserve); stopping this " + "sub-agent so the root agent can finish the scan." + ) diff --git a/strix/core/runner.py b/strix/core/runner.py index 77a4eff1..fdb432d0 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -31,7 +31,7 @@ from strix.core.execution import ( from strix.core.execution import ( spawn_child_agent as start_child_agent, ) -from strix.core.hooks import BudgetExceededError, ReportUsageHooks +from strix.core.hooks import BudgetExceededError, ReportUsageHooks, recomputed_budget_flags from strix.core.inputs import ( DEFAULT_MAX_TURNS, build_root_task, @@ -40,6 +40,7 @@ from strix.core.inputs import ( ) from strix.core.paths import run_dir_for, runtime_state_dir from strix.core.sessions import open_agent_session +from strix.report.state import get_global_report_state from strix.runtime import session_manager from strix.telemetry.logging import set_scan_id, setup_scan_logging from strix.tools.output_store import ( @@ -185,6 +186,18 @@ async def run_strix_scan( f"Cannot resume scan {scan_id}: missing SDK session database at {agents_db}", ) await coordinator.restore(snap) + report_state = get_global_report_state() + if report_state is not None: + budget_stopped, reserve_stopped = recomputed_budget_flags( + report_state.get_total_llm_cost(), + max_budget_usd, + interactive=interactive, + ) + await coordinator.reset_budget_stops( + budget_stopped=budget_stopped, + reserve_stopped=reserve_stopped, + budget_paused=interactive and coordinator.budget_paused, + ) for aid, parent in coordinator.parent_of.items(): if parent is None: root_id = aid @@ -245,7 +258,14 @@ async def run_strix_scan( sandbox=SandboxRunConfig(client=bundle["client"], session=bundle["session"]), trace_include_sensitive_data=False, ) - hooks = ReportUsageHooks(model=resolved_model, max_budget_usd=max_budget_usd) + hooks = ReportUsageHooks( + model=resolved_model, + max_budget_usd=max_budget_usd, + max_turns=max_turns, + interactive=interactive, + ) + if interactive: + coordinator.set_budget_extender(hooks.extend_budget) scope_context = build_scope_context(scan_config) root_context = _merge_root_prompt_context(scope_context, extra_system_prompt_context) diff --git a/strix/interface/cli.py b/strix/interface/cli.py index f5079120..c03e8c86 100644 --- a/strix/interface/cli.py +++ b/strix/interface/cli.py @@ -13,6 +13,7 @@ from rich.panel import Panel from rich.text import Text from strix.config import load_settings +from strix.core.inputs import DEFAULT_MAX_TURNS from strix.core.runner import run_strix_scan from strix.report.state import ReportState, set_global_report_state from strix.runtime import session_manager @@ -184,6 +185,7 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915 local_sources=getattr(args, "local_sources", None) or [], interactive=bool(getattr(args, "interactive", False)), max_budget_usd=getattr(args, "max_budget_usd", None), + max_turns=getattr(args, "max_turns", DEFAULT_MAX_TURNS), ) finally: stop_updates.set() diff --git a/strix/interface/main.py b/strix/interface/main.py index 3791a914..bbb07ee8 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -31,6 +31,7 @@ from strix.config.models import ( is_known_openai_bare_model, is_recommended_or_frontier_model, ) +from strix.core.inputs import DEFAULT_MAX_TURNS from strix.core.paths import run_dir_for, runtime_state_dir from strix.interface.cli import run_cli from strix.interface.tui import run_tui @@ -481,6 +482,16 @@ def _positive_budget(value: str) -> float: return budget +def _positive_int(value: str) -> int: + try: + parsed = int(value) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"invalid int value: {value!r}") from exc + if parsed <= 0: + raise argparse.ArgumentTypeError("must be an integer greater than 0") + return parsed + + def parse_arguments() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Strix Multi-Agent Cybersecurity Penetration Testing Tool", @@ -636,10 +647,27 @@ Examples: ) parser.add_argument( - "--max-budget-usd", + "--max-budget", + dest="max_budget_usd", + metavar="USD", type=_positive_budget, default=None, - help="Maximum LLM cost in USD (> 0). The scan stops cleanly when this limit is reached.", + help=( + "Maximum LLM cost in USD (> 0). The scan stops cleanly when this limit is reached. " + "Graduated wrap-up warnings are sent to all agents as it is approached." + ), + ) + + parser.add_argument( + "--max-turns", + dest="max_turns", + metavar="N", + type=_positive_int, + default=DEFAULT_MAX_TURNS, + help=( + "Maximum turns per agent (> 0, default %(default)s). Each agent is force-stopped " + "when it reaches this limit, with graduated wrap-up warnings as it is approached." + ), ) parser.add_argument( diff --git a/strix/interface/tui/app.py b/strix/interface/tui/app.py index 28d3ef08..cc5a36c4 100644 --- a/strix/interface/tui/app.py +++ b/strix/interface/tui/app.py @@ -34,6 +34,7 @@ from textual.widgets.tree import TreeNode from strix.config import load_settings from strix.config.models import is_recommended_or_frontier_model from strix.core.hooks import BudgetExceededError +from strix.core.inputs import DEFAULT_MAX_TURNS 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 @@ -814,6 +815,7 @@ class StrixTUIApp(App): # type: ignore[misc] self._scan_completed = threading.Event() self._scan_error: BaseException | None = None self._error_noted_agents: set[str] = set() + self._budget_pause_notified = False self._spinner_frame_index: int = 0 self._sweep_num_squares: int = 6 @@ -1046,6 +1048,7 @@ class StrixTUIApp(App): # type: ignore[misc] self.live_view.record_agent_error(agent_id, error) else: self._error_noted_agents.discard(agent_id) + self._notify_budget_pause(statuses) if self._scan_loop is None or self._scan_loop.is_closed(): return @@ -1057,6 +1060,19 @@ class StrixTUIApp(App): # type: ignore[misc] self._agent_graph_sync_future = asyncio.run_coroutine_threadsafe(collect(), self._scan_loop) + def _notify_budget_pause(self, statuses: dict[str, Any]) -> None: + paused = any(status == "budget_paused" for status in statuses.values()) + if paused and not self._budget_pause_notified: + self._budget_pause_notified = True + self.notify( + "Budget limit reached \u2014 agents paused. Send a message to continue " + "(this extends the budget), or ctrl-q to quit.", + severity="warning", + timeout=15, + ) + elif not paused: + self._budget_pause_notified = False + def _update_agent_node(self, agent_id: str, agent_data: dict[str, Any]) -> bool: if agent_id not in self.agent_nodes: return False @@ -1069,6 +1085,7 @@ class StrixTUIApp(App): # type: ignore[misc] status_indicators = { "running": "⚪", "waiting": "⏸", + "budget_paused": "⏸", "completed": "🟢", "failed": "🔴", "crashed": "🔴", @@ -1266,10 +1283,17 @@ class StrixTUIApp(App): # type: ignore[misc] self._stop_dot_animation() return (text, Text(), False) - if status == "waiting": + if status in {"waiting", "budget_paused"}: text = Text() - text.append("Send message to resume", style="dim") - return (text, Text(), False) + keymap = Text() + if status == "budget_paused": + text.append("Budget limit reached", style="yellow") + text.append(" \u00b7 ", style="dim") + text.append("Send a message to continue", style="dim") + keymap = keymap_styled([("ctrl-q", "quit")]) + else: + text.append("Send message to resume", style="dim") + return (text, keymap, False) if status == "running": if self._agent_has_real_activity(agent_id): @@ -1494,6 +1518,7 @@ class StrixTUIApp(App): # type: ignore[misc] coordinator=self.coordinator, interactive=True, max_budget_usd=getattr(self.args, "max_budget_usd", None), + max_turns=getattr(self.args, "max_turns", DEFAULT_MAX_TURNS), event_sink=self._capture_sdk_event, ), ) @@ -1501,10 +1526,7 @@ class StrixTUIApp(App): # type: ignore[misc] 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") + logger.info("Scan stopped: --max-budget limit reached") except (ConnectionError, TimeoutError) as e: logging.exception("Network error during scan") self._scan_error = e @@ -1559,6 +1581,7 @@ class StrixTUIApp(App): # type: ignore[misc] status_indicators = { "running": "⚪", "waiting": "⏸", + "budget_paused": "⏸", "completed": "🟢", "failed": "🔴", "crashed": "🔴", @@ -1605,6 +1628,7 @@ class StrixTUIApp(App): # type: ignore[misc] status_indicators = { "running": "⚪", "waiting": "⏸", + "budget_paused": "⏸", "completed": "🟢", "failed": "🔴", "crashed": "🔴", @@ -1729,7 +1753,10 @@ class StrixTUIApp(App): # type: ignore[misc] message=message, ) if not submitted: - self.notify("Scan loop is not ready; message was not sent", severity="warning") + if self._scan_completed.is_set(): + self.notify("The scan has ended; message was not sent", severity="warning") + else: + self.notify("Scan loop is not ready; message was not sent", severity="warning") return self._displayed_events.clear() diff --git a/strix/tools/finish/tool.py b/strix/tools/finish/tool.py index 21064e71..90edc089 100644 --- a/strix/tools/finish/tool.py +++ b/strix/tools/finish/tool.py @@ -253,6 +253,8 @@ async def finish_scan( parent_id = inner.get("parent_id") if coordinator is not None and parent_id is None and me is not None: active_agents = await coordinator.active_agents_except(me) + if active_agents and coordinator.reserve_stopped: + active_agents = [] else: active_agents = [] diff --git a/tests/test_e2e_budget_lifecycle.py b/tests/test_e2e_budget_lifecycle.py new file mode 100644 index 00000000..f8b9021a --- /dev/null +++ b/tests/test_e2e_budget_lifecycle.py @@ -0,0 +1,357 @@ +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any +from unittest.mock import MagicMock, patch + +import pytest + +from strix.core import execution +from strix.core.agents import AgentCoordinator +from strix.core.execution import _start_child_runner, run_agent_loop +from strix.core.hooks import BudgetExceededError, ReportUsageHooks +from strix.core.sessions import open_agent_session + + +if TYPE_CHECKING: + from collections.abc import AsyncIterator, Callable + from pathlib import Path + + +MAX_BUDGET = 10.0 +COST_PER_CALL = 1.0 + + +class _FakeLedger: + def __init__(self) -> None: + self.cost = 0.0 + self.calls: list[str] = [] + + def record_sdk_usage(self, **_kwargs: Any) -> None: + return + + def get_total_llm_cost(self) -> float: + return self.cost + + +class _FakeStream: + def __init__( + self, + *, + ledger: _FakeLedger, + hooks: ReportUsageHooks, + context: dict[str, Any], + agent: Any, + ) -> None: + self._ledger = ledger + self._hooks = hooks + self._context = context + self._agent = agent + self.run_loop_exception: BaseException | None = None + self.final_output = None + + async def stream_events(self) -> AsyncIterator[Any]: + self._ledger.cost += COST_PER_CALL + self._ledger.calls.append(str(self._context.get("agent_id"))) + ctx_wrapper = MagicMock() + ctx_wrapper.context = self._context + try: + await self._hooks.on_llm_end(ctx_wrapper, self._agent, MagicMock()) + except Exception as exc: # noqa: BLE001 + self.run_loop_exception = exc + items: tuple[Any, ...] = () + for item in items: + yield item + + def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002 + return + + +def _fake_runner(ledger: _FakeLedger) -> Any: + class _FakeRunner: + @staticmethod + def run_streamed( + agent: Any, + input: Any, # noqa: A002, ARG004 + *, + run_config: Any, # noqa: ARG004 + context: dict[str, Any], + max_turns: int, # noqa: ARG004 + session: Any, # noqa: ARG004 + hooks: ReportUsageHooks, + ) -> _FakeStream: + return _FakeStream(ledger=ledger, hooks=hooks, context=context, agent=agent) + + return _FakeRunner + + +async def _noop_compact(*_args: Any, **_kwargs: Any) -> bool: + return False + + +async def _wait_until(predicate: Callable[[], bool], *, timeout: float = 5.0) -> None: + async def _poll() -> None: + while not predicate(): + await asyncio.sleep(0.01) + + await asyncio.wait_for(_poll(), timeout=timeout) + + +@pytest.mark.asyncio +async def test_full_budget_lifecycle_reserve_then_cap( # noqa: PLR0915 + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + ledger = _FakeLedger() + hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET) + monkeypatch.setattr(execution, "Runner", _fake_runner(ledger)) + monkeypatch.setattr(execution, "_compact_session", _noop_compact) + + coordinator = AgentCoordinator() + db_path = tmp_path / "agents.sqlite" + sessions: list[Any] = [] + run_config = MagicMock() + + await coordinator.register("root", "strix", parent_id=None) + root_session = open_agent_session("root", db_path) + sessions.append(root_session) + + root_exc: list[BaseException] = [] + + async def _root_loop() -> None: + try: + await run_agent_loop( + agent=MagicMock(), + initial_input=[], + run_config=run_config, + context={"agent_id": "root", "parent_id": None}, + max_turns=500, + coordinator=coordinator, + agent_id="root", + interactive=True, + session=root_session, + start_parked=True, + hooks=hooks, + ) + except BaseException as exc: + root_exc.append(exc) + raise + + with patch("strix.core.hooks.get_global_report_state", return_value=ledger): + root_task = asyncio.create_task(_root_loop()) + await asyncio.sleep(0.05) + + for child_id in ("child-a", "child-b"): + await coordinator.register(child_id, "recon", parent_id="root") + await _start_child_runner( + parent_ctx={"agent_id": "root", "parent_id": None}, + coordinator=coordinator, + agents_db_path=db_path, + sessions_to_close=sessions, + run_config=run_config, + max_turns=500, + interactive=True, + child_agent=MagicMock(), + child_id=child_id, + name=f"recon-{child_id}", + parent_id="root", + task="probe things", + initial_input=[], + hooks=hooks, + ) + await _wait_until(lambda: ledger.cost >= 2.0) + reserve_before = coordinator.reserve_stopped + assert reserve_before is False + + async def _wait_spend_above(amount: float) -> None: + await _wait_until(lambda: ledger.cost > amount) + + turn = 0 + while ledger.cost < MAX_BUDGET * 0.90 - 1e-9: + target = ("child-a", "child-b")[turn % 2] + spent_before = ledger.cost + assert await coordinator.send(target, {"from": "user", "content": "keep going"}) + await _wait_spend_above(spent_before) + turn += 1 + + await _wait_until(lambda: coordinator.reserve_stopped) + + await _wait_until( + lambda: ( + coordinator.statuses["child-a"] == "stopped" + and coordinator.statuses["child-b"] == "stopped" + ) + ) + + assert coordinator.reserve_stopped is True + + await _wait_until(lambda: coordinator.budget_stopped) + assert ledger.cost == pytest.approx(MAX_BUDGET) + + assert len(ledger.calls) == 10 + assert set(ledger.calls[:9]) == {"child-a", "child-b"} + assert ledger.calls[9] == "root" + + root_items = await root_session.get_items() + notices = [item for item in root_items if "Budget reserve" in str(item)] + assert len(notices) == 1 + + with pytest.raises(BudgetExceededError): + await root_task + assert root_exc and isinstance(root_exc[0], BudgetExceededError) + + assert {aid: str(status) for aid, status in coordinator.statuses.items()} == { + "root": "stopped", + "child-a": "stopped", + "child-b": "stopped", + } + assert coordinator.budget_stopped is True + assert coordinator.reserve_stopped is True + + for session in sessions: + session.close() + + +@pytest.mark.asyncio +async def test_respawned_children_after_reserve_never_spend( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + ledger = _FakeLedger() + ledger.cost = 9.5 + hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET) + monkeypatch.setattr(execution, "Runner", _fake_runner(ledger)) + monkeypatch.setattr(execution, "_compact_session", _noop_compact) + + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child-a", "recon", parent_id="root") + snap = await coordinator.snapshot() + snap["reserve_stopped"] = True + + restored = AgentCoordinator() + await restored.restore(snap) + assert restored.reserve_stopped is True + + sessions: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=ledger): + await _start_child_runner( + parent_ctx={"agent_id": "root", "parent_id": None}, + coordinator=restored, + agents_db_path=tmp_path / "agents.sqlite", + sessions_to_close=sessions, + run_config=MagicMock(), + max_turns=500, + interactive=True, + child_agent=MagicMock(), + child_id="child-a", + name="recon-child-a", + parent_id="root", + task="probe things", + initial_input=[], + hooks=hooks, + ) + await _wait_until(lambda: restored.statuses["child-a"] == "stopped") + + assert ledger.cost == pytest.approx(9.5) + assert ledger.calls == [] + for session in sessions: + session.close() + + +@pytest.mark.asyncio +async def test_resumed_parked_root_after_reserve_is_renotified_and_finalizes( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + ledger = _FakeLedger() + ledger.cost = 9.0 + hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET) + monkeypatch.setattr(execution, "Runner", _fake_runner(ledger)) + monkeypatch.setattr(execution, "_compact_session", _noop_compact) + + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.set_status("root", "waiting") + snap = await coordinator.snapshot() + snap["reserve_stopped"] = True + + restored = AgentCoordinator() + await restored.restore(snap) + assert restored.reserve_stopped is True + + root_session = open_agent_session("root", tmp_path / "agents.sqlite") + with patch("strix.core.hooks.get_global_report_state", return_value=ledger): + root_task = asyncio.create_task( + run_agent_loop( + agent=MagicMock(), + initial_input=[], + run_config=MagicMock(), + context={"agent_id": "root", "parent_id": None}, + max_turns=500, + coordinator=restored, + agent_id="root", + interactive=True, + session=root_session, + start_parked=True, + hooks=hooks, + ) + ) + with pytest.raises(BudgetExceededError): + await asyncio.wait_for(root_task, timeout=5.0) + + assert ledger.calls == ["root"] + assert ledger.cost == pytest.approx(MAX_BUDGET) + root_items = await root_session.get_items() + notices = [item for item in root_items if "Budget reserve" in str(item)] + assert len(notices) == 1 + root_session.close() + + +@pytest.mark.asyncio +async def test_interactive_budget_pause_then_user_message_extends_and_resumes( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + ledger = _FakeLedger() + ledger.cost = 9.0 + hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET, interactive=True) + monkeypatch.setattr(execution, "Runner", _fake_runner(ledger)) + monkeypatch.setattr(execution, "_compact_session", _noop_compact) + + coordinator = AgentCoordinator() + coordinator.set_budget_extender(hooks.extend_budget) + await coordinator.register("root", "strix", parent_id=None) + root_session = open_agent_session("root", tmp_path / "agents.sqlite") + + with patch("strix.core.hooks.get_global_report_state", return_value=ledger): + root_task = asyncio.create_task( + run_agent_loop( + agent=MagicMock(), + initial_input=[], + run_config=MagicMock(), + context={"agent_id": "root", "parent_id": None}, + max_turns=500, + coordinator=coordinator, + agent_id="root", + interactive=True, + session=root_session, + start_parked=True, + hooks=hooks, + ) + ) + await asyncio.sleep(0.05) + + assert await coordinator.send("root", {"from": "user", "content": "go"}) + await _wait_until(lambda: coordinator.budget_paused) + assert coordinator.statuses["root"] == "budget_paused" + assert ledger.cost == pytest.approx(MAX_BUDGET) + assert not root_task.done() + assert coordinator.budget_stopped is False + + assert await coordinator.send("root", {"from": "user", "content": "keep going"}) + await _wait_until(lambda: not coordinator.budget_paused) + await _wait_until(lambda: ledger.cost > MAX_BUDGET) + await _wait_until(lambda: coordinator.statuses["root"] == "waiting") + assert not root_task.done() + + root_task.cancel() + await root_task + + root_session.close() diff --git a/tests/test_execution.py b/tests/test_execution.py index 59a37e65..0350e470 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -3,10 +3,265 @@ from __future__ import annotations import asyncio +import contextlib +import json +from typing import Any import pytest +from agents.memory import SQLiteSession +from agents.tool_context import ToolContext from strix.core.agents import AgentCoordinator +from strix.core.execution import _notify_root_on_budget_reserve +from strix.tools.finish.tool import finish_scan + + +async def _call_finish_scan( + coordinator: AgentCoordinator, agent_id: str, parent_id: str | None +) -> dict[str, Any]: + ctx = ToolContext( + context={"coordinator": coordinator, "agent_id": agent_id, "parent_id": parent_id}, + tool_name="finish_scan", + tool_call_id="call-1", + tool_arguments="{}", + ) + fields = ("executive_summary", "methodology", "technical_analysis", "recommendations") + result: str = await finish_scan.on_invoke_tool(ctx, json.dumps(dict.fromkeys(fields, "x"))) + parsed: dict[str, Any] = json.loads(result) + return parsed + + +@pytest.mark.asyncio +async def test_reserve_stop_notifies_root_once(monkeypatch: pytest.MonkeyPatch) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child-a", "recon", parent_id="root") + await coordinator.register("child-b", "recon", parent_id="root") + + sent: list[tuple[str, dict[str, Any]]] = [] + + async def _record(target_agent_id: str, message: dict[str, Any]) -> bool: + sent.append((target_agent_id, message)) + return True + + monkeypatch.setattr(coordinator, "send", _record) + + await _notify_root_on_budget_reserve(coordinator) + await _notify_root_on_budget_reserve(coordinator) + + assert len(sent) == 1 + target, message = sent[0] + assert target == "root" + assert message["type"] == "budget_reserve_stop" + assert "finish_scan" in str(message["content"]) + + +@pytest.mark.asyncio +async def test_concurrent_reserve_claims_yield_single_root() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + for i in range(12): + await coordinator.register(f"child-{i}", "recon", parent_id="root") + + results = await asyncio.gather(*(coordinator.claim_reserve_notification() for _ in range(12))) + + assert results.count("root") == 1 + assert all(r is None for r in results if r != "root") + + +@pytest.mark.asyncio +async def test_claim_reserve_sets_flag_and_wakes_parked_agents() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + + flag_before = coordinator.reserve_stopped + assert flag_before is False + waiter = asyncio.create_task(coordinator.wait_for_message("child")) + await asyncio.sleep(0) + assert not waiter.done() + + await coordinator.claim_reserve_notification() + + flag_after = coordinator.reserve_stopped + assert flag_after is True + await asyncio.wait_for(waiter, timeout=1.0) + + +@pytest.mark.asyncio +async def test_finish_scan_bypasses_active_agent_guard_after_reserve() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + await coordinator.set_status("child", "running") + + blocked = await _call_finish_scan(coordinator, "root", None) + assert blocked["scan_completed"] is False + assert blocked["active_agents"] + + await coordinator.claim_reserve_notification() + + finished = await _call_finish_scan(coordinator, "root", None) + assert finished["scan_completed"] is True + assert coordinator.statuses["root"] == "completed" + + +@pytest.mark.asyncio +async def test_finish_scan_gate_ignores_sub_agent_caller() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + await coordinator.set_status("child", "running") + + result = await _call_finish_scan(coordinator, "child", "root") + assert "active_agents" not in result + assert result["success"] is False + assert "root" in result["error"] + + +@pytest.mark.asyncio +async def test_reserve_stop_notify_noop_without_root(monkeypatch: pytest.MonkeyPatch) -> None: + coordinator = AgentCoordinator() + await coordinator.register("child", "recon", parent_id="missing") + + sent: list[tuple[str, dict[str, Any]]] = [] + + async def _record(target_agent_id: str, message: dict[str, Any]) -> bool: + sent.append((target_agent_id, message)) + return True + + monkeypatch.setattr(coordinator, "send", _record) + await _notify_root_on_budget_reserve(coordinator) + + assert sent == [] + + +@pytest.mark.asyncio +async def test_snapshot_round_trip_preserves_stop_flags() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.trigger_budget_stop() + await coordinator.claim_reserve_notification() + + snap = await coordinator.snapshot() + assert snap["budget_stopped"] is True + assert snap["reserve_stopped"] is True + + restored = AgentCoordinator() + await restored.restore(snap) + assert restored.budget_stopped is True + assert restored.reserve_stopped is True + + +@pytest.mark.asyncio +async def test_legacy_snapshot_without_stop_flags_defaults_to_false() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + snap = await coordinator.snapshot() + del snap["budget_stopped"] + del snap["reserve_stopped"] + + restored = AgentCoordinator() + await restored.restore(snap) + assert restored.budget_stopped is False + assert restored.reserve_stopped is False + + +@pytest.mark.asyncio +async def test_randomized_reserve_claim_race_many_interleavings() -> None: + for seed in range(25): + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + child_ids = [f"child-{i}" for i in range(8)] + for child_id in child_ids: + await coordinator.register(child_id, "recon", parent_id="root") + + waiters = [asyncio.create_task(coordinator.wait_for_message(cid)) for cid in child_ids] + await asyncio.sleep(0) + + async def _claim(delay: float, coord: AgentCoordinator = coordinator) -> str | None: + await asyncio.sleep(delay) + return await coord.claim_reserve_notification() + + delays = [((seed * 31 + i * 17) % 50) / 10_000 for i in range(len(child_ids))] + results = await asyncio.gather(*(_claim(delay) for delay in delays)) + + assert results.count("root") == 1, f"seed {seed}: expected exactly one winner" + await asyncio.wait_for(asyncio.gather(*waiters), timeout=1.0) + assert coordinator.reserve_stopped is True + + +@pytest.mark.asyncio +async def test_reserve_claim_never_loses_root_wake() -> None: + for _ in range(10): + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + + root_waiter = asyncio.create_task(coordinator.wait_for_message("root")) + await asyncio.sleep(0) + assert not root_waiter.done() + + await coordinator.claim_reserve_notification() + await asyncio.sleep(0) + + async with coordinator._lock: + coordinator.pending_counts["root"] = 1 + coordinator.runtimes["root"].wake.set() + + await asyncio.wait_for(root_waiter, timeout=1.0) + + +@pytest.mark.asyncio +async def test_budget_stop_takes_precedence_over_reserve_for_all_roles() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + await coordinator.claim_reserve_notification() + await coordinator.trigger_budget_stop() + + await asyncio.wait_for(coordinator.wait_for_message("root"), timeout=1.0) + await asyncio.wait_for(coordinator.wait_for_message("child"), timeout=1.0) + assert coordinator.budget_stopped is True + assert coordinator.reserve_stopped is True + + +@pytest.mark.asyncio +async def test_root_not_released_by_reserve_alone() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + + await coordinator.claim_reserve_notification() + + root_waiter = asyncio.create_task(coordinator.wait_for_message("root")) + await asyncio.sleep(0.02) + assert not root_waiter.done() + + root_waiter.cancel() + with contextlib.suppress(asyncio.CancelledError): + await root_waiter + + +@pytest.mark.asyncio +async def test_snapshot_during_concurrent_claims_is_consistent() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + for i in range(6): + await coordinator.register(f"child-{i}", "recon", parent_id="root") + + claims = [asyncio.create_task(coordinator.claim_reserve_notification()) for _ in range(6)] + snap = await coordinator.snapshot() + await asyncio.gather(*claims) + + assert isinstance(snap["reserve_stopped"], bool) + final_snap = await coordinator.snapshot() + assert final_snap["reserve_stopped"] is True + + restored = AgentCoordinator() + await restored.restore(final_snap) + assert restored.reserve_stopped is True + assert await restored.claim_reserve_notification() is None @pytest.mark.asyncio @@ -42,3 +297,133 @@ async def test_wait_for_message_returns_immediately_after_budget_stop() -> None: # No pending messages, but the stop flag short-circuits the wait. await asyncio.wait_for(coordinator.wait_for_message("agent"), timeout=1.0) + + +@pytest.mark.asyncio +async def test_pause_for_budget_sets_flag_and_status() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + await coordinator.pause_for_budget("root") + assert coordinator.budget_paused is True + assert coordinator.statuses["root"] == "budget_paused" + + +@pytest.mark.asyncio +async def test_resume_from_budget_pause_extends_and_nudges( + monkeypatch: pytest.MonkeyPatch, +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child-a", "recon", parent_id="root") + await coordinator.register("child-b", "recon", parent_id="root") + await coordinator.pause_for_budget("root") + await coordinator.pause_for_budget("child-a") + await coordinator.pause_for_budget("child-b") + + extensions: list[int] = [] + coordinator.set_budget_extender(lambda: extensions.append(1)) + + sent: list[tuple[str, dict[str, Any]]] = [] + + async def _record(target_agent_id: str, message: dict[str, Any]) -> bool: + sent.append((target_agent_id, message)) + return True + + monkeypatch.setattr(coordinator, "send", _record) + + await coordinator.resume_from_budget_pause(exclude="root") + + assert coordinator.budget_paused is False + assert len(extensions) == 1 + assert all(coordinator.statuses[aid] == "waiting" for aid in ("root", "child-a", "child-b")) + assert sorted(target for target, _ in sent) == ["child-a", "child-b"] + assert all(message["type"] == "budget_extended" for _, message in sent) + + await coordinator.resume_from_budget_pause(exclude="root") + assert len(extensions) == 1 + + +@pytest.mark.asyncio +async def test_user_send_resumes_budget_pause(tmp_path: Any) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + session = SQLiteSession("root", tmp_path / "agents.db") + await coordinator.attach_runtime("root", session=session) + await coordinator.pause_for_budget("root") + + extensions: list[int] = [] + coordinator.set_budget_extender(lambda: extensions.append(1)) + + delivered = await coordinator.send("root", {"from": "user", "content": "keep going"}) + + assert delivered is True + assert coordinator.budget_paused is False + assert len(extensions) == 1 + assert coordinator.statuses["root"] == "waiting" + assert coordinator.pending_counts["root"] == 1 + session.close() + + +@pytest.mark.asyncio +async def test_non_user_send_does_not_resume_budget_pause(tmp_path: Any) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + session = SQLiteSession("root", tmp_path / "agents.db") + await coordinator.attach_runtime("root", session=session) + await coordinator.pause_for_budget("root") + + extensions: list[int] = [] + coordinator.set_budget_extender(lambda: extensions.append(1)) + + await coordinator.send("root", {"from": "system", "content": "status"}) + + assert coordinator.budget_paused is True + assert extensions == [] + assert coordinator.statuses["root"] == "budget_paused" + session.close() + + +@pytest.mark.asyncio +async def test_reset_budget_stops_clears_pause_and_normalizes_statuses() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.trigger_budget_stop() + await coordinator.claim_reserve_notification() + await coordinator.pause_for_budget("root") + + await coordinator.reset_budget_stops(budget_stopped=False, reserve_stopped=False) + + assert coordinator.budget_stopped is False + assert coordinator.reserve_stopped is False + assert coordinator.budget_paused is False + assert coordinator.statuses["root"] == "waiting" + + +@pytest.mark.asyncio +async def test_reset_budget_stops_can_preserve_pause() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.pause_for_budget("root") + + await coordinator.reset_budget_stops( + budget_stopped=False, reserve_stopped=False, budget_paused=True + ) + + assert coordinator.budget_paused is True + assert coordinator.statuses["root"] == "budget_paused" + + +@pytest.mark.asyncio +async def test_snapshot_round_trip_preserves_budget_pause() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.pause_for_budget("root") + + snap = await coordinator.snapshot() + assert snap["budget_paused"] is True + + restored = AgentCoordinator() + await restored.restore(snap) + assert restored.budget_paused is True + assert restored.statuses["root"] == "budget_paused" diff --git a/tests/test_hooks.py b/tests/test_hooks.py index 5dcbed4f..16fce426 100644 --- a/tests/test_hooks.py +++ b/tests/test_hooks.py @@ -2,11 +2,18 @@ from __future__ import annotations +from typing import Any from unittest.mock import MagicMock, patch import pytest -from strix.core.hooks import BudgetExceededError, ReportUsageHooks +from strix.core.hooks import ( + BudgetExceededError, + BudgetPausedError, + ReportUsageHooks, + SubagentBudgetReservedError, + recomputed_budget_flags, +) def _make_hooks(max_budget: float | None) -> ReportUsageHooks: @@ -20,9 +27,22 @@ def _make_report_state(cost: float) -> MagicMock: return state -def _make_context(agent_id: str = "test-agent") -> MagicMock: +def _make_context(agent_id: str = "test-agent", parent_id: str | None = None) -> MagicMock: ctx: MagicMock = MagicMock() - ctx.context = {"agent_id": agent_id} + ctx.context = {"agent_id": agent_id, "parent_id": parent_id} + return ctx + + +def _make_warn_context( + *, + requests: int, + parent_id: str | None = None, + agent_id: str = "test-agent", +) -> MagicMock: + ctx: MagicMock = MagicMock() + ctx.context = {"agent_id": agent_id, "parent_id": parent_id} + ctx.usage = MagicMock() + ctx.usage.requests = requests return ctx @@ -89,6 +109,127 @@ async def test_error_message_includes_amounts() -> None: assert "7.1234" in str(exc_info.value) +@pytest.mark.asyncio +async def test_subagent_stops_at_reserve() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(9.0) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(SubagentBudgetReservedError), + ): + await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_subagent_below_reserve_does_not_raise() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(8.99) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_subagent_overshoot_to_full_budget_triggers_scan_wide_stop() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(10.5) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetExceededError), + ): + await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_root_keeps_running_inside_reserve() -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(9.5) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_root_hard_stop_stays_at_full_budget() -> 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(parent_id=None), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_budget_warning_mentions_reserve() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0) + state = _make_report_state(7.5) + root_items: list[Any] = [] + sub_items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items + ) + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items + ) + assert "stopped at 90%" in root_items[0]["content"] + assert "stopped at 90%" in sub_items[0]["content"] + assert "root agent's final report" in sub_items[0]["content"] + + +@pytest.mark.asyncio +async def test_subagent_critical_budget_warning_reachable_before_reserve() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0) + state = _make_report_state(8.6) + sub_items: list[Any] = [] + root_items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items + ) + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items + ) + assert "[CRITICAL]" in sub_items[0]["content"] + assert "[URGENT]" in root_items[0]["content"] + + +@pytest.mark.parametrize( + ("parent_id", "cost", "expected"), + [ + ("root-1", 0.0, None), + ("root-1", 8.9999, None), + ("root-1", 9.0, SubagentBudgetReservedError), + ("root-1", 9.0001, SubagentBudgetReservedError), + ("root-1", 9.5, SubagentBudgetReservedError), + ("root-1", 9.9999, SubagentBudgetReservedError), + ("root-1", 10.0, BudgetExceededError), + ("root-1", 10.0001, BudgetExceededError), + ("root-1", 25.0, BudgetExceededError), + (None, 0.0, None), + (None, 8.9999, None), + (None, 9.0, None), + (None, 9.5, None), + (None, 9.9999, None), + (None, 10.0, BudgetExceededError), + (None, 10.0001, BudgetExceededError), + (None, 25.0, BudgetExceededError), + ], +) +@pytest.mark.asyncio +async def test_budget_enforcement_decision_table( + parent_id: str | None, cost: float, expected: type[Exception] | None +) -> None: + hooks = _make_hooks(10.0) + state = _make_report_state(cost) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + if expected is None: + await hooks.on_llm_end(_make_context(parent_id=parent_id), MagicMock(), MagicMock()) + else: + with pytest.raises(expected): + await hooks.on_llm_end(_make_context(parent_id=parent_id), MagicMock(), MagicMock()) + state.record_sdk_usage.assert_called_once() + + @pytest.mark.asyncio async def test_no_raise_when_report_state_none() -> None: hooks = _make_hooks(1.0) @@ -106,3 +247,247 @@ def test_non_positive_budget_rejected(bad_budget: float) -> None: def test_budget_exceeded_error_is_runtime_error() -> None: err = BudgetExceededError("test") assert isinstance(err, RuntimeError) + + +def test_non_positive_max_turns_rejected() -> None: + with pytest.raises(ValueError, match="positive integer"): + ReportUsageHooks(model="test-model", max_turns=0) + + +@pytest.mark.asyncio +async def test_no_turn_warning_below_first_band() -> None: + hooks = ReportUsageHooks(model="test-model", max_turns=100) + items: list[Any] = [] + await hooks.on_llm_start(_make_warn_context(requests=68), MagicMock(), None, items) + assert items == [] + + +@pytest.mark.asyncio +async def test_turn_warning_notice_band() -> None: + hooks = ReportUsageHooks(model="test-model", max_turns=100) + items: list[Any] = [] + await hooks.on_llm_start(_make_warn_context(requests=69), MagicMock(), None, items) + assert len(items) == 1 + content = items[0]["content"] + assert "[NOTICE]" in content + assert "finish_scan" in content + + +@pytest.mark.asyncio +async def test_turn_warning_escalates_and_names_subagent_tool() -> None: + hooks = ReportUsageHooks(model="test-model", max_turns=100) + items: list[Any] = [] + await hooks.on_llm_start( + _make_warn_context(requests=95, parent_id="root-1"), MagicMock(), None, items + ) + assert len(items) == 1 + content = items[0]["content"] + assert "[CRITICAL]" in content + assert "agent_finish" in content + + +@pytest.mark.asyncio +async def test_turn_warning_root_directive_distinct_from_subagent() -> None: + hooks = ReportUsageHooks(model="test-model", max_turns=100) + + root_items: list[Any] = [] + await hooks.on_llm_start( + _make_warn_context(requests=85, parent_id=None), MagicMock(), None, root_items + ) + root = root_items[0]["content"] + + sub_items: list[Any] = [] + await hooks.on_llm_start( + _make_warn_context(requests=85, parent_id="root-1"), MagicMock(), None, sub_items + ) + sub = sub_items[0]["content"] + + assert root != sub + assert "root agent" in root + assert "finish_scan" in root + assert "agent_finish" not in root + assert "whole scan" in root + assert "sub-agent" in sub + assert "agent_finish" in sub + assert "finish_scan" not in sub + assert "confirmed" in sub + + +@pytest.mark.asyncio +async def test_budget_warning_root_directive_distinct_from_subagent() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0) + state = _make_report_state(8.6) + + root_items: list[Any] = [] + sub_items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items + ) + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items + ) + + root = root_items[0]["content"] + sub = sub_items[0]["content"] + assert "finish_scan" in root and "agent_finish" not in root + assert "agent_finish" in sub and "finish_scan" not in sub + assert "confirmed" in sub + + +@pytest.mark.parametrize("parent_id", [None, "root-1"]) +@pytest.mark.asyncio +async def test_turn_warning_directive_escalates_per_stage(parent_id: str | None) -> None: + hooks = ReportUsageHooks(model="test-model", max_turns=100) + contents: dict[str, str] = {} + for label, requests in (("notice", 69), ("urgent", 85), ("critical", 95)): + items: list[Any] = [] + await hooks.on_llm_start( + _make_warn_context(requests=requests, parent_id=parent_id), MagicMock(), None, items + ) + contents[label] = items[0]["content"] + + assert len({contents["notice"], contents["urgent"], contents["critical"]}) == 3 + assert "[NOTICE]" in contents["notice"] and "begin planning" in contents["notice"] + assert "[URGENT]" in contents["urgent"] and "prioritize" in contents["urgent"] + assert "[CRITICAL]" in contents["critical"] and "STOP" in contents["critical"] + + +@pytest.mark.asyncio +async def test_no_turn_warning_when_max_turns_unset() -> None: + hooks = ReportUsageHooks(model="test-model") + items: list[Any] = [] + await hooks.on_llm_start(_make_warn_context(requests=999), MagicMock(), None, items) + assert items == [] + + +@pytest.mark.asyncio +async def test_no_budget_warning_below_first_band() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0) + state = _make_report_state(6.9) + items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start(_make_warn_context(requests=0), MagicMock(), None, items) + assert items == [] + + +@pytest.mark.asyncio +async def test_budget_warning_broadcast_content() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0) + state = _make_report_state(9.6) + items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start(_make_warn_context(requests=0), MagicMock(), None, items) + assert len(items) == 1 + content = items[0]["content"] + assert "[CRITICAL]" in content + assert "shared across every agent" in content + + +@pytest.mark.asyncio +async def test_turn_and_budget_warnings_stack() -> None: + hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0, max_turns=100) + state = _make_report_state(8.6) + items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start(_make_warn_context(requests=89), MagicMock(), None, items) + assert len(items) == 2 + joined = " ".join(i["content"] for i in items) + assert "Turn budget" in joined + assert "cost budget" in joined + + +def _make_interactive_hooks(max_budget: float | None) -> ReportUsageHooks: + return ReportUsageHooks(model="test-model", max_budget_usd=max_budget, interactive=True) + + +@pytest.mark.asyncio +async def test_interactive_at_budget_pauses_instead_of_stopping() -> None: + hooks = _make_interactive_hooks(10.0) + state = _make_report_state(10.0) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetPausedError), + ): + await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_interactive_subagent_has_no_reserve() -> None: + hooks = _make_interactive_hooks(10.0) + state = _make_report_state(9.5) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_interactive_subagent_pauses_at_full_budget() -> None: + hooks = _make_interactive_hooks(10.0) + state = _make_report_state(10.5) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetPausedError), + ): + await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_extend_budget_lifts_the_pause() -> None: + hooks = _make_interactive_hooks(10.0) + state = _make_report_state(10.5) + hooks.extend_budget() + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_extend_budget_adds_original_amount_each_time() -> None: + hooks = _make_interactive_hooks(10.0) + hooks.extend_budget() + hooks.extend_budget() + state = _make_report_state(29.9) + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock()) + state = _make_report_state(30.0) + with ( + patch("strix.core.hooks.get_global_report_state", return_value=state), + pytest.raises(BudgetPausedError), + ): + await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock()) + + +@pytest.mark.asyncio +async def test_interactive_subagent_uses_root_warning_bands() -> None: + hooks = _make_interactive_hooks(10.0) + state = _make_report_state(7.4) + items: list[Any] = [] + with patch("strix.core.hooks.get_global_report_state", return_value=state): + await hooks.on_llm_start( + _make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, items + ) + assert len(items) == 1 + content = items[0]["content"] + assert "[NOTICE]" in content + assert "paused until the user chooses to continue" in content + assert "reserve" not in content.lower() + + +@pytest.mark.parametrize( + ("cost", "max_budget", "interactive", "expected"), + [ + (0.0, None, False, (False, False)), + (100.0, None, False, (False, False)), + (5.0, 10.0, False, (False, False)), + (9.0, 10.0, False, (False, True)), + (10.0, 10.0, False, (True, True)), + (10.0, 20.0, False, (False, False)), + (10.0, 10.0, True, (False, False)), + ], +) +def test_recomputed_budget_flags( + cost: float, + max_budget: float | None, + interactive: bool, + expected: tuple[bool, bool], +) -> None: + assert recomputed_budget_flags(cost, max_budget, interactive=interactive) == expected