mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge b13ad9ce1f into f285229b51
This commit is contained in:
commit
4230eeaf89
6 changed files with 1166 additions and 197 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue