Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/max_iterations_limiter.py: 40%

77 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2Max Iterations Limiter for LiteLLM Proxy. 

3 

4Enforces a per-session cap on the number of LLM calls an agentic loop can make. 

5Callers send a `session_id` with each request (via `x-litellm-session-id` header 

6or `metadata.session_id`), and this hook counts calls per session. When the count 

7exceeds `max_iterations` (configured in agent litellm_params or key metadata), returns 429. 

8 

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

10Follows the same pattern as parallel_request_limiter_v3.py. 

11""" 

12 

13import os 

14from typing import TYPE_CHECKING, Any, Final 

15 

16from litellm import DualCache 

17from litellm._logging import verbose_proxy_logger 

18from litellm.exceptions import RateLimitType 

19from litellm.integrations.custom_logger import CustomLogger 

20from litellm.proxy._types import UserAPIKeyAuth 

21from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError 

22from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit 

23 

24if TYPE_CHECKING: 24 ↛ 25line 24 didn't jump to line 25 because the condition on line 24 was never true

25 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache 

26 

27 InternalUsageCache = _InternalUsageCache 

28else: 

29 InternalUsageCache = Any 

30 

31 

32# Redis Lua script for atomic increment with TTL. 

33# Returns the new count after increment. 

34# Only sets EXPIRE on first increment (when count becomes 1). 

35MAX_ITERATIONS_INCREMENT_SCRIPT: Final = """ 

36local key = KEYS[1] 

37local ttl = tonumber(ARGV[1]) 

38 

39local current = redis.call('INCR', key) 

40if current == 1 then 

41 redis.call('EXPIRE', key, ttl) 

42end 

43 

44return current 

45""" 

46 

47# Default TTL for session iteration counters (1 hour) 

48DEFAULT_MAX_ITERATIONS_TTL: Final = 3600 

49 

50 

51class _PROXY_MaxIterationsHandler(CustomLogger): 

52 """ 

53 Pre-call hook that enforces max_iterations per session. 

54 

55 Configuration: 

56 - max_iterations: set in agent litellm_params (preferred) 

57 e.g. litellm_params={"max_iterations": 25} 

58 Falls back to key metadata max_iterations for backwards compatibility. 

59 - session_id: sent by caller via x-litellm-session-id header or 

60 metadata.session_id in request body 

61 

62 Cache key pattern: 

63 {session_iterations:<session_id>}:count 

64 

65 Multi-instance support: 

66 Uses Redis Lua script for atomic increment (same pattern as 

67 parallel_request_limiter_v3). Falls back to in-memory cache 

68 when Redis is unavailable. 

69 """ 

70 

71 def __init__(self, internal_usage_cache: InternalUsageCache): 

72 self.internal_usage_cache = internal_usage_cache 

73 self.ttl = int(os.getenv("LITELLM_MAX_ITERATIONS_TTL", DEFAULT_MAX_ITERATIONS_TTL)) 

74 

75 # Register Lua script with Redis if available (same pattern as v3 limiter) 

76 if self.internal_usage_cache.dual_cache.redis_cache is not None: 76 ↛ 77line 76 didn't jump to line 77 because the condition on line 76 was never true

77 self.increment_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

78 MAX_ITERATIONS_INCREMENT_SCRIPT 

79 ) 

80 else: 

81 self.increment_script = None 

82 

83 async def async_pre_call_hook( 

84 self, 

85 user_api_key_dict: UserAPIKeyAuth, 

86 cache: DualCache, 

87 data: dict, 

88 call_type: str, 

89 ) -> Exception | str | dict | None: 

90 """ 

91 Check session iteration count before making the API call. 

92 

93 Extracts session_id from request metadata and max_iterations from 

94 agent litellm_params. If the session has exceeded max_iterations, raises 429. 

95 """ 

96 # Extract session_id from request data 

97 session_id: Final = self._get_session_id(data) 

98 if session_id is None: 98 ↛ 101line 98 didn't jump to line 101 because the condition on line 98 was always true

99 return None 

100 

101 max_iterations: Final = self._get_max_iterations(user_api_key_dict) 

102 if max_iterations is None: 

103 return None 

104 

105 verbose_proxy_logger.debug( 

106 "MaxIterationsHandler: session_id=%s, max_iterations=%s", 

107 session_id, 

108 max_iterations, 

109 ) 

110 

111 # Increment and check 

112 cache_key: Final = self._make_cache_key(session_id) 

113 current_count: Final = await self._increment_and_get(cache_key) 

114 

115 if current_count > max_iterations: 

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"Max iterations exceeded for session {session_id}. " 

120 f"Current count: {current_count}, max_iterations: {max_iterations}." 

121 ), 

122 rate_limit_type=RateLimitType.MAX_ITERATIONS, 

123 model=resolved_model, 

124 llm_provider=llm_provider, 

125 ) 

126 

127 verbose_proxy_logger.debug( 

128 "MaxIterationsHandler: session_id=%s, count=%s/%s", 

129 session_id, 

130 current_count, 

131 max_iterations, 

132 ) 

133 

134 return None 

135 

136 def _get_session_id(self, data: dict) -> str | None: 

137 """Extract session_id from request metadata.""" 

138 metadata: Final = data.get("metadata") or {} 

139 session_id = metadata.get("session_id") 

140 if session_id is not None: 140 ↛ 141line 140 didn't jump to line 141 because the condition on line 140 was never true

141 return str(session_id) 

142 

143 # Also check litellm_metadata (used for /thread and /assistant endpoints) 

144 litellm_metadata: Final = data.get("litellm_metadata") or {} 

145 session_id = litellm_metadata.get("session_id") 

146 if session_id is not None: 146 ↛ 147line 146 didn't jump to line 147 because the condition on line 146 was never true

147 return str(session_id) 

148 

149 return None 

150 

151 def _get_max_iterations(self, user_api_key_dict: UserAPIKeyAuth) -> int | None: 

152 """Extract max_iterations from agent litellm_params, with fallback to key metadata.""" 

153 # Try agent litellm_params first 

154 agent_id: Final = user_api_key_dict.agent_id 

155 if agent_id is not None: 

156 from litellm.proxy.agent_endpoints.agent_registry import ( 

157 global_agent_registry, 

158 ) 

159 

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

161 if agent is not None: 

162 litellm_params: Final = agent.litellm_params or {} 

163 max_iterations = litellm_params.get("max_iterations") 

164 if max_iterations is not None: 

165 return int(max_iterations) 

166 

167 # Fallback to key metadata for backwards compatibility 

168 metadata: Final = user_api_key_dict.metadata or {} 

169 max_iterations = metadata.get("max_iterations") 

170 if max_iterations is not None: 

171 return int(max_iterations) 

172 return None 

173 

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

175 """ 

176 Create cache key for session iteration counter. 

177 

178 Uses Redis hash-tag pattern {session_iterations:<session_id>} so all 

179 keys for a session land on the same Redis Cluster slot. 

180 """ 

181 return f"{{session_iterations:{session_id}}}:count" 

182 

183 async def _increment_and_get(self, cache_key: str) -> int: 

184 """ 

185 Atomically increment the session counter and return the new value. 

186 

187 Tries Redis first (via registered Lua script for atomicity across 

188 instances), falls back to in-memory cache. 

189 """ 

190 if self.increment_script is not None: 

191 try: 

192 result: Final = await self.increment_script( 

193 keys=[cache_key], 

194 args=[self.ttl], 

195 ) 

196 return int(result) 

197 except Exception as e: 

198 verbose_proxy_logger.warning( 

199 "MaxIterationsHandler: Redis failed, falling back to in-memory: %s", 

200 str(e), 

201 ) 

202 

203 # Fallback: in-memory cache 

204 return await self._in_memory_increment(cache_key) 

205 

206 async def _in_memory_increment(self, cache_key: str) -> int: 

207 """Increment counter in in-memory cache with TTL.""" 

208 current: Final = await self.internal_usage_cache.async_get_cache( 

209 key=cache_key, 

210 litellm_parent_otel_span=None, 

211 local_only=True, 

212 ) 

213 new_value: Final = (int(current) if current is not None else 0) + 1 

214 await self.internal_usage_cache.async_set_cache( 

215 key=cache_key, 

216 value=new_value, 

217 ttl=self.ttl, 

218 litellm_parent_otel_span=None, 

219 local_only=True, 

220 ) 

221 return new_value