mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
commit
33df5891a9
21 changed files with 1371 additions and 76 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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
4
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue