fix: preserve per-agent session limits on upgrade

This commit is contained in:
AlisinaDevelo 2026-09-27 08:31:18 +02:00
parent fe84cde709
commit ba3ecf9d9e
6 changed files with 704 additions and 153 deletions

View file

@ -14,10 +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
@ -36,21 +38,71 @@ 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 total_new_key = KEYS[2]
local agent_key = KEYS[3]
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)
if redis.call('EXISTS', total_new_key) == 0 then
if redis.call('EXISTS', agent_key) == 1 then
return redis.error_reply('agent session spend exists without migration total')
end
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
redis.call('SET', total_new_key, '0')
if legacy_ttl >= 0 then
redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000)
end
end
return new_val
if redis.call('EXISTS', legacy_key) == 0 then
return redis.error_reply('legacy session spend expired before agent scope')
end
local legacy_value = tonumber(redis.call('INCRBYFLOAT', legacy_key, amount))
local total_new_value = tonumber(redis.call('INCRBYFLOAT', total_new_key, amount))
local agent_existed = redis.call('EXISTS', agent_key)
local agent_value = tonumber(redis.call('INCRBYFLOAT', agent_key, amount))
if agent_existed == 0 then
local migration_ttl = redis.call('PTTL', total_new_key)
if migration_ttl >= 0 then
redis.call('PEXPIRE', agent_key, migration_ttl + 1000)
end
end
return tostring(math.max(legacy_value - total_new_value, 0) + agent_value)
"""
MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT: Final = """
local legacy_value = tonumber(redis.call('GET', KEYS[1])) or 0
local total_new_value = tonumber(redis.call('GET', KEYS[2]))
local agent_value = tonumber(redis.call('GET', KEYS[3])) or 0
if total_new_value == nil then
if redis.call('EXISTS', KEYS[3]) == 1 then
return redis.error_reply('agent session spend exists without migration total')
end
return tostring(math.max(legacy_value, agent_value))
end
if redis.call('EXISTS', KEYS[1]) == 0 then
return redis.error_reply('legacy session spend expired before agent scope')
end
return tostring(math.max(legacy_value - total_new_value, 0) + agent_value)
"""
# Default TTL for session budget counters (1 hour)
@ -65,11 +117,14 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
- max_budget_per_session: dollar cap per agent and session_id
Cache key pattern:
{agent_session_budget:[<agent_id>,<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.ttl = int(
os.getenv(
"LITELLM_MAX_BUDGET_PER_SESSION_TTL",
@ -78,11 +133,18 @@ 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
self.increment_script = cast(
Callable[..., Awaitable[object]],
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_BUDGET_SESSION_INCREMENT_SCRIPT
),
)
self.get_agent_spend_script = cast(
Callable[..., Awaitable[object]],
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT
),
)
else:
self.increment_script = None
async def async_pre_call_hook(
self,
@ -104,8 +166,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
return None
max_budget = float(max_budget)
cache_key: Final = self._make_cache_key(session_id, agent_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",
@ -131,7 +192,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 {}
@ -152,17 +215,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), agent.agent_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",
@ -174,6 +231,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."""
@ -210,65 +268,115 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
def _make_cache_key(self, session_id: str, agent_id: str) -> str:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
scope: Final = json.dumps((global_agent_registry.stable_agent_id(agent_id), session_id), separators=(",", ":"))
return f"{{agent_session_budget:{scope}}}:spend"
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
return f"{{session_budget:{session_id}}}:agent:{stable_agent_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:
def _make_legacy_cache_key(self, session_id: str) -> str:
return f"{{session_budget:{session_id}}}:spend"
def _make_total_new_cache_key(self, session_id: str) -> str:
return f"{{session_budget:{session_id}}}:agent-scope-total"
async def _get_agent_spend(self, session_id: str, agent_id: str) -> float:
legacy_key: Final = self._make_legacy_cache_key(session_id)
total_new_key: Final = self._make_total_new_cache_key(session_id)
agent_key: Final = self._make_cache_key(session_id, agent_id)
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, total_new_key, agent_key],
args=[],
)
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)
total_new_value: Final = await self._get_local_spend(total_new_key)
agent_value: Final = await self._get_local_spend(agent_key)
if legacy_value is None:
if total_new_value is not None:
raise RuntimeError("Legacy session spend expired before agent scope")
return float(agent_value or 0.0)
if total_new_value is None:
return max(float(legacy_value), float(agent_value or 0.0))
return max(float(legacy_value) - float(total_new_value), 0.0) + float(agent_value or 0.0)
async def _get_local_spend(self, cache_key: str) -> float | None:
result: Final[object | None] = cast(
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)
total_new_key: Final = self._make_total_new_cache_key(session_id)
agent_key: Final = self._make_cache_key(session_id, agent_id)
if self.increment_script is not None:
try:
result: Final = await self.increment_script(
keys=[cache_key],
result: Final[object] = await self.increment_script(
keys=[legacy_key, total_new_key, agent_key],
args=[str(amount), self.ttl],
)
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)
total_new_value = await self._get_local_spend(total_new_key)
agent_value = await self._get_local_spend(agent_key)
if legacy_value is None:
if total_new_value is not None:
raise RuntimeError("Legacy session spend expired before agent scope")
legacy_value = 0.0
total_new_value = total_new_value or 0.0
agent_value = agent_value or 0.0
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 if legacy_value == 0.0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=total_new_key,
value=new_total_new,
ttl=self.ttl + 1 if total_new_value == 0.0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=agent_key,
value=new_agent,
ttl=self.ttl + 1 if agent_value == 0.0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
return max(new_legacy - new_total_new, 0.0) + new_agent

View file

@ -10,9 +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
@ -30,19 +32,59 @@ 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 total_new_key = KEYS[2]
local agent_key = KEYS[3]
local ttl = tonumber(ARGV[1])
if redis.call('EXISTS', total_new_key) == 0 then
if redis.call('EXISTS', agent_key) == 1 then
return redis.error_reply('agent session count exists without migration total')
end
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
redis.call('SET', total_new_key, '0')
if legacy_ttl >= 0 then
redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000)
end
end
if redis.call('EXISTS', legacy_key) == 0 then
return redis.error_reply('legacy session count expired before agent scope')
end
local legacy_value = redis.call('INCR', legacy_key)
local total_new_value = redis.call('INCR', total_new_key)
local agent_existed = redis.call('EXISTS', agent_key)
local agent_value = redis.call('INCR', agent_key)
if agent_existed == 0 then
local migration_ttl = redis.call('PTTL', total_new_key)
if migration_ttl >= 0 then
redis.call('PEXPIRE', agent_key, migration_ttl + 1000)
end
end
return math.max(legacy_value - total_new_value, 0) + agent_value
"""
# Default TTL for session iteration counters (1 hour)
@ -61,26 +103,28 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
metadata.session_id in request body
Cache key pattern:
{agent_session_iterations:[<agent_id>,<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.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
self.increment_script = cast(
Callable[..., Awaitable[object]],
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(MAX_ITERATIONS_INCREMENT_SCRIPT),
)
else:
self.increment_script = None
async def async_pre_call_hook(
self,
@ -111,8 +155,10 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
)
# Increment and check
cache_key: Final = self._make_cache_key(session_id, user_api_key_dict.agent_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)
@ -177,52 +223,109 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
"""
Create cache key for session iteration counter.
The Redis hash tag includes both identities when an agent is configured.
Keys without an agent retain the legacy session scope.
Agent-scoped counters share the legacy session hash tag so migration
scripts can atomically update both scopes on Redis Cluster.
"""
if agent_id is None:
return f"{{session_iterations:{session_id}}}:count"
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
scope: Final = json.dumps((global_agent_registry.stable_agent_id(agent_id), session_id), separators=(",", ":"))
return f"{{agent_session_iterations:{scope}}}:count"
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
return f"{{session_iterations:{session_id}}}:agent:{stable_agent_id}:count"
async def _increment_and_get(self, cache_key: str) -> int:
"""
Atomically increment the session counter and return the new value.
def _make_legacy_cache_key(self, session_id: str) -> str:
return f"{{session_iterations:{session_id}}}:count"
Tries Redis first (via registered Lua script for atomicity across
instances), falls back to in-memory cache.
"""
def _make_total_new_cache_key(self, session_id: str) -> str:
return f"{{session_iterations:{session_id}}}:agent-scope-total"
async def _get_local_count(self, cache_key: str) -> int | None:
local_result: Final[object | None] = cast(
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:
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 if current is None else None,
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)
total_new_key: Final = self._make_total_new_cache_key(session_id)
agent_key: Final = self._make_cache_key(session_id, agent_id)
if self.increment_script is not None:
try:
result: Final = await self.increment_script(
keys=[cache_key],
result: Final[object] = await self.increment_script(
keys=[legacy_key, total_new_key, agent_key],
args=[self.ttl],
)
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)
total_new_value = await self._get_local_count(total_new_key)
agent_value = await self._get_local_count(agent_key)
if legacy_value is None:
if total_new_value is not None:
raise RuntimeError("Legacy session count expired before agent scope")
legacy_value = 0
total_new_value = total_new_value or 0
agent_value = agent_value or 0
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 if legacy_value == 0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=total_new_key,
value=new_total_new,
ttl=self.ttl + 1 if total_new_value == 0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=agent_key,
value=new_agent,
ttl=self.ttl + 1 if agent_value == 0 else None,
litellm_parent_otel_span=None,
local_only=True,
)
return max(new_legacy - new_total_new, 0) + new_agent

View file

@ -8,15 +8,17 @@ Tests that session-scoped budget tracking works correctly:
- Requests without agent_id pass through
"""
import logging
import asyncio
import os
import socket
import uuid
from typing import Final
from unittest.mock import patch
from unittest.mock import AsyncMock, patch
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 (
@ -35,6 +37,13 @@ def _make_mock_agent(max_budget_per_session: float, agent_id: str = "agent-budge
)
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():
"""
@ -79,8 +88,7 @@ async def test_budget_per_session_exceeds_budget():
)
session_id = "session-over-budget"
cache_key = handler._make_cache_key(session_id, "agent-budget-123")
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)
@ -113,8 +121,7 @@ async def test_budget_per_session_independent_sessions():
agent_id="agent-budget-123",
)
cache_key_a = handler._make_cache_key("session-A", "agent-budget-123")
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)
@ -245,35 +252,215 @@ async def test_legacy_agent_id_shares_the_registered_agents_budget() -> None:
assert "Current spend: $1.0000" in str(rejected.value.detail)
class _OpenBreakerRedis:
def __init__(self) -> None:
from litellm.caching.redis_cache import RedisCircuitBreaker
@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))
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):
with pytest.raises(HTTPException) as rejected:
await handler.async_pre_call_hook(
UserAPIKeyAuth(agent_id="agent-budget-123"),
cache,
{"metadata": {"session_id": session_id}},
"",
)
@_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)
assert rejected.value.status_code == 429
assert "Current spend: $1.0000" in str(rejected.value.detail)
@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_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 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(kwargs, None, None, None)
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)
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_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:
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(DualCache()))
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)
total_new_key: Final = handler._make_total_new_cache_key(session_id)
agent_a_key: Final = handler._make_cache_key(session_id, "agent-a")
agent_b_key: Final = handler._make_cache_key(session_id, "agent-b")
keys.extend((legacy_key, total_new_key, agent_a_key, agent_b_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[:4]}
assert len(hash_tags) == 1
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
agent_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_a_key))
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
assert current_ttl <= sidecar_ttl <= current_ttl + 1100
assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100
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(total_new_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(total_new_key))
assert after_increment_ttl <= before_increment_ttl + 50
assert after_increment_sidecar_ttl <= before_increment_sidecar_ttl + 50
concurrent_legacy_key: Final = handler._make_legacy_cache_key(concurrent_session)
keys.extend(
(
concurrent_legacy_key,
handler._make_total_new_cache_key(concurrent_session),
handler._make_cache_key(concurrent_session, "agent-a"),
handler._make_cache_key(concurrent_session, "agent-b"),
)
)
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)
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")
await redis_client.delete(redis.check_and_fix_namespace(total_new_key))
with pytest.raises(Exception, match="agent session spend exists without migration total"):
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,8 +6,12 @@ Tests that session-scoped iteration counting works correctly:
- Different sessions have independent counters
"""
import asyncio
import os
import socket
import uuid
from typing import Final
from unittest.mock import patch
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
@ -29,6 +33,13 @@ def _make_mock_agent(max_iterations: int, agent_id: str = "agent-test-123") -> A
)
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():
"""
@ -229,3 +240,144 @@ async def test_legacy_agent_id_shares_the_registered_agents_iteration_limit() ->
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))
total_new_count: Final = await cache.async_get_cache(key=handler._make_total_new_cache_key(session_id))
agent_count: Final = await cache.async_get_cache(key=handler._make_cache_key(session_id, "agent-test-123"))
assert legacy_count - total_new_count + agent_count == 3
@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:
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(DualCache()))
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)
total_new_key: Final = handler._make_total_new_cache_key(session_id)
keys.extend(
(
legacy_key,
total_new_key,
handler._make_cache_key(session_id, "agent-a"),
handler._make_cache_key(session_id, "agent-b"),
)
)
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(total_new_key))
agent_ttl: Final = await redis_client.pttl(
redis.check_and_fix_namespace(handler._make_cache_key(session_id, "agent-a"))
)
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
assert current_ttl <= sidecar_ttl <= current_ttl + 1100
assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100
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(total_new_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(total_new_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))
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")
await redis_client.delete(redis.check_and_fix_namespace(total_new_key))
with pytest.raises(Exception, match="agent session count exists without migration total"):
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

@ -953,7 +953,7 @@ async def test_max_budget_per_session_limiter_populates_provider():
)
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(
@ -989,7 +989,7 @@ async def test_max_budget_per_session_limiter_unknown_model_falls_back():
)
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

@ -841,6 +841,7 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
)
internal_cache = MagicMock()
internal_cache.dual_cache = DualCache()
internal_cache.async_get_cache = AsyncMock(return_value=10.0)
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=internal_cache,