Merge pull request #30888 from BerriAI/litellm_backport_1_89_x_0620

chore(release): backport #30480, #30543, #30542, #30573 to stable/1.89.x and cut 1.89.3
This commit is contained in:
yuneng-jiang 2026-06-20 14:45:08 -07:00 committed by GitHub
commit 33df5891a9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
21 changed files with 1371 additions and 76 deletions

View file

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

View file

@ -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(

View file

@ -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

View file

@ -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]:

View file

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

View file

@ -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",

View file

@ -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"

View file

@ -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

View file

@ -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(

View file

@ -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,

View file

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

View file

@ -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",
]

View file

@ -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")

View file

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

View file

@ -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}"
)

View file

@ -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:

View file

@ -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

View file

@ -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."""

View file

@ -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
# ---------------------------------------------------------------------------

View file

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

4
uv.lock generated
View file

@ -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" },