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/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/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/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/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/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/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/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/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..c1c2479b7b8 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 @@ -2070,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 @@ -2174,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) """ ######################################################### @@ -2190,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/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", ] 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/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/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}" + ) 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/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 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_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 # --------------------------------------------------------------------------- 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 + ) 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" },