Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/max_budget_per_session_limiter.py: 37%
114 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Per-Session Budget Limiter for LiteLLM Proxy.
4Enforces a dollar-amount cap per session (identified by `session_id` /
5`x-litellm-trace-id`). After each successful LLM call the response cost is
6accumulated against the session. When the accumulated spend exceeds
7`max_budget_per_session` (configured in agent litellm_params), subsequent
8requests for that session receive a 429.
10Note: trace-id enforcement (require_trace_id_on_calls_by_agent) is handled
11separately in auth_checks.py at the agent level, not in this hook.
13Works across multiple proxy instances via DualCache (in-memory + Redis).
14Follows the same pattern as max_iterations_limiter.py.
15"""
17import logging
18import os
19from typing import TYPE_CHECKING, Any, Final
21from litellm import DualCache
22from litellm._logging import verbose_proxy_logger
23from litellm.caching.redis_cache import log_redis_failure
24from litellm.exceptions import RateLimitType
25from litellm.integrations.custom_logger import CustomLogger
26from litellm.proxy._types import UserAPIKeyAuth
27from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
28from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
30if TYPE_CHECKING: 30 ↛ 31line 30 didn't jump to line 31 because the condition on line 30 was never true
31 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
33 InternalUsageCache = _InternalUsageCache
34else:
35 InternalUsageCache = Any
38# Redis Lua script for atomic float increment with TTL.
39# INCRBYFLOAT returns the new value as a string.
40# Only sets EXPIRE on first call (when prior value was nil).
41MAX_BUDGET_SESSION_INCREMENT_SCRIPT: Final = """
42local key = KEYS[1]
43local amount = ARGV[1]
44local ttl = tonumber(ARGV[2])
46local existed = redis.call('EXISTS', key)
47local new_val = redis.call('INCRBYFLOAT', key, amount)
48if existed == 0 then
49 redis.call('EXPIRE', key, ttl)
50end
52return new_val
53"""
55# Default TTL for session budget counters (1 hour)
56DEFAULT_MAX_BUDGET_PER_SESSION_TTL: Final = 3600
59class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
60 """
61 Pre-call hook that enforces max_budget_per_session.
63 Configuration (set in agent litellm_params):
64 - max_budget_per_session: dollar cap per session_id
66 Cache key pattern:
67 {session_budget:<session_id>}:spend
68 """
70 def __init__(self, internal_usage_cache: InternalUsageCache):
71 self.internal_usage_cache = internal_usage_cache
72 self.ttl = int(
73 os.getenv(
74 "LITELLM_MAX_BUDGET_PER_SESSION_TTL",
75 DEFAULT_MAX_BUDGET_PER_SESSION_TTL,
76 )
77 )
79 if self.internal_usage_cache.dual_cache.redis_cache is not None: 79 ↛ 80line 79 didn't jump to line 80 because the condition on line 79 was never true
80 self.increment_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
81 MAX_BUDGET_SESSION_INCREMENT_SCRIPT
82 )
83 else:
84 self.increment_script = None
86 async def async_pre_call_hook(
87 self,
88 user_api_key_dict: UserAPIKeyAuth,
89 cache: DualCache,
90 data: dict,
91 call_type: str,
92 ) -> Exception | str | dict | None:
93 """
94 Before each LLM call, check if max_budget_per_session is set and
95 whether accumulated spend exceeds the budget (429 if so).
96 """
97 max_budget = self._get_max_budget_per_session(user_api_key_dict)
99 session_id: Final = self._get_session_id(data)
101 if max_budget is None or session_id is None: 101 ↛ 104line 101 didn't jump to line 104 because the condition on line 101 was always true
102 return None
104 max_budget = float(max_budget)
105 cache_key: Final = self._make_cache_key(session_id)
106 current_spend: Final = await self._get_current_spend(cache_key)
108 verbose_proxy_logger.debug(
109 "MaxBudgetPerSessionHandler: session_id=%s, spend=%.4f, max=%.2f",
110 session_id,
111 current_spend,
112 max_budget,
113 )
115 if current_spend >= max_budget:
116 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(data.get("model") if data else None)
117 raise ProxyRateLimitError(
118 detail=(
119 f"Session budget exceeded for session {session_id}. "
120 f"Current spend: ${current_spend:.4f}, "
121 f"max_budget_per_session: ${max_budget:.2f}."
122 ),
123 rate_limit_type=RateLimitType.BUDGET,
124 model=resolved_model,
125 llm_provider=llm_provider,
126 )
128 return None
130 async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
131 """
132 After a successful LLM call, increment the session spend by the response cost.
133 """
134 try:
135 litellm_params: Final = kwargs.get("litellm_params") or {}
136 metadata: Final = litellm_params.get("metadata") or {}
137 session_id: Final = metadata.get("session_id")
138 if session_id is None: 138 ↛ 141line 138 didn't jump to line 141 because the condition on line 138 was always true
139 return
141 agent_id: Final = metadata.get("agent_id")
142 if agent_id is None:
143 return
145 from litellm.proxy.agent_endpoints.agent_registry import (
146 global_agent_registry,
147 )
149 agent: Final = global_agent_registry.get_agent_by_id(agent_id=str(agent_id))
150 if agent is None:
151 return
153 agent_litellm_params: Final = agent.litellm_params or {}
154 max_budget: Final = agent_litellm_params.get("max_budget_per_session")
155 if max_budget is None:
156 return
158 response_cost: Final = kwargs.get("response_cost") or 0.0
159 if response_cost <= 0:
160 return
162 cache_key: Final = self._make_cache_key(str(session_id))
163 await self._increment_spend(cache_key, float(response_cost))
165 verbose_proxy_logger.debug(
166 "MaxBudgetPerSessionHandler: incremented session %s spend by %.6f",
167 session_id,
168 response_cost,
169 )
170 except Exception as e:
171 verbose_proxy_logger.warning(
172 "MaxBudgetPerSessionHandler: error in async_log_success_event: %s",
173 str(e),
174 )
176 def _get_session_id(self, data: dict) -> str | None:
177 """Extract session_id from request metadata."""
178 metadata: Final = data.get("metadata") or {}
179 session_id = metadata.get("session_id")
180 if session_id is not None: 180 ↛ 181line 180 didn't jump to line 181 because the condition on line 180 was never true
181 return str(session_id)
183 litellm_metadata: Final = data.get("litellm_metadata") or {}
184 session_id = litellm_metadata.get("session_id")
185 if session_id is not None: 185 ↛ 186line 185 didn't jump to line 186 because the condition on line 185 was never true
186 return str(session_id)
188 return None
190 def _get_max_budget_per_session(self, user_api_key_dict: UserAPIKeyAuth) -> float | None:
191 """Extract max_budget_per_session from agent litellm_params."""
192 agent_id: Final = user_api_key_dict.agent_id
193 if agent_id is None: 193 ↛ 196line 193 didn't jump to line 196 because the condition on line 193 was always true
194 return None
196 from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
198 agent: Final = global_agent_registry.get_agent_by_id(agent_id=agent_id)
199 if agent is None:
200 return None
202 litellm_params: Final = agent.litellm_params or {}
203 max_budget: Final = litellm_params.get("max_budget_per_session")
204 if max_budget is not None:
205 return float(max_budget)
206 return None
208 def _make_cache_key(self, session_id: str) -> str:
209 return f"{{session_budget:{session_id}}}:spend"
211 async def _get_current_spend(self, cache_key: str) -> float:
212 """Read current accumulated spend for a session."""
213 if self.internal_usage_cache.dual_cache.redis_cache is not None:
214 try:
215 result = await self.internal_usage_cache.dual_cache.redis_cache.async_get_cache(key=cache_key)
216 if result is not None:
217 return float(result)
218 return 0.0
219 except Exception as e:
220 log_redis_failure(
221 verbose_proxy_logger,
222 logging.WARNING,
223 "MaxBudgetPerSessionHandler: Redis GET failed, falling back to in-memory",
224 e,
225 )
227 result = await self.internal_usage_cache.async_get_cache(
228 key=cache_key,
229 litellm_parent_otel_span=None,
230 local_only=True,
231 )
232 if result is not None:
233 return float(result)
234 return 0.0
236 async def _increment_spend(self, cache_key: str, amount: float) -> float:
237 """Atomically increment the session spend and return the new value."""
238 if self.increment_script is not None:
239 try:
240 result: Final = await self.increment_script(
241 keys=[cache_key],
242 args=[str(amount), self.ttl],
243 )
244 return float(result)
245 except Exception as e:
246 log_redis_failure(
247 verbose_proxy_logger,
248 logging.WARNING,
249 "MaxBudgetPerSessionHandler: Redis INCRBYFLOAT failed, falling back to in-memory",
250 e,
251 )
253 return await self._in_memory_increment_spend(cache_key, amount)
255 async def _in_memory_increment_spend(self, cache_key: str, amount: float) -> float:
256 current: Final = await self.internal_usage_cache.async_get_cache(
257 key=cache_key,
258 litellm_parent_otel_span=None,
259 local_only=True,
260 )
261 new_value: Final = (float(current) if current is not None else 0.0) + amount
262 await self.internal_usage_cache.async_set_cache(
263 key=cache_key,
264 value=new_value,
265 ttl=self.ttl,
266 litellm_parent_otel_span=None,
267 local_only=True,
268 )
269 return new_value