diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index aecb2552b53..3e50fe66039 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import ( record_output_item_chunk, record_output_text_chunk, ) +from litellm.responses.utils import normalize_responses_api_stream_options from litellm.types.llms.openai import ( ChatCompletionAnnotation, ChatCompletionReasoningItem, @@ -320,6 +321,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): responses_api_request["tool_choice"] = ( # type: ignore[assignment] self._normalize_tool_choice_for_responses_api(value) ) + elif key == "stream_options": + stream_options = normalize_responses_api_stream_options(value) + if stream_options is not None: + responses_api_request["stream_options"] = stream_options elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys(): responses_api_request[key] = value # type: ignore elif key == "previous_response_id": @@ -360,8 +365,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): continue if key == "instructions" and instructions: request_data["instructions"] = instructions - elif key == "stream_options" and isinstance(value, dict): - request_data["stream_options"] = value.get("include_obfuscation") elif key == "user" and isinstance(value, str): # OpenAI API requires user param to be max 64 chars - truncate if longer if len(value) <= 64: diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index f639ad49d5e..57b05c9bec8 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -17,7 +17,11 @@ from typing import ( ) from litellm._logging import verbose_logger -from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, + redact_nested_match_and_regex_keys, +) from litellm.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.secret_managers.main import str_to_bool @@ -107,6 +111,8 @@ class CustomGuardrail(CustomLogger): # If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path. use_native_during_call_hook: ClassVar[bool] = False + records_own_guardrail_information: ClassVar[bool] = False + def __init__( self, guardrail_name: Optional[str] = None, @@ -954,17 +960,8 @@ class CustomGuardrail(CustomLogger): # should not happen container[key] = [existing, slg] - if "metadata" in request_data: - if request_data["metadata"] is None: - request_data["metadata"] = {} - _append_guardrail_info(request_data["metadata"]) - elif "litellm_metadata" in request_data: - _append_guardrail_info(request_data["litellm_metadata"]) - else: - # Ensure guardrail info is always logged (e.g. proxy may not have set - # metadata yet). Attach to "metadata" so spend log / standard logging see it. - request_data["metadata"] = {} - _append_guardrail_info(request_data["metadata"]) + _, metadata_bucket = get_or_create_metadata_bucket(request_data) + _append_guardrail_info(metadata_bucket) _guardrail_self_recorded.set(True) @@ -1223,7 +1220,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object) """ if logging_obj is None: return - meta_src = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} slg_info = meta_src.get("standard_logging_guardrail_information") if not slg_info: return @@ -1256,6 +1253,14 @@ def log_guardrail_information(func): so it stays correct when guardrails run concurrently (asyncio copies the context into each gathered task): counting shared entries would let one guardrail's append hide another guardrail's missing record. + + A guardrail that only records an entry when it actually runs (e.g. + ``HeadroomGuardrail``, which returns the inputs untouched on an endpoint + whose payload it cannot act on) sets ``records_own_guardrail_information = + True`` so the auto-record is skipped even on the return paths where it + recorded nothing; otherwise a no-op early return would be logged as an + "allow"/"success" run even though the guardrail did nothing. The exception + branch below still records so a genuine failure is not lost. """ import functools import inspect @@ -1291,7 +1296,7 @@ def log_guardrail_information(func): self_recorded_token = _guardrail_self_recorded.set(False) try: response = await func(*args, **kwargs) - if _guardrail_self_recorded.get(): + if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( response=response, @@ -1333,7 +1338,7 @@ def log_guardrail_information(func): self_recorded_token = _guardrail_self_recorded.set(False) try: response = func(*args, **kwargs) - if _guardrail_self_recorded.get(): + if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( response=response, diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index fea55cd1db4..12465377b51 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -883,8 +883,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): request_data: dict, parent_span: Optional[Any], ) -> None: - """Emit ``guardrail`` spans from ``request_data["metadata"] - ["standard_logging_guardrail_information"]``. + """Emit ``guardrail`` spans from the request's proxy-internal metadata bucket + (``standard_logging_guardrail_information``). Routed through ``_create_guardrail_span`` so the dedupe state in ``_otel_internal`` is honoured — if ``_handle_failure`` already @@ -892,7 +892,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): """ from opentelemetry import trace as _trace - metadata = (request_data or {}).get("metadata") or {} + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + ) + + request_data = request_data or {} + metadata = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} guardrail_information = metadata.get("standard_logging_guardrail_information") if not guardrail_information: return diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 88dddb59cc7..cecc35ee1c1 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -195,6 +195,25 @@ def get_metadata_variable_name_from_kwargs( return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" +def get_or_create_metadata_bucket( + request_data: dict, +) -> tuple[Literal["metadata", "litellm_metadata"], dict]: + """ + Return the proxy-internal metadata bucket for this request, creating it if absent. + + Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI + ``metadata`` field can remain provider-safe (string values only). Every writer and + reader of proxy-internal metadata resolves the bucket through here, so a caller that + supplies its own ``metadata`` field cannot split them across two dicts. + """ + metadata_key = get_metadata_variable_name_from_kwargs(request_data) + metadata_bucket = request_data.get(metadata_key) + if not isinstance(metadata_bucket, dict): + metadata_bucket = {} + request_data[metadata_key] = metadata_bucket + return metadata_key, metadata_bucket + + def get_litellm_metadata_from_kwargs(kwargs: dict): """ Helper to get litellm metadata from all litellm request kwargs diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 7000c20d9c4..90f735707bf 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -600,9 +600,15 @@ class AnthropicMessagesHandler(BaseTranslation): guardrail_inputs["tool_calls"] = tool_calls_list try: + prepared_request_data = self._prepare_request_data( + request_data, + model_response, + user_api_key_dict, + key="response", + ) _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=guardrail_inputs, - request_data=request_data if request_data is not None else {}, + request_data=prepared_request_data, input_type="response", logging_obj=litellm_logging_obj, ) @@ -618,9 +624,15 @@ class AnthropicMessagesHandler(BaseTranslation): string_so_far = self.get_streaming_string_so_far(responses_so_far) try: + prepared_request_data = self._prepare_request_data( + request_data, + responses_so_far, + user_api_key_dict, + key="responses", + ) _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs={"texts": [string_so_far]}, - request_data=request_data if request_data is not None else {}, + request_data=prepared_request_data, input_type="response", logging_obj=litellm_logging_obj, ) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a9c2a12aff7..33bca782e0b 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -1,12 +1,16 @@ import copy import os -from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional +from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional 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.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams from litellm.proxy.common_utils.encrypt_decrypt_utils import ( @@ -406,23 +410,6 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]: return headers -def get_metadata_variable_name_from_kwargs( - kwargs: dict, -) -> Literal["metadata", "litellm_metadata"]: - """ - Helper to return what the "metadata" field should be called in the request data - - - New endpoints return `litellm_metadata` - - Old endpoints return `metadata` - - Context: - - LiteLLM used `metadata` as an internal field for storing metadata - - OpenAI then started using this field for their metadata - - LiteLLM is now moving to using `litellm_metadata` for our metadata - """ - return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" - - LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( { "applied_policies", @@ -450,23 +437,6 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( ) -def _get_or_create_proxy_metadata_bucket( - request_data: Dict, -) -> tuple[Literal["metadata", "litellm_metadata"], dict]: - """ - Return the proxy-internal metadata bucket for this request. - - Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI - ``metadata`` field can remain provider-safe (string values only). - """ - metadata_key = get_metadata_variable_name_from_kwargs(request_data) - metadata_bucket = request_data.get(metadata_key) - if not isinstance(metadata_bucket, dict): - metadata_bucket = {} - request_data[metadata_key] = metadata_bucket - return metadata_key, metadata_bucket - - def sanitize_openai_provider_metadata( metadata: Optional[Dict[str, Any]], ) -> Optional[Dict[str, str]]: @@ -496,7 +466,7 @@ def sanitize_openai_provider_metadata( def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_name: Optional[str]): if guardrail_name is None: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) if "applied_guardrails" in _metadata: if guardrail_name not in _metadata["applied_guardrails"]: _metadata["applied_guardrails"].append(guardrail_name) @@ -513,7 +483,7 @@ def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optio """ if policy_name is None: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) if "applied_policies" in _metadata: if policy_name not in _metadata["applied_policies"]: _metadata["applied_policies"].append(policy_name) @@ -531,7 +501,7 @@ def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str, """ if not policy_sources: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) existing = _metadata.get("policy_sources", {}) if not isinstance(existing, dict): existing = {} diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 2d67c22f0aa..ca3bb0ee361 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -4,7 +4,7 @@ import json import re import time import uuid -from typing import TYPE_CHECKING, Any, List, Literal, Optional +from typing import TYPE_CHECKING, Any, ClassVar, List, Literal, Optional import httpx from fastapi import HTTPException @@ -209,6 +209,8 @@ def _build_responses_followup_items( class HeadroomGuardrail(CustomGuardrail): + records_own_guardrail_information: ClassVar[bool] = True + @classmethod def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: return [ @@ -481,7 +483,21 @@ class HeadroomGuardrail(CustomGuardrail): ) end_time = time.time() + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + if not compression_succeeded: + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response={"error": "headroom compression unavailable; request forwarded uncompressed"}, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER, + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType] self.add_standard_logging_guardrail_information_to_request_data( @@ -493,6 +509,7 @@ class HeadroomGuardrail(CustomGuardrail): end_time=end_time, duration=end_time - start_time, ) + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) hashes = extract_hashes_from_messages(compressed) if not hashes: diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 28b9dec100f..f2c5a95202b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -30,6 +30,10 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( @@ -432,7 +436,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): Override to store only the Model Armor API response, not the entire data dict. This prevents circular references in logging. """ - metadata = (request_data.get("metadata") or {}) if isinstance(request_data, dict) else {} + metadata = ( + request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} + if isinstance(request_data, dict) + else {} + ) guardrail_response = metadata.get("_model_armor_response", {}) # Determine status – default to "success" but prefer the explicit value if present. @@ -471,7 +479,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): blocking, while fail_on_error still governs real Model Armor API errors. """ from litellm.proxy.common_utils.callback_utils import ( - _get_or_create_proxy_metadata_bucket, add_guardrail_to_applied_guardrails_header, ) @@ -491,7 +498,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) # Use the same metadata bucket the header helper writes to, so the logged Model Armor # payload and status land where _process_response reads them on every route. - _, metadata = _get_or_create_proxy_metadata_bucket(data) + _, metadata = get_or_create_metadata_bucket(data) fail_on_error = bool(self.optional_params.get("fail_on_error", True)) if unscannable_references > 0: @@ -607,7 +614,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # overwritten by another coroutine. blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content) if isinstance(data, dict): - metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request + _, metadata = get_or_create_metadata_bucket(data) # ensures metadata exists and is unique per request # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( metadata.get("_model_armor_response"), @@ -702,7 +709,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content) # Store the armor response for logging if isinstance(data, dict): - metadata = data.setdefault("metadata", {}) + _, metadata = get_or_create_metadata_bucket(data) # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( metadata.get("_model_armor_response"), @@ -868,7 +875,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Attach Model Armor response & status to this request's metadata to avoid race conditions if isinstance(request_data, dict): - metadata = request_data.setdefault("metadata", {}) + _, metadata = get_or_create_metadata_bucket(request_data) metadata["_model_armor_response"] = self._build_logging_response(armor_response) metadata["_model_armor_status"] = ( "blocked" if self._should_block_content(armor_response) else "success" diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py index 5c9f93fc2cd..717f5b6c5fe 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal, NoReturn from urllib.parse import urlsplit import httpx -from pydantic import ValidationError +from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -24,11 +24,12 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( STRAIKER_WEBHOOK_SCHEMA_VERSION, StraikerGuardrailConfigModel, @@ -42,7 +43,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( StraikerWebhookStream, StraikerWebhookUsage, ) -from litellm.types.utils import GenericGuardrailAPIInputs, Usage +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -57,6 +58,7 @@ RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504}) UNREACHABLE_STATUS = frozenset({502, 503, 504}) _APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"}) _OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool) +_JSON_DICT_ADAPTER = TypeAdapter(dict[str, object]) @dataclass(frozen=True, slots=True) @@ -137,6 +139,44 @@ def _resolve_destination(request_data: dict) -> str | None: return None +def _route_has_translation(request_data: dict) -> bool: + from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route + from litellm.llms import load_guardrail_translation_mappings + + route = _as_dict(request_data.get("litellm_metadata")).get("user_api_key_request_route") + if not isinstance(route, str) or not route: + return False + mappings = load_guardrail_translation_mappings() + return any(call_type in mappings for call_type in get_call_types_for_route(route) or ()) + + +def _request_structured_messages(request_data: dict) -> list[dict[str, Any]] | None: + messages = request_data.get("messages") + if messages: + return messages if isinstance(messages, list) else None + if not _route_has_translation(request_data): + return None + return resolve_structured_messages(messages=None, request_kwargs=request_data) + + +def _hook_name(value: object) -> str: + return value.value if isinstance(value, GuardrailEventHooks) else str(value) + + +def _configured_modes(event_hook: object) -> list[str] | None: + if isinstance(event_hook, list): + names = [_hook_name(v) for v in event_hook] + elif isinstance(event_hook, (str, GuardrailEventHooks)): + names = [_hook_name(event_hook)] + elif isinstance(event_hook, Mode): + default = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default] + tags = [v for value in event_hook.tags.values() for v in (value if isinstance(value, list) else [value])] + names = [_hook_name(v) for v in (*default, *tags) if v is not None] + else: + return None + return list(dict.fromkeys(names)) or None + + def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str: call_type = ( (getattr(logging_obj, "call_type", None) if logging_obj is not None else None) @@ -146,23 +186,76 @@ def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: d return call_type if isinstance(call_type, str) and call_type else "unknown" +def _jsonable_dict(value: object) -> dict[str, object] | None: + if isinstance(value, BaseModel): + return _JSON_DICT_ADAPTER.validate_python(value.model_dump(mode="json", exclude_none=True)) + if isinstance(value, dict): + return _JSON_DICT_ADAPTER.validate_python(value) + return None + + +def _opaque_dict_list(value: object) -> list[dict[str, object]] | None: + if not isinstance(value, list): + return None + items = tuple(plain for item in value if (plain := _jsonable_dict(item)) is not None) + return list(items) if items else None + + +def _choice_terminal_reason(choice: object) -> str | None: + if isinstance(choice, dict): + return _as_optional_str(choice.get("finish_reason")) or _as_optional_str(choice.get("stop_reason")) + return _as_optional_str(getattr(choice, "finish_reason", None)) or _as_optional_str( + getattr(choice, "stop_reason", None) + ) + + def _response_finish_reason(response: Any) -> str | None: + if response is None: + return None + if isinstance(response, dict): + top = _as_optional_str(response.get("finish_reason")) or _as_optional_str(response.get("stop_reason")) + if top: + return top + choices = response.get("choices") + if not isinstance(choices, list): + return None + for choice in choices: + reason = _choice_terminal_reason(choice) + if reason: + return reason + return None + + top = _as_optional_str(getattr(response, "finish_reason", None)) or _as_optional_str( + getattr(response, "stop_reason", None) + ) + if top: + return top choices = getattr(response, "choices", None) if not isinstance(choices, list): return None for choice in choices: - reason = getattr(choice, "finish_reason", None) - if isinstance(reason, str) and reason: + reason = _choice_terminal_reason(choice) + if reason: return reason return None +def _as_optional_int(value: object) -> int | None: + return value if isinstance(value, int) and not isinstance(value, bool) else None + + +def _usage_token_count(usage: object, openai_key: str, anthropic_key: str) -> int | None: + get = usage.get if isinstance(usage, dict) else lambda key: getattr(usage, key, None) + openai_count = _as_optional_int(get(openai_key)) + return openai_count if openai_count is not None else _as_optional_int(get(anthropic_key)) + + def _build_usage(response: object) -> StraikerWebhookUsage | None: - usage = getattr(response, "usage", None) - if not isinstance(usage, Usage): + usage = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None) + if usage is None: return None - input_tokens = usage.prompt_tokens - output_tokens = usage.completion_tokens + input_tokens = _usage_token_count(usage, "prompt_tokens", "input_tokens") + output_tokens = _usage_token_count(usage, "completion_tokens", "output_tokens") if input_tokens is None and output_tokens is None: return None return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens) @@ -234,6 +327,8 @@ class StraikerGuardrail(CustomGuardrail): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) super().__init__(**kwargs) + self.configured_modes = _configured_modes(self.event_hook) + def _webhook_url(self) -> str: return f"{self.api_base}{WEBHOOK_PATH}" @@ -263,6 +358,7 @@ class StraikerGuardrail(CustomGuardrail): ) -> StraikerWebhookContext: return StraikerWebhookContext( call_surface=_resolve_call_surface(logging_obj, request_data), + mode=self.configured_modes, model=model, model_provider=_resolve_provider(request_data, model), destination=_resolve_destination(request_data), @@ -287,9 +383,9 @@ class StraikerGuardrail(CustomGuardrail): content = StraikerWebhookContent( texts=list(inputs.get("texts") or []), images=list(inputs.get("images") or []), - structured_messages=inputs.get("structured_messages"), - tools=inputs.get("tools"), - tool_calls=inputs.get("tool_calls"), + structured_messages=_opaque_dict_list(inputs.get("structured_messages")), + tools=_opaque_dict_list(inputs.get("tools")), + tool_calls=_opaque_dict_list(inputs.get("tool_calls")), ) if input_type == "request": @@ -305,9 +401,8 @@ class StraikerGuardrail(CustomGuardrail): response_obj = request_data.get("response") content.finish_reason = _response_finish_reason(response_obj) - original_messages = request_data.get("messages") request_content = StraikerWebhookContent( - structured_messages=original_messages if isinstance(original_messages, list) else None, + structured_messages=_opaque_dict_list(_request_structured_messages(request_data)), ) phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none" event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 8e3abfbf159..d4d23cd2e37 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -147,8 +147,10 @@ class UnifiedLLMGuardrails(CustomLogger): litellm_logging_obj=data.get("litellm_logging_obj"), ) - # Add guardrail to applied guardrails header - add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name) + if not guardrail_to_apply.records_own_guardrail_information: + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=guardrail_to_apply.guardrail_name + ) return data async def async_moderation_hook( @@ -274,8 +276,10 @@ class UnifiedLLMGuardrails(CustomLogger): if e.original_response is None: e.original_response = response raise - # Add guardrail to applied guardrails header - add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name) + if not guardrail_to_apply.records_own_guardrail_information: + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=guardrail_to_apply.guardrail_name + ) return response diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index acb2e50c79b..9364d7eae3a 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -38,6 +38,10 @@ from litellm._uuid import uuid from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -668,16 +672,13 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d """ if guardrail_data is None: return - source_metadata = guardrail_data.get("metadata") - if not isinstance(source_metadata, dict): - return + source_key = get_metadata_variable_name_from_kwargs(guardrail_data) + source_metadata = guardrail_data.get(source_key) or {} entries = source_metadata.get("standard_logging_guardrail_information") if not entries: return - metadata = request_data.get("metadata") - if not isinstance(metadata, dict): - metadata = request_data["metadata"] = {} + _, metadata = get_or_create_metadata_bucket(request_data) metadata.setdefault("standard_logging_guardrail_information", list(entries)) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 25fa0819930..1f5aacc2115 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,5 +1,5 @@ import asyncio -from typing import TYPE_CHECKING, Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Mapping, Optional import httpx from fastapi import HTTPException, status @@ -145,6 +145,30 @@ class ProxyModelNotFoundError(HTTPException): super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) +REQUIRED_BODY_PARAM_BY_ROUTE: Mapping[str, str] = { + "acompletion": "messages", + "aembedding": "input", +} + + +class ProxyMissingRequiredParamError(HTTPException): + def __init__(self, route: str, param: str): + detail = {"error": f"{route}: Missing required parameter: '{param}'."} + super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) + self.type = "invalid_request_error" + self.param = param + + +def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None: + required_param = REQUIRED_BODY_PARAM_BY_ROUTE.get(route_type) + if required_param is None or data.get(required_param) is not None: + return + raise ProxyMissingRequiredParamError( + route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type), + param=required_param, + ) + + def get_team_id_from_data(data: dict) -> Optional[str]: """ Get the team id from the data's metadata or litellm_metadata params. @@ -353,6 +377,8 @@ async def route_request( """ Common helper to route the request """ + raise_if_required_body_param_missing(route_type=route_type, data=data) + await add_shared_session_to_data(data) # Strip router-internal mock_testing_* flags. Combined with an diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 12c890ec91d..429ddeef36a 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -5,6 +5,7 @@ from typing import ( Dict, Iterable, List, + Mapping, Optional, Type, Union, @@ -24,6 +25,7 @@ from litellm.types.llms.openai import ( ResponseInputParam, ResponsesAPIOptionalRequestParams, ResponsesAPIResponse, + ResponsesAPIStreamOptions, ResponseText, ) from litellm.types.responses.main import DecodedResponseId @@ -35,6 +37,17 @@ from litellm.types.utils import ( ) +def normalize_responses_api_stream_options( + stream_options: object, +) -> ResponsesAPIStreamOptions | None: + if not isinstance(stream_options, Mapping): + return None + include_obfuscation = stream_options.get("include_obfuscation") + if not isinstance(include_obfuscation, bool): + return None + return ResponsesAPIStreamOptions(include_obfuscation=include_obfuscation) + + class ResponsesAPIRequestUtils: """Helper utils for constructing ResponseAPI requests""" @@ -156,15 +169,19 @@ class ResponsesAPIRequestUtils: drop_params=should_drop_params, ) + stream_options = normalize_responses_api_stream_options(mapped_params.get("stream_options")) + params_with_normalized_stream_options = { + **{key: value for key, value in mapped_params.items() if key != "stream_options"}, + **({} if stream_options is None else {"stream_options": stream_options}), + } + # add any allowed_openai_params to the mapped_params - mapped_params = _apply_openai_param_overrides( - optional_params=mapped_params, + return _apply_openai_param_overrides( + optional_params=params_with_normalized_stream_options, non_default_params=non_default_params, allowed_openai_params=allowed_openai_params or [], ) - return mapped_params - @staticmethod def get_requested_response_api_optional_param( params: Dict[str, Any], diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 9f689a2dd31..314bb653196 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1145,6 +1145,10 @@ class ContextManagementEntry(TypedDict, total=False): """Token threshold at which compaction is triggered for this entry. Minimum 1000.""" +class ResponsesAPIStreamOptions(TypedDict, total=False): + include_obfuscation: bool + + class ResponsesAPIOptionalRequestParams(TypedDict, total=False): """TypedDict for Optional parameters supported by the responses API.""" @@ -1171,7 +1175,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False): max_tool_calls: Optional[int] prompt_cache_key: Optional[str] prompt_cache_retention: Optional[str] - stream_options: Optional[dict] + stream_options: Optional[ResponsesAPIStreamOptions] top_logprobs: Optional[int] partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation context_management: Optional[List[ContextManagementEntry]] diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py index b4375237917..0e816985cb0 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py @@ -4,9 +4,6 @@ from typing import Literal from pydantic import BaseModel, ConfigDict, Field -from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk -from litellm.types.utils import ChatCompletionMessageToolCall - from .base import GuardrailConfigModel StraikerWebhookEventType = Literal["pre_call", "post_call"] @@ -32,9 +29,9 @@ class StraikerWebhookContent(BaseModel): texts: list[str] = Field(default_factory=list) images: list[str] = Field(default_factory=list) - structured_messages: list[AllMessageValues] | None = None + structured_messages: list[dict[str, object]] | None = None tools: list[dict[str, object]] | None = None - tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None + tool_calls: list[dict[str, object]] | None = None finish_reason: str | None = None @@ -45,6 +42,7 @@ class StraikerWebhookUsage(BaseModel): class StraikerWebhookContext(BaseModel): call_surface: str + mode: list[str] | None = None model: str | None = None model_provider: str | None = None destination: str | None = None diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index de5b3b32105..a2e18a62638 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -22,78 +22,6 @@ def redis_no_ping(): yield -@pytest.mark.parametrize("namespace", [None, "test"]) -@pytest.mark.asyncio -async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache(namespace=namespace) - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - assert redis_cache is not None - - expected_key = "test:test" if namespace else "test" - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_set_cache - await redis_cache.async_increment(key=expected_key, value=1) - - # Verify that the set method was called on the mock Redis instance - mock_redis_instance.incrbyfloat.assert_called_once_with( - name=expected_key, amount=1 - ) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl( - monkeypatch, redis_no_ping -): - """With refresh_ttl=True, every increment should call expire() to bump - the TTL, even when the key already has a TTL (counter-style use).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment( - key="spend:team_member:u:t", value=0.05, refresh_ttl=True - ) - - mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl( - monkeypatch, redis_no_ping -): - """Default (refresh_ttl=False) preserves window-style semantics: TTL is - set only on first creation, never refreshed (used by rate-limit windows).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment(key="rate_limit:window", value=1) - - mock_redis_instance.expire.assert_not_awaited() - - @pytest.mark.parametrize("namespace", [None, "litellm"]) @pytest.mark.asyncio async def test_async_delete_cache_applies_namespace( @@ -140,42 +68,6 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping) assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0 -@pytest.mark.asyncio -async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Setup the return value for mget - mock_redis_instance.mget.return_value = [ - b'{"key1": "value1"}', - None, - b'{"key3": "value3"}', - ] - - test_keys = ["key1", "key2", "key3"] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_batch_get_cache - result = await redis_cache.async_batch_get_cache(key_list=test_keys) - - # Verify mget was called with the correct keys - mock_redis_instance.mget.assert_called_once() - - # Check that results were properly decoded - assert result["key1"] == {"key1": "value1"} - assert result["key2"] is None - assert result["key3"] == {"key3": "value3"} - - @pytest.mark.asyncio async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): """Test the helper method that handles LPOP with count for Redis versions < 7.0""" @@ -202,41 +94,6 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): assert mock_pipeline.execute.call_count == 2 -@pytest.mark.asyncio -async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping): - """Verify that multiple rpush ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a", "b"]), - RedisPipelineRpushOperation(key="key2", values=["c"]), - RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - assert result == [3, 5, 1] - assert mock_pipeline.rpush.call_count == 3 - mock_pipeline.rpush.assert_any_call("key1", "a", "b") - mock_pipeline.rpush.assert_any_call("key2", "c") - mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f") - mock_pipeline.execute.assert_called_once() - - @pytest.mark.asyncio async def test_async_rpush_pipeline_empty_list_returns_empty( monkeypatch, redis_no_ping @@ -256,183 +113,6 @@ async def test_async_rpush_pipeline_empty_list_returns_empty( mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping): - """Verify that multiple lpop ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock( - return_value=[ - [b"val1", b"val2"], # key1 results - None, # key2 empty - [b"val3"], # key3 results - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - RedisPipelineLpopOperation(key="key3", count=5), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 3 - assert results[0] == ["val1", "val2"] - assert results[1] is None - assert results[2] == ["val3"] - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results( - monkeypatch, redis_no_ping -): - """Verify Redis < 7 fallback issues individual LPOPs and regroups correctly""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "6.2.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - - # With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands - # Simulate: key1 has 2 values then None, key2 has 1 value then None - mock_pipeline.execute = AsyncMock( - return_value=[ - b"val1", - b"val2", - None, # 3 LPOPs for key1 - b"val3", - None, # 2 LPOPs for key2 - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=3), - RedisPipelineLpopOperation(key="key2", count=2), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 2 - assert results[0] == ["val1", "val2"] # 2 values, None filtered out - assert results[1] == ["val3"] # 1 value, None filtered out - # All 5 individual LPOPs should be queued, but only 1 execute() call - assert mock_pipeline.lpop.call_count == 5 - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - # Simulate: first RPUSH succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a"]), - RedisPipelineRpushOperation(key="key2", values=["b"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in LPOP pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - # Simulate: first LPOP succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[[b"val1"], Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - @pytest.mark.asyncio async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): """Empty lpop_list should return empty list without touching Redis""" @@ -450,111 +130,6 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_lpop_pipeline_propagates_redis_exception( - monkeypatch, redis_no_ping -): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "redis_version", - [ - # Standard cases - "7.0.0", # Standard Redis string version - 7.0, # Valkey/ElastiCache float version (THE BUG this fix addresses) - 7, # Integer version (e.g., from some Redis forks) - # Version < 7 - "6", # String without dots, version < 7 - # Malformed versions (fallback to 7) - "latest", # Non-numeric version - "", # Empty string - -7.0, # Negative float - # Format variations - " 7.0.0 ", # Whitespace (should be stripped) - "7.0.0-rc1", # Version with suffix - "10.0.0", # Double digit major version - ], -) -async def test_async_lpop_with_float_redis_version( - monkeypatch, redis_no_ping, redis_version -): - """ - Test async_lpop with various Redis version formats (especially float). - - This test specifically addresses the issue where AWS ElastiCache Valkey - returns redis_version as a float (e.g., 7.0) instead of a string (e.g., "7.0.0"), - which caused a 'float' object has no attribute 'split' error when trying to - use the Redis transaction buffer feature. - - The fix converts the version to a string and handles edge cases like: - - Floats (7.0) and integers (7) - - Strings with/without dots ("7" vs "7.0.0") - - Malformed versions ("v7.0.0", "latest") - fallback to version 7 - - Whitespace (" 7.0.0 ") - - Negative versions (fallback to version 7) - - Related: Database deadlock issues when use_redis_transaction_buffer is enabled. - """ - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - - # Create RedisCache instance - redis_cache = RedisCache() - redis_cache.redis_version = redis_version # Set the version to test - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Mock lpop to return a test value (Redis >= 7.0 behavior) - mock_redis_instance.lpop.return_value = [b"value1", b"value2"] - - # Mock pipeline for Redis < 7.0 (used when major_version < 7) - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - # Make pipeline() a regular method (not async) that returns the mock - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - # Mock handle_lpop_count_for_older_redis_versions for Redis < 7 - with patch.object( - redis_cache, - "handle_lpop_count_for_older_redis_versions", - return_value=[b"value1", b"value2"], - ): - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_lpop with count - this should not raise AttributeError - result = await redis_cache.async_lpop(key="test_key", count=2) - - # Verify the method completed without error - assert result is not None - - # LIT-3374: the namespace must be applied uniformly across every key-taking # Redis operation, not just get/set/increment. Before the fix these paths wrote # or read raw keys, so with a namespace configured the prefixed keys other diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 6a1de0586dd..6907e4d0d02 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -2853,3 +2853,77 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id(): assert stream_tool_id("fc_unique_abc123", "call_0") == "fc_unique_abc123" assert stream_tool_id("fc_2", "call_tokyo") == "call_tokyo" + + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "stream_options,expected_wire_stream_options", + [ + ({"include_usage": True, "include_obfuscation": False}, {"include_obfuscation": False}), + ({"include_usage": True}, None), + ], +) +async def test_acompletion_bridge_normalizes_stream_options_on_the_wire( + stream_options, expected_wire_stream_options +): + """include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict.""" + from unittest.mock import AsyncMock + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + responses_payload = { + "id": "resp_bridge_stream_options", + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "gpt-5.5", + "output": [ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "temperature": None, + "tool_choice": "auto", + "tools": [], + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "truncation": None, + "user": None, + } + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = json.dumps(responses_payload) + mock_response.headers = httpx.Headers({}) + mock_response.json.return_value = responses_payload + + with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post: + mock_post.return_value = mock_response + + await litellm.acompletion( + model="openai/responses/gpt-5.5", + messages=[{"role": "user", "content": "hi"}], + api_key="fake-api-key", + stream_options=stream_options, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + if expected_wire_stream_options is None: + assert "stream_options" not in request_body + else: + assert request_body["stream_options"] == expected_wire_stream_options diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 4ea79f9e2a4..bac0ae54033 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -3,9 +3,12 @@ from unittest.mock import AsyncMock import pytest -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) from litellm.proxy._types import CallTypes, UserAPIKeyAuth -from litellm.types.utils import GuardrailTracingDetail +from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail class TestCustomGuardrailDeploymentHook: @@ -654,6 +657,53 @@ class TestGuardrailLoggingAggregation: assert len(info) == 2 assert info[1]["guardrail_name"] == "test_guardrail" + def test_caller_metadata_does_not_divert_the_entry_from_the_reader(self): + """A caller-supplied `metadata` field must not send the entry to a bucket the + spend log never reads. Routes in LITELLM_METADATA_ROUTES (/v1/messages, + /v1/responses, batches, files) seed `litellm_metadata`, and Claude Code sends + `metadata.user_id`, so both keys are present on the same request.""" + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"user_api_key_hash": "abc"}, + } + + self._invoke_add_log(request_data) + + assert ( + "standard_logging_guardrail_information" not in request_data["metadata"] + ), "entry landed in the caller's metadata, where the spend log does not read it" + info = request_data["litellm_metadata"][ + "standard_logging_guardrail_information" + ] + assert len(info) == 1 + assert info[0]["guardrail_name"] == "test_guardrail" + + def test_entry_and_applied_guardrails_header_share_one_bucket(self): + """The x-litellm-applied-guardrails writer and the guardrail-info writer must + resolve the same bucket, otherwise the response header and the spend log + disagree about whether the guardrail ran.""" + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {}, + } + + self._invoke_add_log(request_data) + add_guardrail_to_applied_guardrails_header( + request_data=request_data, guardrail_name="test_guardrail" + ) + + buckets = { + key + for key in ("metadata", "litellm_metadata") + for field in ("standard_logging_guardrail_information", "applied_guardrails") + if field in request_data[key] + } + assert buckets == {"litellm_metadata"} + class TestGuardrailOtelSpanEmission: """Recording a guardrail emits its otel span inline, so every guardrail @@ -1947,3 +1997,55 @@ class TestOnlyScanNewMessages: cache.async_set_cache = AsyncMock(side_effect=RuntimeError("redis down")) await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache) + + +def _guardrail_entries(request_data: dict) -> list: + container = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + entries = container.get("standard_logging_guardrail_information") + return entries if isinstance(entries, list) else [] + + +class _NoopGuardrail(CustomGuardrail): + """apply_guardrail that returns the inputs untouched and records nothing.""" + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + return inputs + + +class _NoopSelfLoggingGuardrail(_NoopGuardrail): + records_own_guardrail_information = True + + +class TestRecordsOwnGuardrailInformation: + """The @log_guardrail_information decorator must not synthesize an "allow"/"success" + entry for a no-op apply_guardrail when the guardrail sets + records_own_guardrail_information (LIT-4650).""" + + @pytest.mark.asyncio + async def test_default_noop_apply_guardrail_is_auto_logged(self): + guardrail = _NoopGuardrail(guardrail_name="g1") + request_data: dict = {"model": "gpt-4o"} + + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["x"]), + request_data=request_data, + input_type="request", + ) + + entries = _guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_status"] == "success" + + @pytest.mark.asyncio + async def test_self_logging_noop_apply_guardrail_is_not_logged(self): + guardrail = _NoopSelfLoggingGuardrail(guardrail_name="g2") + request_data: dict = {"model": "gpt-4o"} + + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["x"]), + request_data=request_data, + input_type="request", + ) + + assert _guardrail_entries(request_data) == [] diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/test_litellm/integrations/test_guardrail_logging_sync.py index 5dcd1114b3d..f9e1a3efbd0 100644 --- a/tests/test_litellm/integrations/test_guardrail_logging_sync.py +++ b/tests/test_litellm/integrations/test_guardrail_logging_sync.py @@ -59,8 +59,11 @@ def test_syncs_from_metadata_key(): assert result == [entry] -def test_metadata_wins_over_litellm_metadata(): - """metadata key takes precedence over litellm_metadata when both are present.""" +def test_litellm_metadata_wins_over_caller_metadata(): + """When both keys are present the helper must read the bucket the writer used, + which get_or_create_metadata_bucket resolves to litellm_metadata. Reading the + caller's metadata instead is how a guardrail entry went missing from spend logs + on the routes that seed litellm_metadata.""" entry_meta = _make_slg_entry("from-metadata") entry_lm = _make_slg_entry("from-litellm_metadata") request_data = { @@ -74,7 +77,25 @@ def test_metadata_wins_over_litellm_metadata(): result = logging_obj.litellm_params["metadata"].get( "standard_logging_guardrail_information" ) - assert result == [entry_meta] + assert result == [entry_lm] + + +def test_syncs_when_caller_sends_its_own_metadata(): + """The Claude Code shape: caller metadata present, guardrail entry in the seeded + litellm_metadata bucket. The entry must still reach the spend-log payload.""" + entry = _make_slg_entry() + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"standard_logging_guardrail_information": [entry]}, + } + logging_obj = _FakeLogging() + + _sync_guardrail_info_to_logging_obj(request_data, logging_obj) + + result = logging_obj.litellm_params["metadata"].get( + "standard_logging_guardrail_information" + ) + assert result == [entry] def test_noop_when_no_guardrail_info(): diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py index ace9399cf53..c3e9d67ddad 100644 --- a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py +++ b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py @@ -279,6 +279,48 @@ class TestGuardrailSpanOnViolation(unittest.TestCase): parent_span.context.span_id, ) + def test_post_call_failure_hook_emits_span_when_caller_sends_metadata(self): + """On routes that seed ``litellm_metadata`` the guardrail entry lives there, + not in the caller's own ``metadata`` field. Reading a hard-coded ``metadata`` + key drops the span for exactly the requests that carry both.""" + otel, provider, exporter = _make_otel() + parent_span = provider.get_tracer(__name__).start_span(PROXY_SPAN_NAME) + + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", + parent_otel_span=parent_span, + request_route="/v1/messages", + ) + + request_data = { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": { + "standard_logging_guardrail_information": [ + _slg_entry("guardrail_intervened", _bedrock_block_response()) + ], + }, + } + + _run( + otel.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("guardrail blocked"), + user_api_key_dict=user_api_key_dict, + ) + ) + + guardrail_spans = [ + s for s in exporter.get_finished_spans() if s.name == GUARDRAIL_SPAN_NAME + ] + self.assertEqual( + len(guardrail_spans), + 1, + "the guardrail span must be emitted from the resolved metadata bucket, " + "not from a hard-coded 'metadata' key", + ) + def test_handle_failure_and_post_call_failure_hook_dedupe(self): """When _handle_failure and async_post_call_failure_hook BOTH fire for the same request (the production flow on a guardrail block), diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index b67ea91bb0b..b4f539da286 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -4,12 +4,53 @@ import pytest from litellm.litellm_core_utils.core_helpers import ( _FINISH_REASON_MAP, + get_or_create_metadata_bucket, map_finish_reason, reconstruct_model_name, redact_nested_match_and_regex_keys, ) +class TestGetOrCreateMetadataBucket: + """The single owner every guardrail writer and reader shares, so the response + header and the spend log can never disagree about which dict a record lives in.""" + + def test_prefers_litellm_metadata_when_both_present(self): + request_data = {"metadata": {"user_id": "caller"}, "litellm_metadata": {}} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "litellm_metadata" + assert bucket is request_data["litellm_metadata"] + + def test_uses_metadata_when_litellm_metadata_absent(self): + request_data = {"metadata": {"user_id": "caller"}} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "metadata" + assert bucket is request_data["metadata"] + + def test_creates_the_bucket_in_place_when_missing(self): + request_data: dict = {} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "metadata" + assert request_data["metadata"] is bucket + bucket["k"] = "v" + assert request_data["metadata"]["k"] == "v" + + def test_replaces_a_non_dict_bucket(self): + request_data = {"litellm_metadata": None} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "litellm_metadata" + assert isinstance(request_data["litellm_metadata"], dict) + assert bucket is request_data["litellm_metadata"] + + def test_reconstruct_model_name_prefers_deployment_value(): """Ensure deployment metadata wins when reconstructing the model name.""" diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 9cd1fbb59a6..48acdd348e9 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -57,6 +57,98 @@ class MockDynamicGuardrail(CustomGuardrail): return inputs +class MockRecordingGuardrail(CustomGuardrail): + """Mock guardrail that records the request_data it was handed.""" + + def __init__(self, guardrail_name: str): + super().__init__(guardrail_name=guardrail_name) + self.request_data: Optional[dict] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.request_data = request_data + return inputs + + +class TestAnthropicMessagesHandlerStreamingRequestData: + """Post-call guardrails on streaming /v1/messages receive the response and identity metadata""" + + @pytest.mark.asyncio + async def test_terminal_chunk_passes_assembled_response_and_metadata(self): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.utils import Choices, Message, ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-sonnet-4-5", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Hello world", role="assistant"), + ) + ], + ) + + with ( + patch.object(handler, "_check_streaming_has_ended", return_value=True), + patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ), + ): + await handler.process_output_streaming_response( + responses_so_far=[b"data: some chunk"], + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"), + request_data={"model": "claude-sonnet-4-5"}, + ) + + assert guardrail.request_data is not None + assert guardrail.request_data["response"] is mock_response + assert ( + guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" + ) + + @pytest.mark.asyncio + async def test_mid_stream_chunk_passes_responses_so_far_and_metadata(self): + from litellm.proxy._types import UserAPIKeyAuth + + handler = AnthropicMessagesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + responses_so_far = [b"data: some chunk"] + + with ( + patch.object(handler, "_check_streaming_has_ended", return_value=False), + patch.object( + handler, "get_streaming_string_so_far", return_value="partial text" + ), + ): + await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"), + request_data={"model": "claude-sonnet-4-5"}, + ) + + assert guardrail.request_data is not None + assert guardrail.request_data["responses"] is responses_so_far + assert ( + guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" + ) + + class TestAnthropicMessagesHandlerStreamingOutputProcessing: """Test streaming output processing functionality""" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 3327fc39f73..8875a75e86f 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -2,6 +2,7 @@ import json import os import sys +import httpx import pytest from fastapi.testclient import TestClient @@ -9,6 +10,7 @@ sys.path.insert(0, os.path.abspath("../../../../..")) from unittest.mock import AsyncMock, MagicMock, patch +import litellm from litellm.anthropic_interface import messages from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.utils import Delta, ModelResponse, StreamingChoices @@ -37,6 +39,68 @@ def test_anthropic_experimental_pass_through_messages_handler(): assert mock_responses.call_args.kwargs["api_key"] == "test-api-key" +@pytest.mark.asyncio +async def test_openai_model_does_not_forward_stream_options_to_responses_api(): + """ + Regression test for LIT-4779. `always_include_stream_usage` injects + stream_options={'include_usage': True} into every streaming request, but OpenAI + models on /v1/messages go to the Responses API, which 400s on that param. + """ + responses_payload = { + "id": "resp_stream_options", + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "gpt-5.5", + "output": [ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "temperature": None, + "tool_choice": "auto", + "tools": [], + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "truncation": None, + "user": None, + } + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = json.dumps(responses_payload) + mock_response.headers = httpx.Headers({}) + mock_response.json.return_value = responses_payload + + with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post: + mock_post.return_value = mock_response + + await litellm.anthropic.messages.acreate( + max_tokens=100, + messages=[{"role": "user", "content": "Hello, how are you?"}], + model="openai/gpt-5.5", + api_key="test-api-key", + stream_options={"include_usage": True}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert "stream_options" not in request_body + + def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_and_api_base_and_custom_values(): """ Test that api key, api base, and extra kwargs are forwarded to litellm.completion for Azure models. diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 7bd1d7a6031..87d67e0e8b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -181,28 +181,6 @@ async def test_force_ipv4_transport(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_ssl_context_transport(): - """Test transport creation with SSL context""" - # Create a test SSL context - ssl_context = ssl.create_default_context() - - transport = AsyncHTTPHandler._create_async_transport(ssl_context=ssl_context) - assert transport is not None - - try: - if isinstance(transport, LiteLLMAiohttpTransport): - # Get the client session and verify SSL context is passed through - client_session = transport._get_valid_client_session() - assert isinstance(client_session, ClientSession) - assert isinstance(client_session.connector, TCPConnector) - # Verify the connector has SSL context set by checking if it's using SSL - assert client_session.connector._ssl is not None - finally: - if isinstance(transport, LiteLLMAiohttpTransport): - await transport.aclose() - - @pytest.mark.asyncio async def test_aiohttp_disabled_transport(): """Test transport creation with aiohttp disabled""" @@ -339,44 +317,6 @@ async def test_ssl_context_with_shared_session(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_aiohttp_transport_trust_env_setting(monkeypatch): - """Test that trust_env setting is properly configured in aiohttp transport""" - transports = [] - try: - # Test 1: Default trust_env behavior - transport = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport) - client_session = transport._get_valid_client_session() - - # Default should be False (litellm.aiohttp_trust_env default) - default_trust_env = getattr(litellm, "aiohttp_trust_env", False) - assert client_session._trust_env == default_trust_env - - # Test 2: Environment variable override - monkeypatch.setenv("AIOHTTP_TRUST_ENV", "True") - transport_with_env = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport_with_env) - client_session_with_env = transport_with_env._get_valid_client_session() - - # Should be True when environment variable is set - assert client_session_with_env._trust_env is True - - # Test 3: Verify environment variable with False value - monkeypatch.setenv("AIOHTTP_TRUST_ENV", "False") - transport_with_false_env = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport_with_false_env) - client_session_with_false_env = ( - transport_with_false_env._get_valid_client_session() - ) - - # Should respect the litellm.aiohttp_trust_env setting when env var is False - assert client_session_with_false_env._trust_env == default_trust_env - finally: - for t in transports: - await t.aclose() - - def test_get_ssl_configuration(): """Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle when no environment variables are set.""" @@ -443,36 +383,6 @@ async def test_create_aiohttp_transport_with_shared_session(): assert not callable(transport.client) # Should not be callable -@pytest.mark.asyncio -async def test_create_aiohttp_transport_without_shared_session(): - """Test that _create_aiohttp_transport creates new session when none provided""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Test without shared session - transport = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) - - # Verify the transport uses a lambda function (for backward compatibility) - assert callable(transport.client) # Should be a lambda function - - -@pytest.mark.asyncio -async def test_create_aiohttp_transport_with_closed_session(): - """Test that _create_aiohttp_transport creates new session when shared session is closed""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Create a mock closed session - mock_session = MockClientSession() - mock_session.closed = True - - # Test with closed session - transport = AsyncHTTPHandler._create_aiohttp_transport( - shared_session=mock_session # type: ignore - ) - - # Verify the transport creates a new session (lambda function) - assert callable(transport.client) # Should be a lambda function - - @pytest.mark.asyncio async def test_async_handler_with_shared_session(): """Test AsyncHTTPHandler initialization with shared session""" @@ -622,27 +532,6 @@ async def test_session_reuse_integration(): await client2.close() -@pytest.mark.asyncio -async def test_session_validation(): - """Test that session validation works correctly""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Test with None session - transport1 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) - assert callable(transport1.client) # Should create lambda - - # Test with closed session - mock_closed_session = MockClientSession() - mock_closed_session.closed = True - transport2 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_closed_session) # type: ignore - assert callable(transport2.client) # Should create lambda - - # Test with valid session - mock_valid_session = MockClientSession() - transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore - assert transport3.client is mock_valid_session # Should reuse session - - @pytest.mark.parametrize( "env_curve,litellm_curve,expected_curve,should_call", [ diff --git a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py index f8bd83fc7df..1043c26c6ec 100644 --- a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py @@ -1,15 +1,10 @@ """ -Tests for LlmPassthroughRouteHandler and the guardrail_translation_mappings registry. +Tests for the guardrail_translation_mappings registry. Validates: - allm_passthrough_route is registered in the mappings (regression: this was the bug) -- Bedrock provider is dispatched to BedrockPassthroughGuardrailHandler -- Unknown provider skips apply_guardrail """ -import pytest -from unittest.mock import AsyncMock, MagicMock, patch - from litellm.llms.pass_through.guardrail_translation import ( guardrail_translation_mappings, ) @@ -40,185 +35,3 @@ class TestRegistry: is PassThroughEndpointHandler ) - -def _make_guardrail() -> MagicMock: - g = MagicMock() - g.guardrail_name = "test-guard" - g.apply_guardrail = AsyncMock(return_value={"texts": []}) - g.skip_system_message_in_guardrail = False - g.skip_tool_message_in_guardrail = False - return g - - -class TestLlmPassthroughRouteHandlerInput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - "data": {"messages": [{"role": "user", "content": [{"text": "hi"}]}]}, - } - guardrail = _make_guardrail() - - await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_apply_guardrail(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "some_unknown_provider", - "endpoint": "v1/chat/completions", - "data": {"messages": [{"role": "user", "content": "hi"}]}, - } - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - @pytest.mark.asyncio - async def test_missing_provider_skips(self): - handler = LlmPassthroughRouteHandler() - data = {"endpoint": "foo/bar", "data": {}} - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - -class TestLlmPassthroughRouteHandlerOutput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_output_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - response = { - "output": { - "message": { - "role": "assistant", - "content": [{"text": "hello"}], - } - } - } - request_data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - } - guardrail = _make_guardrail() - - await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_output(self): - handler = LlmPassthroughRouteHandler() - response = {"some": "response"} - request_data = {"custom_llm_provider": "unknown"} - guardrail = _make_guardrail() - - result = await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is response - - -class TestDeAnonymizeEventStream: - @pytest.mark.asyncio - async def test_bedrock_provider_dispatches_to_handler(self): - body = b"original-stream-bytes" - expected = b"de-anonymized-bytes" - proxy_logging_obj = MagicMock() - user_api_key_dict = MagicMock() - - with patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=expected), - ) as mock_handler: - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=proxy_logging_obj, - user_api_key_dict=user_api_key_dict, - data={"custom_llm_provider": "bedrock"}, - ) - - mock_handler.assert_awaited_once() - assert result == expected - - @pytest.mark.asyncio - async def test_unknown_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={"custom_llm_provider": "anthropic"}, - ) - - assert result is body - - @pytest.mark.asyncio - async def test_missing_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={}, - ) - - assert result is body - - -class TestSupportsEventStreamDeAnonymization: - def test_bedrock_converse_stream_is_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream" - ) - is True - ) - - def test_bedrock_invoke_stream_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", - "model/us.amazon.nova-lite-v1:0/invoke-with-response-stream", - ) - is False - ) - - def test_unknown_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "model/foo/converse-stream" - ) - is False - ) - - def test_missing_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - None, "model/foo/converse-stream" - ) - is False - ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index fe6cb98d1f5..9002d1f81a3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -729,10 +729,15 @@ async def test_openai_moderation_post_call_request_data_passthrough(): mock_make_request.assert_called_once() - # Guardrail info in the REAL request_data (not a throwaway) - guardrail_info_list = request_data["metadata"].get( - "standard_logging_guardrail_information" + # Guardrail info in the REAL request_data (not a throwaway). The unified hook + # seeds litellm_metadata, so read the bucket the resolver names rather than + # assuming "metadata"; the spend log reads it the same way. + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, ) + + bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)] + guardrail_info_list = bucket.get("standard_logging_guardrail_information") assert guardrail_info_list is not None assert isinstance(guardrail_info_list[0]["guardrail_response"], dict) assert "results" in guardrail_info_list[0]["guardrail_response"] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py index 0358ca998aa..914af0e2368 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py @@ -259,10 +259,15 @@ async def test_openai_moderation_streaming_end_of_stream_request_data_passthroug ): pass - # Verify guardrail info reached the REAL request_data (not a throwaway) - guardrail_info_list = request_data["metadata"].get( - "standard_logging_guardrail_information" + # Verify guardrail info reached the REAL request_data (not a throwaway). The + # unified hook seeds litellm_metadata, so read the bucket the resolver names + # rather than assuming "metadata"; the spend log reads it the same way. + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, ) + + bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)] + guardrail_info_list = bucket.get("standard_logging_guardrail_information") assert ( guardrail_info_list is not None ), "Guardrail info should be in request_data after streaming" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index 7f412c008ca..776df985d46 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -114,6 +114,24 @@ def guardrail() -> HeadroomGuardrail: return _make_guardrail() +def _recorded_guardrail_entries(request_data: dict) -> list: + for container_key in ("metadata", "litellm_metadata"): + container = request_data.get(container_key) + if isinstance(container, dict): + entries = container.get("standard_logging_guardrail_information") + if isinstance(entries, list): + return entries + return [] + + +def _applied_guardrails(request_data: dict) -> list: + for container_key in ("metadata", "litellm_metadata"): + container = request_data.get(container_key) + if isinstance(container, dict) and isinstance(container.get("applied_guardrails"), list): + return container["applied_guardrails"] + return [] + + @pytest.mark.asyncio async def test_apply_guardrail_compresses_and_returns_structured_messages( guardrail: HeadroomGuardrail, @@ -123,6 +141,7 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages( structured_messages=ORIGINAL_MESSAGES, ) mock_response = _make_compress_response(COMPRESSED_MESSAGES) + request_data = {"model": "gpt-4o"} with patch.object( guardrail.async_handler, @@ -132,12 +151,19 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages( ): result = await guardrail.apply_guardrail( inputs=inputs, - request_data={"model": "gpt-4o"}, + request_data=request_data, input_type="request", ) assert result.get("structured_messages") == COMPRESSED_MESSAGES + entries = _recorded_guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_name"] == "headroom" + assert entries[0]["guardrail_status"] == "success" + assert entries[0]["guardrail_provider"] == "headroom" + assert "headroom" in _applied_guardrails(request_data) + @pytest.mark.asyncio async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present( @@ -719,6 +745,7 @@ async def test_apply_guardrail_bypass_header_skips_compression( mock_post.assert_not_called() assert result.get("structured_messages") == ORIGINAL_MESSAGES + assert _recorded_guardrail_entries(request_data) == [] @pytest.mark.asyncio @@ -729,16 +756,18 @@ async def test_apply_guardrail_response_type_passthrough( texts=["some response text"], structured_messages=ORIGINAL_MESSAGES, ) + request_data: dict = {"model": "gpt-4o"} with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="response", ) mock_post.assert_not_called() assert result is inputs + assert _recorded_guardrail_entries(request_data) == [] @pytest.mark.asyncio @@ -746,16 +775,48 @@ async def test_apply_guardrail_empty_structured_messages_passthrough( guardrail: HeadroomGuardrail, ): inputs = GenericGuardrailAPIInputs(texts=["hello"]) + request_data: dict = {"model": "gpt-4o"} with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="request", ) mock_post.assert_not_called() assert result is inputs + assert _recorded_guardrail_entries(request_data) == [] + assert "headroom" not in _applied_guardrails(request_data) + + +@pytest.mark.asyncio +async def test_passthrough_handler_does_not_log_headroom_as_run( + guardrail: HeadroomGuardrail, +): + """Regression for LIT-4650. + + A passthrough request drives headroom through PassThroughEndpointHandler, which + only supplies `texts` (no `structured_messages`). Headroom cannot compress that + shape and no-ops, so it must not appear in the spend log's + standard_logging_guardrail_information as a successful run. + """ + from litellm.llms.pass_through.guardrail_translation.handler import ( + PassThroughEndpointHandler, + ) + + data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]} + + with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: + await PassThroughEndpointHandler().process_input_messages( + data=data, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + mock_post.assert_not_called() + + assert _recorded_guardrail_entries(data) == [] + assert "headroom" not in _applied_guardrails(data) @pytest.mark.asyncio @@ -884,6 +945,7 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed() texts=["hello"], structured_messages=ORIGINAL_MESSAGES, ) + request_data = {"model": "gpt-4o"} with patch.object( guardrail.async_handler, @@ -893,12 +955,18 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed() ): result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="request", ) assert result["structured_messages"] == ORIGINAL_MESSAGES + entries = _recorded_guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_name"] == "headroom" + assert entries[0]["guardrail_status"] == "guardrail_failed_to_respond" + assert "headroom" in _applied_guardrails(request_data) + @pytest.mark.asyncio async def test_apply_guardrail_http_error_fail_open_forwards_uncompressed(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 18b5bd92411..89b6af27719 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -3502,6 +3502,44 @@ async def test_single_scan_response_stays_a_dict(): assert isinstance(request_data["metadata"]["_model_armor_response"], dict) +@pytest.mark.asyncio +async def test_scan_result_reaches_the_logger_on_a_seeded_route(): + """On routes that seed `litellm_metadata` the scan result must land in that bucket + and be found by `_process_response`. Writing the file-scan result through the shared + resolver while the text-scan writers and the reader used a hard-coded `metadata` key + split the record in two, so the logged guardrail payload came back empty.""" + guardrail = _make_guardrail() + pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8") + request_data = { + "model": "claude-haiku", + "messages": [_file_message(pdf_b64)], + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(return_value=_armor_response(blocked=False)), + ): + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + + assert "_model_armor_response" not in request_data["metadata"] + assert "_model_armor_response" in request_data["litellm_metadata"] + + before = len(request_data["litellm_metadata"].get("standard_logging_guardrail_information", [])) + guardrail._process_response(response=None, request_data=request_data) + + logged = request_data["litellm_metadata"]["standard_logging_guardrail_information"] + assert len(logged) == before + 1 + assert logged[-1]["guardrail_response"], "the logger recorded an empty Model Armor payload" + + @pytest.mark.asyncio async def test_pre_call_blocks_supported_document_with_undecodable_base64(): """A supported document whose inline base64 will not decode cannot be scanned, so it fails closed.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py index ca57118ee9d..36a2e205ea7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py @@ -1,4 +1,5 @@ import json +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import httpx @@ -8,6 +9,9 @@ from litellm.exceptions import GuardrailRaisedException, ModifyResponseException from litellm.proxy.guardrails.guardrail_hooks.straiker import initialize_guardrail from litellm.proxy.guardrails.guardrail_hooks.straiker.straiker import ( StraikerGuardrail, + _build_usage, + _request_structured_messages, + _response_finish_reason, ) from litellm.proxy.guardrails.guardrail_registry import ( guardrail_class_registry, @@ -17,7 +21,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( StraikerGuardrailConfigModel, StraikerGuardrailConfigModelOptionalParams, ) -from litellm.types.utils import Choices, Message, ModelResponse, Usage +from litellm.types.utils import ( + ChatCompletionMessageToolCall, + Choices, + Function, + Message, + ModelResponse, + Usage, +) def _mock_response(action: str, turn_id: str = "turn-1", schema_version: str = "1", **extra) -> MagicMock: @@ -208,7 +219,9 @@ async def test_request_envelope_transport_and_shape(): "metadata": {"user_api_key_alias": "team-key", "agent_id": "chatbot-app", "app_name": "Chatbot"}, } - out = await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj()) + out = await g.apply_guardrail( + inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj() + ) assert out is inputs url = g.async_handler.post.call_args.args[0] @@ -232,6 +245,35 @@ async def test_request_envelope_transport_and_shape(): assert "metadata" not in payload +@pytest.mark.asyncio +async def test_request_envelope_ignores_unsupported_opaque_items(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + + await g.apply_guardrail( + inputs={ + "texts": ["hello"], + "tools": [ + object(), + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object"}}, + }, + ], + }, + request_data={"model": "m", "messages": [{"role": "user", "content": "hello"}]}, + input_type="request", + logging_obj=_logging_obj(), + ) + + assert _posted_payload(g)["request"]["tools"] == [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object"}}, + } + ] + + @pytest.mark.asyncio async def test_webhook_metadata_session_id_and_opaque_passthrough(): g = _make_guardrail() @@ -331,6 +373,64 @@ async def test_context_session_id_from_request_metadata(): assert "metadata" not in payload +@pytest.mark.asyncio +async def test_context_mode_from_string_event_hook(): + g = _make_guardrail(event_hook="pre_call") + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["pre_call"] + + +@pytest.mark.asyncio +async def test_context_mode_from_list_event_hook(): + from litellm.types.guardrails import GuardrailEventHooks + + g = _make_guardrail(event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["pre_call", "post_call"] + + +@pytest.mark.asyncio +async def test_context_mode_from_tagged_mode_is_flattened_and_deduped(): + from litellm.types.guardrails import Mode + + g = _make_guardrail( + event_hook=Mode(tags={"team-a": "pre_call", "team-b": ["post_call", "pre_call"]}, default="post_call") + ) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["post_call", "pre_call"] + + +@pytest.mark.asyncio +async def test_context_mode_omitted_when_event_hook_absent(): + g = _make_guardrail(event_hook=None) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert "mode" not in _posted_payload(g)["context"] + + @pytest.mark.asyncio async def test_identity_key_and_team_coalesce_alias_over_id(): g = _make_guardrail() @@ -426,6 +526,7 @@ async def test_application_source_from_agent_id(): ) assert _posted_payload(g)["application"] == {"source": "analytics-app", "name": "Analytics"} + @pytest.mark.asyncio async def test_request_block_raises_guardrail_exception_with_reason(): g = _make_guardrail() @@ -560,6 +661,34 @@ async def test_response_envelope_and_block_replaces_response(): assert payload["request"]["structured_messages"] == [{"role": "user", "content": "original prompt"}] +@pytest.mark.asyncio +async def test_post_call_resolves_request_from_responses_input_when_messages_absent(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + response = ModelResponse( + choices=[Choices(finish_reason="stop", index=0, message=Message(content="answer", role="assistant"))], + model="gpt-4o-mini", + ) + request_data = { + "model": "gpt-4o-mini", + "input": "responses-surface prompt", + "response": response, + "litellm_metadata": {"user_api_key_request_route": "/v1/responses"}, + } + + await g.apply_guardrail( + inputs={"texts": ["answer"], "model": "gpt-4o-mini"}, + request_data=request_data, + input_type="response", + logging_obj=_logging_obj(), + ) + + payload = _posted_payload(g) + assert payload["event"]["type"] == "post_call" + messages = payload["request"]["structured_messages"] + assert any(m.get("content") == "responses-surface prompt" for m in messages) + + @pytest.mark.asyncio async def test_post_call_fail_closed_raises_modify_response_exception(): g = _make_guardrail(unreachable_fallback="fail_closed") @@ -731,3 +860,210 @@ async def test_unreachable_http_status_fail_closed_blocks(): await g.apply_guardrail( inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj() ) + + +@pytest.mark.asyncio +async def test_post_call_preserves_anthropic_tool_blocks_in_request_messages(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + anthropic_messages = [ + {"role": "user", "content": "What's the weather in Paris?"}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Paris"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": "18C, cloudy", + } + ], + }, + ] + response = { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Mild and cloudy."}], + "stop_reason": "end_turn", + "model": "claude-sonnet-5", + } + await g.apply_guardrail( + inputs={"texts": ["Mild and cloudy."], "model": "claude-sonnet-5"}, + request_data={ + "model": "claude-sonnet-5", + "messages": anthropic_messages, + "response": response, + }, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["request"]["structured_messages"] == anthropic_messages + assert payload["response"]["finish_reason"] == "end_turn" + + +@pytest.mark.asyncio +async def test_pre_call_preserves_anthropic_tool_blocks_in_structured_messages(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + anthropic_messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Paris"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": "18C", + } + ], + }, + ] + await g.apply_guardrail( + inputs={"structured_messages": anthropic_messages, "model": "claude-sonnet-5"}, + request_data={"model": "claude-sonnet-5", "messages": anthropic_messages}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["request"]["structured_messages"] == anthropic_messages + + +@pytest.mark.asyncio +async def test_response_finish_reason_from_openai_choices_still_works(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + response = ModelResponse( + choices=[Choices(finish_reason="tool_calls", index=0, message=Message(content=None, role="assistant"))], + model="gpt-4o-mini", + ) + await g.apply_guardrail( + inputs={ + "texts": [], + "tool_calls": [ + ChatCompletionMessageToolCall( + id="c1", + type="function", + function=Function(name="f", arguments="{}"), + ) + ], + }, + request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}], "response": response}, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["response"]["finish_reason"] == "tool_calls" + assert payload["response"]["tool_calls"] == [ + {"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}} + ] + + +@pytest.mark.parametrize( + ("response", "expected"), + [ + (None, None), + ({"choices": "invalid"}, None), + ({"choices": [{"finish_reason": "length"}]}, "length"), + ({"choices": [{"stop_reason": "end_turn"}]}, "end_turn"), + ({"choices": [{}]}, None), + (SimpleNamespace(stop_reason="end_turn"), "end_turn"), + ], +) +def test_response_finish_reason_handles_supported_shapes(response, expected): + assert _response_finish_reason(response) == expected + + +@pytest.mark.parametrize( + "request_data", + [ + {"input": ["ssn 123-45-6789"], "litellm_metadata": {"user_api_key_request_route": "/vllm/v1/embeddings"}}, + {"input": [[1, 2, 3]], "litellm_metadata": {}}, + {"input": "confidential memo", "litellm_metadata": {}}, + {"input": "confidential memo"}, + ], +) +def test_request_messages_not_resolved_for_unmapped_surfaces(request_data): + """Bodies from surfaces without a translation handler yield no messages, and never raise.""" + assert _request_structured_messages(request_data) is None + + +@pytest.mark.parametrize( + ("request_data", "expected"), + [ + ( + {"messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}}, + [{"role": "user", "content": "hi"}], + ), + ( + { + "input": [{"role": "user", "content": "weather in Paris?"}], + "litellm_metadata": {"user_api_key_request_route": "/v1/responses"}, + }, + [{"role": "user", "content": "weather in Paris?"}], + ), + ], +) +def test_request_messages_resolved_for_mapped_surfaces(request_data, expected): + assert _request_structured_messages(request_data) == expected + + +@pytest.mark.parametrize( + ("response", "expected"), + [ + ({"usage": {"input_tokens": 10, "output_tokens": 5}}, (10, 5)), + ({"usage": {"prompt_tokens": 7, "completion_tokens": 3}}, (7, 3)), + (SimpleNamespace(usage=Usage(prompt_tokens=7, completion_tokens=3)), (7, 3)), + ({"usage": {"prompt_tokens": 0, "input_tokens": 99}}, (0, None)), + ({"usage": {}}, None), + ({}, None), + ], +) +def test_build_usage_handles_openai_and_anthropic_shapes(response, expected): + usage = _build_usage(response) + if expected is None: + assert usage is None + else: + assert (usage.input_tokens, usage.output_tokens) == expected + + +@pytest.mark.asyncio +async def test_anthropic_non_streaming_response_reports_usage(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={ + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "response": { + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + }, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["usage"] == {"input_tokens": 10, "output_tokens": 5} + assert payload["response"]["finish_reason"] == "end_turn" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e84e9b74201..bf904dbe394 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -4,7 +4,10 @@ import pytest import litellm from litellm.caching import DualCache -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.llms.base_llm.guardrail_translation.utils import ( effective_skip_system_message_for_guardrail, @@ -1490,3 +1493,109 @@ class TestStreamingTransform: # None holdback treated as 0: full text emitted, no crash. assert "".join(_delta_text(i) for i in out) == "ABCDEF" + + +def _applied_guardrails(data: dict) -> list: + for key in ("metadata", "litellm_metadata"): + meta = data.get(key) + if isinstance(meta, dict) and isinstance(meta.get("applied_guardrails"), list): + return meta["applied_guardrails"] + return [] + + +class _TextsOnlyTranslation(BaseTranslation): + """Mimics a passthrough handler: hands the guardrail only `texts`, never + structured_messages, so a structured_messages-based guardrail no-ops.""" + + async def process_input_messages(self, data, guardrail_to_apply, litellm_logging_obj=None): # type: ignore[override] + await guardrail_to_apply.apply_guardrail( + inputs={"texts": ["payload"]}, + request_data=data, + input_type="request", + logging_obj=litellm_logging_obj, + ) + return data + + async def process_output_response( # type: ignore[override] + self, + response, + guardrail_to_apply, + litellm_logging_obj=None, + user_api_key_dict=None, + request_data=None, + ): + return response + + +class _SelfLoggingGuardrail(CustomGuardrail): + records_own_guardrail_information = True + + def __init__(self, *, self_add: bool): + super().__init__(guardrail_name="self-logging") + self._self_add = self_add + + def should_run_guardrail(self, data, event_type): # type: ignore[override] + return True + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + if self._self_add: + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) + return inputs + + +class _AutoLoggingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="auto-logging") + + def should_run_guardrail(self, data, event_type): # type: ignore[override] + return True + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + return inputs + + +class TestAppliedGuardrailsReflectsExecution: + """The unified hook must not auto-mark a self-logging guardrail + (records_own_guardrail_information) as applied; such a guardrail owns that + decision and marks itself only when it actually ran (LIT-4650). Ordinary + guardrails are still auto-marked by the hook after dispatch.""" + + @staticmethod + def _data(guardrail): + return { + "guardrail_to_apply": guardrail, + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello world"}], + } + + async def _run(self, guardrail): + unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation} + data = self._data(guardrail) + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=None, + cache=DualCache(), + data=data, + call_type=CallTypes.pass_through.value, + ) + return data + + @pytest.mark.asyncio + async def test_self_logging_guardrail_is_not_auto_marked_applied(self): + data = await self._run(_SelfLoggingGuardrail(self_add=False)) + assert "self-logging" not in _applied_guardrails(data) + + @pytest.mark.asyncio + async def test_self_logging_guardrail_that_self_marks_is_applied(self): + data = await self._run(_SelfLoggingGuardrail(self_add=True)) + assert "self-logging" in _applied_guardrails(data) + + @pytest.mark.asyncio + async def test_ordinary_guardrail_is_auto_marked_applied(self): + data = await self._run(_AutoLoggingGuardrail()) + assert "auto-logging" in _applied_guardrails(data) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bad76864ca7..087aaec9215 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1504,50 +1504,6 @@ def test_team_info_masking(): assert "public-test-key" not in str(exc_info.value) -def test_embedding_input_array_of_tokens(client_no_auth): - """ - Test to bypass decoding input as array of tokens for selected providers - - Ref: https://github.com/BerriAI/litellm/issues/10113 - """ - from litellm.proxy import proxy_server - - # The client_no_auth fixture should initialize the router - # Assert this to catch any router initialization regressions - assert proxy_server.llm_router is not None, ( - "llm_router is None after client_no_auth fixture initialized. " - "This indicates a router initialization issue that should be investigated." - ) - - try: - with mock.patch.object( - proxy_server.llm_router, - "aembedding", - return_value=example_embedding_result, - ) as mock_aembedding: - test_data = { - "model": "vllm_embed_model", - "input": [[2046, 13269, 158208]], - } - - response = client_no_auth.post("/v1/embeddings", json=test_data) - - # Assert that aembedding was called, and that input was not modified - mock_aembedding.assert_called_once() - call_args, call_kwargs = mock_aembedding.call_args - assert call_kwargs["model"] == "vllm_embed_model" - assert call_kwargs["input"] == [[2046, 13269, 158208]] - - assert response.status_code == 200 - result = response.json() - print(len(result["data"][0]["embedding"])) - assert ( - len(result["data"][0]["embedding"]) > 10 - ) # this usually has len==1536 so - except Exception as e: - pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") - - @pytest.mark.asyncio async def test_get_all_team_models(): """ diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index f506b9665a6..93b3ef1cce8 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -12,24 +12,25 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque @pytest.mark.parametrize( - "route_type", + "route_type, required_body_params", [ - "atext_completion", - "acompletion", - "aembedding", - "aimage_generation", - "aspeech", - "atranscription", - "amoderation", - "arerank", + ("atext_completion", {}), + ("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}), + ("aembedding", {"input": "Hello"}), + ("aimage_generation", {}), + ("aspeech", {}), + ("atranscription", {}), + ("amoderation", {}), + ("arerank", {}), ], ) @pytest.mark.asyncio -async def test_route_request_dynamic_credentials(route_type): +async def test_route_request_dynamic_credentials(route_type, required_body_params): data = { "model": "openai/gpt-4o-mini-2024-07-18", "api_key": "my-bad-key", "api_base": "https://api.openai.com/v1 ", + **required_body_params, } llm_router = MagicMock() # Ensure that the dynamic method exists on the llm_router mock. @@ -887,3 +888,59 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value(): call_kwargs = llm_router.acompletion.call_args[1] assert call_kwargs["enable_tag_filtering"] is True + + +@pytest.mark.parametrize( + "route_type, param, route", + [ + ("acompletion", "messages", "/chat/completions"), + ("aembedding", "input", "/embeddings"), + ], +) +@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None}]) +def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra): + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra}) + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == param + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.detail == {"error": f"{route}: Missing required parameter: '{param}'."} + + +@pytest.mark.parametrize( + "route_type, data", + [ + ("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}), + ("acompletion", {"model": "gpt-4o", "messages": []}), + ("atext_completion", {"model": "gpt-4o"}), + ("aembedding", {"model": "text-embedding-3-small", "input": "hi"}), + ("arerank", {"model": "rerank-model"}), + ("aimage_generation", {"model": "dall-e-3"}), + ], +) +def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data): + from litellm.proxy.route_llm_request import raise_if_required_body_param_missing + + raise_if_required_body_param_missing(route_type=route_type, data=data) + + +@pytest.mark.asyncio +async def test_route_request_rejects_chat_completion_without_messages(): + """A /chat/completions body without `messages` used to splat into + Router.acompletion() and surface the resulting TypeError as a 500.""" + from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError + + llm_router = MagicMock() + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + await route_request({"model": "gpt-4o"}, llm_router, None, "acompletion") + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == "messages" + llm_router.acompletion.assert_not_called() diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index 44dfa240d42..83b9c34636e 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -198,6 +198,54 @@ async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error(): assert "not supported" in str(excinfo.value).lower() +@pytest.mark.asyncio +async def test_aresponses_drops_stream_options(): + """The Responses API rejects include_usage, so include_usage-only stream_options must never reach the wire.""" + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_stream_options_test", "gpt-5.5"), 200 + ) + + await litellm.aresponses( + model="openai/gpt-5.5", + api_key="fake-api-key", + input="hi", + stream_options={"include_usage": True}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert "stream_options" not in request_body + + +@pytest.mark.asyncio +async def test_aresponses_keeps_include_obfuscation_in_stream_options(): + """include_obfuscation is a valid Responses API stream option and must survive the include_usage strip.""" + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_stream_options_obfuscation", "gpt-5.5"), 200 + ) + + await litellm.aresponses( + model="openai/gpt-5.5", + api_key="fake-api-key", + input="hi", + stream_options={"include_usage": True, "include_obfuscation": False}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert request_body["stream_options"] == {"include_obfuscation": False} + + @pytest.mark.asyncio async def test_aresponses_request_level_drop_params_drops_bedrock_mantle_service_tier( monkeypatch, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index cbf3a2cc1f6..89c38c6274f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -25,8 +25,17 @@ vi.mock("@/components/networking", () => ({ // Mock the child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
{data.date}
+Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}
+Total Requests: {data.metrics.api_requests}
+Successful: {data.metrics.successful_requests}
+Failed: {data.metrics.failed_requests}
+Total Tokens: {data.metrics.total_tokens}
++ Total {capitalizedEntityLabel}s: {entityCount} +
+Spend by {capitalizedEntityLabel}:
+ {Object.entries(data.breakdown.entities || {}) + .sort(([, a], [, b]) => { + const spendA = (a as EntityMetrics).metrics.spend; + const spendB = (b as EntityMetrics).metrics.spend; + return spendB - spendA; + }) + .slice(0, 5) + .map(([entity, entityData]) => { + const metrics = entityData as EntityMetrics; + return ( ++ {getEntityLabel(entity, metrics.metadata)}: $ + {formatNumberWithCommas(metrics.metrics.spend, 2)} +
+ ); + })} + {entityCount > 5 &&...and {entityCount - 5} more
} +{data.metadata.alias}
+Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}
+Requests: {data.metrics.api_requests.toLocaleString()}
++ Successful: {data.metrics.successful_requests.toLocaleString()} +
+Failed: {data.metrics.failed_requests.toLocaleString()}
+Tokens: {data.metrics.total_tokens.toLocaleString()}
+{data.date}
-- Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)} -
-Total Requests: {data.metrics.api_requests}
-Successful: {data.metrics.successful_requests}
-Failed: {data.metrics.failed_requests}
-Total Tokens: {data.metrics.total_tokens}
-- Total {capitalizedEntityLabel}s: {entityCount} -
-Spend by {capitalizedEntityLabel}:
- {Object.entries(data.breakdown.entities || {}) - .sort(([, a], [, b]) => { - const spendA = (a as EntityMetrics).metrics.spend; - const spendB = (b as EntityMetrics).metrics.spend; - return spendB - spendA; - }) - .slice(0, 5) - .map(([entity, entityData]) => { - const metrics = entityData as EntityMetrics; - return ( -- {getEntityLabel(entity, metrics.metadata)}: $ - {formatNumberWithCommas(metrics.metrics.spend, 2)} -
- ); - })} - {entityCount > 5 && ( -...and {entityCount - 5} more
- )} -{data.metadata.alias}
-Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}
-Requests: {data.metrics.api_requests.toLocaleString()}
-- Successful: {data.metrics.successful_requests.toLocaleString()} -
-Failed: {data.metrics.failed_requests.toLocaleString()}
-Tokens: {data.metrics.total_tokens.toLocaleString()}
-