mirror of
https://github.com/usestrix/strix.git
synced 2026-10-01 02:03:55 +00:00
feat(budget): budget_policy=pause parks every agent at the limit until the operator resumes
Adds budget_policy: stop | pause to run_strix_scan / ReportUsageHooks / AgentCoordinator, independent of interactive mode. Under pause the agents get no budget warnings and no sub-agent reserve; each agent parks before its next LLM call once spent >= limit or the scan is paused, sessions and sandbox stay alive, and coordinator.resume_budget(max_budget_usd=...) replaces the limit and wakes every parked agent without adding anything to any session. coordinator.pause_budget() parks a running scan the same way. In-flight calls are never cancelled, so spent may end above the limit. Parked agents count as active for stop_agent.
This commit is contained in:
parent
0ff9f8c324
commit
d355838ea0
6 changed files with 751 additions and 41 deletions
|
|
@ -24,6 +24,8 @@ logger = logging.getLogger(__name__)
|
|||
|
||||
Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed", "budget_paused"]
|
||||
|
||||
BudgetPolicy = Literal["stop", "pause"]
|
||||
|
||||
TERMINAL_STATUSES: frozenset[str] = frozenset({"completed", "stopped", "crashed", "failed"})
|
||||
|
||||
# Why an agent parked. The user can message any agent, so this - not the agent's
|
||||
|
|
@ -68,7 +70,10 @@ class AgentCoordinator:
|
|||
self._budget_stopped = False
|
||||
self._reserve_stopped = False
|
||||
self._budget_paused = False
|
||||
self._resume_epoch = 0
|
||||
self._budget_policy: BudgetPolicy = "stop"
|
||||
self._extend_budget: Callable[[], None] | None = None
|
||||
self._set_budget_limit: Callable[[float | None], None] | None = None
|
||||
|
||||
def set_snapshot_path(self, path: Path) -> None:
|
||||
self._snapshot_path = path
|
||||
|
|
@ -95,15 +100,89 @@ class AgentCoordinator:
|
|||
def budget_paused(self) -> bool:
|
||||
return self._budget_paused
|
||||
|
||||
@property
|
||||
def resume_epoch(self) -> int:
|
||||
"""Bumped by every ``resume_budget``; an agent parks against the value it read."""
|
||||
return self._resume_epoch
|
||||
|
||||
@property
|
||||
def budget_policy(self) -> BudgetPolicy:
|
||||
return self._budget_policy
|
||||
|
||||
def set_budget_policy(self, policy: BudgetPolicy) -> None:
|
||||
self._budget_policy = policy
|
||||
|
||||
def set_budget_extender(self, extend: Callable[[], None]) -> None:
|
||||
self._extend_budget = extend
|
||||
|
||||
def set_budget_limit_setter(self, setter: Callable[[float | None], None]) -> None:
|
||||
self._set_budget_limit = setter
|
||||
|
||||
async def pause_for_budget(self, agent_id: str) -> None:
|
||||
async with self._lock:
|
||||
self._budget_paused = True
|
||||
await self.set_status(agent_id, "budget_paused")
|
||||
|
||||
async def park_for_budget(self, agent_id: str) -> None:
|
||||
"""Record that ``agent_id`` parked before an LLM call (pause policy)."""
|
||||
await self.set_status(agent_id, "budget_paused")
|
||||
|
||||
async def pause_budget(self) -> None:
|
||||
"""Operator pause: every agent parks before its next LLM call.
|
||||
|
||||
Agents mid-call or mid-tool finish that step first, so their spend still
|
||||
lands; nothing is cancelled.
|
||||
"""
|
||||
async with self._lock:
|
||||
self._budget_paused = True
|
||||
logger.info("scan paused by the operator")
|
||||
await self._maybe_snapshot()
|
||||
|
||||
async def resume_budget(self, *, max_budget_usd: float | None = None) -> list[str]:
|
||||
"""Lift the pause and wake every parked agent; returns the woken agent ids.
|
||||
|
||||
With ``max_budget_usd`` the scan's limit is replaced first (``None`` keeps
|
||||
the current one). Agents continue with the LLM call they parked on; no
|
||||
message is added to any session. An agent that parks again on its next
|
||||
call (the new limit is already spent) is not an error.
|
||||
"""
|
||||
if max_budget_usd is not None and self._set_budget_limit is not None:
|
||||
self._set_budget_limit(max_budget_usd)
|
||||
async with self._lock:
|
||||
self._budget_paused = False
|
||||
self._resume_epoch += 1
|
||||
woken = [aid for aid, status in self.statuses.items() if status == "budget_paused"]
|
||||
for aid in woken:
|
||||
self.runtimes.setdefault(aid, AgentRuntime()).wake.set()
|
||||
logger.info("scan resumed; woke %d parked agent(s)", len(woken))
|
||||
await self._maybe_snapshot()
|
||||
return woken
|
||||
|
||||
async def wait_for_budget_resume(self, agent_id: str, *, parked_epoch: int) -> None:
|
||||
"""Block until a resume newer than ``parked_epoch``, a scan-wide stop, or
|
||||
the agent itself being stopped while parked.
|
||||
|
||||
``parked_epoch`` is the ``resume_epoch`` the agent read when it decided to
|
||||
park, so a resume that lands between that decision and this wait is not
|
||||
missed. ``agent_id`` is ``running`` again on return unless it was stopped.
|
||||
"""
|
||||
while True:
|
||||
async with self._lock:
|
||||
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
|
||||
if (
|
||||
self._budget_stopped
|
||||
or self._resume_epoch != parked_epoch
|
||||
or self.statuses.get(agent_id) != "budget_paused"
|
||||
):
|
||||
break
|
||||
wake = runtime.wake
|
||||
wake.clear()
|
||||
await wake.wait()
|
||||
if not self._budget_stopped and self.statuses.get(agent_id) == "budget_paused":
|
||||
await self.set_status(agent_id, "running")
|
||||
|
||||
async def resume_from_budget_pause(self, *, exclude: str | None = None) -> None:
|
||||
"""Legacy interactive resume: extend by the original budget and nudge agents."""
|
||||
async with self._lock:
|
||||
if not self._budget_paused:
|
||||
return
|
||||
|
|
@ -308,7 +387,7 @@ class AgentCoordinator:
|
|||
unknown, or it is terminal and its loop does not park for wake-ups.
|
||||
"""
|
||||
from_user = message.get("from") == "user"
|
||||
if from_user and self._budget_paused:
|
||||
if from_user and self._budget_paused and self._budget_policy != "pause":
|
||||
await self.resume_from_budget_pause(exclude=target_agent_id)
|
||||
async with self._lock:
|
||||
if target_agent_id not in self.statuses:
|
||||
|
|
|
|||
|
|
@ -522,33 +522,51 @@ async def _run_until_lifecycle(
|
|||
await coordinator.set_status(agent_id, "stopped")
|
||||
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
|
||||
|
||||
if interactive:
|
||||
result = await _run_cycle_parked(
|
||||
agent,
|
||||
coordinator,
|
||||
agent_id,
|
||||
input_data=input_data,
|
||||
run_config=run_config,
|
||||
context=context,
|
||||
max_turns=max_turns,
|
||||
session=session,
|
||||
event_sink=event_sink,
|
||||
hooks=hooks,
|
||||
)
|
||||
else:
|
||||
result = await _run_cycle(
|
||||
agent,
|
||||
coordinator,
|
||||
agent_id,
|
||||
input_data=input_data,
|
||||
run_config=run_config,
|
||||
context=context,
|
||||
max_turns=max_turns,
|
||||
session=session,
|
||||
interactive=False,
|
||||
event_sink=event_sink,
|
||||
hooks=hooks,
|
||||
)
|
||||
try:
|
||||
if interactive:
|
||||
result = await _run_cycle_parked(
|
||||
agent,
|
||||
coordinator,
|
||||
agent_id,
|
||||
input_data=input_data,
|
||||
run_config=run_config,
|
||||
context=context,
|
||||
max_turns=max_turns,
|
||||
session=session,
|
||||
event_sink=event_sink,
|
||||
hooks=hooks,
|
||||
)
|
||||
else:
|
||||
result = await _run_cycle(
|
||||
agent,
|
||||
coordinator,
|
||||
agent_id,
|
||||
input_data=input_data,
|
||||
run_config=run_config,
|
||||
context=context,
|
||||
max_turns=max_turns,
|
||||
session=session,
|
||||
interactive=False,
|
||||
event_sink=event_sink,
|
||||
hooks=hooks,
|
||||
)
|
||||
except BudgetPausedError as exc:
|
||||
if coordinator.budget_policy != "pause":
|
||||
raise
|
||||
# The agent parked right before an LLM call; everything up to that
|
||||
# point is already in its session. Once resumed, the same call goes
|
||||
# out with nothing added to the conversation.
|
||||
await coordinator.wait_for_budget_resume(agent_id, parked_epoch=exc.resume_epoch)
|
||||
if (
|
||||
not coordinator.budget_stopped
|
||||
and await _agent_status(coordinator, agent_id) != "running"
|
||||
):
|
||||
# Stopped while parked (operator stop or a parent's stop_agent).
|
||||
await coordinator.reset_recovery(agent_id)
|
||||
return result
|
||||
if session is not None:
|
||||
input_data = []
|
||||
continue
|
||||
|
||||
status = await _agent_status(coordinator, agent_id)
|
||||
if status != "running":
|
||||
|
|
@ -759,7 +777,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||
await coordinator.detach_stream(agent_id, stream)
|
||||
except BudgetPausedError as exc:
|
||||
logger.info("agent %s paused at the scan budget limit: %s", agent_id, exc)
|
||||
await coordinator.pause_for_budget(agent_id)
|
||||
if coordinator.budget_policy == "pause":
|
||||
await coordinator.park_for_budget(agent_id)
|
||||
else:
|
||||
await coordinator.pause_for_budget(agent_id)
|
||||
raise
|
||||
except SubagentBudgetReservedError as exc:
|
||||
logger.info("sub-agent %s stopped at the budget reserve: %s", agent_id, exc)
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any
|
|||
|
||||
from agents.lifecycle import RunHooks
|
||||
|
||||
from strix.core.agents import BudgetPolicy, coordinator_from_context
|
||||
from strix.report.state import get_global_report_state
|
||||
|
||||
|
||||
|
|
@ -22,6 +23,26 @@ logger = logging.getLogger(__name__)
|
|||
|
||||
LLM_TURN_KEY = "llm_turn"
|
||||
|
||||
# ``BudgetPolicy`` decides what happens when the accumulated LLM cost reaches
|
||||
# ``max_budget_usd``.
|
||||
#
|
||||
# ``stop``: the agents are warned as the limit approaches, sub-agents are cut at
|
||||
# a reserve so the root can write its report, and the scan ends at the limit.
|
||||
#
|
||||
# ``pause``: the agents are never told a limit exists. Every agent parks right
|
||||
# before its next LLM call once the limit is reached (or an operator pauses the
|
||||
# scan), keeping its session, context and sandbox alive, and continues with that
|
||||
# same call when the operator raises the limit or resumes.
|
||||
__all__ = [
|
||||
"LLM_TURN_KEY",
|
||||
"BudgetExceededError",
|
||||
"BudgetPausedError",
|
||||
"BudgetPolicy",
|
||||
"ReportUsageHooks",
|
||||
"SubagentBudgetReservedError",
|
||||
"recomputed_budget_flags",
|
||||
]
|
||||
|
||||
_STAGE_LABELS: tuple[str, ...] = ("NOTICE", "URGENT", "CRITICAL")
|
||||
_TURN_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
|
||||
_ROOT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
|
||||
|
|
@ -38,7 +59,20 @@ class SubagentBudgetReservedError(RuntimeError):
|
|||
|
||||
|
||||
class BudgetPausedError(RuntimeError):
|
||||
"""Raised to park one agent when an interactive scan reaches its budget."""
|
||||
"""Raised to park one agent until the scan budget is raised or the pause lifted.
|
||||
|
||||
``resume_epoch`` is the coordinator's ``resume_epoch`` at the moment the agent
|
||||
decided to park; the agent waits for a resume newer than that.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, *, resume_epoch: int = 0) -> None:
|
||||
super().__init__(message)
|
||||
self.resume_epoch = resume_epoch
|
||||
|
||||
|
||||
def _validate_budget(max_budget_usd: float | None) -> None:
|
||||
if max_budget_usd is not None and (not math.isfinite(max_budget_usd) or max_budget_usd <= 0):
|
||||
raise ValueError("max_budget_usd must be a finite number greater than 0")
|
||||
|
||||
|
||||
def recomputed_budget_flags(
|
||||
|
|
@ -46,11 +80,12 @@ def recomputed_budget_flags(
|
|||
max_budget_usd: float | None,
|
||||
*,
|
||||
interactive: bool,
|
||||
budget_policy: BudgetPolicy = "stop",
|
||||
) -> tuple[bool, bool]:
|
||||
"""Return the (budget_stopped, reserve_stopped) flags a resumed scan should carry."""
|
||||
if max_budget_usd is None:
|
||||
return False, False
|
||||
if interactive:
|
||||
if interactive or budget_policy == "pause":
|
||||
return False, False
|
||||
budget_stopped = cost >= max_budget_usd
|
||||
reserve_stopped = cost >= max_budget_usd * _SUBAGENT_BUDGET_RESERVE
|
||||
|
|
@ -121,18 +156,32 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||
max_budget_usd: float | None = None,
|
||||
max_turns: int | None = None,
|
||||
interactive: bool = False,
|
||||
budget_policy: BudgetPolicy = "stop",
|
||||
) -> None:
|
||||
if max_budget_usd is not None and (
|
||||
not math.isfinite(max_budget_usd) or max_budget_usd <= 0
|
||||
):
|
||||
raise ValueError("max_budget_usd must be a finite number greater than 0")
|
||||
_validate_budget(max_budget_usd)
|
||||
if max_turns is not None and max_turns <= 0:
|
||||
raise ValueError("max_turns must be a positive integer")
|
||||
if budget_policy not in ("stop", "pause"):
|
||||
raise ValueError(f"unknown budget_policy: {budget_policy!r}")
|
||||
self._model = model
|
||||
self._max_budget_usd = max_budget_usd
|
||||
self._budget_increment = max_budget_usd
|
||||
self._max_turns = max_turns
|
||||
self._interactive = interactive
|
||||
self._budget_policy: BudgetPolicy = budget_policy
|
||||
|
||||
@property
|
||||
def max_budget_usd(self) -> float | None:
|
||||
return self._max_budget_usd
|
||||
|
||||
@property
|
||||
def budget_policy(self) -> BudgetPolicy:
|
||||
return self._budget_policy
|
||||
|
||||
def set_max_budget_usd(self, max_budget_usd: float | None) -> None:
|
||||
"""Replace the scan's cost limit; ``None`` removes it."""
|
||||
_validate_budget(max_budget_usd)
|
||||
self._max_budget_usd = max_budget_usd
|
||||
|
||||
def extend_budget(self) -> None:
|
||||
if self._max_budget_usd is None or self._budget_increment is None:
|
||||
|
|
@ -146,6 +195,8 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||
system_prompt: str | None, # noqa: ARG002
|
||||
input_items: list[TResponseInputItem],
|
||||
) -> None:
|
||||
if self._budget_policy == "pause":
|
||||
self._pause_if_limited(context)
|
||||
context.context[LLM_TURN_KEY] = int(context.context.get(LLM_TURN_KEY, 0)) + 1
|
||||
try:
|
||||
self._maybe_warn_turns(context, input_items)
|
||||
|
|
@ -153,6 +204,32 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||
except Exception:
|
||||
logger.exception("budget/turn warning injection failed")
|
||||
|
||||
def _pause_if_limited(self, context: RunContextWrapper[dict[str, Any]]) -> None:
|
||||
"""Park the agent before a paid call when the scan is at its limit or paused.
|
||||
|
||||
Only calls that have already returned are counted, so calls in flight on
|
||||
other agents still land and are paid for: ``spent`` may end up above the
|
||||
limit, which is expected and never an error under this policy.
|
||||
"""
|
||||
coordinator = coordinator_from_context(context.context)
|
||||
epoch = coordinator.resume_epoch if coordinator is not None else 0
|
||||
if coordinator is not None and coordinator.budget_paused:
|
||||
raise BudgetPausedError(
|
||||
"scan paused; waiting for the operator to resume", resume_epoch=epoch
|
||||
)
|
||||
if self._max_budget_usd is None:
|
||||
return
|
||||
report_state = get_global_report_state()
|
||||
if report_state is None:
|
||||
return
|
||||
cost = report_state.get_total_llm_cost()
|
||||
if cost >= self._max_budget_usd:
|
||||
raise BudgetPausedError(
|
||||
f"Scan budget of ${self._max_budget_usd:.2f} reached (spent ${cost:.4f}); "
|
||||
"pausing until the operator raises the limit",
|
||||
resume_epoch=epoch,
|
||||
)
|
||||
|
||||
def _maybe_warn_turns(
|
||||
self,
|
||||
context: RunContextWrapper[dict[str, Any]],
|
||||
|
|
@ -182,7 +259,7 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||
context: RunContextWrapper[dict[str, Any]],
|
||||
input_items: list[TResponseInputItem],
|
||||
) -> None:
|
||||
if self._max_budget_usd is None:
|
||||
if self._max_budget_usd is None or self._budget_policy == "pause":
|
||||
return
|
||||
report_state = get_global_report_state()
|
||||
if report_state is None:
|
||||
|
|
@ -250,6 +327,11 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||
except Exception:
|
||||
logger.exception("failed to record SDK usage for agent %s", agent_id)
|
||||
|
||||
if self._budget_policy == "pause":
|
||||
# The finished call is paid for and its tool calls still run for free;
|
||||
# the agent parks before its next call, in ``on_llm_start``.
|
||||
return
|
||||
|
||||
if self._max_budget_usd is not None:
|
||||
cost = report_state.get_total_llm_cost()
|
||||
if cost >= self._max_budget_usd:
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ from strix.config.models import (
|
|||
uses_chat_completions_tool_schema,
|
||||
)
|
||||
from strix.config.settings import DEFAULT_MAX_TURNS
|
||||
from strix.core.agents import AgentCoordinator
|
||||
from strix.core.agents import AgentCoordinator, BudgetPolicy
|
||||
from strix.core.execution import (
|
||||
respawn_subagents,
|
||||
run_agent_loop,
|
||||
|
|
@ -189,6 +189,7 @@ async def run_strix_scan(
|
|||
interactive: bool = False,
|
||||
max_turns: int = DEFAULT_MAX_TURNS,
|
||||
max_budget_usd: float | None = None,
|
||||
budget_policy: BudgetPolicy = "stop",
|
||||
model: str | None = None,
|
||||
cleanup_on_exit: bool = True,
|
||||
event_sink: StreamEventSink | None = None,
|
||||
|
|
@ -208,6 +209,12 @@ async def run_strix_scan(
|
|||
``extra_system_prompt_context`` is merged into the root agent's scan
|
||||
context before prompt rendering. Child agents keep the standard scan prompt
|
||||
and context.
|
||||
``budget_policy`` decides what happens when the LLM spend reaches
|
||||
``max_budget_usd``: ``"stop"`` warns the agents as the limit approaches and
|
||||
ends the scan at it; ``"pause"`` tells the agents nothing and parks every
|
||||
agent before its next LLM call until the caller resumes the scan through
|
||||
``coordinator.resume_budget()`` (optionally with a higher limit) or cancels
|
||||
it. ``coordinator.pause_budget()`` parks a running scan the same way.
|
||||
``mcp_connection_requests`` supplies the run's MCP connections from any
|
||||
source: when given, the engine connects those requests; when ``None`` (the
|
||||
command-line default) it reads ``~/.strix/mcp-servers.json`` itself. Either
|
||||
|
|
@ -256,9 +263,12 @@ async def run_strix_scan(
|
|||
if not strict_tool_schemas:
|
||||
logger.info("Sending non-strict tool schemas: %s caps strict tools", resolved_model)
|
||||
|
||||
if budget_policy not in ("stop", "pause"):
|
||||
raise ValueError(f"unknown budget_policy: {budget_policy!r}")
|
||||
if coordinator is None:
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.set_snapshot_path(agents_path)
|
||||
coordinator.set_budget_policy(budget_policy)
|
||||
|
||||
from strix.tools.coverage.tools import hydrate_coverage_from_disk
|
||||
from strix.tools.notes.tools import hydrate_notes_from_disk
|
||||
|
|
@ -289,11 +299,17 @@ async def run_strix_scan(
|
|||
report_state.get_total_llm_cost(),
|
||||
max_budget_usd,
|
||||
interactive=interactive,
|
||||
budget_policy=budget_policy,
|
||||
)
|
||||
# Under the pause policy the hooks re-park at the first call if the
|
||||
# spend is still at the limit, so a restored pause flag would only
|
||||
# hold agents back after the limit was raised.
|
||||
await coordinator.reset_budget_stops(
|
||||
budget_stopped=budget_stopped,
|
||||
reserve_stopped=reserve_stopped,
|
||||
budget_paused=interactive and coordinator.budget_paused,
|
||||
budget_paused=(
|
||||
interactive and budget_policy != "pause" and coordinator.budget_paused
|
||||
),
|
||||
)
|
||||
for aid, parent in coordinator.parent_of.items():
|
||||
if parent is None:
|
||||
|
|
@ -372,8 +388,10 @@ async def run_strix_scan(
|
|||
max_budget_usd=max_budget_usd,
|
||||
max_turns=max_turns,
|
||||
interactive=interactive,
|
||||
budget_policy=budget_policy,
|
||||
)
|
||||
if interactive:
|
||||
coordinator.set_budget_limit_setter(hooks.set_max_budget_usd)
|
||||
if interactive and budget_policy != "pause":
|
||||
coordinator.set_budget_extender(hooks.extend_budget)
|
||||
|
||||
scope_context = build_scope_context(scan_config)
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from strix.report.state import get_global_report_state
|
|||
from strix.skills import validate_requested_skills
|
||||
|
||||
|
||||
_ACTIVE_STATUSES: frozenset[str] = frozenset({"running", "waiting"})
|
||||
_ACTIVE_STATUSES: frozenset[str] = frozenset({"running", "waiting", "budget_paused"})
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
|
@ -816,7 +816,7 @@ async def stop_agent(
|
|||
"success": False,
|
||||
"error": (
|
||||
f"Agent {target_agent_id} is already '{current_status}'; "
|
||||
"stop_agent only acts on running/waiting agents — use "
|
||||
"stop_agent only acts on running/waiting/paused agents — use "
|
||||
"view_agent_graph to find still-active descendants and "
|
||||
"stop them individually, or send_message_to_agent if you "
|
||||
"want to wake this one with new instructions"
|
||||
|
|
|
|||
510
tests/test_budget_pause_policy.py
Normal file
510
tests/test_budget_pause_policy.py
Normal file
|
|
@ -0,0 +1,510 @@
|
|||
"""``budget_policy="pause"``: agents park before a paid call and never hear about budgets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.core import execution
|
||||
from strix.core.agents import AgentCoordinator
|
||||
from strix.core.execution import _start_child_runner, run_agent_loop
|
||||
from strix.core.hooks import (
|
||||
BudgetExceededError,
|
||||
BudgetPausedError,
|
||||
ReportUsageHooks,
|
||||
recomputed_budget_flags,
|
||||
)
|
||||
from strix.core.sessions import open_agent_session
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
COST_PER_CALL = 1.0
|
||||
_CALL_LATENCY_S = 0.005
|
||||
|
||||
|
||||
class _FakeLedger:
|
||||
def __init__(self) -> None:
|
||||
self.cost = 0.0
|
||||
self.calls: list[str] = []
|
||||
self.remaining: dict[str, int] = {}
|
||||
self.warned_inputs: list[list[Any]] = []
|
||||
self.gate: asyncio.Event | None = None
|
||||
self.in_flight = 0
|
||||
|
||||
def record_sdk_usage(self, **_kwargs: Any) -> None:
|
||||
return
|
||||
|
||||
def get_total_llm_cost(self) -> float:
|
||||
return self.cost
|
||||
|
||||
|
||||
class _FakeStream:
|
||||
"""One ``Runner.run_streamed`` call: several LLM turns, each guarded by the hooks.
|
||||
|
||||
Mirrors the SDK's ordering: ``on_llm_start`` runs before the paid request,
|
||||
``on_llm_end`` after it. A ``BudgetPausedError`` from ``on_llm_start`` ends
|
||||
the stream without spending, exactly like the SDK surfacing a hook error.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ledger: _FakeLedger,
|
||||
hooks: ReportUsageHooks,
|
||||
context: dict[str, Any],
|
||||
agent: Any,
|
||||
coordinator: AgentCoordinator,
|
||||
) -> None:
|
||||
self._ledger = ledger
|
||||
self._hooks = hooks
|
||||
self._context = context
|
||||
self._agent = agent
|
||||
self._coordinator = coordinator
|
||||
self.run_loop_exception: BaseException | None = None
|
||||
self.final_output = None
|
||||
|
||||
async def stream_events(self) -> AsyncIterator[Any]:
|
||||
agent_id = str(self._context.get("agent_id"))
|
||||
ctx_wrapper = MagicMock()
|
||||
ctx_wrapper.context = self._context
|
||||
while self._ledger.remaining.get(agent_id, 0) > 0:
|
||||
input_items: list[Any] = []
|
||||
try:
|
||||
await self._hooks.on_llm_start(ctx_wrapper, self._agent, None, input_items)
|
||||
except BudgetPausedError as exc:
|
||||
self.run_loop_exception = exc
|
||||
return
|
||||
self._ledger.warned_inputs.append(input_items)
|
||||
if self._ledger.gate is not None:
|
||||
self._ledger.in_flight += 1
|
||||
await self._ledger.gate.wait()
|
||||
self._ledger.in_flight -= 1
|
||||
self._ledger.cost += COST_PER_CALL
|
||||
self._ledger.calls.append(agent_id)
|
||||
self._ledger.remaining[agent_id] -= 1
|
||||
await self._hooks.on_llm_end(ctx_wrapper, self._agent, MagicMock())
|
||||
await asyncio.sleep(_CALL_LATENCY_S)
|
||||
if self._coordinator.statuses.get(agent_id) == "running":
|
||||
await self._coordinator.set_status(agent_id, "completed")
|
||||
items: tuple[Any, ...] = ()
|
||||
for item in items:
|
||||
yield item
|
||||
|
||||
def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002
|
||||
return
|
||||
|
||||
|
||||
def _fake_runner(ledger: _FakeLedger, coordinator: AgentCoordinator) -> Any:
|
||||
class _FakeRunner:
|
||||
@staticmethod
|
||||
def run_streamed(
|
||||
agent: Any,
|
||||
input: Any, # noqa: A002, ARG004
|
||||
*,
|
||||
run_config: Any, # noqa: ARG004
|
||||
context: dict[str, Any],
|
||||
max_turns: int, # noqa: ARG004
|
||||
session: Any, # noqa: ARG004
|
||||
hooks: ReportUsageHooks,
|
||||
) -> _FakeStream:
|
||||
return _FakeStream(
|
||||
ledger=ledger,
|
||||
hooks=hooks,
|
||||
context=context,
|
||||
agent=agent,
|
||||
coordinator=coordinator,
|
||||
)
|
||||
|
||||
return _FakeRunner
|
||||
|
||||
|
||||
async def _noop_compact(*_args: Any, **_kwargs: Any) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
async def _wait_until(predicate: Callable[[], bool], *, timeout: float = 5.0) -> None:
|
||||
async def _poll() -> None:
|
||||
while not predicate():
|
||||
await asyncio.sleep(0.001)
|
||||
|
||||
await asyncio.wait_for(_poll(), timeout=timeout)
|
||||
|
||||
|
||||
def _all_parked(coordinator: AgentCoordinator, *agent_ids: str) -> bool:
|
||||
return all(coordinator.statuses.get(aid) == "budget_paused" for aid in agent_ids)
|
||||
|
||||
|
||||
class _Scan:
|
||||
"""Root + children driven through the real non-interactive loops."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
max_budget_usd: float | None,
|
||||
) -> None:
|
||||
self.ledger = _FakeLedger()
|
||||
self.hooks = ReportUsageHooks(
|
||||
model="test-model", max_budget_usd=max_budget_usd, budget_policy="pause"
|
||||
)
|
||||
self.coordinator = AgentCoordinator()
|
||||
self.coordinator.set_budget_policy("pause")
|
||||
self.coordinator.set_budget_limit_setter(self.hooks.set_max_budget_usd)
|
||||
monkeypatch.setattr(execution, "Runner", _fake_runner(self.ledger, self.coordinator))
|
||||
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
|
||||
self.db_path = tmp_path / "agents.sqlite"
|
||||
self.sessions: list[Any] = []
|
||||
self.run_config = MagicMock()
|
||||
self.root_ctx: dict[str, Any] = {
|
||||
"agent_id": "root",
|
||||
"parent_id": None,
|
||||
"coordinator": self.coordinator,
|
||||
}
|
||||
self.root_task: asyncio.Task[Any] | None = None
|
||||
|
||||
async def start_root(self, *, calls: int) -> None:
|
||||
self.ledger.remaining["root"] = calls
|
||||
await self.coordinator.register("root", "strix", parent_id=None)
|
||||
session = open_agent_session("root", self.db_path)
|
||||
self.sessions.append(session)
|
||||
self.root_task = asyncio.create_task(
|
||||
run_agent_loop(
|
||||
agent=MagicMock(),
|
||||
initial_input=[],
|
||||
run_config=self.run_config,
|
||||
context=self.root_ctx,
|
||||
max_turns=500,
|
||||
coordinator=self.coordinator,
|
||||
agent_id="root",
|
||||
interactive=False,
|
||||
session=session,
|
||||
hooks=self.hooks,
|
||||
)
|
||||
)
|
||||
|
||||
async def start_child(self, child_id: str, *, calls: int) -> None:
|
||||
self.ledger.remaining[child_id] = calls
|
||||
await self.coordinator.register(child_id, "recon", parent_id="root")
|
||||
await _start_child_runner(
|
||||
parent_ctx=self.root_ctx,
|
||||
coordinator=self.coordinator,
|
||||
agents_db_path=self.db_path,
|
||||
sessions_to_close=self.sessions,
|
||||
run_config=self.run_config,
|
||||
max_turns=500,
|
||||
interactive=False,
|
||||
child_agent=MagicMock(),
|
||||
child_id=child_id,
|
||||
name=f"recon-{child_id}",
|
||||
parent_id="root",
|
||||
task="probe things",
|
||||
initial_input=[],
|
||||
hooks=self.hooks,
|
||||
)
|
||||
|
||||
def tasks(self) -> list[asyncio.Task[Any]]:
|
||||
tasks = [self.root_task] if self.root_task is not None else []
|
||||
tasks.extend(rt.task for rt in self.coordinator.runtimes.values() if rt.task is not None)
|
||||
return tasks
|
||||
|
||||
async def teardown(self) -> None:
|
||||
for task in self.tasks():
|
||||
task.cancel()
|
||||
await asyncio.gather(*self.tasks(), return_exceptions=True)
|
||||
for session in self.sessions:
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pause_policy_parks_every_agent_at_the_limit_and_resumes_in_place(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=5.0)
|
||||
agents = ("root", "child-a", "child-b")
|
||||
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
|
||||
await scan.start_root(calls=100)
|
||||
await scan.start_child("child-a", calls=100)
|
||||
await scan.start_child("child-b", calls=100)
|
||||
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
|
||||
assert scan.ledger.cost == pytest.approx(5.0)
|
||||
assert len(scan.ledger.calls) == 5
|
||||
assert scan.coordinator.budget_paused is False
|
||||
assert scan.coordinator.budget_stopped is False
|
||||
assert scan.coordinator.reserve_stopped is False
|
||||
assert all(not task.done() for task in scan.tasks())
|
||||
|
||||
woken = await scan.coordinator.resume_budget(max_budget_usd=8.0)
|
||||
assert sorted(woken) == sorted(agents)
|
||||
assert scan.hooks.max_budget_usd == 8.0
|
||||
await _wait_until(lambda: scan.ledger.cost >= 8.0)
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
|
||||
assert scan.ledger.cost == pytest.approx(8.0)
|
||||
|
||||
await scan.coordinator.resume_budget(max_budget_usd=9.0)
|
||||
await _wait_until(lambda: scan.ledger.cost >= 9.0)
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
|
||||
assert scan.ledger.cost == pytest.approx(9.0)
|
||||
assert len(scan.ledger.calls) == 9
|
||||
assert all(not task.done() for task in scan.tasks())
|
||||
|
||||
# The model never saw a budget message: no warning band, no resume note.
|
||||
assert scan.ledger.warned_inputs
|
||||
assert all(items == [] for items in scan.ledger.warned_inputs)
|
||||
for session in scan.sessions:
|
||||
assert await session.get_items() == []
|
||||
|
||||
# Stop while parked: an individual stop wakes that loop and it exits.
|
||||
child_a_task = scan.coordinator.runtimes["child-a"].task
|
||||
assert child_a_task is not None
|
||||
await scan.coordinator.request_stop("child-a")
|
||||
await asyncio.wait_for(child_a_task, timeout=5.0)
|
||||
assert scan.coordinator.statuses["child-a"] == "stopped"
|
||||
assert scan.coordinator.statuses["root"] == "budget_paused"
|
||||
assert scan.coordinator.statuses["child-b"] == "budget_paused"
|
||||
|
||||
# A scan-wide cancel while parked tears the rest down cleanly.
|
||||
await scan.teardown()
|
||||
assert all(task.done() for task in scan.tasks())
|
||||
assert scan.ledger.cost == pytest.approx(9.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pause_policy_operator_pause_and_resume_without_a_limit(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=None)
|
||||
agents = ("root", "child-a", "child-b")
|
||||
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
|
||||
await scan.start_root(calls=6)
|
||||
await scan.start_child("child-a", calls=6)
|
||||
await scan.start_child("child-b", calls=6)
|
||||
|
||||
await _wait_until(lambda: scan.ledger.cost >= 3.0)
|
||||
await scan.coordinator.pause_budget()
|
||||
spent_at_pause = scan.ledger.cost
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, *agents))
|
||||
await _wait_until(lambda: scan.coordinator.budget_paused, timeout=0.1)
|
||||
assert scan.ledger.cost == pytest.approx(spent_at_pause)
|
||||
assert all(not task.done() for task in scan.tasks())
|
||||
|
||||
await asyncio.sleep(0.05)
|
||||
assert scan.ledger.cost == pytest.approx(spent_at_pause)
|
||||
|
||||
woken = await scan.coordinator.resume_budget()
|
||||
assert sorted(woken) == sorted(agents)
|
||||
await _wait_until(lambda: not scan.coordinator.budget_paused, timeout=0.1)
|
||||
await asyncio.wait_for(asyncio.gather(*scan.tasks(), return_exceptions=True), timeout=5.0)
|
||||
assert scan.ledger.cost == pytest.approx(18.0)
|
||||
assert {aid: str(s) for aid, s in scan.coordinator.statuses.items()} == {
|
||||
"root": "completed",
|
||||
"child-a": "completed",
|
||||
"child-b": "completed",
|
||||
}
|
||||
assert all(items == [] for items in scan.ledger.warned_inputs)
|
||||
|
||||
for session in scan.sessions:
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pause_policy_keeps_in_flight_overshoot(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
scan = _Scan(tmp_path, monkeypatch, max_budget_usd=1.0)
|
||||
scan.ledger.gate = asyncio.Event()
|
||||
await scan.coordinator.register("root", "strix", parent_id=None)
|
||||
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=scan.ledger):
|
||||
await scan.start_child("child-a", calls=100)
|
||||
await scan.start_child("child-b", calls=100)
|
||||
|
||||
# Both calls were dispatched under the limit; neither is cancelled.
|
||||
await _wait_until(lambda: scan.ledger.in_flight == 2)
|
||||
assert scan.ledger.cost == 0.0
|
||||
scan.ledger.gate.set()
|
||||
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, "child-a", "child-b"))
|
||||
assert scan.ledger.cost == pytest.approx(2.0)
|
||||
assert scan.hooks.max_budget_usd is not None
|
||||
assert scan.ledger.cost > scan.hooks.max_budget_usd
|
||||
assert scan.coordinator.budget_stopped is False
|
||||
assert all(not task.done() for task in scan.tasks())
|
||||
|
||||
# Resuming at a limit that is already spent parks again without a call.
|
||||
await scan.coordinator.resume_budget(max_budget_usd=1.5)
|
||||
await asyncio.sleep(0.05)
|
||||
await _wait_until(lambda: _all_parked(scan.coordinator, "child-a", "child-b"))
|
||||
assert scan.ledger.cost == pytest.approx(2.0)
|
||||
|
||||
await scan.teardown()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resume_between_park_decision_and_wait_is_not_missed() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.set_budget_policy("pause")
|
||||
await coordinator.register("a", "strix", parent_id=None)
|
||||
|
||||
parked_epoch = coordinator.resume_epoch
|
||||
await coordinator.park_for_budget("a")
|
||||
assert _all_parked(coordinator, "a")
|
||||
await coordinator.resume_budget()
|
||||
|
||||
await asyncio.wait_for(
|
||||
coordinator.wait_for_budget_resume("a", parked_epoch=parked_epoch), timeout=1.0
|
||||
)
|
||||
assert coordinator.statuses["a"] == "running"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parked_wait_returns_on_stop_signals() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.set_budget_policy("pause")
|
||||
await coordinator.register("a", "strix", parent_id=None)
|
||||
await coordinator.register("b", "recon", parent_id="a")
|
||||
await coordinator.park_for_budget("a")
|
||||
await coordinator.park_for_budget("b")
|
||||
|
||||
wait_a = asyncio.create_task(
|
||||
coordinator.wait_for_budget_resume("a", parked_epoch=coordinator.resume_epoch)
|
||||
)
|
||||
wait_b = asyncio.create_task(
|
||||
coordinator.wait_for_budget_resume("b", parked_epoch=coordinator.resume_epoch)
|
||||
)
|
||||
await asyncio.sleep(0.02)
|
||||
assert not wait_a.done()
|
||||
assert not wait_b.done()
|
||||
|
||||
await coordinator.request_stop("b")
|
||||
await asyncio.wait_for(wait_b, timeout=1.0)
|
||||
assert coordinator.statuses["b"] == "stopped"
|
||||
assert not wait_a.done()
|
||||
|
||||
await coordinator.trigger_budget_stop()
|
||||
await asyncio.wait_for(wait_a, timeout=1.0)
|
||||
assert coordinator.statuses["a"] == "budget_paused"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resume_budget_replaces_the_limit_and_validates_it() -> None:
|
||||
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.set_budget_policy("pause")
|
||||
coordinator.set_budget_limit_setter(hooks.set_max_budget_usd)
|
||||
|
||||
await coordinator.resume_budget(max_budget_usd=25.0)
|
||||
assert hooks.max_budget_usd == 25.0
|
||||
await coordinator.resume_budget()
|
||||
assert hooks.max_budget_usd == 25.0
|
||||
with pytest.raises(ValueError, match="greater than 0"):
|
||||
await coordinator.resume_budget(max_budget_usd=0.0)
|
||||
with pytest.raises(ValueError, match="finite"):
|
||||
await coordinator.resume_budget(max_budget_usd=float("inf"))
|
||||
assert hooks.max_budget_usd == 25.0
|
||||
|
||||
|
||||
def _ctx(coordinator: AgentCoordinator | None, *, parent_id: str | None = None) -> MagicMock:
|
||||
wrapper = MagicMock()
|
||||
wrapper.context = {"agent_id": "x", "parent_id": parent_id}
|
||||
if coordinator is not None:
|
||||
wrapper.context["coordinator"] = coordinator
|
||||
return wrapper
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pause_hooks_never_warn_and_park_only_at_the_limit() -> None:
|
||||
ledger = _FakeLedger()
|
||||
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.set_budget_policy("pause")
|
||||
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
|
||||
for cost in (7.0, 8.5, 9.5, 9.99):
|
||||
ledger.cost = cost
|
||||
for parent_id in (None, "root"):
|
||||
items: list[Any] = []
|
||||
await hooks.on_llm_start(
|
||||
_ctx(coordinator, parent_id=parent_id), MagicMock(), None, items
|
||||
)
|
||||
assert items == []
|
||||
await hooks.on_llm_end(
|
||||
_ctx(coordinator, parent_id=parent_id), MagicMock(), MagicMock()
|
||||
)
|
||||
|
||||
ledger.cost = 10.0
|
||||
await hooks.on_llm_end(_ctx(coordinator, parent_id="root"), MagicMock(), MagicMock())
|
||||
with pytest.raises(BudgetPausedError) as at_limit:
|
||||
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, [])
|
||||
assert at_limit.value.resume_epoch == coordinator.resume_epoch
|
||||
|
||||
ledger.cost = 13.7
|
||||
with pytest.raises(BudgetPausedError):
|
||||
await hooks.on_llm_start(_ctx(coordinator, parent_id="root"), MagicMock(), None, [])
|
||||
|
||||
hooks.set_max_budget_usd(20.0)
|
||||
items = []
|
||||
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, items)
|
||||
assert items == []
|
||||
|
||||
await coordinator.pause_budget()
|
||||
with pytest.raises(BudgetPausedError, match="paused"):
|
||||
await hooks.on_llm_start(_ctx(coordinator), MagicMock(), None, [])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pause_hooks_do_not_count_a_parked_turn() -> None:
|
||||
ledger = _FakeLedger()
|
||||
ledger.cost = 10.0
|
||||
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0, budget_policy="pause")
|
||||
coordinator = AgentCoordinator()
|
||||
ctx = _ctx(coordinator)
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
|
||||
with pytest.raises(BudgetPausedError):
|
||||
await hooks.on_llm_start(ctx, MagicMock(), None, [])
|
||||
assert "llm_turn" not in ctx.context
|
||||
hooks.set_max_budget_usd(11.0)
|
||||
await hooks.on_llm_start(ctx, MagicMock(), None, [])
|
||||
assert ctx.context["llm_turn"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_policy_is_unchanged() -> None:
|
||||
ledger = _FakeLedger()
|
||||
hooks = ReportUsageHooks(model="m", max_budget_usd=10.0)
|
||||
assert hooks.budget_policy == "stop"
|
||||
|
||||
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
|
||||
ledger.cost = 7.0
|
||||
items: list[Any] = []
|
||||
await hooks.on_llm_start(_ctx(None), MagicMock(), None, items)
|
||||
assert len(items) == 1
|
||||
assert "Scan cost budget" in str(items[0])
|
||||
|
||||
ledger.cost = 10.0
|
||||
with pytest.raises(BudgetExceededError):
|
||||
await hooks.on_llm_end(_ctx(None), MagicMock(), MagicMock())
|
||||
|
||||
assert recomputed_budget_flags(10.0, 10.0, interactive=False, budget_policy="stop") == (
|
||||
True,
|
||||
True,
|
||||
)
|
||||
assert recomputed_budget_flags(10.0, 10.0, interactive=False, budget_policy="pause") == (
|
||||
False,
|
||||
False,
|
||||
)
|
||||
|
||||
|
||||
def test_budget_policy_is_validated() -> None:
|
||||
with pytest.raises(ValueError, match="budget_policy"):
|
||||
ReportUsageHooks(model="m", budget_policy="later") # type: ignore[arg-type]
|
||||
Loading…
Add table
Reference in a new issue