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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Max Iterations Limiter for LiteLLM Proxy.
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.
9Works across multiple proxy instances via DualCache (in-memory + Redis).
10Follows the same pattern as parallel_request_limiter_v3.py.
11"""
13import os
14from typing import TYPE_CHECKING, Any, Final
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
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
27 InternalUsageCache = _InternalUsageCache
28else:
29 InternalUsageCache = Any
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])
39local current = redis.call('INCR', key)
40if current == 1 then
41 redis.call('EXPIRE', key, ttl)
42end
44return current
45"""
47# Default TTL for session iteration counters (1 hour)
48DEFAULT_MAX_ITERATIONS_TTL: Final = 3600
51class _PROXY_MaxIterationsHandler(CustomLogger):
52 """
53 Pre-call hook that enforces max_iterations per session.
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
62 Cache key pattern:
63 {session_iterations:<session_id>}:count
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 """
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))
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
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.
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
101 max_iterations: Final = self._get_max_iterations(user_api_key_dict)
102 if max_iterations is None:
103 return None
105 verbose_proxy_logger.debug(
106 "MaxIterationsHandler: session_id=%s, max_iterations=%s",
107 session_id,
108 max_iterations,
109 )
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)
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 )
127 verbose_proxy_logger.debug(
128 "MaxIterationsHandler: session_id=%s, count=%s/%s",
129 session_id,
130 current_count,
131 max_iterations,
132 )
134 return None
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)
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)
149 return None
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 )
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)
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
174 def _make_cache_key(self, session_id: str) -> str:
175 """
176 Create cache key for session iteration counter.
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"
183 async def _increment_and_get(self, cache_key: str) -> int:
184 """
185 Atomically increment the session counter and return the new value.
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 )
203 # Fallback: in-memory cache
204 return await self._in_memory_increment(cache_key)
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