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)