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

1""" 

2Per-Session Budget Limiter for LiteLLM Proxy. 

3 

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. 

9 

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. 

12 

13Works across multiple proxy instances via DualCache (in-memory + Redis). 

14Follows the same pattern as max_iterations_limiter.py. 

15""" 

16 

17import logging 

18import os 

19from typing import TYPE_CHECKING, Any, Final 

20 

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 

29 

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 

32 

33 InternalUsageCache = _InternalUsageCache 

34else: 

35 InternalUsageCache = Any 

36 

37 

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

45 

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 

51 

52return new_val 

53""" 

54 

55# Default TTL for session budget counters (1 hour) 

56DEFAULT_MAX_BUDGET_PER_SESSION_TTL: Final = 3600 

57 

58 

59class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): 

60 """ 

61 Pre-call hook that enforces max_budget_per_session. 

62 

63 Configuration (set in agent litellm_params): 

64 - max_budget_per_session: dollar cap per session_id 

65 

66 Cache key pattern: 

67 {session_budget:<session_id>}:spend 

68 """ 

69 

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 ) 

78 

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 

85 

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) 

98 

99 session_id: Final = self._get_session_id(data) 

100 

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 

103 

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) 

107 

108 verbose_proxy_logger.debug( 

109 "MaxBudgetPerSessionHandler: session_id=%s, spend=%.4f, max=%.2f", 

110 session_id, 

111 current_spend, 

112 max_budget, 

113 ) 

114 

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 ) 

127 

128 return None 

129 

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 

140 

141 agent_id: Final = metadata.get("agent_id") 

142 if agent_id is None: 

143 return 

144 

145 from litellm.proxy.agent_endpoints.agent_registry import ( 

146 global_agent_registry, 

147 ) 

148 

149 agent: Final = global_agent_registry.get_agent_by_id(agent_id=str(agent_id)) 

150 if agent is None: 

151 return 

152 

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 

157 

158 response_cost: Final = kwargs.get("response_cost") or 0.0 

159 if response_cost <= 0: 

160 return 

161 

162 cache_key: Final = self._make_cache_key(str(session_id)) 

163 await self._increment_spend(cache_key, float(response_cost)) 

164 

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 ) 

175 

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) 

182 

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) 

187 

188 return None 

189 

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 

195 

196 from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry 

197 

198 agent: Final = global_agent_registry.get_agent_by_id(agent_id=agent_id) 

199 if agent is None: 

200 return None 

201 

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 

207 

208 def _make_cache_key(self, session_id: str) -> str: 

209 return f"{{session_budget:{session_id}}}:spend" 

210 

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 ) 

226 

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 

235 

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 ) 

252 

253 return await self._in_memory_increment_spend(cache_key, amount) 

254 

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