Spaces:
Running
Running
File size: 6,626 Bytes
69e310f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 | """Loop safety: the hard iteration cap and the token/spend budget guard.
Both brakes are proved to actually stop a run, not merely to exist.
"""
from __future__ import annotations
from typing import Any
import pytest
from langgraph.graph import END, START, StateGraph
from app.core.budget import BudgetExceededError, BudgetGuard, Usage
from app.core.settings import Settings
from app.graph.context import RunContext, use_context
from app.graph.state import RunState, initial_state
from app.graph.supervisor import route_from_supervisor, supervisor_node
from app.models.run import RunStatus, TokenSpend
def _state(settings: Settings, iterations: int = 0) -> RunState:
state = initial_state(
run_id="run_loop",
watchlist_key="AAPL",
tickers=["AAPL"],
session_date="2026-01-02",
days=settings.price_history_days,
news_limit=settings.news_limit,
)
state["iterations"] = iterations
return state
class TestIterationCap:
async def test_supervisor_aborts_at_the_cap(
self, run_context: RunContext, settings: Settings
) -> None:
state = _state(settings, iterations=settings.max_iterations)
with use_context(run_context):
update = await supervisor_node(state)
assert update["status"] == RunStatus.ITERATION_ABORT
assert "iteration cap" in (update["abort_reason"] or "")
assert update["plan"]["dispatch"] == []
def test_routing_ends_the_graph_on_abort(self) -> None:
aborted = RunState(status=RunStatus.ITERATION_ABORT, plan={"dispatch": ["data_agent"]})
assert route_from_supervisor(aborted) == END
async def test_forced_loop_terminates(
self, run_context: RunContext, settings: Settings
) -> None:
"""A worker that never completes its ticker must still terminate the run.
This is the mandated forced-loop scenario: the workers are wired to return
nothing at all, so the watchlist can never become complete. The run must
stop at the iteration cap rather than spin forever.
"""
async def never_completes(state: RunState) -> dict[str, Any]:
# Records an attempt but never produces metrics or sentiment.
return {"agents_completed": ["stub"]}
graph: StateGraph[RunState, RunContext, RunState, RunState] = StateGraph(
RunState, context_schema=RunContext
)
graph.add_node("supervisor", supervisor_node)
graph.add_node("data_agent", never_completes)
graph.add_node("news_agent", never_completes)
graph.add_node("writer", never_completes)
graph.add_edge(START, "supervisor")
graph.add_conditional_edges(
"supervisor",
route_from_supervisor,
["data_agent", "news_agent", "writer", END],
)
graph.add_edge("data_agent", "supervisor")
graph.add_edge("news_agent", "supervisor")
graph.add_edge("writer", END)
compiled = graph.compile()
result = await compiled.ainvoke(
_state(settings),
{"recursion_limit": 200},
context=run_context,
)
# It terminated, and it terminated for the documented reason.
assert result["iterations"] <= settings.max_iterations + 1
assert result["status"] == RunStatus.ITERATION_ABORT
assert "iteration cap" in result["abort_reason"]
class TestBudgetGuard:
def test_charge_accumulates_and_trips(self) -> None:
guard = BudgetGuard(Settings(token_budget_usd=0.01, token_budget_tokens=1_000_000))
snapshot = guard.charge("claude-haiku-4-5", Usage(input_tokens=1000, output_tokens=100))
assert snapshot.calls == 1
assert snapshot.spent_tokens == 1100
with pytest.raises(BudgetExceededError):
for _ in range(200):
guard.charge("claude-sonnet-4-6", Usage(input_tokens=100_000, output_tokens=10_000))
def test_token_ceiling_trips_independently_of_usd(self) -> None:
guard = BudgetGuard(Settings(token_budget_usd=1000.0, token_budget_tokens=500))
with pytest.raises(BudgetExceededError) as exc:
guard.charge("claude-haiku-4-5", Usage(input_tokens=600))
assert exc.value.limit_tokens == 500
def test_preflight_refuses_a_call_that_cannot_fit(self) -> None:
guard = BudgetGuard(Settings(token_budget_usd=0.0001, token_budget_tokens=1_000_000))
with pytest.raises(BudgetExceededError):
guard.ensure_headroom("claude-sonnet-4-6", projected_output_tokens=4096)
def test_unknown_model_prices_at_the_most_expensive_tier(self) -> None:
"""A mis-configured model must never under-report spend."""
guard = BudgetGuard(Settings(token_budget_usd=1000.0, token_budget_tokens=10_000_000))
known = guard.charge("claude-haiku-4-5", Usage(output_tokens=1_000_000)).spent_usd
guard2 = BudgetGuard(Settings(token_budget_usd=1000.0, token_budget_tokens=10_000_000))
unknown = guard2.charge("totally-made-up-model", Usage(output_tokens=1_000_000)).spent_usd
assert unknown > known
async def test_run_aborts_with_budget_status(
self, run_context: RunContext, settings: Settings
) -> None:
"""A supervisor whose budget cannot cover one call aborts with BUDGET_ABORT."""
starved = settings.model_copy(update={"token_budget_usd": 0.000001})
context = RunContext(
run_id=run_context.run_id,
settings=starved,
engine=run_context.engine,
mcp=run_context.mcp,
bus=run_context.bus,
tracer=run_context.tracer,
budget=BudgetGuard(starved),
)
with use_context(context):
update = await supervisor_node(_state(starved))
assert update["status"] == RunStatus.BUDGET_ABORT
assert "budget" in (update["abort_reason"] or "").lower()
def test_snapshot_is_serialisable(self) -> None:
guard = BudgetGuard(Settings())
guard.charge("claude-haiku-4-5", Usage(input_tokens=5, output_tokens=5))
payload = guard.snapshot().to_dict()
assert payload["calls"] == 1
assert payload["remaining_usd"] >= 0
def test_spend_addition_is_lossless(self) -> None:
total = TokenSpend()
for _ in range(10):
total = total.plus(TokenSpend(input_tokens=1, output_tokens=2, cost_usd=0.1, calls=1))
assert total.input_tokens == 10
assert total.output_tokens == 20
assert total.calls == 10
assert round(total.cost_usd, 6) == 1.0
|