From 311be7f2cddcaadf657756861a18e15220d9c53d Mon Sep 17 00:00:00 2001 From: Shivam Rawat Date: Mon, 15 Jun 2026 20:50:38 -0700 Subject: [PATCH 1/6] fix(integrations): cap Anthropic cache_control injection at 4 blocks (#30480) * fix(integrations): cap Anthropic cache_control injection at 4 blocks Respect Anthropic's 4 cache_control breakpoint limit by counting client-supplied blocks, skipping messages that already carry cache_control, and stopping further auto-injection once the limit is reached. Co-authored-by: Cursor * fix(integrations): reserve cache slot for tool_config and short-circuit cap Address review feedback on the cache_control cap: break out of the injection loop before resolving target indices once the limit is reached, and reserve one of the four breakpoint slots when a tool_config injection point is present so the cachePoint appended by the Bedrock transform does not push the total past Anthropic's limit. Co-authored-by: Cursor --------- Co-authored-by: Cursor (cherry picked from commit fc9d789d24bc4bbed4512c5da60e0d988866890c) --- .../anthropic_cache_control_hook.py | 155 ++++++-- .../test_anthropic_cache_control_hook.py | 354 ++++++++++++++++++ 2 files changed, 476 insertions(+), 33 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 213622cb43a..296bfb6fc85 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -27,6 +27,11 @@ else: LiteLLMLoggingObj = Any +# Anthropic (and Bedrock Claude) reject requests with more than 4 cache_control +# breakpoints: "A maximum of 4 blocks with cache_control may be provided." +MAX_CACHE_CONTROL_BLOCKS = 4 + + class AnthropicCacheControlHook(CustomPromptManagement): def get_chat_completion_prompt( self, @@ -61,16 +66,30 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = copy.deepcopy(messages) # Separate message-level and non-message-level injection points - remaining_points = [] + message_points: List[CacheControlMessageInjectionPoint] = [] + remaining_points: List[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": - point = cast(CacheControlMessageInjectionPoint, point) - processed_messages = self._process_message_injection( - point=point, messages=processed_messages - ) + message_points.append(cast(CacheControlMessageInjectionPoint, point)) else: remaining_points.append(point) + # Non-message points (currently Bedrock tool_config) are handled in the + # provider transform, where each tool_config point appends at most one + # cachePoint to the tools. That block also counts toward Anthropic's + # limit, so reserve a slot for it here to leave room. + reserved_blocks = ( + 1 + if any(p.get("location") == "tool_config" for p in remaining_points) + else 0 + ) + + processed_messages = self._apply_message_injections( + points=message_points, + messages=processed_messages, + max_blocks=MAX_CACHE_CONTROL_BLOCKS - reserved_blocks, + ) + # Pass through non-message injection points for provider-specific handling if remaining_points: non_default_params["cache_control_injection_points"] = remaining_points @@ -78,14 +97,71 @@ class AnthropicCacheControlHook(CustomPromptManagement): return model, processed_messages, non_default_params @staticmethod - def _process_message_injection( - point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] + def _apply_message_injections( + points: List[CacheControlMessageInjectionPoint], + messages: List[AllMessageValues], + max_blocks: int, ) -> List[AllMessageValues]: - """Process message-level cache control injection.""" - control: ChatCompletionCachedContent = point.get( - "control", None - ) or ChatCompletionCachedContent(type="ephemeral") + """Apply message-level cache control injection points in order. + Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control + breakpoints per request. Client-supplied breakpoints count toward that + limit, so we never inject onto a message that already carries + cache_control (preserving the client's TTL) and we stop injecting once + ``max_blocks`` is reached. Injection points are honored in config order, + so earlier points win when slots are scarce. + """ + used_blocks = sum( + AnthropicCacheControlHook._count_cache_control_blocks(msg) + for msg in messages + ) + + limit_reached = False + for point in points: + if used_blocks >= max_blocks: + limit_reached = True + break + + control: ChatCompletionCachedContent = point.get( + "control", None + ) or ChatCompletionCachedContent(type="ephemeral") + + for target_index in AnthropicCacheControlHook._resolve_target_indices( + point=point, messages=messages + ): + if used_blocks >= max_blocks: + limit_reached = True + break + + if AnthropicCacheControlHook._message_has_cache_control( + messages[target_index] + ): + # Client already marked this message; don't overwrite it. + continue + + messages[target_index] = ( + AnthropicCacheControlHook._safe_insert_cache_control_in_message( + messages[target_index], control + ) + ) + used_blocks += 1 + + if limit_reached: + break + + if limit_reached: + verbose_logger.warning( + f"AnthropicCacheControlHook: Reached the Anthropic limit of " + f"{MAX_CACHE_CONTROL_BLOCKS} cache_control blocks. Skipping further injection." + ) + + return messages + + @staticmethod + def _resolve_target_indices( + point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] + ) -> List[int]: + """Resolve which message indices an injection point targets.""" _targetted_index: Optional[Union[int, str]] = point.get("index", None) targetted_index: Optional[int] = None if isinstance(_targetted_index, str): @@ -96,36 +172,49 @@ class AnthropicCacheControlHook(CustomPromptManagement): else: targetted_index = _targetted_index - targetted_role = point.get("role", None) - # Case 1: Target by specific index if targetted_index is not None: original_index = targetted_index - # Handle negative indices (convert to positive) if targetted_index < 0: targetted_index += len(messages) if 0 <= targetted_index < len(messages): - messages[targetted_index] = ( - AnthropicCacheControlHook._safe_insert_cache_control_in_message( - messages[targetted_index], control - ) - ) - else: - verbose_logger.warning( - f"AnthropicCacheControlHook: Provided index {original_index} is out of bounds for message list of length {len(messages)}. " - f"Targeted index was {targetted_index}. Skipping cache control injection for this point." - ) + return [targetted_index] + + verbose_logger.warning( + f"AnthropicCacheControlHook: Provided index {original_index} is out of bounds for message list of length {len(messages)}. " + f"Targeted index was {targetted_index}. Skipping cache control injection for this point." + ) + return [] + # Case 2: Target by role - elif targetted_role is not None: - for msg in messages: - if msg.get("role") == targetted_role: - msg = ( - AnthropicCacheControlHook._safe_insert_cache_control_in_message( - message=msg, control=control - ) - ) - return messages + targetted_role = point.get("role", None) + if targetted_role is not None: + return [ + idx + for idx, msg in enumerate(messages) + if msg.get("role") == targetted_role + ] + + return [] + + @staticmethod + def _count_cache_control_blocks(message: AllMessageValues) -> int: + """Count cache_control breakpoints on a message (message + content level).""" + count = 0 + if message.get("cache_control") is not None: + count += 1 + content = message.get("content") + if isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("cache_control") is not None: + count += 1 + return count + + @staticmethod + def _message_has_cache_control(message: AllMessageValues) -> bool: + """Return True if the message already carries any cache_control.""" + return AnthropicCacheControlHook._count_cache_control_blocks(message) > 0 @staticmethod def _safe_insert_cache_control_in_message( diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py index 1a4d03528e7..6afe5efc54d 100644 --- a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py +++ b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py @@ -1087,3 +1087,357 @@ async def test_anthropic_cache_control_hook_string_negative_index(): f"Expected cachePoint in last message content, got: {last_message_content}. " "String index '-1' was not parsed correctly (str.isdigit() returns False for negative strings)." ) + + +def _count_cache_control(messages: List[AllMessageValues]) -> int: + """Count cache_control breakpoints across messages (message + content level).""" + count = 0 + for message in messages: + if message.get("cache_control") is not None: + count += 1 + content = message.get("content") + if isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("cache_control") is not None: + count += 1 + return count + + +def _build_injection_points(): + return [ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + { + "location": "message", + "index": -1, + "control": {"type": "ephemeral", "ttl": "5m"}, + }, + ] + + +def test_cache_control_hook_caps_at_four_blocks_with_client_cache_control(): + """Regression for LIT-3667 / Anthropic 'A maximum of 4 blocks ... Found 5'. + + A Hermes-style request already carries 4 client cache_control breakpoints on + its system messages. With both auto-inject points configured the hook must + NOT add a 5th breakpoint, and must NOT overwrite the client's existing + breakpoints (TTL must be preserved). + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": f"System block {i}", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + } + for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": _build_injection_points() + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert ( + _count_cache_control(processed) == 4 + ), "Hook must cap cache_control at Anthropic's limit of 4 blocks" + + # Client TTL on system blocks must be preserved (not overwritten by config). + for i in range(4): + assert processed[i]["content"][-1]["cache_control"] == { + "type": "ephemeral", + "ttl": "1h", + } + + # The last (user) message must not receive a 5th breakpoint. + user_message = processed[-1] + assert user_message.get("cache_control") is None + user_content = user_message.get("content") + if isinstance(user_content, list): + assert all( + block.get("cache_control") is None + for block in user_content + if isinstance(block, dict) + ) + + +def test_cache_control_hook_caps_at_four_blocks_without_client_cache_control(): + """Four plain system messages + role:system + index:-1 must stay at 4 blocks. + + role:system fills all four slots, so the index:-1 point is skipped. + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + {"role": "system", "content": f"System {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": _build_injection_points() + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert _count_cache_control(processed) == 4 + # All four system messages cached; user message skipped (limit reached). + assert all(processed[i].get("cache_control") is not None for i in range(4)) + assert processed[-1].get("cache_control") is None + + +def test_cache_control_hook_does_not_overwrite_existing_cache_control(): + """If a targeted message already has client cache_control, do not inject.""" + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Cached by client", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + }, + {"role": "user", "content": "hello"}, + ] + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + # Target the already-cached system message with a different TTL. + non_default_params={ + "cache_control_injection_points": [ + { + "location": "message", + "index": 0, + "control": {"type": "ephemeral", "ttl": "5m"}, + } + ] + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + # Client's 1h TTL must be preserved, not replaced by the config's 5m. + assert processed[0]["content"][-1]["cache_control"] == { + "type": "ephemeral", + "ttl": "1h", + } + assert _count_cache_control(processed) == 1 + + +@pytest.mark.asyncio +async def test_cache_control_hook_bedrock_payload_caps_cachepoints_at_four(): + """End-to-end: outgoing Bedrock payload must not exceed 4 cachePoint blocks. + + Reproduces the customer report where 4 client cache_control system blocks + plus auto-inject produced 5 cachePoint blocks and Bedrock returned 400. + """ + with patch.dict( + os.environ, + { + "AWS_ACCESS_KEY_ID": "fake_access_key_id", + "AWS_SECRET_ACCESS_KEY": "fake_secret_access_key", + "AWS_REGION_NAME": "us-east-1", + }, + ): + litellm.callbacks = [AnthropicCacheControlHook()] + + mock_response = MagicMock() + mock_response.json.return_value = { + "output": {"message": {"role": "assistant", "content": "ok"}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 100, "outputTokens": 4, "totalTokens": 104}, + } + mock_response.status_code = 200 + + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=mock_response) as mock_post: + messages = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": f"System block {i}", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + } + for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + await litellm.acompletion( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + max_tokens=32, + cache_control_injection_points=_build_injection_points(), + client=client, + ) + + request_body = json.loads(mock_post.call_args.kwargs["data"]) + + cache_points = sum( + 1 + for block in request_body.get("system", []) + if isinstance(block, dict) and "cachePoint" in block + ) + for msg in request_body.get("messages", []): + content = msg.get("content", []) + if isinstance(content, list): + cache_points += sum( + 1 + for block in content + if isinstance(block, dict) and "cachePoint" in block + ) + + assert cache_points <= 4, ( + f"Bedrock payload exceeded Anthropic's 4 cache_control block limit: " + f"found {cache_points} cachePoint blocks" + ) + + +def test_cache_control_hook_reserves_slot_for_tool_config_point(): + """A tool_config injection point consumes one of the 4 slots downstream. + + With role:system targeting 4 system messages plus a tool_config point, the + hook must inject at most 3 message-level blocks so the tool_config cachePoint + appended by the Bedrock transform keeps the total at 4, not 5. + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + {"role": "system", "content": f"System {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, non_default_params = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": [ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + {"location": "tool_config"}, + ] + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert _count_cache_control(processed) == 3 + # The tool_config point is passed through for the provider transform. + assert non_default_params["cache_control_injection_points"] == [ + {"location": "tool_config"} + ] + + +@pytest.mark.asyncio +async def test_cache_control_hook_bedrock_payload_caps_with_tool_config_point(): + """End-to-end: message + tool_config injection must not exceed 4 cachePoints.""" + with patch.dict( + os.environ, + { + "AWS_ACCESS_KEY_ID": "fake_access_key_id", + "AWS_SECRET_ACCESS_KEY": "fake_secret_access_key", + "AWS_REGION_NAME": "us-east-1", + }, + ): + litellm.callbacks = [AnthropicCacheControlHook()] + + mock_response = MagicMock() + mock_response.json.return_value = { + "output": {"message": {"role": "assistant", "content": "ok"}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 100, "outputTokens": 4, "totalTokens": 104}, + } + mock_response.status_code = 200 + + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=mock_response) as mock_post: + messages = [ + {"role": "system", "content": f"System block {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "What is the weather?"}) + + await litellm.acompletion( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + max_tokens=32, + tools=[ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + }, + } + ], + cache_control_injection_points=[ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + {"location": "tool_config"}, + ], + client=client, + ) + + request_body = json.loads(mock_post.call_args.kwargs["data"]) + + cache_points = sum( + 1 + for block in request_body.get("system", []) + if isinstance(block, dict) and "cachePoint" in block + ) + for msg in request_body.get("messages", []): + content = msg.get("content", []) + if isinstance(content, list): + cache_points += sum( + 1 + for block in content + if isinstance(block, dict) and "cachePoint" in block + ) + for tool in request_body.get("toolConfig", {}).get("tools", []): + if isinstance(tool, dict) and "cachePoint" in tool: + cache_points += 1 + + assert cache_points <= 4, ( + f"Bedrock payload exceeded Anthropic's 4 cache_control block limit " + f"when mixing message and tool_config injection: found {cache_points}" + ) From 3ab500c79e9958a8c594c2d1d5cc0427f9cc1afc Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 16 Jun 2026 11:17:03 -0700 Subject: [PATCH 2/6] fix(guardrails): run pre_call hook once for model-level guardrails (#30543) * fix(guardrails): run pre_call hook once for model-level guardrails A CustomGuardrail attached to a deployment via litellm_params.guardrails gets its async_pre_call_hook invoked twice per request: once by the proxy pre-call loop and again by async_pre_call_deployment_hook after the router spreads the model-level guardrails into the top-level request kwargs. Record in request metadata that the proxy pre-call loop already ran a given guardrail, and have the deployment hook skip it when the marker is present. Direct-SDK usage never runs the proxy loop, so the deployment hook stays the sole invocation there and still fires exactly once. The marker key is stripped from untrusted caller metadata so a request body cannot suppress a model-only guardrail by pre-seeding it. * fix(guardrails): mark pre_call dedup on the post-hook request data Record the exactly-once marker after async_pre_call_hook runs, on the data object that flows downstream, rather than before it. A guardrail whose hook returns a brand-new request dict (instead of mutating or spreading the one it received) would otherwise discard the marker, letting the deployment hook re-run the guardrail a second time. (cherry picked from commit 4faeabc2541912bfec38c2d755cbad0ae3394671) --- litellm/constants.py | 4 + litellm/integrations/custom_guardrail.py | 54 +++++++ litellm/proxy/common_utils/callback_utils.py | 2 + litellm/proxy/litellm_pre_call_utils.py | 2 + .../proxy/policy_engine/pipeline_executor.py | 4 + litellm/proxy/utils.py | 2 + .../integrations/test_custom_guardrail.py | 106 ++++++++++++ .../proxy/test_model_level_guardrails.py | 152 +++++++++++++++++- 8 files changed, 325 insertions(+), 1 deletion(-) diff --git a/litellm/constants.py b/litellm/constants.py index 711c8a413d3..3041c498c1d 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -190,6 +190,10 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int( # Override with LITELLM_MAX_CALLBACKS env var for large deployments (e.g., many teams with guardrails) MAX_CALLBACKS = get_env_int("LITELLM_MAX_CALLBACKS", 100) +# Metadata key recording which pre_call guardrails the proxy loop already ran, +# so the deployment-level hook does not re-run them for the same request +PRE_CALL_EXECUTED_GUARDRAILS_KEY = "_pre_call_executed_guardrails" + # Generic fallback for unknown models DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index fc5f0429b63..059658991fb 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1,3 +1,4 @@ +import secrets from datetime import datetime from typing import ( TYPE_CHECKING, @@ -43,6 +44,7 @@ if TYPE_CHECKING: dc = DualCache() +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.exceptions import ( BlockedPiiEntityError, GuardrailRaisedException, @@ -50,6 +52,12 @@ from litellm.exceptions import ( SensitiveDataRouteException, ) +# Per-process secret tagging each recorded marker. The deployment hook only +# honors markers carrying this token, so a caller cannot forge the metadata +# field to suppress a guardrail on the direct-SDK path that never reaches the +# proxy's metadata sanitizer. +_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16) + def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]: """Extract session_id from request data (litellm_session_id or metadata).""" @@ -458,6 +466,49 @@ class CustomGuardrail(CustomLogger): return False + def _pre_call_marker(self) -> Optional[str]: + name = self.guardrail_name + if not name: + return None + return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}" + + def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None: + """ + Record that this guardrail's ``async_pre_call_hook`` already ran for this + request, so the deployment-level hook does not run it a second time. + + The proxy runs pre-call guardrails in ``ProxyLogging.pre_call_hook``. The + router later spreads a deployment's model-level ``guardrails`` into the + top-level request kwargs, which would otherwise re-trigger the same hook + from ``async_pre_call_deployment_hook``. + """ + marker = self._pre_call_marker() + if marker is None: + return + for meta_key in ("metadata", "litellm_metadata"): + meta = data.get(meta_key) + if isinstance(meta, dict): + executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY) + if isinstance(executed, list): + if marker not in executed: + executed.append(marker) + else: + meta[PRE_CALL_EXECUTED_GUARDRAILS_KEY] = [marker] + return + data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]} + + def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool: + marker = self._pre_call_marker() + if marker is None: + return False + for meta_key in ("metadata", "litellm_metadata"): + meta = data.get(meta_key) + if isinstance(meta, dict): + executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY) + if isinstance(executed, list) and marker in executed: + return True + return False + async def async_pre_call_deployment_hook( self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] ) -> Optional[dict]: @@ -468,6 +519,9 @@ class CustomGuardrail(CustomLogger): if litellm_guardrails is None or not isinstance(litellm_guardrails, list): return kwargs + if self._pre_call_hook_already_ran(kwargs): + return kwargs + if ( self.should_run_guardrail( data=kwargs, event_type=GuardrailEventHooks.pre_call diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a65e737f248..ffd2ac68c54 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, import litellm from litellm import get_secret from litellm._logging import verbose_proxy_logger +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams @@ -490,6 +491,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( "guardrail_config", "_guardrail_pipelines", "_pipeline_managed_guardrails", + PRE_CALL_EXECUTED_GUARDRAILS_KEY, "disable_global_guardrails", "disable_global_guardrail", "opted_out_global_guardrails", diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 7666b23f2af..22657f6250c 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -13,6 +13,7 @@ from starlette.datastructures import Headers import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host @@ -161,6 +162,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = ( "secret_fields", "_guardrail_pipelines", "_pipeline_managed_guardrails", + PRE_CALL_EXECUTED_GUARDRAILS_KEY, ) _UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS = frozenset( diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index 3c5a1d67be4..e46e3e1dc9f 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -171,6 +171,10 @@ class PipelineExecutor: data=data, call_type=call_type, # type: ignore ) + if isinstance(callback, CustomGuardrail): + callback.mark_pre_call_hook_ran(data) + if isinstance(response, dict): + callback.mark_pre_call_hook_ran(response) elif mode == "post_call": response = await target.async_post_call_success_hook( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d86cd38a51c..2c41b8155b9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1155,6 +1155,8 @@ class ProxyLogging: response=response, data=data, call_type=call_type ) + callback.mark_pre_call_hook_ran(data) + except SensitiveDataRouteException: status = "intervened" raise diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index f0bc7b8ebed..57fb0fe6714 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -84,6 +84,112 @@ class TestCustomGuardrailDeploymentHook: assert result["messages"] == mock_result["messages"] assert result["messages"] != original_messages + @pytest.mark.asyncio + async def test_deployment_hook_skips_when_pre_call_already_ran(self): + """The deployment hook must not re-run async_pre_call_hook once the proxy + pre-call loop has already run it for this request.""" + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {}, + } + + guardrail.mark_pre_call_hook_ran(kwargs) + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 0 + + @pytest.mark.asyncio + async def test_deployment_hook_runs_when_not_marked(self): + """Without the proxy marker (direct-SDK usage) the deployment hook is the + only execution path and must still run the guardrail exactly once.""" + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {}, + } + + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 1 + + def test_mark_pre_call_hook_ran_uses_litellm_metadata(self): + """The marker is recorded in litellm_metadata when that is the metadata + bucket in use, and is then visible to the skip check.""" + from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY + + guardrail = CustomGuardrail(guardrail_name="g1") + kwargs = {"litellm_metadata": {}} + + guardrail.mark_pre_call_hook_ran(kwargs) + + assert kwargs["litellm_metadata"][PRE_CALL_EXECUTED_GUARDRAILS_KEY] + assert guardrail._pre_call_hook_already_ran(kwargs) is True + + @pytest.mark.asyncio + async def test_deployment_hook_ignores_forged_caller_marker(self): + """A direct-SDK caller controls request metadata but cannot know the + per-process token, so a hand-crafted marker must not suppress a + requested guardrail in async_pre_call_deployment_hook.""" + from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {PRE_CALL_EXECUTED_GUARDRAILS_KEY: ["g1"]}, + } + + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 1 + class TestCustomGuardrailShouldRunGuardrail: diff --git a/tests/test_litellm/proxy/test_model_level_guardrails.py b/tests/test_litellm/proxy/test_model_level_guardrails.py index 3d74edd772b..9a79fa7f496 100644 --- a/tests/test_litellm/proxy/test_model_level_guardrails.py +++ b/tests/test_litellm/proxy/test_model_level_guardrails.py @@ -19,7 +19,6 @@ from litellm.proxy.utils import ( _merge_guardrails_with_existing, ) - # --------------------------------------------------------------------------- # Unit tests for _check_and_merge_model_level_guardrails # --------------------------------------------------------------------------- @@ -159,6 +158,157 @@ class TestCheckAndMergeModelLevelGuardrails: assert "existing" in result["metadata"]["guardrails"] +# --------------------------------------------------------------------------- +# Regression test: pre_call hook must run exactly once with model-level guardrails +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_pre_call_hook_runs_once_with_model_level_guardrails(): + """ + A guardrail attached at the model level (litellm_params.guardrails) is + spread into the top-level request kwargs by the router. The proxy pre-call + loop (async_pre_call_hook) and the deployment-level hook + (async_pre_call_deployment_hook) must together invoke async_pre_call_hook + exactly once, not twice. + """ + from litellm.caching.caching import DualCache + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy._types import CallTypes, UserAPIKeyAuth + from litellm.proxy.utils import ProxyLogging + from litellm.types.guardrails import GuardrailEventHooks + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="counting-guardrail", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + self.pre_call_count = 0 + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + + with patch("litellm.callbacks", [guardrail]): + ProxyLogging._callback_capabilities_cache.clear() + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hello"}], + "metadata": {}, + } + + # Path A: proxy pre-call loop runs the guardrail and records that it ran + data = await proxy_logging.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=data, + call_type="acompletion", + ) + + # Path B: the router spreads the deployment's model-level guardrails into + # the top-level kwargs, then litellm.acompletion fires the deployment hook + data["guardrails"] = ["counting-guardrail"] + await guardrail.async_pre_call_deployment_hook(data, CallTypes.acompletion) + + assert guardrail.pre_call_count == 1 + + +@pytest.mark.asyncio +async def test_pre_call_hook_runs_once_when_hook_returns_fresh_dict(): + """ + async_pre_call_hook may return a brand-new request dict instead of mutating + or spreading the one it received. The exactly-once marker must live on the + data that flows downstream, so the deployment hook still skips the guardrail + even when the proxy loop swapped in a fresh dict that never carried it. + """ + from litellm.caching.caching import DualCache + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy._types import CallTypes, UserAPIKeyAuth + from litellm.proxy.utils import ProxyLogging + from litellm.types.guardrails import GuardrailEventHooks + + class FreshDictGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="counting-guardrail", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + self.pre_call_count = 0 + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.pre_call_count += 1 + return {"model": data["model"], "messages": data["messages"]} + + guardrail = FreshDictGuardrail() + + with patch("litellm.callbacks", [guardrail]): + ProxyLogging._callback_capabilities_cache.clear() + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hello"}], + "metadata": {}, + } + + data = await proxy_logging.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=data, + call_type="acompletion", + ) + + data["guardrails"] = ["counting-guardrail"] + await guardrail.async_pre_call_deployment_hook(data, CallTypes.acompletion) + + assert guardrail.pre_call_count == 1 + + +@pytest.mark.asyncio +async def test_deployment_hook_runs_pre_call_without_proxy_loop(): + """ + Direct-SDK usage (litellm.acompletion(..., guardrails=[...]) without the + proxy) never runs the proxy pre-call loop, so the deployment hook is the + only place the guardrail executes and it must still run exactly once. + """ + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy._types import CallTypes + from litellm.types.guardrails import GuardrailEventHooks + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="counting-guardrail", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + self.pre_call_count = 0 + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hello"}], + "guardrails": ["counting-guardrail"], + "metadata": {}, + } + + await guardrail.async_pre_call_deployment_hook(data, CallTypes.acompletion) + + assert guardrail.pre_call_count == 1 + + # --------------------------------------------------------------------------- # Integration test: post_call_success_hook with model-level guardrails # --------------------------------------------------------------------------- From 9683808366a3ee4348395f025da82263b824cfb7 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 16 Jun 2026 11:17:49 -0700 Subject: [PATCH 3/6] fix(guardrails): stop re-initializing DB guardrails on every poll (#30542) * fix(guardrails): stop re-initializing DB guardrails on every poll InMemoryGuardrailHandler._has_guardrail_params_changed compared the in-memory LitellmParams against the raw dict loaded from the DB. The in-memory side carries every field default and coerces enums via model_dump(), while the DB side only holds the keys originally stored, so the two shapes never compared equal and the guardrail was rebuilt on every poll cycle. Each rebuild created a fresh instance, but delete_in_memory_guardrail only removed the old callback from litellm.callbacks. Request handling promotes guardrail callbacks into the success/failure/async lists, so the previous instance stayed referenced there and instances accumulated. Normalize both sides through LitellmParams(...).model_dump() before diffing, and purge the callback from every callback list on delete. * refactor(guardrails): narrow params-normalization fallback to ValidationError The comparison normalizer caught a bare Exception and silently fell back to the raw dict, which hid the cause and quietly degraded the affected guardrail back to re-initializing on every poll. Catch only the ValidationError that LitellmParams construction can raise, log a warning so the offending row is diagnosable, and let any other error surface instead of being swallowed. * refactor(callbacks): add remove_callback_from_all_lists helper to manager Move the knowledge of which callback lists a callback can be promoted into out of the guardrail registry and into LoggingCallbackManager, where the rest of the callback-list bookkeeping already lives. delete_in_memory_guardrail now delegates to the new helper instead of iterating the lists itself. (cherry picked from commit 9fa74ad8b4a3cf206c847d92d20c1bd20daa2b69) --- .../logging_callback_manager.py | 16 ++ .../proxy/guardrails/guardrail_registry.py | 64 +++++-- .../test_logging_callback_manager.py | 23 +++ .../guardrails/test_guardrail_registry.py | 169 ++++++++++++++++++ 4 files changed, 253 insertions(+), 19 deletions(-) diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 6c749118dec..b7adda3a9a4 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -394,6 +394,22 @@ class LoggingCallbackManager: + litellm._async_failure_callback ) + def remove_callback_from_all_lists(self, obj, require_self=False) -> None: + """ + Remove a callback object from every callback list it may have been + promoted into, so a re-initialized callback leaves no stale instance behind. + """ + for callback_list in ( + litellm.callbacks, + litellm.success_callback, + litellm.failure_callback, + litellm._async_success_callback, + litellm._async_failure_callback, + ): + self.remove_callback_from_list_by_object( + callback_list, obj, require_self=require_self + ) + def get_active_additional_logging_utils_from_custom_logger( self, ) -> Set[AdditionalLoggingUtils]: diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index aafcc5f1819..d589217a7ea 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -5,6 +5,8 @@ import os from datetime import datetime, timezone from typing import Any, Dict, List, Literal, Optional, Set, Type, cast +from pydantic import ValidationError + import litellm from litellm import Router from litellm._logging import verbose_proxy_logger @@ -598,21 +600,25 @@ class InMemoryGuardrailHandler: def delete_in_memory_guardrail(self, guardrail_id: str) -> None: """ Delete a guardrail in memory and remove from litellm callbacks. + + The callback is purged from every callback list, not just + litellm.callbacks: request handling promotes guardrail callbacks into the + success/failure/async lists, so removing it from only litellm.callbacks + leaves the old instance stranded in those lists on every re-initialization. """ # Remove from in-memory storage self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None) self._sources.pop(guardrail_id, None) - # Remove the callback from litellm.callbacks custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.pop( guardrail_id, None ) - if custom_guardrail_callback: - litellm.logging_callback_manager.remove_callback_from_list_by_object( - callback_list=litellm.callbacks, - obj=custom_guardrail_callback, - require_self=False, - ) + if custom_guardrail_callback is None: + return + + litellm.logging_callback_manager.remove_callback_from_all_lists( + custom_guardrail_callback + ) def list_in_memory_guardrails(self) -> List[Guardrail]: """ @@ -654,6 +660,34 @@ class InMemoryGuardrailHandler: self.delete_in_memory_guardrail(guardrail_id) return stale_ids + @staticmethod + def _normalize_litellm_params_for_comparison( + params: Optional[Any], + ) -> Optional[Dict[str, Any]]: + """ + Render litellm_params to a canonical dict so an in-memory LitellmParams and + the raw dict loaded from the DB compare equal when they describe the same + config. The in-memory side is a LitellmParams whose model_dump() carries + every field default and coerces enums, while the DB side is the raw stored + dict holding only the keys originally provided. Comparing those two shapes + directly never matches, so each DB poll would re-initialize the guardrail + forever; normalizing both through LitellmParams keeps the diff meaningful. + """ + if params is None: + return None + if isinstance(params, LitellmParams): + return params.model_dump() + if isinstance(params, dict): + try: + return LitellmParams(**params).model_dump() + except ValidationError as e: + verbose_proxy_logger.warning( + f"Could not normalize guardrail litellm_params for comparison; " + f"treating the guardrail as changed. Error: {e}" + ) + return params + return params + def _has_guardrail_params_changed( self, guardrail_id: str, new_guardrail: Guardrail ) -> bool: @@ -670,19 +704,11 @@ class InMemoryGuardrailHandler: return True # Compare litellm_params - existing_params = existing.get("litellm_params") - new_params = new_guardrail.get("litellm_params") - - # Convert to dicts for comparison - existing_dict = ( - existing_params.model_dump() - if isinstance(existing_params, LitellmParams) - else existing_params + existing_dict = self._normalize_litellm_params_for_comparison( + existing.get("litellm_params") ) - new_dict = ( - new_params.model_dump() - if isinstance(new_params, LitellmParams) - else new_params + new_dict = self._normalize_litellm_params_for_comparison( + new_guardrail.get("litellm_params") ) # Compare and identify specific differences diff --git a/tests/litellm_utils_tests/test_logging_callback_manager.py b/tests/litellm_utils_tests/test_logging_callback_manager.py index d9540f8f850..d9bfca425e4 100644 --- a/tests/litellm_utils_tests/test_logging_callback_manager.py +++ b/tests/litellm_utils_tests/test_logging_callback_manager.py @@ -192,6 +192,29 @@ def test_remove_callback_from_list_by_object(): assert len(litellm._async_failure_callback) == 0 +def test_remove_callback_from_all_lists(): + manager = LoggingCallbackManager() + manager._reset_all_callbacks() + + class TestLogger(CustomLogger): + pass + + obj = TestLogger() + manager.add_litellm_callback(obj) + manager.add_litellm_success_callback(obj) + manager.add_litellm_failure_callback(obj) + manager.add_litellm_async_success_callback(obj) + manager.add_litellm_async_failure_callback(obj) + + manager.remove_callback_from_all_lists(obj) + + assert obj not in litellm.callbacks + assert obj not in litellm.success_callback + assert obj not in litellm.failure_callback + assert obj not in litellm._async_success_callback + assert obj not in litellm._async_failure_callback + + def test_reset_callbacks(callback_manager): # Add various callbacks callback_manager.add_litellm_callback("test") diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 9f7173383b0..0ef9ad857f9 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -180,3 +180,172 @@ def test_sync_guardrail_from_db_marks_source_db_when_unchanged(): handler.sync_guardrail_from_db(g) assert handler.get_source("collide") == "db" + + +def _db_litellm_params() -> dict: + """ + Shape produced by GuardrailRegistry.get_all_guardrails_from_db: litellm_params + is a raw dict (not a LitellmParams), holding only the keys originally stored, + a non-schema extra key, and plain-string enum values. + """ + return { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "default_on": True, + "version": 2, + "blocked_words": [{"keyword": "secret", "action": "BLOCK"}], + } + + +def test_unchanged_db_params_do_not_register_as_changed(): + """ + A DB poll returns litellm_params as a raw dict while the in-memory copy is a + LitellmParams whose model_dump() fills every field default and coerces enums. + The two shapes must compare equal when the config is identical; otherwise + every poll cycle re-initializes the guardrail indefinitely. + """ + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "11111111-1111-1111-1111-111111111111" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=dict(raw)) + assert handler._has_guardrail_params_changed(gid, new) is False + + +def test_changed_db_params_register_as_changed(): + """Normalizing both sides must still surface a genuine config change.""" + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "22222222-2222-2222-2222-222222222222" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + changed = {**raw, "blocked_words": [{"keyword": "different", "action": "BLOCK"}]} + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=changed) + assert handler._has_guardrail_params_changed(gid, new) is True + + +def test_unnormalizable_db_params_register_as_changed_without_raising(): + """ + A DB row whose litellm_params fail LitellmParams validation must not crash the + poll loop. The comparison falls back to treating the guardrail as changed so it + re-initializes (and surfaces the bad row in logs) rather than propagating the + validation error up through the polling cycle. + """ + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "55555555-5555-5555-5555-555555555555" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + malformed = {**raw, "default_on": "not-a-bool-xyz"} + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=malformed) + assert handler._has_guardrail_params_changed(gid, new) is True + + +def _all_callback_lists(): + import litellm + + return [ + litellm.callbacks, + litellm.success_callback, + litellm.failure_callback, + litellm._async_success_callback, + litellm._async_failure_callback, + ] + + +def test_delete_in_memory_guardrail_removes_callback_from_all_lists(): + """ + Request handling promotes guardrail callbacks from litellm.callbacks into the + success/failure/async lists. delete_in_memory_guardrail must purge the callback + from every list, otherwise a re-initialized guardrail leaves its old instance + stranded in those lists and instances accumulate. + """ + handler = InMemoryGuardrailHandler() + callback = CustomGuardrail( + guardrail_name="cf-delete", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + ) + gid = "33333333-3333-3333-3333-333333333333" + handler.IN_MEMORY_GUARDRAILS[gid] = _make_guardrail(gid, "cf-delete") + handler._sources[gid] = "db" + handler.guardrail_id_to_custom_guardrail[gid] = callback + + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + try: + for cb_list in lists: + cb_list.append(callback) + + handler.delete_in_memory_guardrail(gid) + + for cb_list in lists: + assert callback not in cb_list + finally: + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot + + +def test_repeated_db_sync_does_not_accumulate_runner_instances(): + """ + End-to-end regression for the OOM: across repeated DB polls (with the config + genuinely changing each cycle to force re-initialization), exactly one live + guardrail instance must exist across all callback lists. On the unfixed code + the stale instance lingers in the success/failure lists and the distinct count + climbs above one. + """ + import litellm + + handler = InMemoryGuardrailHandler() + gid = "44444444-4444-4444-4444-444444444444" + name = "cf-accum" + + def db_guardrail(word: str) -> Guardrail: + params = { + **_db_litellm_params(), + "blocked_words": [{"keyword": word, "action": "BLOCK"}], + } + return Guardrail(guardrail_id=gid, guardrail_name=name, litellm_params=params) + + def promote_into_request_lists() -> None: + manager = litellm.logging_callback_manager + for callback in list(litellm.callbacks): + manager.add_litellm_success_callback(callback) + manager.add_litellm_failure_callback(callback) + manager.add_litellm_async_success_callback(callback) + manager.add_litellm_async_failure_callback(callback) + + def distinct_runner_instances() -> int: + seen = set() + for callback in litellm.logging_callback_manager._get_all_callbacks(): + if ( + isinstance(callback, CustomGuardrail) + and getattr(callback, "guardrail_name", None) == name + ): + seen.add(id(callback)) + return len(seen) + + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + try: + for cycle in range(5): + handler.sync_guardrail_from_db(db_guardrail(f"word-{cycle}")) + promote_into_request_lists() + + assert distinct_runner_instances() == 1 + finally: + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot From ccf892d339ff3b8d38247811eaa7c9a50f0f7867 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 16 Jun 2026 18:56:14 -0700 Subject: [PATCH 4/6] fix(guardrails): return 400 not 500 when AIM blocks a request (#30573) * fix(guardrails): return 400 not 500 when AIM blocks a request AIM guardrail blocks raised a bare HTTPException whose type and param serialized as the literal string "None", which broke OpenAI-SDK error parsing for downstream consumers. Switching AIM to raise a ProxyException surfaced a second bug: the shared error funnel re-derived the HTTP status from a nonexistent status_code attribute and downgraded the 400 to a 500. The funnel now honors an already-normalized ProxyException rather than rebuilding it, and ProxyException is excluded from llm_exceptions alerting so a content-policy block no longer pages on-call as an LLM API failure Resolves LIT-3751 * fix(guardrails): route all AIM rejection paths through ProxyException The block-action fix left two AIM rejection paths raising a bare HTTPException: the multimodal anonymize rejection and the output-side block. Both serialized type and param as the literal string "None", the same malformed shape the block fix removed. Funnel all three through a shared _rejection helper so they return a conformant OpenAI error body. The output block carries content_policy_violation; the multimodal rejection stays a plain invalid_request_error because it is a usage error, not a policy violation Resolves LIT-3751 * fix(guardrails): record AIM ProxyException blocks in failure logs Switching AIM blocks from HTTPException to ProxyException made _is_proxy_only_llm_api_error return False for them, so _handle_logging_proxy_only_error was skipped and the blocked prompt was dropped from the configured failure loggers. Classify ProxyException as a proxy-only error alongside HTTPException so guardrail blocks are recorded again, matching the prior behavior. The llm_exceptions alert suppression is a separate check and stays in place Resolves LIT-3751 * style(guardrails): use str | None over Optional[str] in AIM _rejection * style(guardrails): collapse AIM _rejection signature per black (cherry picked from commit b5fcd859bec1388267f6f1f9affc125190555525) --- litellm/proxy/common_request_processing.py | 7 + .../guardrails/guardrail_hooks/aim/aim.py | 34 +++-- litellm/proxy/utils.py | 5 +- tests/local_testing/test_aim_guardrails.py | 135 +++++++++++++++++- .../proxy/test_common_request_processing.py | 35 +++++ tests/test_litellm/proxy/test_proxy_utils.py | 116 ++++++++++++++- 6 files changed, 313 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6558543370d..6020d528515 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1907,6 +1907,13 @@ class ProxyBaseLLMRequestProcessing: except Exception: pass + if isinstance(e, ProxyException): + e.headers = { + **e.headers, + **{k: v if isinstance(v, str) else str(v) for k, v in headers.items()}, + } + raise e + if isinstance(e, HTTPException): raw_detail = getattr(e, "detail", str(e)) message, structured_fields = _serialize_http_exception_detail(raw_detail) diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 5b5f91195e7..d70c8e4f310 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -9,7 +9,6 @@ import json import os from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union -from fastapi import HTTPException from pydantic import BaseModel from websockets.asyncio.client import ClientConnection, connect @@ -21,7 +20,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.guardrails._content_utils import ( apply_redacted_messages_back, build_inspection_messages, @@ -129,6 +128,16 @@ class AimGuardrail(CustomGuardrail): verbose_proxy_logger.error(f"Aim: {action_type} action") return data + @staticmethod + def _rejection(message: str, *, openai_code: str | None = None) -> ProxyException: + return ProxyException( + message=message, + type="invalid_request_error", + param=None, + code=400, + openai_code=openai_code, + ) + def _handle_block_action(self, analysis_result: Any, required_action: Any) -> None: detection_message = required_action.get("detection_message", None) verbose_proxy_logger.info( @@ -136,7 +145,7 @@ class AimGuardrail(CustomGuardrail): policies=list(analysis_result["policy_drill_down"].keys()), ), ) - raise HTTPException(status_code=400, detail=detection_message) + raise self._rejection(detection_message, openai_code="content_policy_violation") def _anonymize_request(self, res: Any, data: dict) -> dict: verbose_proxy_logger.info("Aim: anonymize action") @@ -148,14 +157,11 @@ class AimGuardrail(CustomGuardrail): # parts from a multimodal request — degrade to block so the # multimodal payload is never silently rewritten. if has_non_string_content(data): - raise HTTPException( - status_code=400, - detail=( - "Aim: anonymize action requested for multimodal input " - "but mask-in-place would drop non-text parts. Send the " - "request with plain string content to use anonymize, " - "or rely on block-mode policies." - ), + raise self._rejection( + "Aim: anonymize action requested for multimodal input " + "but mask-in-place would drop non-text parts. Send the " + "request with plain string content to use anonymize, " + "or rely on block-mode policies." ) redacted_messages = [ { @@ -287,9 +293,9 @@ class AimGuardrail(CustomGuardrail): if aim_output_guardrail_result and aim_output_guardrail_result.get( "detection_message" ): - raise HTTPException( - status_code=400, - detail=aim_output_guardrail_result.get("detection_message"), + raise self._rejection( + aim_output_guardrail_result.get("detection_message"), + openai_code="content_policy_violation", ) if aim_output_guardrail_result and aim_output_guardrail_result.get( "redacted_output" diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2c41b8155b9..c1c2479b7b8 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2072,7 +2072,7 @@ class ProxyLogging: litellm_call_id=request_data.get("litellm_call_id", ""), status="fail" ) if AlertType.llm_exceptions in self.alert_types and not isinstance( - original_exception, HTTPException + original_exception, (HTTPException, ProxyException) ): """ Just alert on LLM API exceptions. Do not alert on user errors @@ -2176,6 +2176,7 @@ class ProxyLogging: e.g should only return True for: - Authentication Errors from user_api_key_auth - HTTP HTTPException (rate limit errors) + - ProxyException (guardrail blocks, budget / rate-limit errors) """ ######################################################### @@ -2192,7 +2193,7 @@ class ProxyLogging: ): return False - return isinstance(original_exception, HTTPException) or ( + return isinstance(original_exception, (HTTPException, ProxyException)) or ( error_type == ProxyErrorTypes.auth_error ) diff --git a/tests/local_testing/test_aim_guardrails.py b/tests/local_testing/test_aim_guardrails.py index 31416c565c1..2cb7f9cd357 100644 --- a/tests/local_testing/test_aim_guardrails.py +++ b/tests/local_testing/test_aim_guardrails.py @@ -6,10 +6,10 @@ import sys from unittest.mock import AsyncMock, patch, call import pytest -from fastapi.exceptions import HTTPException from httpx import Request, Response from litellm import DualCache +from litellm.proxy._types import ProxyException from litellm.proxy.guardrails.guardrail_hooks.aim.aim import ( AimGuardrail, AimGuardrailMissingSecrets, @@ -101,7 +101,7 @@ async def test_block_callback(mode: str): ], } - with pytest.raises(HTTPException, match="Jailbreak detected"): + with pytest.raises(ProxyException, match="Jailbreak detected") as exc_info: with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", return_value=Response( @@ -135,6 +135,137 @@ async def test_block_callback(mode: str): call_type="completion", ) + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code == "content_policy_violation" + + +@pytest.mark.asyncio +async def test_output_block_raises_proxy_exception(): + """An output-side block is a content-policy violation, like the input block: + it must surface a conformant ProxyException, not a bare HTTPException whose + type/param serialize as the literal string "None". Regression for LIT-3751.""" + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "gibberish-guard", + "litellm_params": { + "guardrail": "aim", + "mode": "post_call", + "api_key": "hs-aim-key", + }, + }, + ], + config_file_path="", + ) + aim_guardrails = [ + callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail) + ] + assert len(aim_guardrails) == 1 + aim_guardrail = aim_guardrails[0] + + block_on_output = Response( + json={ + "analysis_result": {"policy_drill_down": {"PII": {}}}, + "required_action": { + "action_type": "block_action", + "detection_message": "Output blocked: leaked secret", + "policy_name": "blocking policy", + }, + }, + status_code=200, + request=Request(method="POST", url="http://aim"), + ) + response = ModelResponse( + choices=[ + { + "finish_reason": "stop", + "index": 0, + "message": {"content": "here is the secret", "role": "assistant"}, + } + ] + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=block_on_output, + ): + with pytest.raises(ProxyException, match="Output blocked") as exc_info: + await aim_guardrail.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "tell me a secret"}]}, + response=response, + user_api_key_dict=UserAPIKeyAuth(), + ) + + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code == "content_policy_violation" + + +@pytest.mark.asyncio +async def test_anonymize_multimodal_rejection_raises_proxy_exception(): + """Anonymize on multimodal input degrades to a 400 because mask-in-place would + drop non-text parts. That is a usage error, not a content-policy violation, so + it must raise a conformant ProxyException WITHOUT the content_policy_violation + code. Regression for LIT-3751.""" + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "gibberish-guard", + "litellm_params": { + "guardrail": "aim", + "mode": "pre_call", + "api_key": "hs-aim-key", + }, + }, + ], + config_file_path="", + ) + aim_guardrails = [ + callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail) + ] + assert len(aim_guardrails) == 1 + aim_guardrail = aim_guardrails[0] + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hi my name is Brian"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ], + }, + ], + } + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=response_with_detections, + ): + with pytest.raises( + ProxyException, match="anonymize action requested for multimodal" + ) as exc_info: + await aim_guardrail.async_pre_call_hook( + data=data, + cache=DualCache(), + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code != "content_policy_violation" + @pytest.mark.asyncio @pytest.mark.parametrize("mode", ["pre_call", "during_call"]) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 0f5a0cbe4b6..478cc3c1403 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2269,6 +2269,41 @@ class TestHandleLLMApiExceptionDictDetail: assert proxy_exc.message == "Content blocked by guardrail" assert proxy_exc.provider_specific_fields is None + async def test_already_normalized_proxy_exception_is_honored(self): + """A ProxyException raised mid-request (e.g. a guardrail block) is already + the OpenAI wire format. The funnel must re-raise it untouched instead of + re-deriving the status from a (nonexistent) status_code attribute and + defaulting to 500. Regression for LIT-3751.""" + from litellm.proxy._types import ProxyException + + exc = ProxyException( + message='"Leroy Jenkins" detected as name', + type="invalid_request_error", + param=None, + code=400, + openai_code="content_policy_violation", + ) + proxy_exc = await self._invoke(exc) + assert proxy_exc is exc + assert proxy_exc.code == "400" + assert proxy_exc.type == "invalid_request_error" + assert proxy_exc.param is None + assert proxy_exc.openai_code == "content_policy_violation" + assert proxy_exc.message == '"Leroy Jenkins" detected as name' + + # The body the OpenAI-SDK client actually receives. The HTTP status line + # comes from int(exc.code) == 400; the wire ``code`` stays the status + # string. ``openai_code`` ("content_policy_violation") is intentionally + # NOT serialized here - to_dict() emits only ``code`` - so this asserts + # the real contract rather than the write-only attribute. + assert int(proxy_exc.code) == 400 + assert proxy_exc.to_dict() == { + "message": '"Leroy Jenkins" detected as name', + "type": "invalid_request_error", + "param": None, + "code": "400", + } + class TestAsyncStreamingDataGeneratorFastPath: """Fast/slow path branching in async_streaming_data_generator.""" diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index f0015d9df0d..ccbcbef212e 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -15,7 +15,7 @@ sys.path.insert( ) # Adds the parent directory to the system path -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from litellm.proxy.utils import get_custom_url, join_paths @@ -368,3 +368,117 @@ class TestPostCallFailureHookLiftsFirstApiCallStartTime: await self._run(request_data) assert "first_api_call_start_time" not in request_data assert "litellm_logging_obj" not in request_data + + +class TestPostCallFailureHookLLMExceptionAlerting: + """The llm_exceptions alert is for infra / LLM-API failures, not user + errors (https://github.com/BerriAI/litellm/issues/3395). Already-normalized + client errors must be excluded so a guardrail content-policy block never + pages on-call. ProxyException is such an error; before LIT-3751 only + HTTPException was excluded, so AIM blocks paged as if the LLM API failed.""" + + async def _alerted(self, exc) -> bool: + import asyncio + from unittest.mock import AsyncMock + + from litellm.proxy._types import AlertType, UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [AlertType.llm_exceptions] + alerting_handler = AsyncMock() + with ( + patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()), + patch.object(proxy_logging_obj, "alerting_handler", new=alerting_handler), + ): + await proxy_logging_obj.post_call_failure_hook( + request_data={}, + original_exception=exc, + user_api_key_dict=UserAPIKeyAuth(), + ) + await asyncio.sleep(0) # let the fire-and-forget alert task run + return alerting_handler.called + + @pytest.mark.asyncio + async def test_proxy_exception_does_not_alert(self): + from litellm.proxy._types import ProxyException + + exc = ProxyException( + message="content blocked", + type="invalid_request_error", + param=None, + code=400, + openai_code="content_policy_violation", + ) + assert await self._alerted(exc) is False + + @pytest.mark.asyncio + async def test_http_exception_does_not_alert(self): + assert ( + await self._alerted(HTTPException(status_code=400, detail="blocked")) + is False + ) + + @pytest.mark.asyncio + async def test_genuine_llm_api_error_still_alerts(self): + assert await self._alerted(Exception("upstream 503")) is True + + +class TestPostCallFailureHookProxyExceptionLogging: + """A guardrail block raises a ProxyException; on an LLM route it must still + drive proxy-only failure logging (_handle_logging_proxy_only_error) so the + blocked request is recorded, exactly as the old HTTPException did. Before + LIT-3751 the classifier only matched HTTPException, so switching AIM to + ProxyException silently dropped the rejected prompt from failure logs.""" + + async def _logged(self, exc, *, request_route) -> bool: + from unittest.mock import AsyncMock + + from litellm.proxy._types import UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + handle_mock = AsyncMock() + with ( + patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()), + patch.object( + proxy_logging_obj, + "_handle_logging_proxy_only_error", + new=handle_mock, + ), + ): + await proxy_logging_obj.post_call_failure_hook( + request_data={}, + original_exception=exc, + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-test", request_route=request_route + ), + ) + return handle_mock.await_count > 0 + + def _block(self): + from litellm.proxy._types import ProxyException + + return ProxyException( + message="content blocked", + type="invalid_request_error", + param=None, + code=400, + openai_code="content_policy_violation", + ) + + @pytest.mark.asyncio + async def test_proxy_exception_on_llm_route_is_logged(self): + assert ( + await self._logged(self._block(), request_route="/v1/chat/completions") + is True + ) + + @pytest.mark.asyncio + async def test_generic_exception_on_llm_route_is_not_logged(self): + # A raw provider/unknown exception is logged by the LLM call path, not here. + assert ( + await self._logged( + Exception("upstream 503"), request_route="/v1/chat/completions" + ) + is False + ) From d5963763112560c551436aece6de7e5634c18be4 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 20 Jun 2026 12:05:10 -0700 Subject: [PATCH 5/6] =?UTF-8?q?bump:=20version=201.89.2=20=E2=86=92=201.89?= =?UTF-8?q?.3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index b5ef93f033c..b338220bf51 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm" -version = "1.89.2" +version = "1.89.3" description = "Library to easily interface with LLM API providers" readme = "README.md" requires-python = ">=3.10, <3.14" @@ -264,7 +264,7 @@ source-exclude = [ profile = "black" [tool.commitizen] -version = "1.89.2" +version = "1.89.3" version_files = [ "pyproject.toml:^version", ] From b4bf258bff8978a97556154f466d6622b9edcc8b Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 20 Jun 2026 12:05:34 -0700 Subject: [PATCH 6/6] chore: refresh uv.lock for 1.89.3 --- uv.lock | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/uv.lock b/uv.lock index efda9444fed..8bd9645ba90 100644 --- a/uv.lock +++ b/uv.lock @@ -9,7 +9,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-06-15T02:02:32.823508Z" +exclude-newer = "2026-06-17T19:05:26.692494Z" exclude-newer-span = "P3D" [manifest] @@ -3280,7 +3280,7 @@ wheels = [ [[package]] name = "litellm" -version = "1.89.2" +version = "1.89.3" source = { editable = "." } dependencies = [ { name = "aiohttp" },