From 9feee0f8c66a8675bba863bdbb43d77f366d557e Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sat, 26 Sep 2026 21:25:57 +0200 Subject: [PATCH 1/7] fix(proxy): scope session limits to each agent --- .../hooks/max_budget_per_session_limiter.py | 21 ++-- litellm/proxy/hooks/max_iterations_limiter.py | 21 ++-- .../test_max_budget_per_session_limiter.py | 118 +++++++++++++++--- .../hooks/test_max_iterations_limiter.py | 101 +++++++++++++-- 4 files changed, 213 insertions(+), 48 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index e07b96e5773..0ac1547693e 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -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,6 +14,7 @@ Works across multiple proxy instances via DualCache (in-memory + Redis). Follows the same pattern as max_iterations_limiter.py. """ +import json import logging import os from typing import TYPE_CHECKING, Any, Final @@ -61,10 +62,10 @@ 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:}:spend + {agent_session_budget:[,]}:spend """ def __init__(self, internal_usage_cache: InternalUsageCache): @@ -97,12 +98,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) + cache_key: Final = self._make_cache_key(session_id, agent_id) current_spend: Final = await self._get_current_spend(cache_key) verbose_proxy_logger.debug( @@ -159,7 +161,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): if response_cost <= 0: return - cache_key: Final = self._make_cache_key(str(session_id)) + cache_key: Final = self._make_cache_key(str(session_id), agent.agent_id) await self._increment_spend(cache_key, float(response_cost)) verbose_proxy_logger.debug( @@ -205,8 +207,11 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return float(max_budget) return None - def _make_cache_key(self, session_id: str) -> str: - return f"{{session_budget:{session_id}}}:spend" + 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" async def _get_current_spend(self, cache_key: str) -> float: """Read current accumulated spend for a session.""" diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 93697afa3c6..eb27f65fcc4 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -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,6 +10,7 @@ Works across multiple proxy instances via DualCache (in-memory + Redis). Follows the same pattern as parallel_request_limiter_v3.py. """ +import json import os from typing import TYPE_CHECKING, Any, Final @@ -60,7 +61,8 @@ class _PROXY_MaxIterationsHandler(CustomLogger): metadata.session_id in request body Cache key pattern: - {session_iterations:}:count + {agent_session_iterations:[,]}:count + Without an agent, retains {session_iterations:}:count. Multi-instance support: Uses Redis Lua script for atomic increment (same pattern as @@ -109,7 +111,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): ) # Increment and check - cache_key: Final = self._make_cache_key(session_id) + 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 current_count > max_iterations: @@ -171,14 +173,19 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return int(max_iterations) return None - def _make_cache_key(self, session_id: str) -> str: + def _make_cache_key(self, session_id: str, agent_id: str | None = None) -> str: """ Create cache key for session iteration counter. - Uses Redis hash-tag pattern {session_iterations:} so all - keys for a session land on the same Redis Cluster slot. + The Redis hash tag includes both identities when an agent is configured. + Keys without an agent retain the legacy session scope. """ - return f"{{session_iterations:{session_id}}}:count" + 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" async def _increment_and_get(self, cache_key: str) -> int: """ diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py index a1b3f313814..de4f581d4e7 100644 --- a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py @@ -8,15 +8,17 @@ Tests that session-scoped budget tracking works correctly: - Requests without agent_id pass through """ +import logging +from typing import Final from unittest.mock import 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,9 +26,9 @@ 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"}, @@ -49,11 +51,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 +79,14 @@ async def test_budget_per_session_exceeds_budget(): ) session_id = "session-over-budget" - cache_key = handler._make_cache_key(session_id) + cache_key = handler._make_cache_key(session_id, "agent-budget-123") await handler._increment_spend(cache_key, 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 +113,14 @@ async def test_budget_per_session_independent_sessions(): agent_id="agent-budget-123", ) - cache_key_a = handler._make_cache_key("session-A") + cache_key_a = handler._make_cache_key("session-A", "agent-budget-123") await handler._increment_spend(cache_key_a, 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,6 +163,88 @@ async def test_no_agent_id_passes(): 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_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}} + + 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 + + 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_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 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) + + class _OpenBreakerRedis: def __init__(self) -> None: from litellm.caching.redis_cache import RedisCircuitBreaker diff --git a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py index 20928ef46d5..9d512859f1b 100644 --- a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py @@ -6,21 +6,23 @@ Tests that session-scoped iteration counting works correctly: - Different sessions have independent counters """ -from unittest.mock import MagicMock, patch +from typing import Final +from unittest.mock import 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"}, @@ -46,11 +48,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 +91,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 +152,80 @@ 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) From 57a9211b1b7d21824d6991126f396eba4a8fbe95 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 06:06:43 +0200 Subject: [PATCH 2/7] fix(proxy): fall back to authenticated agent ids --- litellm/proxy/hooks/max_budget_per_session_limiter.py | 4 +++- litellm/proxy/hooks/max_iterations_limiter.py | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 0ac1547693e..77a6e021c81 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -210,7 +210,9 @@ 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=(",", ":")) + stable_agent_id: Final = global_agent_registry.stable_agent_id(agent_id) + canonical_agent_id: Final = stable_agent_id if isinstance(stable_agent_id, str) else agent_id + scope: Final = json.dumps((canonical_agent_id, session_id), separators=(",", ":")) return f"{{agent_session_budget:{scope}}}:spend" async def _get_current_spend(self, cache_key: str) -> float: diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index eb27f65fcc4..1b26c59fc83 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -184,7 +184,9 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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=(",", ":")) + stable_agent_id: Final = global_agent_registry.stable_agent_id(agent_id) + canonical_agent_id: Final = stable_agent_id if isinstance(stable_agent_id, str) else agent_id + scope: Final = json.dumps((canonical_agent_id, session_id), separators=(",", ":")) return f"{{agent_session_iterations:{scope}}}:count" async def _increment_and_get(self, cache_key: str) -> int: From 0d0df83bfd82b91bbf443bfc0ceb9d4cbd8d96b4 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 06:24:53 +0200 Subject: [PATCH 3/7] test(proxy): configure agent registry test doubles --- litellm/proxy/hooks/max_budget_per_session_limiter.py | 4 +--- litellm/proxy/hooks/max_iterations_limiter.py | 4 +--- .../proxy/hooks/test_proxy_rate_limit_provider_field.py | 7 +++++-- 3 files changed, 7 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 77a6e021c81..0ac1547693e 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -210,9 +210,7 @@ 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 - stable_agent_id: Final = global_agent_registry.stable_agent_id(agent_id) - canonical_agent_id: Final = stable_agent_id if isinstance(stable_agent_id, str) else agent_id - scope: Final = json.dumps((canonical_agent_id, session_id), separators=(",", ":")) + scope: Final = json.dumps((global_agent_registry.stable_agent_id(agent_id), session_id), separators=(",", ":")) return f"{{agent_session_budget:{scope}}}:spend" async def _get_current_spend(self, cache_key: str) -> float: diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 1b26c59fc83..eb27f65fcc4 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -184,9 +184,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return f"{{session_iterations:{session_id}}}:count" from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry - stable_agent_id: Final = global_agent_registry.stable_agent_id(agent_id) - canonical_agent_id: Final = stable_agent_id if isinstance(stable_agent_id, str) else agent_id - scope: Final = json.dumps((canonical_agent_id, session_id), separators=(",", ":")) + scope: Final = json.dumps((global_agent_registry.stable_agent_id(agent_id), session_id), separators=(",", ":")) return f"{{agent_session_iterations:{scope}}}:count" async def _increment_and_get(self, cache_key: str) -> int: diff --git a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py index 49bbd498cb9..879f2b1d889 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -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,6 +951,7 @@ 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) ): @@ -985,6 +987,7 @@ 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) ): From 3d42321bcce8b16b01b20d915d42362adb771629 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 06:41:07 +0200 Subject: [PATCH 4/7] test(proxy): configure stable agent IDs in error tests --- tests/unit/test_rate_limit_error_unification.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py index e5acba938c7..0e0e5273e6d 100644 --- a/tests/unit/test_rate_limit_error_unification.py +++ b/tests/unit/test_rate_limit_error_unification.py @@ -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, @@ -854,6 +855,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, From ba3ecf9d9e1d461fe6cca355b8f3df68d6aa62f3 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 08:31:18 +0200 Subject: [PATCH 5/7] fix: preserve per-agent session limits on upgrade --- .../hooks/max_budget_per_session_limiter.py | 236 +++++++++++----- litellm/proxy/hooks/max_iterations_limiter.py | 211 +++++++++++---- .../test_max_budget_per_session_limiter.py | 251 +++++++++++++++--- .../hooks/test_max_iterations_limiter.py | 154 ++++++++++- .../test_proxy_rate_limit_provider_field.py | 4 +- .../unit/test_rate_limit_error_unification.py | 1 + 6 files changed, 704 insertions(+), 153 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 0ac1547693e..6a3e66ff842 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -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:[,]}:spend + {session_budget:}:agent::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 diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index eb27f65fcc4..9288a71c611 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -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:[,]}:count + {session_iterations:}:agent::count Without an agent, retains {session_iterations:}: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 diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py index de4f581d4e7..5c5b84c6ed8 100644 --- a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py @@ -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)) diff --git a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py index 9d512859f1b..7839b3f0cf8 100644 --- a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py @@ -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)) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py index 879f2b1d889..556a3f1068d 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -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( diff --git a/tests/unit/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py index 0e0e5273e6d..abef8cf512e 100644 --- a/tests/unit/test_rate_limit_error_unification.py +++ b/tests/unit/test_rate_limit_error_unification.py @@ -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, From d21a5ec97f1abe891e775e4b914f3743965aa526 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 08:59:36 +0200 Subject: [PATCH 6/7] fix: keep session limiter migration state together --- .../hooks/max_budget_per_session_limiter.py | 222 ++++++++++++------ litellm/proxy/hooks/max_iterations_limiter.py | 161 ++++++++----- .../test_max_budget_per_session_limiter.py | 120 ++++++++-- .../hooks/test_max_iterations_limiter.py | 117 +++++++-- .../unit/test_rate_limit_error_unification.py | 2 +- 5 files changed, 435 insertions(+), 187 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 6a3e66ff842..0850fe0c0f4 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -43,16 +43,15 @@ else: # Cluster. Old proxy instances continue to update the aggregate key. MAX_BUDGET_SESSION_INCREMENT_SCRIPT: Final = """ local legacy_key = KEYS[1] -local total_new_key = KEYS[2] -local agent_key = KEYS[3] +local agent_scope_key = KEYS[2] local amount = tonumber(ARGV[1]) local ttl = tonumber(ARGV[2]) +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 -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', 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) @@ -62,10 +61,13 @@ if redis.call('EXISTS', total_new_key) == 0 then 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('SET', total_new_key, '0') + redis.call('HSET', agent_scope_key, '__total_new', '0') if legacy_ttl >= 0 then - redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000) + redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1) end end @@ -73,35 +75,60 @@ 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 +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 -return tostring(math.max(legacy_value - total_new_value, 0) + agent_value) +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_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 +local legacy_key = KEYS[1] +local agent_scope_key = KEYS[2] +local legacy_value = tonumber(redis.call('GET', legacy_key)) 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)) +if redis.call('EXISTS', agent_scope_key) == 0 then + return tostring(legacy_value) end -if redis.call('EXISTS', KEYS[1]) == 0 then + +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) """ @@ -125,6 +152,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 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", @@ -132,19 +160,27 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): ) ) - if self.internal_usage_cache.dual_cache.redis_cache is not None: - 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 - ), - ) + 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( + Callable[..., Awaitable[object]], + redis_cache.async_register_script(MAX_BUDGET_SESSION_INCREMENT_SCRIPT), + ) + self.get_agent_spend_script = cast( + 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, @@ -265,27 +301,52 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return float(max_budget) return None - def _make_cache_key(self, session_id: str, agent_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"{{session_budget:{session_id}}}:agent:{stable_agent_id}:spend" + return f"agent:{stable_agent_id}" 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_local_scope(self, cache_key: str) -> dict[str, object] | 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 None: + return None + if isinstance(result, dict) and all(isinstance(key, str) for key in result): + return cast(dict[str, object], result) + 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) - total_new_key: Final = self._make_total_new_cache_key(session_id) - agent_key: Final = self._make_cache_key(session_id, agent_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: Final[object] = await self.get_agent_spend_script( - keys=[legacy_key, total_new_key, agent_key], - args=[], + keys=[legacy_key, agent_scope_key], + args=[agent_field], ) if isinstance(result, (int, float, str, bytes)): return float(result) @@ -300,15 +361,19 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): raise 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) + agent_scope: Final = await self._get_local_scope(agent_scope_key) if legacy_value is None: - if total_new_value is not None: + if agent_scope 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) + 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( @@ -325,13 +390,14 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 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) + 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[object] = await self.increment_script( - keys=[legacy_key, total_new_key, agent_key], - args=[str(amount), self.ttl], + keys=[legacy_key, agent_scope_key], + args=[str(amount), self.ttl, agent_field], ) if isinstance(result, (int, float, str, bytes)): return float(result) @@ -347,35 +413,39 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 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) + agent_scope = await self._get_local_scope(agent_scope_key) if legacy_value is None: - if total_new_value is not None: + if agent_scope 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 + 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 if legacy_value == 0.0 else None, + ttl=self.ttl, 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, + 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, ) diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 9288a71c611..112b8bd4c00 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -45,14 +45,10 @@ if #KEYS == 1 then return current end -local total_new_key = KEYS[2] -local agent_key = KEYS[3] +local agent_scope_key = KEYS[2] 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 - +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) @@ -62,10 +58,13 @@ if redis.call('EXISTS', total_new_key) == 0 then 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('SET', total_new_key, '0') + redis.call('HSET', agent_scope_key, '__total_new', '0') if legacy_ttl >= 0 then - redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000) + redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1) end end @@ -73,16 +72,30 @@ 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 +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 """ @@ -117,14 +130,25 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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 = cast( - Callable[..., Awaitable[object]], - self.internal_usage_cache.dual_cache.redis_cache.async_register_script(MAX_ITERATIONS_INCREMENT_SCRIPT), - ) + 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( + 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, @@ -219,25 +243,32 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return int(max_iterations) return None - def _make_cache_key(self, session_id: str, agent_id: str | None = None) -> str: - """ - Create cache key for session iteration counter. - - 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 - - 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" - def _make_legacy_cache_key(self, session_id: str) -> str: return f"{{session_iterations:{session_id}}}:count" - def _make_total_new_cache_key(self, session_id: str) -> str: - return f"{{session_iterations:{session_id}}}:agent-scope-total" + def _make_agent_scope_cache_key(self, session_id: str) -> str: + return f"{{session_iterations:{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}" + + async def _get_local_scope(self, cache_key: str) -> dict[str, object] | 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 None: + return None + if isinstance(result, dict) and all(isinstance(key, str) for key in result): + return cast(dict[str, object], result) + 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( @@ -253,6 +284,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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], @@ -268,7 +300,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): await self.internal_usage_cache.async_set_cache( key=cache_key, value=new_value, - ttl=self.ttl if current is None else None, + ttl=self.ttl, litellm_parent_otel_span=None, local_only=True, ) @@ -276,13 +308,14 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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) + 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[object] = await self.increment_script( - keys=[legacy_key, total_new_key, agent_key], - args=[self.ttl], + keys=[legacy_key, agent_scope_key], + args=[self.ttl, agent_field], ) if isinstance(result, (int, float, str, bytes)): return int(result) @@ -296,35 +329,41 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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) + agent_scope = await self._get_local_scope(agent_scope_key) if legacy_value is None: - if total_new_value is not None: + if agent_scope 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 + 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 if legacy_value == 0 else None, + ttl=self.ttl, 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, + 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, ) diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py index 5c5b84c6ed8..5b4885f96fd 100644 --- a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py @@ -295,6 +295,71 @@ async def test_agent_budget_carries_existing_spend_into_its_counter() -> None: 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() @@ -358,7 +423,11 @@ async def test_agent_spend_is_recorded_before_a_budget_is_configured() -> None: @pytest.mark.asyncio async def test_agent_budget_redis_errors_are_not_retried_or_read_locally() -> None: - handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(DualCache())) + 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 @@ -405,10 +474,8 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N 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)) + 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) @@ -420,32 +487,32 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N 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]} + 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(total_new_key)) - agent_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_a_key)) + 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 current_ttl <= sidecar_ttl <= current_ttl + 1100 - assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100 + 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(total_new_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(total_new_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) - 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"), - ) - ) + 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)), @@ -454,13 +521,18 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N 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") - - 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)) diff --git a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py index 7839b3f0cf8..1cef0943dbc 100644 --- a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py @@ -264,9 +264,71 @@ async def test_agent_iteration_limit_keeps_count_from_an_existing_session() -> N 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 + 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 @@ -292,7 +354,11 @@ async def test_agent_iteration_counter_tracks_old_and_new_pods_independently() - @pytest.mark.asyncio async def test_agent_iteration_redis_error_is_not_retried_or_counted_locally() -> None: - handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(DualCache())) + 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: @@ -328,15 +394,8 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() - 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"), - ) - ) + 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) @@ -349,20 +408,18 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() - 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")) - ) + 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 current_ttl <= sidecar_ttl <= current_ttl + 1100 - assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100 + 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(total_new_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(total_new_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 @@ -371,13 +428,23 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() - ) 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") - - 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)) diff --git a/tests/unit/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py index abef8cf512e..3c7b597d8fc 100644 --- a/tests/unit/test_rate_limit_error_unification.py +++ b/tests/unit/test_rate_limit_error_unification.py @@ -842,7 +842,7 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: internal_cache = MagicMock() internal_cache.dual_cache = DualCache() - internal_cache.async_get_cache = AsyncMock(return_value=10.0) + internal_cache.async_get_cache = AsyncMock(side_effect=[10.0, None]) handler = _PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=internal_cache, ) From b13ad9ce1f316c7f269f06de85517d634e6fbb99 Mon Sep 17 00:00:00 2001 From: AlisinaDevelo Date: Sun, 27 Sep 2026 09:09:26 +0200 Subject: [PATCH 7/7] fix: satisfy type-discipline gate --- litellm/proxy/hooks/max_budget_per_session_limiter.py | 10 +++++----- litellm/proxy/hooks/max_iterations_limiter.py | 8 ++++---- 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 0850fe0c0f4..8c9f5161dda 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -172,11 +172,11 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): if redis_cache is self._registered_redis_cache: return - self.increment_script = cast( + 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( + 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), ) @@ -314,7 +314,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return f"{{session_budget:{session_id}}}:spend" async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None: - result: Final[object | None] = cast( + 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, @@ -325,7 +325,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 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) + 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 @@ -376,7 +376,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 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( + 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, diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 112b8bd4c00..10bf8e40372 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -144,7 +144,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): if redis_cache is self._registered_redis_cache: return - self.increment_script = cast( + 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), ) @@ -256,7 +256,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger): return f"agent:{stable_agent_id}" async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None: - result: Final[object | None] = cast( + 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, @@ -267,11 +267,11 @@ class _PROXY_MaxIterationsHandler(CustomLogger): 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) + 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( + 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,