fix(proxy): scope session limits to each agent

This commit is contained in:
AlisinaDevelo 2026-09-26 21:25:57 +02:00
parent 41070b1363
commit 9feee0f8c6
4 changed files with 213 additions and 48 deletions

View file

@ -1,7 +1,7 @@
"""
Per-Session Budget Limiter for LiteLLM Proxy.
Enforces a dollar-amount cap per session (identified by `session_id` /
Enforces a dollar-amount cap per agent and session (identified by `session_id` /
`x-litellm-trace-id`). After each successful LLM call the response cost is
accumulated against the session. When the accumulated spend exceeds
`max_budget_per_session` (configured in agent litellm_params), subsequent
@ -14,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:<session_id>}:spend
{agent_session_budget:[<agent_id>,<session_id>]}: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."""

View file

@ -1,7 +1,7 @@
"""
Max Iterations Limiter for LiteLLM Proxy.
Enforces a per-session cap on the number of LLM calls an agentic loop can make.
Enforces a per-agent, per-session cap on the number of LLM calls an agentic loop can make.
Callers send a `session_id` with each request (via `x-litellm-session-id` header
or `metadata.session_id`), and this hook counts calls per session. When the count
exceeds `max_iterations` (configured in agent litellm_params or key metadata), returns 429.
@ -10,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:<session_id>}:count
{agent_session_iterations:[<agent_id>,<session_id>]}:count
Without an agent, retains {session_iterations:<session_id>}: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:<session_id>} 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:
"""

View file

@ -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

View file

@ -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)