mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): scope session limits to each agent
This commit is contained in:
parent
41070b1363
commit
9feee0f8c6
4 changed files with 213 additions and 48 deletions
|
|
@ -1,7 +1,7 @@
|
|||
"""
|
||||
Per-Session Budget Limiter for LiteLLM Proxy.
|
||||
|
||||
Enforces a dollar-amount cap per session (identified by `session_id` /
|
||||
Enforces a dollar-amount cap per agent and session (identified by `session_id` /
|
||||
`x-litellm-trace-id`). After each successful LLM call the response cost is
|
||||
accumulated against the session. When the accumulated spend exceeds
|
||||
`max_budget_per_session` (configured in agent litellm_params), subsequent
|
||||
|
|
@ -14,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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue