mirror of
https://github.com/usestrix/strix.git
synced 2026-08-22 11:02:08 +02:00
feat(runtime): graduated wrap-up warnings, budget reserve, and interactive budget pause/continue (#893)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
This commit is contained in:
co-authored by
Ahmed Allam
parent
47617969d3
commit
c55a8fa4ba
+388
-3
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user