This commit is contained in:
Alisina Karimi 2026-09-30 22:54:14 +02:00 • committed by GitHub
commit 4230eeaf89
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1166 additions and 197 deletions

View file

@ -1,7 +1,7 @@
"""
Per-Session Budget Limiter for LiteLLM Proxy.
Enforces a dollar-amount cap per session (identified by `session_id` /
Enforces a dollar-amount cap per agent and session (identified by `session_id` /
`x-litellm-trace-id`). After each successful LLM call the response cost is
accumulated against the session. When the accumulated spend exceeds
`max_budget_per_session` (configured in agent litellm_params), subsequent
@ -14,9 +14,12 @@ Works across multiple proxy instances via DualCache (in-memory + Redis).
Follows the same pattern as max_iterations_limiter.py.
"""
import asyncio
import json
import logging
import os
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Final, cast
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
@ -35,21 +38,98 @@ else:
InternalUsageCache = Any
# Redis Lua script for atomic float increment with TTL.
# INCRBYFLOAT returns the new value as a string.
# Only sets EXPIRE on first call (when prior value was nil).
# Redis Lua scripts keep the legacy aggregate and new per-agent counters in
# sync. All keys use the legacy session hash tag so this also works on Redis
# Cluster. Old proxy instances continue to update the aggregate key.
MAX_BUDGET_SESSION_INCREMENT_SCRIPT: Final = """
local key = KEYS[1]
local amount = ARGV[1]
local legacy_key = KEYS[1]
local agent_scope_key = KEYS[2]
local amount = tonumber(ARGV[1])
local ttl = tonumber(ARGV[2])
local existed = redis.call('EXISTS', key)
local new_val = redis.call('INCRBYFLOAT', key, amount)
if existed == 0 then
redis.call('EXPIRE', key, ttl)
local agent_field = ARGV[3]
if amount == nil or amount <= 0 or amount ~= amount or math.abs(amount) == math.huge then
return redis.error_reply('agent session spend increment is invalid')
end
return new_val
if redis.call('EXISTS', agent_scope_key) == 0 then
if redis.call('EXISTS', legacy_key) == 0 then
redis.call('SET', legacy_key, '0')
redis.call('PEXPIRE', legacy_key, ttl * 1000)
end
local legacy_ttl = redis.call('PTTL', legacy_key)
if legacy_ttl == -2 then
return redis.error_reply('legacy session spend expired during migration')
end
if legacy_ttl >= 0 and legacy_ttl <= 1 then
return redis.error_reply('legacy session spend is expiring before agent scope')
end
redis.call('HSET', agent_scope_key, '__total_new', '0')
if legacy_ttl >= 0 then
redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1)
end
end
if redis.call('EXISTS', legacy_key) == 0 then
return redis.error_reply('legacy session spend expired before agent scope')
end
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
if total_new_raw == false then
return redis.error_reply('agent session scope is missing its migration total')
end
local total_new_value = tonumber(total_new_raw)
if total_new_value == nil or total_new_value ~= total_new_value or math.abs(total_new_value) == math.huge then
return redis.error_reply('agent session migration total is not numeric')
end
local agent_raw = redis.call('HGET', agent_scope_key, agent_field)
local agent_value = tonumber(agent_raw or '0')
if agent_value == nil or agent_value ~= agent_value or math.abs(agent_value) == math.huge then
return redis.error_reply('agent session counter is not numeric')
end
if math.abs(total_new_value + amount) == math.huge or math.abs(agent_value + amount) == math.huge then
return redis.error_reply('agent session spend increment is out of range')
end
local legacy_value = tonumber(redis.call('INCRBYFLOAT', legacy_key, amount))
local next_total_new = tonumber(redis.call('HINCRBYFLOAT', agent_scope_key, '__total_new', amount))
local next_agent_value = tonumber(redis.call('HINCRBYFLOAT', agent_scope_key, agent_field, amount))
return tostring(math.max(legacy_value - next_total_new, 0) + next_agent_value)
"""
MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT: Final = """
local legacy_key = KEYS[1]
local agent_scope_key = KEYS[2]
local legacy_value = tonumber(redis.call('GET', legacy_key)) or 0
if redis.call('EXISTS', agent_scope_key) == 0 then
return tostring(legacy_value)
end
if redis.call('EXISTS', legacy_key) == 0 then
return redis.error_reply('legacy session spend expired before agent scope')
end
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
if total_new_raw == false then
return redis.error_reply('agent session scope is missing its migration total')
end
local total_new_value = tonumber(total_new_raw)
if total_new_value == nil or total_new_value ~= total_new_value or math.abs(total_new_value) == math.huge then
return redis.error_reply('agent session migration total is not numeric')
end
local agent_raw = redis.call('HGET', agent_scope_key, ARGV[1])
local agent_value = 0
if agent_raw ~= false then
agent_value = tonumber(agent_raw)
if agent_value == nil or agent_value ~= agent_value or math.abs(agent_value) == math.huge then
return redis.error_reply('agent session counter is not numeric')
end
end
return tostring(math.max(legacy_value - total_new_value, 0) + agent_value)
"""
# Default TTL for session budget counters (1 hour)
@ -61,14 +141,18 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
Pre-call hook that enforces max_budget_per_session.
Configuration (set in agent litellm_params):
- max_budget_per_session: dollar cap per session_id
- max_budget_per_session: dollar cap per agent and session_id
Cache key pattern:
{session_budget:<session_id>}:spend
{session_budget:<session_id>}:agent:<agent_id>:spend
"""
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
self._local_lock = asyncio.Lock()
self.increment_script: Callable[..., Awaitable[object]] | None = None
self.get_agent_spend_script: Callable[..., Awaitable[object]] | None = None
self._registered_redis_cache: object | None = None
self.ttl = int(
os.getenv(
"LITELLM_MAX_BUDGET_PER_SESSION_TTL",
@ -76,12 +160,27 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
)
)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
self.increment_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_BUDGET_SESSION_INCREMENT_SCRIPT
)
else:
self._ensure_redis_scripts()
def _ensure_redis_scripts(self) -> None:
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
if redis_cache is None:
self.increment_script = None
self.get_agent_spend_script = None
self._registered_redis_cache = None
return
if redis_cache is self._registered_redis_cache:
return
self.increment_script = cast( # cast-ok: Redis registration returns the callable invoked below
Callable[..., Awaitable[object]],
redis_cache.async_register_script(MAX_BUDGET_SESSION_INCREMENT_SCRIPT),
)
self.get_agent_spend_script = cast( # cast-ok: Redis registration returns the callable invoked below
Callable[..., Awaitable[object]],
redis_cache.async_register_script(MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT),
)
self._registered_redis_cache = redis_cache
async def async_pre_call_hook(
self,
@ -97,13 +196,13 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
max_budget = self._get_max_budget_per_session(user_api_key_dict)
session_id: Final = self._get_session_id(data)
agent_id: Final = user_api_key_dict.agent_id
if max_budget is None or session_id is None:
if max_budget is None or session_id is None or agent_id is None:
return None
max_budget = float(max_budget)
cache_key: Final = self._make_cache_key(session_id)
current_spend: Final = await self._get_current_spend(cache_key)
current_spend: Final = await self._get_agent_spend(session_id, agent_id)
verbose_proxy_logger.debug(
"MaxBudgetPerSessionHandler: session_id=%s, spend=%.4f, max=%.2f",
@ -129,7 +228,9 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
After a successful LLM call, increment the session spend by the response cost.
Record every successful agent call so limits added later still see the
session's accumulated spend. The pre-call hook enforces a cap only when
one is configured.
"""
try:
litellm_params: Final = kwargs.get("litellm_params") or {}
@ -150,17 +251,11 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
if agent is None:
return
agent_litellm_params: Final = agent.litellm_params or {}
max_budget: Final = agent_litellm_params.get("max_budget_per_session")
if max_budget is None:
return
response_cost: Final = kwargs.get("response_cost") or 0.0
if response_cost <= 0:
return
cache_key: Final = self._make_cache_key(str(session_id))
await self._increment_spend(cache_key, float(response_cost))
await self._increment_agent_spend(str(session_id), agent.agent_id, float(response_cost))
verbose_proxy_logger.debug(
"MaxBudgetPerSessionHandler: incremented session %s spend by %.6f",
@ -172,6 +267,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
"MaxBudgetPerSessionHandler: error in async_log_success_event: %s",
str(e),
)
raise
def _get_session_id(self, data: dict) -> str | None:
"""Extract session_id from request metadata."""
@ -205,65 +301,152 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
return float(max_budget)
return None
def _make_cache_key(self, session_id: str) -> str:
def _make_agent_scope_cache_key(self, session_id: str) -> str:
return f"{{session_budget:{session_id}}}:agent-scope"
def _make_agent_scope_field(self, agent_id: str) -> str:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
return f"agent:{stable_agent_id}"
def _make_legacy_cache_key(self, session_id: str) -> str:
return f"{{session_budget:{session_id}}}:spend"
async def _get_current_spend(self, cache_key: str) -> float:
"""Read current accumulated spend for a session."""
if self.internal_usage_cache.dual_cache.redis_cache is not None:
async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None:
result: Final[object | None] = cast( # cast-ok: cache API returns Any; validate value before use
object | None,
await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
),
)
if result is None:
return None
if isinstance(result, dict) and all(isinstance(key, str) for key in result):
return cast(dict[str, object], result) # cast-ok: keys are checked and values remain opaque
raise RuntimeError("Agent session scope cache has an invalid value")
@staticmethod
def _get_scope_spend(scope: dict[str, object], field: str) -> float:
value = scope.get(field)
if value is None:
return 0.0
if isinstance(value, (int, float, str, bytes)):
return float(value)
raise RuntimeError("Agent session scope has an invalid counter")
async def _get_agent_spend(self, session_id: str, agent_id: str) -> float:
legacy_key: Final = self._make_legacy_cache_key(session_id)
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
agent_field: Final = self._make_agent_scope_field(agent_id)
self._ensure_redis_scripts()
if self.get_agent_spend_script is not None:
try:
result = await self.internal_usage_cache.dual_cache.redis_cache.async_get_cache(key=cache_key)
if result is not None:
result: Final[object] = await self.get_agent_spend_script(
keys=[legacy_key, agent_scope_key],
args=[agent_field],
)
if isinstance(result, (int, float, str, bytes)):
return float(result)
return 0.0
raise TypeError(f"Unexpected Redis spend result: {type(result).__name__}")
except Exception as e:
log_redis_failure(
verbose_proxy_logger,
logging.WARNING,
"MaxBudgetPerSessionHandler: Redis GET failed, falling back to in-memory",
"MaxBudgetPerSessionHandler: Redis agent spend read failed",
e,
)
raise
result = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
legacy_value: Final = await self._get_local_spend(legacy_key)
agent_scope: Final = await self._get_local_scope(agent_scope_key)
if legacy_value is None:
if agent_scope is not None:
raise RuntimeError("Legacy session spend expired before agent scope")
return 0.0
if agent_scope is None:
return float(legacy_value)
raw_total_new = agent_scope.get("__total_new")
if not isinstance(raw_total_new, (int, float, str, bytes)):
raise RuntimeError("Agent session scope is missing its migration total")
total_new_value = float(raw_total_new)
agent_value = self._get_scope_spend(agent_scope, agent_field)
return max(float(legacy_value) - total_new_value, 0.0) + agent_value
async def _get_local_spend(self, cache_key: str) -> float | None:
result: Final[object | None] = cast( # cast-ok: cache API returns Any; validate value before use
object | None,
await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
),
)
if result is not None:
if isinstance(result, (int, float, str, bytes)):
return float(result)
return 0.0
return None
async def _increment_spend(self, cache_key: str, amount: float) -> float:
"""Atomically increment the session spend and return the new value."""
async def _increment_agent_spend(self, session_id: str, agent_id: str, amount: float) -> float:
legacy_key: Final = self._make_legacy_cache_key(session_id)
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
agent_field: Final = self._make_agent_scope_field(agent_id)
self._ensure_redis_scripts()
if self.increment_script is not None:
try:
result: Final = await self.increment_script(
keys=[cache_key],
args=[str(amount), self.ttl],
result: Final[object] = await self.increment_script(
keys=[legacy_key, agent_scope_key],
args=[str(amount), self.ttl, agent_field],
)
return float(result)
if isinstance(result, (int, float, str, bytes)):
return float(result)
raise TypeError(f"Unexpected Redis spend result: {type(result).__name__}")
except Exception as e:
log_redis_failure(
verbose_proxy_logger,
logging.WARNING,
"MaxBudgetPerSessionHandler: Redis INCRBYFLOAT failed, falling back to in-memory",
"MaxBudgetPerSessionHandler: Redis migration increment failed; refusing an unsafe retry",
e,
)
raise
return await self._in_memory_increment_spend(cache_key, amount)
async def _in_memory_increment_spend(self, cache_key: str, amount: float) -> float:
current: Final = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
new_value: Final = (float(current) if current is not None else 0.0) + amount
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=new_value,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return new_value
async with self._local_lock:
legacy_value = await self._get_local_spend(legacy_key)
agent_scope = await self._get_local_scope(agent_scope_key)
if legacy_value is None:
if agent_scope is not None:
raise RuntimeError("Legacy session spend expired before agent scope")
legacy_value = 0.0
if agent_scope is None:
total_new_value = 0.0
agent_value = 0.0
else:
raw_total_new = agent_scope.get("__total_new")
if not isinstance(raw_total_new, (int, float, str, bytes)):
raise RuntimeError("Agent session scope is missing its migration total")
total_new_value = float(raw_total_new)
agent_value = self._get_scope_spend(agent_scope, agent_field)
new_legacy: Final = float(legacy_value) + amount
new_total_new: Final = float(total_new_value) + amount
new_agent: Final = float(agent_value) + amount
await self.internal_usage_cache.async_set_cache(
key=legacy_key,
value=new_legacy,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=agent_scope_key,
value={
**(agent_scope or {}),
"__total_new": new_total_new,
agent_field: new_agent,
},
# Keep this map on a shorter rolling TTL than the aggregate fallback.
ttl=max(self.ttl - 1, 0),
litellm_parent_otel_span=None,
local_only=True,
)
return max(new_legacy - new_total_new, 0.0) + new_agent

View file

@ -1,7 +1,7 @@
"""
Max Iterations Limiter for LiteLLM Proxy.
Enforces a per-session cap on the number of LLM calls an agentic loop can make.
Enforces a per-agent, per-session cap on the number of LLM calls an agentic loop can make.
Callers send a `session_id` with each request (via `x-litellm-session-id` header
or `metadata.session_id`), and this hook counts calls per session. When the count
exceeds `max_iterations` (configured in agent litellm_params or key metadata), returns 429.
@ -10,8 +10,11 @@ Works across multiple proxy instances via DualCache (in-memory + Redis).
Follows the same pattern as parallel_request_limiter_v3.py.
"""
import asyncio
import json
import os
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Final, cast
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
@ -29,19 +32,72 @@ else:
InternalUsageCache = Any
# Redis Lua script for atomic increment with TTL.
# Returns the new count after increment.
# Only sets EXPIRE on first increment (when count becomes 1).
# Redis Lua scripts keep the legacy aggregate and new per-agent counters in
# sync. All keys use the legacy session hash tag so this also works on Redis
# Cluster. Old proxy instances continue to update the aggregate key.
MAX_ITERATIONS_INCREMENT_SCRIPT: Final = """
local key = KEYS[1]
local ttl = tonumber(ARGV[1])
local current = redis.call('INCR', key)
if current == 1 then
redis.call('EXPIRE', key, ttl)
local legacy_key = KEYS[1]
if #KEYS == 1 then
local current = redis.call('INCR', legacy_key)
if current == 1 then
redis.call('EXPIRE', legacy_key, tonumber(ARGV[1]))
end
return current
end
return current
local agent_scope_key = KEYS[2]
local ttl = tonumber(ARGV[1])
local agent_field = ARGV[2]
if redis.call('EXISTS', agent_scope_key) == 0 then
if redis.call('EXISTS', legacy_key) == 0 then
redis.call('SET', legacy_key, '0')
redis.call('PEXPIRE', legacy_key, ttl * 1000)
end
local legacy_ttl = redis.call('PTTL', legacy_key)
if legacy_ttl == -2 then
return redis.error_reply('legacy session count expired during migration')
end
if legacy_ttl >= 0 and legacy_ttl <= 1 then
return redis.error_reply('legacy session count is expiring before agent scope')
end
redis.call('HSET', agent_scope_key, '__total_new', '0')
if legacy_ttl >= 0 then
redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1)
end
end
if redis.call('EXISTS', legacy_key) == 0 then
return redis.error_reply('legacy session count expired before agent scope')
end
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
if total_new_raw == false then
return redis.error_reply('agent session scope is missing its migration total')
end
if string.match(total_new_raw, '^%d+$') == nil then
return redis.error_reply('agent session migration total is not an integer')
end
local validated_total_new = tonumber(total_new_raw)
if validated_total_new == nil or validated_total_new >= 9223372036854774784 then
return redis.error_reply('agent session migration total is out of range')
end
local agent_raw = redis.call('HGET', agent_scope_key, agent_field)
if agent_raw ~= false and string.match(agent_raw, '^%d+$') == nil then
return redis.error_reply('agent session counter is not an integer')
end
local validated_agent_value = tonumber(agent_raw or '0')
if validated_agent_value >= 9223372036854774784 then
return redis.error_reply('agent session counter is out of range')
end
local legacy_value = redis.call('INCR', legacy_key)
local total_new_value = redis.call('HINCRBY', agent_scope_key, '__total_new', 1)
local agent_value = redis.call('HINCRBY', agent_scope_key, agent_field, 1)
return math.max(legacy_value - total_new_value, 0) + agent_value
"""
# Default TTL for session iteration counters (1 hour)
@ -60,25 +116,39 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
metadata.session_id in request body
Cache key pattern:
{session_iterations:<session_id>}:count
{session_iterations:<session_id>}:agent:<agent_id>:count
Without an agent, retains {session_iterations:<session_id>}:count.
Multi-instance support:
Uses Redis Lua script for atomic increment (same pattern as
parallel_request_limiter_v3). Falls back to in-memory cache
when Redis is unavailable.
Uses Redis Lua scripts for atomic increments when Redis is configured.
Uses process-local memory only when Redis is not configured; Redis
errors propagate so a failed shared counter cannot silently bypass the
limit.
"""
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
self._local_lock = asyncio.Lock()
self.increment_script: Callable[..., Awaitable[object]] | None = None
self._registered_redis_cache: object | None = None
self.ttl = int(os.getenv("LITELLM_MAX_ITERATIONS_TTL", DEFAULT_MAX_ITERATIONS_TTL))
# Register Lua script with Redis if available (same pattern as v3 limiter)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
self.increment_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_ITERATIONS_INCREMENT_SCRIPT
)
else:
self._ensure_redis_scripts()
def _ensure_redis_scripts(self) -> None:
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
if redis_cache is None:
self.increment_script = None
self._registered_redis_cache = None
return
if redis_cache is self._registered_redis_cache:
return
self.increment_script = cast( # cast-ok: Redis registration returns the callable invoked below
Callable[..., Awaitable[object]],
redis_cache.async_register_script(MAX_ITERATIONS_INCREMENT_SCRIPT),
)
self._registered_redis_cache = redis_cache
async def async_pre_call_hook(
self,
@ -109,8 +179,10 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
)
# Increment and check
cache_key: Final = self._make_cache_key(session_id)
current_count: Final = await self._increment_and_get(cache_key)
if user_api_key_dict.agent_id is None:
current_count = await self._increment_legacy_and_get(self._make_legacy_cache_key(session_id))
else:
current_count = await self._increment_agent_and_get(session_id, user_api_key_dict.agent_id)
if current_count > max_iterations:
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(data.get("model") if data else None)
@ -171,51 +243,128 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
return int(max_iterations)
return None
def _make_cache_key(self, session_id: str) -> str:
"""
Create cache key for session iteration counter.
Uses Redis hash-tag pattern {session_iterations:<session_id>} so all
keys for a session land on the same Redis Cluster slot.
"""
def _make_legacy_cache_key(self, session_id: str) -> str:
return f"{{session_iterations:{session_id}}}:count"
async def _increment_and_get(self, cache_key: str) -> int:
"""
Atomically increment the session counter and return the new value.
def _make_agent_scope_cache_key(self, session_id: str) -> str:
return f"{{session_iterations:{session_id}}}:agent-scope"
Tries Redis first (via registered Lua script for atomicity across
instances), falls back to in-memory cache.
"""
def _make_agent_scope_field(self, agent_id: str) -> str:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
return f"agent:{stable_agent_id}"
async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None:
result: Final[object | None] = cast( # cast-ok: cache API returns Any; validate value before use
object | None,
await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
),
)
if result is None:
return None
if isinstance(result, dict) and all(isinstance(key, str) for key in result):
return cast(dict[str, object], result) # cast-ok: keys are checked and values remain opaque
raise RuntimeError("Agent session scope cache has an invalid value")
async def _get_local_count(self, cache_key: str) -> int | None:
local_result: Final[object | None] = cast( # cast-ok: cache API returns Any; validate value below
object | None,
await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
),
)
if isinstance(local_result, (int, float, str, bytes)):
return int(local_result)
return None
async def _increment_legacy_and_get(self, cache_key: str) -> int:
self._ensure_redis_scripts()
if self.increment_script is not None:
result: Final[object] = await self.increment_script(
keys=[cache_key],
args=[self.ttl],
)
if isinstance(result, (int, float, str, bytes)):
return int(result)
raise TypeError(f"Unexpected Redis iteration result: {type(result).__name__}")
async with self._local_lock:
current: Final = await self._get_local_count(cache_key)
new_value: Final = (current or 0) + 1
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=new_value,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return new_value
async def _increment_agent_and_get(self, session_id: str, agent_id: str) -> int:
legacy_key: Final = self._make_legacy_cache_key(session_id)
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
agent_field: Final = self._make_agent_scope_field(agent_id)
self._ensure_redis_scripts()
if self.increment_script is not None:
try:
result: Final = await self.increment_script(
keys=[cache_key],
args=[self.ttl],
result: Final[object] = await self.increment_script(
keys=[legacy_key, agent_scope_key],
args=[self.ttl, agent_field],
)
return int(result)
if isinstance(result, (int, float, str, bytes)):
return int(result)
raise TypeError(f"Unexpected Redis iteration result: {type(result).__name__}")
except Exception as e:
verbose_proxy_logger.warning(
"MaxIterationsHandler: Redis failed, falling back to in-memory: %s",
"MaxIterationsHandler: Redis migration increment failed; refusing an unsafe retry: %s",
str(e),
)
raise
# Fallback: in-memory cache
return await self._in_memory_increment(cache_key)
async def _in_memory_increment(self, cache_key: str) -> int:
"""Increment counter in in-memory cache with TTL."""
current: Final = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
new_value: Final = (int(current) if current is not None else 0) + 1
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=new_value,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return new_value
async with self._local_lock:
legacy_value = await self._get_local_count(legacy_key)
agent_scope = await self._get_local_scope(agent_scope_key)
if legacy_value is None:
if agent_scope is not None:
raise RuntimeError("Legacy session count expired before agent scope")
legacy_value = 0
total_new_value = 0
agent_value = 0
if agent_scope is not None:
raw_total_new = agent_scope.get("__total_new")
if not isinstance(raw_total_new, (int, float, str, bytes)):
raise RuntimeError("Agent session scope is missing its migration total")
raw_agent_value = agent_scope.get(agent_field, 0)
if not isinstance(raw_agent_value, (int, float, str, bytes)):
raise RuntimeError("Agent session scope has an invalid counter")
total_new_value = int(raw_total_new)
agent_value = int(raw_agent_value)
new_legacy: Final = legacy_value + 1
new_total_new: Final = total_new_value + 1
new_agent: Final = agent_value + 1
await self.internal_usage_cache.async_set_cache(
key=legacy_key,
value=new_legacy,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=agent_scope_key,
value={
**(agent_scope or {}),
"__total_new": new_total_new,
agent_field: new_agent,
},
# Keep this map on a shorter rolling TTL than the aggregate fallback.
ttl=max(self.ttl - 1, 0),
litellm_parent_otel_span=None,
local_only=True,
)
return max(new_legacy - new_total_new, 0) + new_agent

View file

@ -8,15 +8,19 @@ Tests that session-scoped budget tracking works correctly:
- Requests without agent_id pass through
"""
from unittest.mock import patch
import asyncio
import os
import socket
import uuid
from typing import Final
from unittest.mock import AsyncMock, patch
import logging
import pytest
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import _redis_circuit_breaker_guard
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.hooks.max_budget_per_session_limiter import (
_PROXY_MaxBudgetPerSessionHandler,
)
@ -24,15 +28,22 @@ from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
def _make_mock_agent(max_budget_per_session: float) -> AgentResponse:
def _make_mock_agent(max_budget_per_session: float, agent_id: str = "agent-budget-123") -> AgentResponse:
return AgentResponse(
agent_id="agent-budget-123",
agent_id=agent_id,
agent_name="budget-agent",
litellm_params={"max_budget_per_session": max_budget_per_session},
agent_card_params={"name": "budget-agent", "version": "1.0.0"},
)
def _redis_port_for_migration_test() -> int | None:
port = int(os.getenv("LITELLM_TEST_REDIS_PORT", "6379"))
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.2)
return port if sock.connect_ex(("127.0.0.1", port)) == 0 else None
@pytest.mark.asyncio
async def test_budget_per_session_under_budget_passes():
"""
@ -49,11 +60,9 @@ async def test_budget_per_session_under_budget_passes():
mock_agent = _make_mock_agent(max_budget_per_session=5.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
registry: Final = AgentRegistry()
registry.register_agent(mock_agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
@ -79,16 +88,13 @@ async def test_budget_per_session_exceeds_budget():
)
session_id = "session-over-budget"
cache_key = handler._make_cache_key(session_id)
await handler._increment_spend(cache_key, 1.50)
await handler._increment_agent_spend(session_id, "agent-budget-123", 1.50)
mock_agent = _make_mock_agent(max_budget_per_session=1.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
registry: Final = AgentRegistry()
registry.register_agent(mock_agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -115,16 +121,13 @@ async def test_budget_per_session_independent_sessions():
agent_id="agent-budget-123",
)
cache_key_a = handler._make_cache_key("session-A")
await handler._increment_spend(cache_key_a, 3.0)
await handler._increment_agent_spend("session-A", "agent-budget-123", 3.0)
mock_agent = _make_mock_agent(max_budget_per_session=2.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
registry: Final = AgentRegistry()
registry.register_agent(mock_agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
# Session A should be blocked
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
@ -167,35 +170,369 @@ async def test_no_agent_id_passes():
assert result is None
class _OpenBreakerRedis:
def __init__(self) -> None:
from litellm.caching.redis_cache import RedisCircuitBreaker
@pytest.mark.asyncio
@pytest.mark.parametrize(
"researcher_id,orchestrator_id,researcher_session,orchestrator_session",
[
("researcher", "orchestrator", "shared-trace", "shared-trace"),
("parent:child", "parent", "trace", "child:trace"),
],
)
async def test_agent_session_budget_counters_do_not_mix_usage(
researcher_id: str, orchestrator_id: str, researcher_session: str, orchestrator_session: str
) -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(3.0, researcher_id))
registry.register_agent(_make_mock_agent(1.0, orchestrator_id))
researcher: Final = UserAPIKeyAuth(agent_id=researcher_id)
orchestrator: Final = UserAPIKeyAuth(agent_id=orchestrator_id)
researcher_data: Final = {"metadata": {"session_id": researcher_session}}
orchestrator_data: Final = {"metadata": {"session_id": orchestrator_session}}
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
for _ in range(3):
self._circuit_breaker.record_failure()
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler.async_log_success_event(
{
"litellm_params": {"metadata": {"session_id": researcher_session, "agent_id": researcher_id}},
"response_cost": 2.0,
},
None,
None,
None,
)
assert await handler.async_pre_call_hook(researcher, cache, researcher_data, "") is None
assert await handler.async_pre_call_hook(orchestrator, cache, orchestrator_data, "") is None
@_redis_circuit_breaker_guard
async def async_get_cache(self, key, **kwargs):
raise AssertionError("never reached")
def async_register_script(self, script):
@_redis_circuit_breaker_guard
async def refused(_self, keys, args):
raise AssertionError("never reached")
return lambda keys, args: refused(self, keys, args)
await handler.async_log_success_event(
{
"litellm_params": {"metadata": {"session_id": orchestrator_session, "agent_id": orchestrator_id}},
"response_cost": 1.0,
},
None,
None,
None,
)
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(orchestrator, cache, orchestrator_data, "")
assert rejected.value.status_code == 429
assert "Current spend: $1.0000" in str(rejected.value.detail)
assert await handler.async_pre_call_hook(researcher, cache, researcher_data, "") is None
@pytest.mark.asyncio
async def test_an_open_circuit_breaker_reads_session_spend_locally_without_a_warning(caplog):
cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
handler = _PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache))
caplog.clear()
async def test_legacy_agent_id_shares_the_registered_agents_budget() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
registry: Final = AgentRegistry()
registry.load_agents_from_config(
(
{
"agent_name": "configured-agent",
"agent_card_params": {"name": "configured-agent", "version": "1"},
"litellm_params": {"max_budget_per_session": 1.0},
},
)
)
legacy_id, agent_id = next(iter(registry.config_agent_legacy_ids.items()))
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
spend = await handler._get_current_spend("{session_budget:quiet}:spend")
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler.async_log_success_event(
{"litellm_params": {"metadata": {"session_id": "trace", "agent_id": legacy_id}}, "response_cost": 1.0},
None,
None,
None,
)
for identity in (agent_id, legacy_id):
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(
UserAPIKeyAuth(agent_id=identity), cache, {"metadata": {"session_id": "trace"}}, ""
)
assert rejected.value.status_code == 429
assert "Current spend: $1.0000" in str(rejected.value.detail)
assert spend == 0.0
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
@pytest.mark.asyncio
async def test_agent_budget_keeps_spend_from_an_existing_session() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "existing-session"
legacy_key: Final = handler._make_legacy_cache_key(session_id)
await cache.async_set_cache(key=legacy_key, value=1.0)
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(1.0))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(
UserAPIKeyAuth(agent_id="agent-budget-123"),
cache,
{"metadata": {"session_id": session_id}},
"",
)
assert rejected.value.status_code == 429
assert "Current spend: $1.0000" in str(rejected.value.detail)
@pytest.mark.asyncio
async def test_agent_budget_carries_existing_spend_into_its_counter() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "existing-session"
await cache.async_set_cache(key=handler._make_legacy_cache_key(session_id), value=0.75)
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(1.0))
kwargs: Final = {
"litellm_params": {"metadata": {"session_id": session_id, "agent_id": "agent-budget-123"}},
"response_cost": 0.1,
}
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler.async_log_success_event(kwargs, None, None, None)
spend: Final = await handler._get_agent_spend(session_id, "agent-budget-123")
assert spend == pytest.approx(0.85)
@pytest.mark.asyncio
async def test_agent_budget_recovers_conservatively_after_scope_eviction() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "scope-eviction-session"
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
registry.register_agent(_make_mock_agent(10.0, "agent-b"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler._increment_agent_spend(session_id, "agent-a", 0.4)
await cache.async_delete_cache(key=handler._make_agent_scope_cache_key(session_id))
# Lost per-agent detail falls back to the aggregate, rather than resetting spend.
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.4)
assert await handler._increment_agent_spend(session_id, "agent-a", 0.1) == pytest.approx(0.5)
@pytest.mark.asyncio
async def test_agent_budget_fails_closed_when_scope_loses_migration_total() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "inconsistent-scope-session"
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
scope_key: Final = handler._make_agent_scope_cache_key(session_id)
await cache.async_set_cache(key=handler._make_legacy_cache_key(session_id), value=0.4)
await cache.async_set_cache(key=scope_key, value={})
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
with pytest.raises(RuntimeError, match="missing its migration total"):
await handler._increment_agent_spend(session_id, "agent-a", 0.1)
@pytest.mark.asyncio
async def test_agent_budget_registers_redis_scripts_when_redis_is_attached_late() -> None:
class FakeRedisCache:
def __init__(self) -> None:
self.scripts: list[str] = []
self.calls: list[dict[str, object]] = []
def async_register_script(self, source: str):
self.scripts.append(source)
async def call(**kwargs: object) -> object:
self.calls.append(kwargs)
return "0.1"
return call
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
redis_cache: Final = FakeRedisCache()
cache.redis_cache = redis_cache
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler._increment_agent_spend("late-redis", "agent-a", 0.1)
assert await handler._get_agent_spend("late-redis", "agent-a") == pytest.approx(0.1)
assert len(redis_cache.scripts) == 2
assert len(redis_cache.calls) == 2
assert all(len(call["keys"]) == 2 for call in redis_cache.calls)
@pytest.mark.asyncio
async def test_agent_budget_tracks_old_and_new_pods_without_mixing_new_agent_spend() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "mixed-rollout-session"
legacy_key: Final = handler._make_legacy_cache_key(session_id)
await cache.async_set_cache(key=legacy_key, value=0.5)
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(3.0, "agent-a"))
registry.register_agent(_make_mock_agent(3.0, "agent-b"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler._increment_agent_spend(session_id, "agent-a", 0.4)
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(0.9)
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.5)
# An old pod still writes only the legacy aggregate.
await cache.async_set_cache(key=legacy_key, value=1.1)
await handler._increment_agent_spend(session_id, "agent-a", 0.1)
await handler._increment_agent_spend(session_id, "agent-b", 0.1)
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(1.2)
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.8)
@pytest.mark.asyncio
async def test_agent_spend_is_recorded_before_a_budget_is_configured() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
session_id: Final = "unlimited-session"
agent: Final = _make_mock_agent(1.0)
agent.litellm_params = {}
registry: Final = AgentRegistry()
registry.register_agent(agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
await handler.async_log_success_event(
{
"litellm_params": {
"metadata": {"session_id": session_id, "agent_id": agent.agent_id},
},
"response_cost": 1.25,
},
None,
None,
None,
)
agent.litellm_params["max_budget_per_session"] = 1.0
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(
UserAPIKeyAuth(agent_id=agent.agent_id),
cache,
{"metadata": {"session_id": session_id}},
"",
)
assert rejected.value.status_code == 429
assert "Current spend: $1.2500" in str(rejected.value.detail)
@pytest.mark.asyncio
async def test_agent_budget_redis_errors_are_not_retried_or_read_locally() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
redis_cache: Final = object()
cache.redis_cache = redis_cache
handler._registered_redis_cache = redis_cache
increment_calls = 0
read_calls = 0
async def fail_increment(**_kwargs: object) -> object:
nonlocal increment_calls
increment_calls += 1
raise TimeoutError("reply timed out after Redis may have applied the script")
async def fail_read(**_kwargs: object) -> object:
nonlocal read_calls
read_calls += 1
raise TimeoutError("Redis read timed out")
handler.increment_script = fail_increment
handler.get_agent_spend_script = fail_read
with patch.object(handler, "_get_local_spend", new=AsyncMock(side_effect=AssertionError("must fail closed"))):
with pytest.raises(TimeoutError):
await handler._increment_agent_spend("uncertain-session", "agent-a", 0.1)
with pytest.raises(TimeoutError):
await handler._get_agent_spend("uncertain-session", "agent-a")
assert increment_calls == 1
assert read_calls == 1
@pytest.mark.asyncio
@pytest.mark.skipif(_redis_port_for_migration_test() is None, reason="requires local Redis for Lua migration path")
async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> None:
from litellm.caching.redis_cache import RedisCache
port: Final = _redis_port_for_migration_test()
assert port is not None
redis: Final = RedisCache(host="127.0.0.1", port=port)
redis_client: Final = redis.init_async_client()
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(DualCache(redis_cache=redis)))
handler.ttl = 30
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
registry.register_agent(_make_mock_agent(10.0, "agent-b"))
session_id: Final = f"budget-migration-{uuid.uuid4().hex}"
concurrent_session: Final = f"budget-concurrent-{uuid.uuid4().hex}"
keys: Final = []
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
try:
legacy_key: Final = handler._make_legacy_cache_key(session_id)
agent_scope_key: Final = handler._make_agent_scope_cache_key(session_id)
keys.extend((legacy_key, agent_scope_key))
await redis.async_set_cache(key=legacy_key, value=0.5, ttl=30)
legacy_redis_key: Final = redis.check_and_fix_namespace(legacy_key)
initial_ttl: Final = await redis_client.pttl(legacy_redis_key)
await handler._increment_agent_spend(session_id, "agent-a", 0.4)
await redis.async_increment(legacy_key, 0.2, ttl=30) # Old pod updates only the aggregate.
await handler._increment_agent_spend(session_id, "agent-a", 0.1)
await handler._increment_agent_spend(session_id, "agent-b", 0.1)
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(1.2)
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.8)
hash_tags: Final = {key[key.index("{") : key.index("}") + 1] for key in keys[:2]}
assert len(hash_tags) == 1
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
assert 0 <= sidecar_ttl <= current_ttl
assert current_ttl <= initial_ttl
await asyncio.sleep(0.25)
before_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
before_increment_sidecar_ttl: Final = await redis_client.pttl(
redis.check_and_fix_namespace(agent_scope_key)
)
await handler._increment_agent_spend(session_id, "agent-a", 0.01)
after_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
assert after_increment_ttl <= before_increment_ttl + 50
assert after_increment_sidecar_ttl <= before_increment_sidecar_ttl + 50
# Losing the hash evicts migration and agent counters together; the aggregate remains a safe baseline.
await redis_client.delete(redis.check_and_fix_namespace(agent_scope_key))
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(1.31)
await handler._increment_agent_spend(session_id, "agent-a", 0.01)
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(1.32)
concurrent_legacy_key: Final = handler._make_legacy_cache_key(concurrent_session)
concurrent_scope_key: Final = handler._make_agent_scope_cache_key(concurrent_session)
keys.extend((concurrent_legacy_key, concurrent_scope_key))
await redis.async_set_cache(key=concurrent_legacy_key, value=5.0, ttl=30)
await asyncio.gather(
*(handler._increment_agent_spend(concurrent_session, "agent-a", 0.01) for _ in range(100)),
*(handler._increment_agent_spend(concurrent_session, "agent-b", 0.01) for _ in range(50)),
)
assert await handler._get_agent_spend(concurrent_session, "agent-a") == pytest.approx(6.0)
assert await handler._get_agent_spend(concurrent_session, "agent-b") == pytest.approx(5.5)
saved_total_new: Final = await redis_client.hget(
redis.check_and_fix_namespace(agent_scope_key), "__total_new"
)
await redis_client.hdel(redis.check_and_fix_namespace(agent_scope_key), "__total_new")
with pytest.raises(Exception, match="agent session scope is missing its migration total"):
await handler._get_agent_spend(session_id, "agent-a")
assert saved_total_new is not None
await redis_client.hset(redis.check_and_fix_namespace(agent_scope_key), "__total_new", saved_total_new)
await redis_client.pexpire(legacy_redis_key, 100)
await asyncio.sleep(0.15)
with pytest.raises(Exception, match="legacy session spend expired before agent scope"):
await handler._get_agent_spend(session_id, "agent-a")
finally:
await redis_client.delete(*(redis.check_and_fix_namespace(key) for key in keys))

View file

@ -6,27 +6,40 @@ Tests that session-scoped iteration counting works correctly:
- Different sessions have independent counters
"""
from unittest.mock import MagicMock, patch
import asyncio
import os
import socket
import uuid
from typing import Final
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
def _make_mock_agent(max_iterations: int) -> AgentResponse:
def _make_mock_agent(max_iterations: int, agent_id: str = "agent-test-123") -> AgentResponse:
return AgentResponse(
agent_id="agent-test-123",
agent_id=agent_id,
agent_name="test-agent",
litellm_params={"max_iterations": max_iterations},
agent_card_params={"name": "test-agent", "version": "1.0.0"},
)
def _redis_port_for_migration_test() -> int | None:
port = int(os.getenv("LITELLM_TEST_REDIS_PORT", "6379"))
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.2)
return port if sock.connect_ex(("127.0.0.1", port)) == 0 else None
@pytest.mark.asyncio
async def test_max_iterations_basic_enforcement():
"""
@ -46,11 +59,9 @@ async def test_max_iterations_basic_enforcement():
mock_agent = _make_mock_agent(max_iterations=3)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
registry: Final = AgentRegistry()
registry.register_agent(mock_agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
# First 3 requests should succeed
for i in range(3):
await handler.async_pre_call_hook(
@ -91,11 +102,9 @@ async def test_max_iterations_different_sessions_independent():
mock_agent = _make_mock_agent(max_iterations=2)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
registry: Final = AgentRegistry()
registry.register_agent(mock_agent)
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
# Session A: 2 calls succeed
for _ in range(2):
await handler.async_pre_call_hook(
@ -154,3 +163,288 @@ async def test_max_iterations_no_agent_id_passes():
call_type="",
)
assert result is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"researcher_id,orchestrator_id,researcher_session,orchestrator_session",
[
("researcher", "orchestrator", "shared-trace", "shared-trace"),
("parent:child", "parent", "trace", "child:trace"),
],
)
async def test_agent_session_iteration_counters_do_not_mix_usage(
researcher_id: str, orchestrator_id: str, researcher_session: str, orchestrator_session: str
) -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(4, researcher_id))
registry.register_agent(_make_mock_agent(2, orchestrator_id))
researcher: Final = UserAPIKeyAuth(agent_id=researcher_id)
orchestrator: Final = UserAPIKeyAuth(agent_id=orchestrator_id)
researcher_data: Final = {"metadata": {"session_id": researcher_session}}
orchestrator_data: Final = {"metadata": {"session_id": orchestrator_session}}
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
for _ in range(3):
assert await handler.async_pre_call_hook(researcher, cache, researcher_data, "") is None
for _ in range(2):
assert await handler.async_pre_call_hook(orchestrator, cache, orchestrator_data, "") is None
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(orchestrator, cache, orchestrator_data, "")
assert rejected.value.status_code == 429
assert "Current count: 3" in str(rejected.value.detail)
assert await handler.async_pre_call_hook(researcher, cache, researcher_data, "") is None
with pytest.raises(HTTPException) as researcher_rejected:
await handler.async_pre_call_hook(researcher, cache, researcher_data, "")
assert researcher_rejected.value.status_code == 429
assert "Current count: 5" in str(researcher_rejected.value.detail)
@pytest.mark.asyncio
async def test_key_metadata_iteration_limit_keeps_existing_session_count() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
key: Final = UserAPIKeyAuth(metadata={"max_iterations": 2})
await cache.async_set_cache(key="{session_iterations:existing}:count", value=2)
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(key, cache, {"metadata": {"session_id": "existing"}}, "")
assert rejected.value.status_code == 429
assert "Current count: 3" in str(rejected.value.detail)
@pytest.mark.asyncio
async def test_legacy_agent_id_shares_the_registered_agents_iteration_limit() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
registry: Final = AgentRegistry()
registry.load_agents_from_config(
(
{
"agent_name": "configured-agent",
"agent_card_params": {"name": "configured-agent", "version": "1"},
"litellm_params": {"max_iterations": 2},
},
)
)
legacy_id, agent_id = next(iter(registry.config_agent_legacy_ids.items()))
data: Final = {"metadata": {"session_id": "trace"}}
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
for identity in (legacy_id, agent_id):
assert await handler.async_pre_call_hook(UserAPIKeyAuth(agent_id=identity), cache, data, "") is None
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(UserAPIKeyAuth(agent_id=legacy_id), cache, data, "")
assert rejected.value.status_code == 429
assert "Current count: 3" in str(rejected.value.detail)
@pytest.mark.asyncio
async def test_agent_iteration_limit_keeps_count_from_an_existing_session() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
session_id: Final = "existing-session"
legacy_key: Final = handler._make_legacy_cache_key(session_id)
await cache.async_set_cache(key=legacy_key, value=2)
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(2))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(
UserAPIKeyAuth(agent_id="agent-test-123"),
cache,
{"metadata": {"session_id": session_id}},
"",
)
assert rejected.value.status_code == 429
assert "Current count: 3" in str(rejected.value.detail)
legacy_count: Final = await cache.async_get_cache(key=handler._make_legacy_cache_key(session_id))
scope: Final = await cache.async_get_cache(key=handler._make_agent_scope_cache_key(session_id))
agent_field: Final = handler._make_agent_scope_field("agent-test-123")
assert isinstance(scope, dict)
assert legacy_count - scope["__total_new"] + scope[agent_field] == 3
@pytest.mark.asyncio
async def test_agent_iteration_recovers_conservatively_after_scope_eviction() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
session_id: Final = "scope-eviction-session"
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(20, "agent-a"))
registry.register_agent(_make_mock_agent(20, "agent-b"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
assert await handler._increment_agent_and_get(session_id, "agent-a") == 1
await cache.async_delete_cache(key=handler._make_agent_scope_cache_key(session_id))
# With the migration scope gone, prior usage becomes a conservative shared baseline.
assert await handler._increment_agent_and_get(session_id, "agent-b") == 2
assert await handler._increment_agent_and_get(session_id, "agent-a") == 2
@pytest.mark.asyncio
async def test_agent_iteration_fails_closed_when_scope_loses_migration_total() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
session_id: Final = "inconsistent-scope-session"
scope_key: Final = handler._make_agent_scope_cache_key(session_id)
await cache.async_set_cache(key=handler._make_legacy_cache_key(session_id), value=1)
await cache.async_set_cache(key=scope_key, value={'agent:"agent-a"': 1})
with pytest.raises(RuntimeError, match="missing its migration total"):
await handler._increment_agent_and_get(session_id, "agent-a")
@pytest.mark.asyncio
async def test_agent_iteration_registers_redis_scripts_when_redis_is_attached_late() -> None:
class FakeRedisCache:
def __init__(self) -> None:
self.scripts: list[str] = []
self.calls: list[dict[str, object]] = []
def async_register_script(self, source: str):
self.scripts.append(source)
async def call(**kwargs: object) -> object:
self.calls.append(kwargs)
return 1
return call
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
redis_cache: Final = FakeRedisCache()
cache.redis_cache = redis_cache
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(10, "agent-a"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
assert await handler._increment_agent_and_get("late-redis", "agent-a") == 1
assert len(redis_cache.scripts) == 1
assert len(redis_cache.calls) == 1
assert len(redis_cache.calls[0]["keys"]) == 2
@pytest.mark.asyncio
async def test_agent_iteration_counter_tracks_old_and_new_pods_independently() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
session_id: Final = "mixed-rollout-session"
legacy_key: Final = handler._make_legacy_cache_key(session_id)
await cache.async_set_cache(key=legacy_key, value=3)
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(20, "agent-a"))
registry.register_agent(_make_mock_agent(20, "agent-b"))
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
assert await handler._increment_agent_and_get(session_id, "agent-a") == 4
assert await handler._increment_agent_and_get(session_id, "agent-b") == 4
# An old pod still writes only the legacy aggregate.
await cache.async_set_cache(key=legacy_key, value=6)
assert await handler._increment_agent_and_get(session_id, "agent-a") == 6
assert await handler._increment_agent_and_get(session_id, "agent-b") == 6
@pytest.mark.asyncio
async def test_agent_iteration_redis_error_is_not_retried_or_counted_locally() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
redis_cache: Final = object()
cache.redis_cache = redis_cache
handler._registered_redis_cache = redis_cache
calls = 0
async def fail_after_attempt(**_kwargs: object) -> object:
nonlocal calls
calls += 1
raise TimeoutError("reply timed out after Redis may have applied the script")
handler.increment_script = fail_after_attempt
with patch.object(handler, "_get_local_count", new=AsyncMock(side_effect=AssertionError("must fail closed"))):
with pytest.raises(TimeoutError):
await handler._increment_agent_and_get("uncertain-session", "agent-a")
assert calls == 1
@pytest.mark.asyncio
@pytest.mark.skipif(_redis_port_for_migration_test() is None, reason="requires local Redis for Lua migration path")
async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() -> None:
from litellm.caching.redis_cache import RedisCache
port: Final = _redis_port_for_migration_test()
assert port is not None
redis: Final = RedisCache(host="127.0.0.1", port=port)
redis_client: Final = redis.init_async_client()
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(DualCache(redis_cache=redis)))
handler.ttl = 30
registry: Final = AgentRegistry()
registry.register_agent(_make_mock_agent(200, "agent-a"))
registry.register_agent(_make_mock_agent(200, "agent-b"))
session_id: Final = f"iterations-migration-{uuid.uuid4().hex}"
keys: Final = []
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
try:
legacy_key: Final = handler._make_legacy_cache_key(session_id)
agent_scope_key: Final = handler._make_agent_scope_cache_key(session_id)
keys.extend((legacy_key, agent_scope_key))
await redis.async_set_cache(key=legacy_key, value=3, ttl=30)
legacy_redis_key: Final = redis.check_and_fix_namespace(legacy_key)
initial_ttl: Final = await redis_client.pttl(legacy_redis_key)
hash_tags: Final = {key[key.index("{") : key.index("}") + 1] for key in keys}
assert len(hash_tags) == 1
assert await handler._increment_agent_and_get(session_id, "agent-a") == 4
assert await handler._increment_agent_and_get(session_id, "agent-b") == 4
await redis_client.incr(legacy_redis_key) # Old pod updates only the aggregate.
assert await handler._increment_agent_and_get(session_id, "agent-a") == 6
assert await handler._increment_agent_and_get(session_id, "agent-b") == 6
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
assert 0 <= sidecar_ttl <= current_ttl
assert current_ttl <= initial_ttl
await asyncio.sleep(0.25)
before_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
before_increment_sidecar_ttl: Final = await redis_client.pttl(
redis.check_and_fix_namespace(agent_scope_key)
)
assert await handler._increment_agent_and_get(session_id, "agent-a") == 7
after_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
assert after_increment_ttl <= before_increment_ttl + 50
assert after_increment_sidecar_ttl <= before_increment_sidecar_ttl + 50
values: Final = await asyncio.gather(
*(handler._increment_agent_and_get(session_id, "agent-a") for _ in range(100))
)
assert sorted(values) == list(range(8, 108))
# Losing the hash evicts migration and agent counters together; the aggregate remains a safe baseline.
await redis_client.delete(redis.check_and_fix_namespace(agent_scope_key))
assert await handler._increment_agent_and_get(session_id, "agent-b") == 110
assert await handler._increment_agent_and_get(session_id, "agent-a") == 110
saved_total_new: Final = await redis_client.hget(
redis.check_and_fix_namespace(agent_scope_key), "__total_new"
)
await redis_client.hdel(redis.check_and_fix_namespace(agent_scope_key), "__total_new")
with pytest.raises(Exception, match="agent session scope is missing its migration total"):
await handler._increment_agent_and_get(session_id, "agent-a")
assert saved_total_new is not None
await redis_client.hset(redis.check_and_fix_namespace(agent_scope_key), "__total_new", saved_total_new)
await redis_client.pexpire(legacy_redis_key, 100)
await asyncio.sleep(0.15)
with pytest.raises(Exception, match="legacy session count expired before agent scope"):
await handler._increment_agent_and_get(session_id, "agent-a")
finally:
await redis_client.delete(*(redis.check_and_fix_namespace(key) for key in keys))

View file

@ -42,6 +42,7 @@ import litellm
from litellm.caching.caching import DualCache
from litellm.exceptions import RateLimitError
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.batch_rate_limiter import (
BatchFileUsage,
_PROXY_BatchRateLimiter,
@ -60,7 +61,6 @@ from litellm.proxy.hooks.parallel_request_limiter import (
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.rate_limiter_utils import (
PROXY_LLM_PROVIDER_FALLBACK,
resolve_llm_provider_for_rate_limit,
@ -68,7 +68,6 @@ from litellm.proxy.hooks.rate_limiter_utils import (
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
# ---------------------------------------------------------------------------
# Helper class itself
# ---------------------------------------------------------------------------
@ -855,6 +854,7 @@ async def test_max_iterations_limiter_populates_provider():
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
mock_registry.stable_agent_id.return_value = "agent-iter"
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -896,6 +896,7 @@ async def test_max_iterations_limiter_unknown_model_falls_back():
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
mock_registry.stable_agent_id.return_value = "agent-iter"
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -950,8 +951,9 @@ async def test_max_budget_per_session_limiter_populates_provider():
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
max_budget=1.0
)
mock_registry.stable_agent_id.return_value = "agent-session-budget"
with patch.object(
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
handler, "_get_agent_spend", new=AsyncMock(return_value=5.0)
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
@ -985,8 +987,9 @@ async def test_max_budget_per_session_limiter_unknown_model_falls_back():
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
max_budget=1.0
)
mock_registry.stable_agent_id.return_value = "agent-session-budget"
with patch.object(
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
handler, "_get_agent_spend", new=AsyncMock(return_value=5.0)
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(

View file

@ -498,6 +498,7 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = agent
mock_registry.stable_agent_id.return_value = "agent-iter-1"
# First call within budget.
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -840,7 +841,8 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
)
internal_cache = MagicMock()
internal_cache.async_get_cache = AsyncMock(return_value=10.0)
internal_cache.dual_cache = DualCache()
internal_cache.async_get_cache = AsyncMock(side_effect=[10.0, None])
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=internal_cache,
)
@ -854,6 +856,7 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = agent
mock_registry.stable_agent_id.return_value = "agent-session-1"
with pytest.raises(ProxyRateLimitError) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,