diff --git a/litellm/__init__.py b/litellm/__init__.py index b9da0524095..c868ae55b4f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -416,6 +416,7 @@ custom_prometheus_metadata_labels: List[str] = [] custom_prometheus_tags: List[str] = [] prometheus_metrics_config: Optional[List] = None prometheus_emit_stream_label: bool = False +prometheus_user_budget_label_include_email_alias: bool = False prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000 prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0 prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0 diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 3cf1d911d7f..3f4e54382c9 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -87,6 +87,16 @@ class CachingHandlerResponse(BaseModel): in_memory_cache_obj = InMemoryCache() +def _is_chat_completion_cached_dict(cached_result: dict) -> bool: + cached_id = cached_result.get("id") + if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"): + return True + obj = cached_result.get("object") + if isinstance(obj, str): + return obj.startswith("chat.completion") + return "choices" in cached_result + + def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bool: """ When stream=True, do not run success callbacks at cache-hit time. @@ -861,27 +871,47 @@ class LLMCachingHandler: elif (call_type == "aresponses" or call_type == "responses") and isinstance( cached_result, dict ): - from litellm.responses.streaming_iterator import ( - CachedResponsesAPIStreamingIterator, - ) - - response_obj = ResponsesAPIResponse(**cached_result) - if ( - hasattr(response_obj, "_hidden_params") - and response_obj._hidden_params is not None - and isinstance(response_obj._hidden_params, dict) - ): - response_obj._hidden_params["cache_hit"] = True - - if kwargs.get("stream", False) is True: - cached_result = CachedResponsesAPIStreamingIterator( - response=response_obj, - logging_obj=logging_obj, - request_data=kwargs, - call_type=call_type, - ) + use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result) + if use_chat_completion_cache: + if kwargs.get("stream", False) is True: + bridge_call_type = ( + CallTypes.acompletion.value + if call_type == "aresponses" + else CallTypes.completion.value + ) + cached_result = self._convert_cached_stream_response( + cached_result=cached_result, + call_type=bridge_call_type, + logging_obj=logging_obj, + model=model, + ) + else: + cached_result = convert_to_model_response_object( + response_object=cached_result, + model_response_object=ModelResponse(), + ) else: - cached_result = response_obj + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + + response_obj = ResponsesAPIResponse(**cached_result) + if ( + hasattr(response_obj, "_hidden_params") + and response_obj._hidden_params is not None + and isinstance(response_obj._hidden_params, dict) + ): + response_obj._hidden_params["cache_hit"] = True + + if kwargs.get("stream", False) is True: + cached_result = CachedResponsesAPIStreamingIterator( + response=response_obj, + logging_obj=logging_obj, + request_data=kwargs, + call_type=call_type, + ) + else: + cached_result = response_obj if ( hasattr(cached_result, "_hidden_params") diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index ce398ee8288..2de7bda6467 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -37,6 +37,15 @@ class ResponsesToCompletionBridgeHandler: stream = litellm_params.get("stream", False) return bool(stream) + @staticmethod + def _is_preformatted_cached_chat_stream(result: Any) -> bool: + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + + return ( + isinstance(result, CustomStreamWrapper) + and result.custom_llm_provider == "cached_response" + ) + @staticmethod def _coerce_response_object( response_obj: Any, @@ -177,6 +186,8 @@ class ResponsesToCompletionBridgeHandler: **request_data, ) + from litellm.types.utils import ModelResponse + stream = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( @@ -192,6 +203,8 @@ class ResponsesToCompletionBridgeHandler: api_key=kwargs.get("api_key"), json_mode=kwargs.get("json_mode"), ) + elif isinstance(result, ModelResponse): + return result elif not stream: responses_api_response = self._collect_response_from_stream(result) return self.transformation_handler.transform_response( @@ -208,6 +221,10 @@ class ResponsesToCompletionBridgeHandler: json_mode=kwargs.get("json_mode"), ) else: + if self._is_preformatted_cached_chat_stream(result): + return self._apply_post_stream_processing( + result, model, custom_llm_provider + ) completion_stream = self.transformation_handler.get_model_response_iterator( streaming_response=result, # type: ignore sync_stream=True, @@ -256,6 +273,8 @@ class ResponsesToCompletionBridgeHandler: aresponses=True, ) + from litellm.types.utils import ModelResponse + stream = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( @@ -271,6 +290,8 @@ class ResponsesToCompletionBridgeHandler: api_key=kwargs.get("api_key"), json_mode=kwargs.get("json_mode"), ) + elif isinstance(result, ModelResponse): + return result elif not stream: responses_api_response = await self._collect_response_from_stream_async( result @@ -289,6 +310,10 @@ class ResponsesToCompletionBridgeHandler: json_mode=kwargs.get("json_mode"), ) else: + if self._is_preformatted_cached_chat_stream(result): + return self._apply_post_stream_processing( + result, model, custom_llm_provider + ) completion_stream = self.transformation_handler.get_model_response_iterator( streaming_response=result, # type: ignore sync_stream=False, diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 32423f23314..e3cbf422e5d 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1141,6 +1141,14 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): event_type = parsed_chunk.get("type") if isinstance(event_type, ResponsesAPIStreamEvents): event_type = event_type.value + + if parsed_chunk.get("object") == "chat.completion.chunk" or ( + event_type is None + and isinstance(parsed_chunk.get("choices"), list) + and parsed_chunk.get("choices") + ): + return ModelResponseStream(**parsed_chunk) + verbose_logger.debug(f"Chat provider: Processing event type: {event_type}") if event_type == "response.created": diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 30af0dcb8ed..2c63455565c 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -3540,6 +3540,10 @@ class PrometheusLogger(CustomLogger): user_object.budget_reset_at = user_info.budget_reset_at if user_object.max_budget is None and user_info.max_budget is not None: user_object.max_budget = user_info.max_budget + if user_info.user_email is not None: + user_object.user_email = user_info.user_email + if user_info.user_alias is not None: + user_object.user_alias = user_info.user_alias return user_object @@ -3556,6 +3560,8 @@ class PrometheusLogger(CustomLogger): """ enum_values = UserAPIKeyLabelValues( user=user.user_id, + user_email=user.user_email or "", + user_alias=user.user_alias or "", ) _labels = prometheus_label_factory( diff --git a/litellm/interactions/litellm_responses_transformation/streaming_iterator.py b/litellm/interactions/litellm_responses_transformation/streaming_iterator.py index 72a3afbc3c5..567e9b523e8 100644 --- a/litellm/interactions/litellm_responses_transformation/streaming_iterator.py +++ b/litellm/interactions/litellm_responses_transformation/streaming_iterator.py @@ -2,7 +2,7 @@ Streaming iterator for transforming Responses API stream to Interactions API stream. """ -from typing import Any, AsyncIterator, Dict, Iterator, Optional, cast +from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, cast from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, @@ -15,6 +15,7 @@ from litellm.types.interactions import ( InteractionsAPIStreamingResponse, ) from litellm.types.llms.openai import ( + ContentPartAddedEvent, OutputTextDeltaEvent, ResponseCompletedEvent, ResponseCreatedEvent, @@ -51,6 +52,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: self.collected_text = "" self.sent_interaction_start = False self.sent_content_start = False + self._pending_events: List[InteractionsAPIStreamingResponse] = [] def _transform_responses_chunk_to_interactions_chunk( self, @@ -80,7 +82,49 @@ class LiteLLMResponsesInteractionsStreamingIterator: ) self.collected_text += delta_text - # Send interaction.start if not sent + # Fallback: emit interaction.start, and queue content.start carrying this + # delta so the first token is preserved in the stream. + if not self.sent_interaction_start: + self.sent_interaction_start = True + self.sent_content_start = True + self._pending_events.append( + InteractionsAPIStreamingResponse( + event_type="content.start", + id=getattr(responses_chunk, "item_id", None), + object="content", + delta={"type": "text", "text": delta_text}, + ) + ) + return InteractionsAPIStreamingResponse( + event_type="interaction.start", + id=getattr(responses_chunk, "item_id", None) + or f"interaction_{id(self)}", + object="interaction", + status="in_progress", + model=self.model, + ) + + # Fallback: emit content.start if ContentPartAddedEvent never arrived + if not self.sent_content_start: + self.sent_content_start = True + return InteractionsAPIStreamingResponse( + event_type="content.start", + id=getattr(responses_chunk, "item_id", None), + object="content", + delta={"type": "text", "text": delta_text}, + ) + + # Normal path: emit content.delta with type field + return InteractionsAPIStreamingResponse( + event_type="content.delta", + id=getattr(responses_chunk, "item_id", None), + object="content", + delta={"type": "text", "text": delta_text}, + ) + + # Handle ContentPartAddedEvent -> content.start (arrives before text deltas) + if isinstance(responses_chunk, ContentPartAddedEvent): + # Fallback: emit interaction.start if ResponseCreatedEvent never arrived if not self.sent_interaction_start: self.sent_interaction_start = True return InteractionsAPIStreamingResponse( @@ -91,8 +135,6 @@ class LiteLLMResponsesInteractionsStreamingIterator: status="in_progress", model=self.model, ) - - # Send content.start if not sent if not self.sent_content_start: self.sent_content_start = True return InteractionsAPIStreamingResponse( @@ -101,14 +143,7 @@ class LiteLLMResponsesInteractionsStreamingIterator: object="content", delta={"type": "text", "text": ""}, ) - - # Send content.delta - return InteractionsAPIStreamingResponse( - event_type="content.delta", - id=getattr(responses_chunk, "item_id", None), - object="content", - delta={"text": delta_text}, - ) + return None # Handle ResponseCreatedEvent or ResponseInProgressEvent -> interaction.start if isinstance(responses_chunk, (ResponseCreatedEvent, ResponseInProgressEvent)): @@ -172,6 +207,10 @@ class LiteLLMResponsesInteractionsStreamingIterator: delattr(self, "_pending_interaction_complete") return pending + # Drain events queued from a prior chunk (e.g. content.start emitted alongside + # the interaction.start fallback for the first OutputTextDeltaEvent). + if self._pending_events: + return self._pending_events.pop(0) # Use a loop instead of recursion to avoid stack overflow sync_iterator = cast( SyncResponsesAPIStreamingIterator, self.responses_stream_iterator @@ -237,6 +276,10 @@ class LiteLLMResponsesInteractionsStreamingIterator: delattr(self, "_pending_interaction_complete") return pending + # Drain events queued from a prior chunk (e.g. content.start emitted alongside + # the interaction.start fallback for the first OutputTextDeltaEvent). + if self._pending_events: + return self._pending_events.pop(0) # Use a loop instead of recursion to avoid stack overflow async_iterator = cast( ResponsesAPIStreamingIterator, self.responses_stream_iterator diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index 0602b1c2f62..620bc91732d 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -5,6 +5,7 @@ from typing import Any, Dict, List, Literal, Optional, Union, cast from httpx import Headers, Response +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret_str @@ -263,9 +264,32 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): cancelling_at=None, cancelled_at=None, request_counts=None, - metadata=original_request.get("metadata", {}), + metadata=self._get_openai_compatible_batch_metadata( + original_request.get("metadata", {}) + ), ) + @staticmethod + def _get_openai_compatible_batch_metadata(metadata: Any) -> Dict[str, str]: + """ + OpenAI Batch metadata only accepts string values. + """ + if not isinstance(metadata, dict): + return {} + + sanitized_metadata: Dict[str, str] = {} + for key, value in metadata.items(): + if key == "standard_logging_guardrail_information" or value is None: + continue + + str_key = str(key) + if isinstance(value, str): + sanitized_metadata[str_key] = value + else: + sanitized_metadata[str_key] = safe_dumps(value) + + return sanitized_metadata + def transform_retrieve_batch_request( self, batch_id: str, diff --git a/litellm/llms/bedrock/chat/invoke_agent/transformation.py b/litellm/llms/bedrock/chat/invoke_agent/transformation.py index e4072c24557..c88fa32b6a0 100644 --- a/litellm/llms/bedrock/chat/invoke_agent/transformation.py +++ b/litellm/llms/bedrock/chat/invoke_agent/transformation.py @@ -299,9 +299,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): ) def _get_response_stream_shape(self): - from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE + from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape - return BEDROCK_RESPONSE_STREAM_SHAPE + return get_bedrock_response_stream_shape() def _extract_response_content(self, events: InvokeAgentEventList) -> str: """Extract the final response content from parsed events.""" diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 92ca75db95b..7a9916f1f31 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -68,9 +68,9 @@ from litellm.utils import CustomStreamWrapper, get_secret from ..base_aws_llm import BaseAWSLLM from ..common_utils import ( - BEDROCK_RESPONSE_STREAM_SHAPE, BedrockError, ModelResponseIterator, + get_bedrock_response_stream_shape, get_bedrock_tool_name, ) @@ -1828,7 +1828,8 @@ class AWSEventStreamDecoder: yield self._chunk_parser(chunk_data=_data) def _parse_message_from_event(self, event) -> Optional[str]: - if BEDROCK_RESPONSE_STREAM_SHAPE is None: + response_stream_shape = get_bedrock_response_stream_shape() + if response_stream_shape is None: raise BedrockError( status_code=500, message=( @@ -1837,9 +1838,7 @@ class AWSEventStreamDecoder: ), ) response_dict = event.to_response_dict() - parsed_response = self.parser.parse( - response_dict, BEDROCK_RESPONSE_STREAM_SHAPE - ) + parsed_response = self.parser.parse(response_dict, response_stream_shape) if response_dict["status_code"] != 200: decoded_body = response_dict["body"].decode() diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 0256d5d4b95..4f4729e4019 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -4,6 +4,7 @@ from __future__ import annotations Common utilities used across bedrock chat/embedding/image generation """ +import functools import json import os from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union @@ -963,10 +964,8 @@ def _load_bedrock_response_stream_shape(): """ Load the ResponseStream shape from botocore's bundled bedrock-runtime schema. - Called once at module import time; the result is stored in - ``BEDROCK_RESPONSE_STREAM_SHAPE`` and reused for the process lifetime. Returns ``None`` if botocore is unavailable or the service model cannot be - loaded, so the module still imports cleanly. + loaded. """ try: from botocore.loaders import Loader @@ -977,15 +976,22 @@ def _load_bedrock_response_stream_shape(): return ServiceModel(service_dict).shape_for("ResponseStream") except Exception as e: verbose_logger.warning( - "litellm: could not pre-load bedrock-runtime response stream shape " + "litellm: could not load bedrock-runtime response stream shape " "— Bedrock event-stream decoding will be unavailable. Error: %s", e, ) return None -# Eagerly resolved once per process — avoids per-instance or per-request disk I/O. -BEDROCK_RESPONSE_STREAM_SHAPE = _load_bedrock_response_stream_shape() +@functools.lru_cache(maxsize=1) +def get_bedrock_response_stream_shape(): + """ + Lazily load and cache the bedrock-runtime ResponseStream shape for the process. + + Avoids importing botocore (and logging warnings) unless Bedrock event-stream + decoding is actually needed. + """ + return _load_bedrock_response_stream_shape() class BedrockEventStreamDecoderBase: @@ -999,7 +1005,8 @@ class BedrockEventStreamDecoderBase: self.parser = EventStreamJSONParser() def _parse_message_from_event(self, event) -> Optional[str]: - if BEDROCK_RESPONSE_STREAM_SHAPE is None: + response_stream_shape = get_bedrock_response_stream_shape() + if response_stream_shape is None: raise BedrockError( status_code=500, message=( @@ -1008,9 +1015,7 @@ class BedrockEventStreamDecoderBase: ), ) response_dict = event.to_response_dict() - parsed_response = self.parser.parse( - response_dict, BEDROCK_RESPONSE_STREAM_SHAPE - ) + parsed_response = self.parser.parse(response_dict, response_stream_shape) if response_dict["status_code"] != 200: decoded_body = response_dict["body"].decode() diff --git a/litellm/llms/bedrock/embed/cohere_transformation.py b/litellm/llms/bedrock/embed/cohere_transformation.py index d00cb74aae0..2c0dc834144 100644 --- a/litellm/llms/bedrock/embed/cohere_transformation.py +++ b/litellm/llms/bedrock/embed/cohere_transformation.py @@ -22,7 +22,7 @@ class BedrockCohereEmbeddingConfig: ) -> dict: for k, v in non_default_params.items(): if k == "encoding_format": - optional_params["embedding_types"] = v + optional_params["embedding_types"] = v if isinstance(v, list) else [v] elif k == "dimensions": optional_params["output_dimension"] = v return optional_params diff --git a/litellm/llms/deepseek/messages/transformation.py b/litellm/llms/deepseek/messages/transformation.py new file mode 100644 index 00000000000..ad60478960e --- /dev/null +++ b/litellm/llms/deepseek/messages/transformation.py @@ -0,0 +1,133 @@ +""" +DeepSeek Anthropic-compatible messages transformation config. +""" + +from typing import Any, Dict, List, Optional, Tuple + +import litellm +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams + + +class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): + """ + DeepSeek exposes an Anthropic-compatible Messages API at + https://api.deepseek.com/anthropic. + + It accepts the native Anthropic Messages conversation shape, including + thinking blocks in assistant history, but rejects Anthropic's explicit + custom-tool discriminator (`{"type": "custom"}`). + """ + + @property + def custom_llm_provider(self) -> Optional[str]: + return "deepseek" + + @staticmethod + def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + return api_key or get_secret_str("DEEPSEEK_API_KEY") or litellm.api_key + + @staticmethod + def get_api_base(api_base: Optional[str] = None) -> str: + return ( + api_base + or get_secret_str("DEEPSEEK_ANTHROPIC_API_BASE") + or get_secret_str("DEEPSEEK_API_BASE") + or "https://api.deepseek.com/anthropic" + ) + + def validate_anthropic_messages_environment( + self, + headers: dict, + model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> Tuple[dict, Optional[str]]: + dynamic_api_key = self.get_api_key(api_key=api_key) + + if ( + "x-api-key" not in headers + and "authorization" not in headers + and dynamic_api_key is not None + ): + headers["x-api-key"] = dynamic_api_key + + if "anthropic-version" not in headers: + headers["anthropic-version"] = "2023-06-01" + if "content-type" not in headers: + headers["content-type"] = "application/json" + + headers = self._update_headers_with_anthropic_beta( + headers=headers, + optional_params=optional_params, + custom_llm_provider=self.custom_llm_provider or "deepseek", + ) + + return headers, api_base + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + base_url = self.get_api_base(api_base=api_base).rstrip("/") + + if base_url.endswith("/v1/messages") and "/anthropic/" in base_url: + return base_url + if base_url.endswith("/v1/messages"): + base_url = base_url[: -len("/v1/messages")] + if base_url.endswith("/v1"): + base_url = base_url[: -len("/v1")] + if base_url.endswith("/beta"): + base_url = base_url[: -len("/beta")] + + if not base_url.endswith("/anthropic") and "/anthropic/" not in base_url: + base_url = f"{base_url}/anthropic" + + return f"{base_url}/v1/messages" + + @staticmethod + def _sanitize_tools_for_deepseek(tools: Any) -> Any: + if not isinstance(tools, list): + return tools + + sanitized_tools = [] + for tool in tools: + if isinstance(tool, dict) and tool.get("type") == "custom": + sanitized_tool = dict(tool) + sanitized_tool.pop("type", None) + sanitized_tools.append(sanitized_tool) + else: + sanitized_tools.append(tool) + return sanitized_tools + + def transform_anthropic_messages_request( + self, + model: str, + messages: List[Dict], + anthropic_messages_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + anthropic_messages_request = super().transform_anthropic_messages_request( + model=model, + messages=messages, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + if "tools" in anthropic_messages_request: + anthropic_messages_request["tools"] = self._sanitize_tools_for_deepseek( + anthropic_messages_request["tools"] + ) + return anthropic_messages_request diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index 50c8ee4220e..6c15d642f8c 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -1,3 +1,4 @@ +import functools import json from typing import AsyncIterator, Iterator, List, Optional, Union @@ -22,14 +23,22 @@ def _load_sagemaker_response_stream_shape(): ) except Exception as e: verbose_logger.warning( - "litellm: could not pre-load sagemaker-runtime response stream shape " + "litellm: could not load sagemaker-runtime response stream shape " "— SageMaker event-stream decoding will be unavailable. Error: %s", e, ) return None -SAGEMAKER_RESPONSE_STREAM_SHAPE = _load_sagemaker_response_stream_shape() +@functools.lru_cache(maxsize=1) +def get_sagemaker_response_stream_shape(): + """ + Lazily load and cache the sagemaker-runtime stream shape for the process. + + Avoids importing botocore (and logging warnings) unless SageMaker event-stream + decoding is actually needed. + """ + return _load_sagemaker_response_stream_shape() class SagemakerError(BaseLLMException): @@ -207,7 +216,8 @@ class AWSEventStreamDecoder: verbose_logger.error(f"Final error parsing accumulated JSON: {e}") def _parse_message_from_event(self, event) -> Optional[str]: - if SAGEMAKER_RESPONSE_STREAM_SHAPE is None: + response_stream_shape = get_sagemaker_response_stream_shape() + if response_stream_shape is None: raise SagemakerError( status_code=500, message=( @@ -216,9 +226,7 @@ class AWSEventStreamDecoder: ), ) response_dict = event.to_response_dict() - parsed_response = self.parser.parse( - response_dict, SAGEMAKER_RESPONSE_STREAM_SHAPE - ) + parsed_response = self.parser.parse(response_dict, response_stream_shape) if response_dict["status_code"] != 200: raise ValueError(f"Bad response code, expected 200: {response_dict}") diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index a43d15a580f..dc27e87726a 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -1,6 +1,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from fastapi import HTTPException, status +from pydantic import BaseModel from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache @@ -53,6 +54,37 @@ def require_caller_user_id_for_non_admin( return user_api_key_dict.user_id +def _check_passthrough_routes_caller_permission( + data: BaseModel, + user_api_key_dict: UserAPIKeyAuth, + *, + entity: str = "key", +) -> None: + """ + Only proxy admins may set `allowed_passthrough_routes` (top-level or under + `metadata`) — it short-circuits the role-based route gate, so keys and teams + must be gated identically. + """ + # view-only admins excluded by design; blocked upstream from writes anyway + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: + return + if getattr(data, "allowed_passthrough_routes", None): + raise HTTPException( + status_code=403, + detail={ + "error": f"Only proxy admins can set `allowed_passthrough_routes` on a {entity}." + }, + ) + metadata = getattr(data, "metadata", None) + if isinstance(metadata, dict) and metadata.get("allowed_passthrough_routes"): + raise HTTPException( + status_code=403, + detail={ + "error": f"Only proxy admins can set `metadata.allowed_passthrough_routes` on a {entity}." + }, + ) + + def _is_user_team_admin( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable ) -> bool: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 7ff706eb91e..2ab147043d9 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -55,6 +55,7 @@ from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_k from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.management_endpoints.common_utils import ( + _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, _is_user_team_admin, _set_object_metadata_field, @@ -548,36 +549,6 @@ def _check_allowed_routes_caller_permission( ) -def _check_passthrough_routes_caller_permission( - data: BaseModel, - user_api_key_dict: UserAPIKeyAuth, -) -> None: - """ - Only proxy admins may set `allowed_passthrough_routes` on a key, either at - the top level of the request or nested under `metadata`. - - The route gate evaluates passthrough access ahead of the standard role - gate, so the field is restricted to admins to keep that ordering safe. - """ - if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: - return - if getattr(data, "allowed_passthrough_routes", None): - raise HTTPException( - status_code=403, - detail={ - "error": "Only proxy admins can set `allowed_passthrough_routes` on a key." - }, - ) - metadata = getattr(data, "metadata", None) - if isinstance(metadata, dict) and metadata.get("allowed_passthrough_routes"): - raise HTTPException( - status_code=403, - detail={ - "error": "Only proxy admins can set `metadata.allowed_passthrough_routes` on a key." - }, - ) - - async def validate_team_id_used_in_service_account_request( team_id: Optional[str], prisma_client: Optional[PrismaClient], diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 35e3d196e9e..86c4d6dcd9a 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -73,6 +73,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_utils import ( + _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, _is_user_team_admin, _set_object_metadata_field, @@ -1049,6 +1050,10 @@ async def new_team( # noqa: PLR0915 Member(role="admin", user_id=user_api_key_dict.user_id) ) + _check_passthrough_routes_caller_permission( + data, user_api_key_dict, entity="team" + ) + ## ADD TO MODEL TABLE _model_id = None if data.model_aliases is not None and isinstance(data.model_aliases, dict): @@ -1646,6 +1651,10 @@ async def update_team( # noqa: PLR0915 user_api_key_dict=user_api_key_dict, ) + _check_passthrough_routes_caller_permission( + data, user_api_key_dict, entity="team" + ) + if data.soft_budget is not None: max_budget_to_check = ( data.max_budget diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d89d3fa9c5..879914e5ac6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6917,6 +6917,15 @@ async def async_data_generator( # noqa: PLR0915 if isinstance(chunk, BaseModel): chunk = _serialize_streaming_chunk(chunk) + elif isinstance(chunk, bytes): + # Some upstream streaming iterators (e.g. AsyncGoogleGenAIGenerateContentStreamingIterator + # for /v1beta/.../streamGenerateContent) yield raw SSE bytes from Gemini. + # Decode to str so the f-string below does not emit a Python b'...' literal, + # and pass already-formatted SSE through unchanged to avoid double "data:" prefix. + chunk = chunk.decode("utf-8", errors="replace") + if chunk.startswith(("data:", "event:", ":")): + yield chunk if chunk.endswith("\n\n") else chunk + "\n\n" + continue elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk break diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 43a287f29bc..7b5c5ab2969 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -160,6 +160,7 @@ class UserAPIKeyLabelNames(Enum): END_USER = "end_user" USER = "user" USER_EMAIL = "user_email" + USER_ALIAS = "user_alias" API_KEY_HASH = "hashed_api_key" API_KEY_ALIAS = "api_key_alias" TEAM = "team" @@ -533,17 +534,9 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, ] - litellm_user_max_budget_metric = [ - UserAPIKeyLabelNames.USER.value, - ] + litellm_user_max_budget_metric = litellm_remaining_user_budget_metric - litellm_user_budget_remaining_hours_metric = [ - UserAPIKeyLabelNames.USER.value, - ] - - litellm_user_budget_remaining_hours_metric = [ - UserAPIKeyLabelNames.USER.value, - ] + litellm_user_budget_remaining_hours_metric = litellm_remaining_user_budget_metric litellm_remaining_api_key_requests_for_model = [ UserAPIKeyLabelNames.API_KEY_HASH.value, @@ -730,6 +723,22 @@ class PrometheusMetricLabels: ): custom_labels.append(UserAPIKeyLabelNames.STREAM.value) + _user_budget_metrics = { + "litellm_remaining_user_budget_metric", + "litellm_user_max_budget_metric", + "litellm_user_budget_remaining_hours_metric", + } + if ( + label_name in _user_budget_metrics + and litellm.prometheus_user_budget_label_include_email_alias is True + ): + for label in [ + UserAPIKeyLabelNames.USER_EMAIL.value, + UserAPIKeyLabelNames.USER_ALIAS.value, + ]: + if label not in default_labels and label not in custom_labels: + custom_labels.append(label) + if label_name in PrometheusMetricLabels._org_label_metrics: for label in [ UserAPIKeyLabelNames.ORG_ID.value, @@ -759,6 +768,7 @@ class UserAPIKeyLabelValues: end_user: Optional[str] = None user: Optional[str] = None user_email: Optional[str] = None + user_alias: Optional[str] = None hashed_api_key: Optional[str] = None api_key_alias: Optional[str] = None team: Optional[str] = None diff --git a/litellm/utils.py b/litellm/utils.py index cfe2dff822f..75c6d2dcf0f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8539,6 +8539,12 @@ class ProviderConfigManager: ) return MinimaxMessagesConfig() + elif litellm.LlmProviders.DEEPSEEK == provider: + from litellm.llms.deepseek.messages.transformation import ( + DeepSeekAnthropicMessagesConfig, + ) + + return DeepSeekAnthropicMessagesConfig() return None @staticmethod diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py index a179a21ba69..cb43f1abbdd 100644 --- a/tests/_vcr_conftest_common.py +++ b/tests/_vcr_conftest_common.py @@ -36,6 +36,75 @@ SAFE_BODY_MATCHER_NAME = "safe_body" KEY_FINGERPRINT_MATCHER_NAME = "key_fingerprint" KEY_FINGERPRINT_HEADER = "x-litellm-key-fp" +VCR_DIAG_DIR_ENV = "LITELLM_VCR_DIAG_DIR" +VCR_DIAG_DIR_DEFAULT = "test-results/vcr-diagnostics" + + +def _vcr_diag_dir() -> str: + return os.environ.get(VCR_DIAG_DIR_ENV) or VCR_DIAG_DIR_DEFAULT + + +def vcr_diag_write_line(msg: str) -> None: + try: + directory = _vcr_diag_dir() + os.makedirs(directory, exist_ok=True) + path = os.path.join(directory, f"{os.getpid()}.log") + with open(path, "a", encoding="utf-8") as fh: + fh.write(msg.rstrip("\n") + "\n") + except OSError: + pass + + +def reset_vcr_diag_dir() -> None: + if os.environ.get("PYTEST_XDIST_WORKER"): + return + directory = _vcr_diag_dir() + if not os.path.isdir(directory): + return + try: + names = os.listdir(directory) + except OSError: + return + for name in names: + if name.endswith(".log"): + try: + os.remove(os.path.join(directory, name)) + except OSError: + pass + + +def emit_vcr_diagnostic_log(terminalreporter) -> None: + directory = _vcr_diag_dir() + if not os.path.isdir(directory): + return + try: + files = sorted(f for f in os.listdir(directory) if f.endswith(".log")) + except OSError: + return + if not files: + return + terminalreporter.write_sep("=", "VCR DIAGNOSTIC LOG", bold=True) + terminalreporter.write_line( + f" source dir: {directory} (also archived as a CI artifact)" + ) + for name in files: + path = os.path.join(directory, name) + try: + with open(path, "r", encoding="utf-8") as fh: + content = fh.read() + except OSError as exc: + terminalreporter.write_line( + f" [failed to read {name}: {type(exc).__name__}: {exc}]" + ) + continue + if not content.strip(): + continue + terminalreporter.write_sep("-", name, bold=False) + for line in content.splitlines(): + terminalreporter.write_line(line) + terminalreporter.write_sep("=", bold=True) + + # Intentionally narrower than ``FILTERED_REQUEST_HEADERS``: AWS SigV4 headers # carry secrets but their values rotate on every call, so fingerprinting them # would defeat caching. @@ -91,6 +160,32 @@ VCR_IMAGE_B64_PLACEHOLDER = "dGVzdA==" VCR_FIXED_MULTIPART_BOUNDARY = "vcr-static-boundary" +def pin_httpx_multipart_boundary(monkeypatch) -> None: + try: + import httpx._multipart as _httpx_multipart + except ImportError: + return + + _original_init = _httpx_multipart.MultipartStream.__init__ + + def _init_with_fixed_boundary(self, data, files, boundary=None, **kwargs): + if boundary is None: + boundary = VCR_FIXED_MULTIPART_BOUNDARY.encode("ascii") + return _original_init(self, data=data, files=files, boundary=boundary, **kwargs) + + monkeypatch.setattr( + _httpx_multipart.MultipartStream, "__init__", _init_with_fixed_boundary + ) + + +@pytest.fixture(scope="session", autouse=True) +def _pin_multipart_boundary(): + monkeypatch = pytest.MonkeyPatch() + pin_httpx_multipart_boundary(monkeypatch) + yield + monkeypatch.undo() + + def _scrub_response(response): if not isinstance(response, dict): return response @@ -139,9 +234,17 @@ def _strip_image_b64_payloads(response): preserves all those checks while shrinking cassettes by ~99%. """ if not isinstance(response, dict): + vcr_diag_write_line( + f"[vcr-strip-b64] response is {type(response).__name__!r}, not " + "dict; skipping b64 scrub" + ) return response body = response.get("body") if not isinstance(body, dict): + vcr_diag_write_line( + f"[vcr-strip-b64] response['body'] is {type(body).__name__!r}, " + "not dict; skipping b64 scrub" + ) return response raw = body.get("string") if raw is None: @@ -151,12 +254,20 @@ def _strip_image_b64_payloads(response): try: text = bytes(raw).decode("utf-8") except UnicodeDecodeError: + vcr_diag_write_line( + "[vcr-strip-b64] response body bytes are not valid UTF-8; " + "skipping b64 scrub" + ) return response was_bytes = True elif isinstance(raw, str): text = raw was_bytes = False else: + vcr_diag_write_line( + f"[vcr-strip-b64] response['body']['string'] is " + f"{type(raw).__name__!r}, not bytes/str; skipping b64 scrub" + ) return response try: @@ -186,6 +297,35 @@ def _before_record_response(response): return filter_non_2xx_response(_scrub_response(_strip_image_b64_payloads(response))) +def _canonical_body(request) -> tuple[bytes, str]: + pre_type = type(getattr(request, "body", None)).__name__ + _materialize_iterable_body(request) + body = getattr(request, "body", None) + if body is None: + return b"", pre_type + if isinstance(body, bytes): + return body, pre_type + if isinstance(body, bytearray): + return bytes(body), pre_type + if isinstance(body, str): + return body.encode("utf-8"), pre_type + if isinstance(body, (dict, list)): + try: + return ( + json.dumps(body, sort_keys=True, separators=(",", ":")).encode("utf-8"), + pre_type, + ) + except (TypeError, ValueError): + pass + method = getattr(request, "method", "?") + uri = getattr(request, "uri", getattr(request, "url", "?")) + vcr_diag_write_line( + f"[vcr-canonical-body] FALLBACK: {method} {uri} body type " + f"{type(body).__name__!r} not coerced to bytes; comparing as b''" + ) + return b"", pre_type + + def _safe_body_matcher(r1, r2) -> None: """Compare request bodies as bytes; never invokes ``json.loads``. @@ -195,27 +335,47 @@ def _safe_body_matcher(r1, r2) -> None: This matcher is strictly more conservative — the only equivalence it gives up vs. the default is "JSON key order doesn't matter". """ - body1 = getattr(r1, "body", None) - body2 = getattr(r2, "body", None) + body1, pre1 = _canonical_body(r1) + body2, pre2 = _canonical_body(r2) if body1 == body2: return - - def _to_bytes(b): - if b is None: - return b"" - if isinstance(b, bytes): - return b - if isinstance(b, str): - return b.encode("utf-8") - return None - - n1 = _to_bytes(body1) - n2 = _to_bytes(body2) - if n1 is not None and n2 is not None and n1 == n2: - return + _emit_body_mismatch_diagnostic(r1, r2, body1, body2, pre1, pre2) raise AssertionError("request bodies differ") +def _emit_body_mismatch_diagnostic(r1, r2, body1, body2, pre1, pre2) -> None: + def _describe(label, asbytes, pre_type): + return ( + f" {label}: pre_canonical_type={pre_type!r} length={len(asbytes)} " + f"sha256={hashlib.sha256(asbytes).hexdigest()} " + f"preview={asbytes[:120]!r}" + ) + + method_a = getattr(r1, "method", "?") + method_b = getattr(r2, "method", "?") + url_a = getattr(r1, "uri", getattr(r1, "url", "?")) + url_b = getattr(r2, "uri", getattr(r2, "url", "?")) + lines = [ + "[vcr-safe-body-matcher] request body mismatch", + f" request[a]: {method_a} {url_a}", + f" request[b]: {method_b} {url_b}", + _describe("body[a]", body1, pre1), + _describe("body[b]", body2, pre2), + ] + if body1 != body2: + offset = next( + (i for i in range(min(len(body1), len(body2))) if body1[i] != body2[i]), + min(len(body1), len(body2)), + ) + start = max(0, offset - 100) + end_a = min(len(body1), offset + 100) + end_b = min(len(body2), offset + 100) + lines.append(f" first divergent byte offset: {offset}") + lines.append(f" window[a] @ {start}..{end_a}: {body1[start:end_a]!r}") + lines.append(f" window[b] @ {start}..{end_b}: {body2[start:end_b]!r}") + vcr_diag_write_line("\n".join(lines)) + + def _iter_header_values(headers, name: str): if headers is None: return @@ -271,6 +431,13 @@ def _compute_key_fingerprint(request) -> str: stable = _stable_key_value(header_name, text) parts.append(f"{header_name}={stable}") if not parts: + method = getattr(request, "method", "?") + uri = getattr(request, "uri", getattr(request, "url", "?")) + vcr_diag_write_line( + f"[vcr-key-fingerprint] no API key header found on {method} " + f"{uri}; falling back to 'no-key'. If this request should have " + "carried auth, something earlier in the pipeline stripped it." + ) return "no-key" digest = hashlib.sha256("\n".join(parts).encode("utf-8")).hexdigest() return digest[:16] @@ -360,6 +527,13 @@ def _normalize_multipart_boundary(request) -> None: elif isinstance(body, str): new_body = body.replace(current_boundary, VCR_FIXED_MULTIPART_BOUNDARY) else: + vcr_diag_write_line( + f"[vcr-multipart-normalize] body normalization SKIPPED: " + f"body type {type(body).__name__!r} is not bytes/bytearray/str. " + f"content-type={content_type_value!r}. " + f"Recorded body will retain the random boundary substring " + f"and the safe_body matcher will miss on the next run." + ) return try: @@ -389,6 +563,7 @@ def _before_record_request(request): headers = getattr(request, "headers", None) if headers is None: return request + _materialize_iterable_body(request) if not any(_iter_header_values(headers, KEY_FINGERPRINT_HEADER)): fingerprint = _compute_key_fingerprint(request) try: @@ -400,6 +575,56 @@ def _before_record_request(request): return request +def _materialize_iterable_body(request) -> None: + body = getattr(request, "body", None) + if body is None or isinstance(body, (bytes, bytearray, str)): + return + if not hasattr(body, "__next__"): + return + try: + chunks = list(body) + except TypeError: + return + + out = _coalesce_chunks_to_bytes(chunks) + if out is None: + method = getattr(request, "method", "?") + uri = getattr(request, "uri", getattr(request, "url", "?")) + first_type = type(chunks[0]).__name__ if chunks else "empty" + vcr_diag_write_line( + f"[vcr-materialize] FALLBACK: {method} {uri} chunk type " + f"{first_type!r} not coerced to bytes; storing b''" + ) + out = b"" + + try: + request.body = out + except (AttributeError, TypeError): + pass + + for attr in ("_was_iter", "_was_file"): + try: + setattr(request, attr, False) + except (AttributeError, TypeError): + pass + + +def _coalesce_chunks_to_bytes(chunks): + if not chunks: + return b"" + first = chunks[0] + try: + if isinstance(first, int): + return bytes(chunks) + if isinstance(first, (bytes, bytearray)): + return b"".join(c if isinstance(c, bytes) else bytes(c) for c in chunks) + if isinstance(first, str): + return "".join(chunks).encode("utf-8") + except (TypeError, ValueError): + return None + return None + + def _key_fingerprint_matcher(r1, r2) -> None: def _fp(req): for value in _iter_header_values( @@ -410,7 +635,17 @@ def _key_fingerprint_matcher(r1, r2) -> None: return value if isinstance(value, str) else str(value) return "no-key" - if _fp(r1) != _fp(r2): + fp1, fp2 = _fp(r1), _fp(r2) + if fp1 != fp2: + method_a = getattr(r1, "method", "?") + method_b = getattr(r2, "method", "?") + url_a = getattr(r1, "uri", getattr(r1, "url", "?")) + url_b = getattr(r2, "uri", getattr(r2, "url", "?")) + vcr_diag_write_line( + "[vcr-key-fingerprint-matcher] API key fingerprints differ\n" + f" request[a]: {method_a} {url_a} fingerprint={fp1!r}\n" + f" request[b]: {method_b} {url_b} fingerprint={fp2!r}" + ) raise AssertionError("API key fingerprints differ") diff --git a/tests/_vcr_redis_persister.py b/tests/_vcr_redis_persister.py index 7fdb7267a38..373cb66696a 100644 --- a/tests/_vcr_redis_persister.py +++ b/tests/_vcr_redis_persister.py @@ -159,9 +159,20 @@ def make_redis_persister( raise CassetteNotFoundError() from exc if data is None: raise CassetteNotFoundError() - if isinstance(data, bytes): - data = data.decode("utf-8") - return deserialize(data, serializer) + try: + if isinstance(data, bytes): + data = data.decode("utf-8") + return deserialize(data, serializer) + except Exception as exc: + _record_cache_failure("load", exc) + msg = ( + f"VCR redis load failed for {cassette_path}; cached " + f"payload is corrupt, treating as cache miss: " + f"{type(exc).__name__}: {exc}" + ) + _log.warning(msg) + warnings.warn(msg, VCRCassetteCacheWarning, stacklevel=2) + raise CassetteNotFoundError() from exc @staticmethod def save_cassette(cassette_path, cassette_dict, serializer): diff --git a/tests/audio_tests/conftest.py b/tests/audio_tests/conftest.py index ff47853d494..c4ff576e5bd 100644 --- a/tests/audio_tests/conftest.py +++ b/tests/audio_tests/conftest.py @@ -5,14 +5,17 @@ import pytest sys.path.insert(0, os.path.abspath("../..")) -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -44,6 +47,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -57,3 +61,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/audio_tests/test_whisper.py b/tests/audio_tests/test_whisper.py index cdf079f8cb4..243d27614b1 100644 --- a/tests/audio_tests/test_whisper.py +++ b/tests/audio_tests/test_whisper.py @@ -23,12 +23,21 @@ pwd = os.path.dirname(os.path.realpath(__file__)) print(pwd) file_path = os.path.join(pwd, "gettysburg.wav") - -audio_file = open(file_path, "rb") - - file2_path = os.path.join(pwd, "eagle.wav") -audio_file2 = open(file2_path, "rb") + +with open(file_path, "rb") as _f: + _GETTYSBURG_BYTES = _f.read() +with open(file2_path, "rb") as _f: + _EAGLE_BYTES = _f.read() + + +def _audio_file(): + return ("gettysburg.wav", _GETTYSBURG_BYTES, "audio/wav") + + +def _audio_file2(): + return ("eagle.wav", _EAGLE_BYTES, "audio/wav") + load_dotenv() @@ -44,7 +53,7 @@ async def _run_transcription( ): transcript = await litellm.atranscription( model=model, - file=audio_file, + file=_audio_file(), api_key=api_key, api_base=api_base, response_format=response_format, @@ -101,7 +110,7 @@ async def test_transcription_caching(): response_1 = await litellm.atranscription( model="whisper-1", - file=audio_file, + file=_audio_file(), ) await asyncio.sleep(5) @@ -110,7 +119,7 @@ async def test_transcription_caching(): response_2 = await litellm.atranscription( model="whisper-1", - file=audio_file, + file=_audio_file(), ) print("response_1", response_1) @@ -122,7 +131,7 @@ async def test_transcription_caching(): response_3 = await litellm.atranscription( model="whisper-1", - file=audio_file2, + file=_audio_file2(), ) print("response_3", response_3) print("response3 hidden params", response_3._hidden_params) @@ -146,7 +155,7 @@ async def test_whisper_log_pre_call(): with patch.object(custom_logger, "log_pre_api_call") as mock_log_pre_call: await litellm.atranscription( model="whisper-1", - file=audio_file, + file=_audio_file(), ) mock_log_pre_call.assert_called_once() @@ -165,7 +174,7 @@ async def test_whisper_log_pre_call(): with patch.object(custom_logger, "log_pre_api_call") as mock_log_pre_call: await litellm.atranscription( model="whisper-1", - file=audio_file, + file=_audio_file(), ) mock_log_pre_call.assert_called_once() @@ -177,7 +186,7 @@ async def test_gpt_4o_transcribe(): from unittest.mock import patch, MagicMock await litellm.atranscription( - model="openai/gpt-4o-transcribe", file=audio_file, response_format="json" + model="openai/gpt-4o-transcribe", file=_audio_file(), response_format="json" ) @@ -187,7 +196,9 @@ async def test_gpt_4o_transcribe_model_mapping(): # Test GPT-4o mini transcribe response = await litellm.atranscription( - model="openai/gpt-4o-mini-transcribe", file=audio_file, response_format="json" + model="openai/gpt-4o-mini-transcribe", + file=_audio_file(), + response_format="json", ) # Check that the response contains the correct model in hidden params @@ -198,7 +209,7 @@ async def test_gpt_4o_transcribe_model_mapping(): # Test GPT-4o transcribe response2 = await litellm.atranscription( - model="openai/gpt-4o-transcribe", file=audio_file, response_format="json" + model="openai/gpt-4o-transcribe", file=_audio_file(), response_format="json" ) # Check that the response contains the correct model in hidden params @@ -209,7 +220,7 @@ async def test_gpt_4o_transcribe_model_mapping(): # Test traditional whisper-1 still works response3 = await litellm.atranscription( - model="openai/whisper-1", file=audio_file, response_format="json" + model="openai/whisper-1", file=_audio_file(), response_format="json" ) # Check that the response contains the correct model in hidden params @@ -262,7 +273,7 @@ async def test_azure_transcribe_model_mapping(): # Make the transcription call response = await litellm.atranscription( model="azure/whisper-1", - file=audio_file, + file=_audio_file(), response_format="json", api_key="test-api-key", api_base="https://my-endpoint-europe-berri-992.openai.azure.com/", diff --git a/tests/guardrails_tests/conftest.py b/tests/guardrails_tests/conftest.py index eb563699b2b..f2f65645c3d 100644 --- a/tests/guardrails_tests/conftest.py +++ b/tests/guardrails_tests/conftest.py @@ -16,14 +16,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -55,6 +58,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -160,3 +164,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/image_gen_tests/conftest.py b/tests/image_gen_tests/conftest.py index 93dec98e708..9f808c11161 100644 --- a/tests/image_gen_tests/conftest.py +++ b/tests/image_gen_tests/conftest.py @@ -9,14 +9,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm # noqa: E402,F401 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -58,6 +61,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -71,3 +75,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 656b8a69117..ca8ec3bbe32 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -103,12 +103,6 @@ class BaseLLMImageEditTest(ABC): pwd = os.path.dirname(os.path.realpath(__file__)) -# Image fixtures must be regenerated per access — module-level -# ``open(...)`` handles get consumed after a single multipart upload, leaving -# subsequent tests in the same process to send empty bodies. That non-determinism -# (a) blows the recorded cassette past ``MAX_EPISODES_PER_CASSETTE`` so the -# persister refuses to save (see ``tests/_vcr_redis_persister.py``), and -# (b) re-bills the live image edit endpoint on every CI run. def _read_image_bytes(filename: str) -> bytes: with open(os.path.join(pwd, filename), "rb") as f: return f.read() @@ -119,32 +113,20 @@ _LITELLM_SITE_BYTES = _read_image_bytes("litellm_site.png") def _make_test_images() -> list: - """Return a fresh pair of image streams seeded with the fixture bytes. + return [_ISHAAN_GITHUB_BYTES, _LITELLM_SITE_BYTES] - Use this everywhere you'd previously have used the module-level - ``TEST_IMAGES``. Each call returns brand new ``BytesIO`` objects whose - file pointers start at 0, so multipart uploads encode the full image - bytes on every test invocation. Parametrized and ``flaky``-retried - test methods call ``get_base_image_edit_call_args`` once per - invocation, so a fresh stream per call is sufficient — the factory - must not auto-rewind on EOF or the SDK's multipart writer will read - the same bytes forever (worker OOM). - """ + +def _make_single_test_image() -> bytes: + return _ISHAAN_GITHUB_BYTES + + +def get_test_images_as_bytesio(): return [ BytesIO(_ISHAAN_GITHUB_BYTES), BytesIO(_LITELLM_SITE_BYTES), ] -def _make_single_test_image() -> BytesIO: - return BytesIO(_ISHAAN_GITHUB_BYTES) - - -def get_test_images_as_bytesio(): - """Helper function to get test images as BytesIO objects""" - return _make_test_images() - - class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest): """ Concrete implementation of BaseLLMImageEditTest for OpenAI image edits. @@ -710,10 +692,9 @@ async def test_multiple_image_edit_with_different_formats(): try: prompt = "Create a cohesive artistic style across all images" - # Test with mixed BytesIO and file objects mixed_images = [ - _make_single_test_image(), # File object - get_test_images_as_bytesio()[1], # BytesIO object + _make_single_test_image(), + get_test_images_as_bytesio()[1], ] result = await aimage_edit( diff --git a/tests/litellm_utils_tests/conftest.py b/tests/litellm_utils_tests/conftest.py index 08745c99c07..418ee76a399 100644 --- a/tests/litellm_utils_tests/conftest.py +++ b/tests/litellm_utils_tests/conftest.py @@ -12,14 +12,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm # noqa: E402,F401 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -86,6 +89,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -116,3 +120,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/llm_responses_api_testing/conftest.py b/tests/llm_responses_api_testing/conftest.py index 2a08db57149..1928b540dad 100644 --- a/tests/llm_responses_api_testing/conftest.py +++ b/tests/llm_responses_api_testing/conftest.py @@ -13,14 +13,17 @@ sys.path.insert( import litellm # noqa: E402 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -52,6 +55,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -116,3 +120,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/llm_translation/conftest.py b/tests/llm_translation/conftest.py index 5fcd31aa32d..d346dae4308 100644 --- a/tests/llm_translation/conftest.py +++ b/tests/llm_translation/conftest.py @@ -18,14 +18,17 @@ sys.path.insert( import litellm # noqa: E402 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -73,6 +76,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -82,6 +86,7 @@ def pytest_runtest_logreport(report): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) # --------------------------------------------------------------------------- diff --git a/tests/llm_translation/realtime/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py index c5f77de6beb..fc9f938b4cd 100644 --- a/tests/llm_translation/realtime/test_openai_realtime.py +++ b/tests/llm_translation/realtime/test_openai_realtime.py @@ -101,7 +101,9 @@ async def test_openai_realtime_direct_call_no_intent(): try: await litellm._arealtime( - model="openai/gpt-4o-realtime-preview", + # OpenAI shut down the gpt-4o-realtime-preview family (incl. the + # undated alias) on 2026-05-07; gpt-realtime is the GA successor. + model="openai/gpt-realtime", websocket=websocket_client, api_key=os.environ.get("OPENAI_API_KEY"), timeout=60, @@ -249,14 +251,16 @@ async def test_openai_realtime_direct_call_with_intent(): websocket_client = RealTimeWebSocketClient() caught_exception = None + # OpenAI shut down the gpt-4o-realtime-preview family (incl. the undated + # alias) on 2026-05-07; gpt-realtime is the GA successor. query_params: RealtimeQueryParams = { - "model": "openai/gpt-4o-realtime-preview", + "model": "openai/gpt-realtime", "intent": "chat", } try: await litellm._arealtime( - model="openai/gpt-4o-realtime-preview", + model="openai/gpt-realtime", websocket=websocket_client, api_key=os.environ.get("OPENAI_API_KEY"), query_params=query_params, diff --git a/tests/llm_translation/realtime/test_openai_realtime_simple.py b/tests/llm_translation/realtime/test_openai_realtime_simple.py index 5522d843e42..073c1ce11af 100644 --- a/tests/llm_translation/realtime/test_openai_realtime_simple.py +++ b/tests/llm_translation/realtime/test_openai_realtime_simple.py @@ -21,7 +21,10 @@ class TestOpenAIRealtime(BaseRealtimeTest): """ def get_model(self) -> str: - return "gpt-4o-realtime-preview" + # OpenAI shut down the entire gpt-4o-realtime-preview family + # (including the undated alias) on 2026-05-07. gpt-realtime is the + # current GA realtime model. + return "gpt-realtime" def get_api_key_env_var(self) -> str: return "OPENAI_API_KEY" diff --git a/tests/llm_translation/realtime/test_realtime_guardrails_openai.py b/tests/llm_translation/realtime/test_realtime_guardrails_openai.py index 170440f6b9f..50cedba2ac0 100644 --- a/tests/llm_translation/realtime/test_realtime_guardrails_openai.py +++ b/tests/llm_translation/realtime/test_realtime_guardrails_openai.py @@ -26,9 +26,7 @@ from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming from litellm.types.guardrails import GuardrailEventHooks OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY") -OPENAI_REALTIME_URL = ( - "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview-2024-12-17" -) +OPENAI_REALTIME_URL = "wss://api.openai.com/v1/realtime?model=gpt-realtime" pytestmark = pytest.mark.skipif( not OPENAI_API_KEY, @@ -192,10 +190,35 @@ async def test_text_message_blocked_by_guardrail_no_ai_response(): len(transcript_deltas) >= 1 ), f"Expected guardrail message in transcript delta, got: {event_types}" - # 3. No *real* AI response should have been generated. - # The guardrail may produce its own response (e.g. "Content blocked: ...") - # via response.cancel + conversation.item.create + response.create. - # We allow the guardrail's own block message but NOT original AI content. + # 3. No *real* AI response to the blocked content should have been + # generated. The original user message is blocked BEFORE it is + # forwarded to OpenAI, so the only thing the model ever sees is the + # guardrail's "say exactly: " prompt + # (see realtime_streaming.py). Two safe outcomes are possible: + # - the model voices the block message verbatim (older realtime + # snapshots did this -> text contains "blocked"), or + # - the model declines to repeat it (gpt-realtime tends to refuse + # verbatim-repeat instructions, e.g. "I'm sorry, but I can't + # repeat that message."). + # Both mean the blocked prompt itself was never answered, so we + # accept either. The hard invariant is that the blocked phrase must + # never leak into AI output, and the model must not have produced a + # normal answer to the user (which would have neither a block nor a + # refusal marker). + safe_markers = ( + "block", + "guardrail", + "content filter", + "policy", + "can't repeat", + "cannot repeat", + "won't repeat", + "can't assist", + "can't help", + "unable to", + "i'm sorry", + "i am sorry", + ) done_events = [e for e in client_events if e.get("type") == "response.done"] for done in done_events: output = done.get("response", {}).get("output", []) @@ -205,11 +228,19 @@ async def test_text_message_blocked_by_guardrail_no_ai_response(): for c in item.get("content", []) ] real_ai_text = " ".join(ai_texts).strip() - # Allow guardrail-generated block messages (contain "Content blocked" or "blocked") if real_ai_text: assert ( - "blocked" in real_ai_text.lower() - or "guardrail" in real_ai_text.lower() + BLOCKED_PHRASE not in real_ai_text + ), f"Blocked phrase leaked into AI response: {real_ai_text!r}" + normalized_ai_text = ( + real_ai_text.lower() + .replace("\u2019", "'") + .replace("\u2018", "'") + .replace("\u201c", '"') + .replace("\u201d", '"') + ) + assert any( + marker in normalized_ai_text for marker in safe_markers ), f"AI responded with non-guardrail content even though message was blocked: {real_ai_text!r}" finally: diff --git a/tests/llm_translation/test_nvidia_nim.py b/tests/llm_translation/test_nvidia_nim.py index 469516407c8..80e764147bb 100644 --- a/tests/llm_translation/test_nvidia_nim.py +++ b/tests/llm_translation/test_nvidia_nim.py @@ -262,3 +262,44 @@ class TestNvidiaNim(BaseLLMRerankTest): def get_expected_cost(self) -> float: """Nvidia NIM rerank models are free (cost = 0.0)""" return 0.0 + + @pytest.mark.asyncio() + @pytest.mark.parametrize("sync_mode", [True, False]) + async def test_basic_rerank(self, sync_mode, monkeypatch): + """ + Override the base live rerank test with a mocked HTTP layer. + + NVIDIA reached end-of-life for the hosted + nvidia/llama-3.2-nv-rerankqa-1b-v2 rerank API on 2026-05-18 and + published no replacement model, so a live call now returns HTTP 410 + ("Gone"). NVIDIA's hosted catalog rotates on a schedule, so pointing + at another live model would only defer the same failure. Mock the + transport instead (same pattern as + test_nvidia_nim_rerank_ranking_endpoint above) so the request/response + transformation and cost calculation stay covered offline. + """ + monkeypatch.setenv("NVIDIA_NIM_API_KEY", "fake-api-key") + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {} + mock_response.text = "" + mock_response.json.return_value = { + "rankings": [ + {"index": 0, "logit": 0.95}, + {"index": 1, "logit": 0.75}, + ], + "usage": {"total_tokens": 7}, + } + + with ( + patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=mock_response, + ), + patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=mock_response, + ), + ): + await super().test_basic_rerank(sync_mode=sync_mode) diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index 0ff7dff668a..6a746041f15 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -22,14 +22,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -84,6 +87,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -93,6 +97,7 @@ def pytest_runtest_logreport(report): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) # --------------------------------------------------------------------------- diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index 2b6712cbaa3..0f4539162a2 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -25,6 +25,7 @@ from unittest.mock import AsyncMock, patch, MagicMock from litellm.caching.caching_handler import ( LLMCachingHandler, CachingHandlerResponse, + _is_chat_completion_cached_dict, _should_defer_streaming_cache_hit_callbacks, ) from litellm.caching.caching import LiteLLMCacheType @@ -40,6 +41,7 @@ from litellm.types.utils import ( from litellm.types.llms.openai import ResponsesAPIResponse from datetime import timedelta, datetime from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm._logging import verbose_logger import logging @@ -1072,6 +1074,70 @@ def test_convert_cached_streaming_responses_result_to_iterator(): ) +def test_is_chat_completion_cached_dict(): + assert _is_chat_completion_cached_dict( + {"id": "chatcmpl-abc", "object": "chat.completion", "choices": []} + ) + assert _is_chat_completion_cached_dict( + {"id": "other", "object": "chat.completion.chunk", "choices": []} + ) + assert not _is_chat_completion_cached_dict( + {"id": "resp_abc", "object": "response", "output": []} + ) + + +def test_convert_cached_aresponses_bridge_chat_completion_stream(): + """ + openai/responses chat-completions bridge caches ModelResponse JSON on aresponses + cache keys; replay must not call ResponsesAPIResponse(**chatcmpl_dict). + """ + caching_handler = LLMCachingHandler( + original_function=aresponses, request_kwargs={}, start_time=datetime.now() + ) + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-5.4", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + cached_result = { + "id": "chatcmpl-bridge-cache-test", + "object": "chat.completion", + "created": int(time.time()), + "model": "gpt-5.4", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hi!"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 7, + "completion_tokens": 11, + "total_tokens": 18, + }, + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.aresponses.value, + kwargs={ + "model": "gpt-5.4", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + logging_obj=logging_obj, + model="gpt-5.4", + args=(), + ) + + assert isinstance(result, CustomStreamWrapper) + + def test_convert_cached_streaming_reasoning_result_to_iterator(): caching_handler = LLMCachingHandler( original_function=responses, request_kwargs={}, start_time=datetime.now() diff --git a/tests/logging_callback_tests/conftest.py b/tests/logging_callback_tests/conftest.py index cdb9200bc83..6dde85f2ca7 100644 --- a/tests/logging_callback_tests/conftest.py +++ b/tests/logging_callback_tests/conftest.py @@ -19,14 +19,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -79,6 +82,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -229,3 +233,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/ocr_tests/conftest.py b/tests/ocr_tests/conftest.py index 66970b8579f..94790bd7aa3 100644 --- a/tests/ocr_tests/conftest.py +++ b/tests/ocr_tests/conftest.py @@ -12,14 +12,17 @@ import pytest sys.path.insert(0, os.path.abspath("../..")) -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -51,6 +54,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -64,3 +68,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/otel_tests/test_prometheus.py b/tests/otel_tests/test_prometheus.py index c9490af07cb..90c71037609 100644 --- a/tests/otel_tests/test_prometheus.py +++ b/tests/otel_tests/test_prometheus.py @@ -610,22 +610,18 @@ def extract_user_budget_metrics(metrics_text: str, user_id: str) -> Dict[str, fl # Escape user_id for regex pattern matching escaped_user_id = re.escape(user_id) - # Get remaining budget - remaining_pattern = ( - f'litellm_remaining_user_budget_metric{{user="{escaped_user_id}"}} ([0-9.]+)' - ) + # Get remaining budget (user_email and user_alias may also be present as labels) + remaining_pattern = rf'litellm_remaining_user_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' remaining_match = re.search(remaining_pattern, metrics_text) metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None # Get total budget - total_pattern = ( - f'litellm_user_max_budget_metric{{user="{escaped_user_id}"}} ([0-9.]+)' - ) + total_pattern = rf'litellm_user_max_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' total_match = re.search(total_pattern, metrics_text) metrics["total"] = float(total_match.group(1)) if total_match else None # Get remaining hours - hours_pattern = f'litellm_user_budget_remaining_hours_metric{{user="{escaped_user_id}"}} ([0-9.]+)' + hours_pattern = rf'litellm_user_budget_remaining_hours_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' hours_match = re.search(hours_pattern, metrics_text) metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None diff --git a/tests/pass_through_unit_tests/conftest.py b/tests/pass_through_unit_tests/conftest.py index 42a95343eb7..390e14b7f11 100644 --- a/tests/pass_through_unit_tests/conftest.py +++ b/tests/pass_through_unit_tests/conftest.py @@ -5,14 +5,17 @@ import pytest sys.path.insert(0, os.path.abspath("../..")) -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -56,6 +59,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -71,3 +75,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/router_unit_tests/conftest.py b/tests/router_unit_tests/conftest.py index fe976515c92..6a8f3e589f4 100644 --- a/tests/router_unit_tests/conftest.py +++ b/tests/router_unit_tests/conftest.py @@ -12,14 +12,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm # noqa: E402,F401 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -97,6 +100,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -123,3 +127,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/search_tests/conftest.py b/tests/search_tests/conftest.py index e06d3e95eee..78ba19a7724 100644 --- a/tests/search_tests/conftest.py +++ b/tests/search_tests/conftest.py @@ -13,14 +13,17 @@ import pytest sys.path.insert(0, os.path.abspath("../..")) -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -52,6 +55,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -65,3 +69,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/test_callbacks_on_proxy.py b/tests/test_callbacks_on_proxy.py index 0b55d820532..17c0db9260f 100644 --- a/tests/test_callbacks_on_proxy.py +++ b/tests/test_callbacks_on_proxy.py @@ -9,12 +9,155 @@ import pytest import asyncio import aiohttp import os +import re import dotenv +from collections import Counter from dotenv import load_dotenv import pytest load_dotenv() +# A *leak* is sustained, monotonic growth of one callback TYPE across the whole +# sampling window. A one-time bump that then plateaus is benign pollution from +# other tests sharing this proxy (this suite runs `pytest -n 4` against a single +# proxy container, so other workers legitimately add team/key-scoped callbacks +# while this test sleeps). We therefore sample N times and only flag a type +# whose normalized count never decreases, grows in >=2 distinct intervals, and +# nets >= LEAK_MIN_NET_GROWTH overall. +NUM_SAMPLES = 4 +SAMPLE_INTERVAL_SECONDS = 20 +LEAK_MIN_NET_GROWTH = 5 +LEAK_MIN_GROWING_INTERVALS = 2 +# A routing-strategy switch / alerting config is a *known, bounded, one-time* +# registration (CCI diagnostic 2026-05-16: total 85->95 on the first interval +# after switching to latency-based-routing, then flat at 95 for 2.5 min under +# load). We absorb that step by settling before the baseline sample, so only +# growth *after* the deliberate perturbation can count as a leak. +SETTLE_SECONDS = 30 + +# Strip instance-identity noise so N leaked instances of one class collapse to +# one rising counter instead of N opaque, unrelated-looking strings. +_ADDR_RE = re.compile(r" at 0x[0-9a-fA-F]+") +_OBJ_RE = re.compile(r"<([\w.]+) object") + + +def _normalize_callback(cb_str: str) -> str: + """Reduce a callback's str() to a stable type key (drops 0x… addresses).""" + s = _ADDR_RE.sub("", cb_str) + m = _OBJ_RE.search(s) + if m: + return m.group(1).split(".")[-1] + # bound methods: ">" -> "Cls.m" + bm = re.search(r"bound method ([\w.]+)", s) + if bm: + return bm.group(1) + return s.strip() + + +def _summarize(all_litellm_callbacks) -> Counter: + return Counter(_normalize_callback(str(c)) for c in all_litellm_callbacks) + + +def _detect_leaks(samples): + """ + samples: list[Counter] taken in time order. + + Returns {callback_type: [counts across samples]} for types that grew + monotonically (never decreased), in >=LEAK_MIN_GROWING_INTERVALS intervals, + and netted >=LEAK_MIN_NET_GROWTH overall — i.e. a real leak, not a one-shot + step from a parallel test. + """ + leaks = {} + all_types = set().union(*[set(s) for s in samples]) if samples else set() + for t in all_types: + series = [s.get(t, 0) for s in samples] + deltas = [b - a for a, b in zip(series, series[1:])] + net = series[-1] - series[0] + non_decreasing = all(d >= 0 for d in deltas) + growing_intervals = sum(1 for d in deltas if d > 0) + if ( + non_decreasing + and net >= LEAK_MIN_NET_GROWTH + and growing_intervals >= LEAK_MIN_GROWING_INTERVALS + ): + leaks[t] = series + return leaks + + +def _terminal_suspects(samples): + """ + Types whose net growth clears the threshold monotonically but is confined + to the *final* interval — `growing_intervals == 1` with that one growing + interval being the last. `_detect_leaks`' `>= 2` guard silently passes + these, so a real leak that accumulates entirely in the last sampled window + is indistinguishable from a one-time terminal step *without one more + sample*. Returns the set of such types so the caller can re-confirm. + """ + suspects = set() + all_types = set().union(*[set(s) for s in samples]) if samples else set() + for t in all_types: + series = [s.get(t, 0) for s in samples] + deltas = [b - a for a, b in zip(series, series[1:])] + if not deltas: + continue + net = series[-1] - series[0] + non_decreasing = all(d >= 0 for d in deltas) + growing = [i for i, d in enumerate(deltas) if d > 0] + if ( + non_decreasing + and net >= LEAK_MIN_NET_GROWTH + and growing == [len(deltas) - 1] + ): + suspects.add(t) + return suspects + + +async def _detect_leaks_confirmed(session, samples): + """ + `_detect_leaks`, plus a single confirmation sample when growth is confined + to the final interval (see `_terminal_suspects`). A genuine ongoing leak + keeps climbing -> now grows in >= 2 intervals -> flagged; a one-time + terminal registration plateaus -> still 1 growing interval -> ignored. + Returns `(leaks, samples)` (samples may have one extra entry appended). + """ + leaks = _detect_leaks(samples) + if not leaks and _terminal_suspects(samples): + await asyncio.sleep(SAMPLE_INTERVAL_SECONDS) + _, _, all_cb = await get_active_callbacks(session=session) + samples = samples + [_summarize(all_cb)] + leaks = _detect_leaks(samples) + return leaks, samples + + +def _format_report(samples, leaks) -> str: + lines = ["Callback count per type across samples (time order):"] + all_types = sorted(set().union(*[set(s) for s in samples])) + for t in all_types: + series = [s.get(t, 0) for s in samples] + marker = " <-- LEAK" if t in leaks else "" + lines.append(f" {t}: {series}{marker}") + totals = [sum(s.values()) for s in samples] + lines.append(f"TOTAL callbacks per sample: {totals}") + if leaks: + lines.append( + "Leaking callback types (sustained monotonic growth): " + + ", ".join(sorted(leaks)) + ) + return "\n".join(lines) + + +async def _sample_callbacks(session, num_samples, interval): + """Take `num_samples` callback snapshots `interval`s apart.""" + samples = [] + alerts = [] + for i in range(num_samples): + if i > 0: + await asyncio.sleep(interval) + num_cb, num_alert, all_cb = await get_active_callbacks(session=session) + samples.append(_summarize(all_cb)) + alerts.append(num_alert) + return samples, alerts + async def config_update(session, routing_strategy=None): url = "http://0.0.0.0:4000/config/update" @@ -97,105 +240,65 @@ async def get_current_routing_strategy(session): @pytest.mark.asyncio @pytest.mark.order1 +@pytest.mark.flaky(reruns=2, reruns_delay=5) async def test_check_num_callbacks(): """ - Test 1: num callbacks should NOT increase over time - -> check current callbacks - -> sleep for 30 seconds - -> check current callbacks - -> sleep for 30 seconds - -> check current callbacks + PROD invariant: no callback TYPE should grow without bound over time. + + This suite runs `pytest -n 4` against one shared proxy, so the raw count is + noisy — other workers legitimately add team/key-scoped callbacks that then + plateau. We settle first, then sample several times, and only fail on + *sustained, monotonic* per-type growth (a genuine leak), naming the type. """ - from litellm._uuid import uuid - async with aiohttp.ClientSession() as session: - await asyncio.sleep(30) - num_callbacks_1, _, all_litellm_callbacks_1 = await get_active_callbacks( - session=session - ) - assert num_callbacks_1 > 0 - await asyncio.sleep(30) + # Absorb proxy warmup / in-flight parallel registration before baseline. + await asyncio.sleep(SETTLE_SECONDS) - num_callbacks_2, _, all_litellm_callbacks_2 = await get_active_callbacks( - session=session + samples, _ = await _sample_callbacks( + session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS ) - print("all_litellm_callbacks_1", all_litellm_callbacks_1) + assert sum(samples[0].values()) > 0, "expected some callbacks registered" - print( - "diff in callbacks=", - set(all_litellm_callbacks_1) - set(all_litellm_callbacks_2), - ) - - assert abs(num_callbacks_1 - num_callbacks_2) <= 4 - - await asyncio.sleep(30) - - num_callbacks_3, _, all_litellm_callbacks_3 = await get_active_callbacks( - session=session - ) - - print( - "diff in callbacks = all_litellm_callbacks3 - all_litellm_callbacks2 ", - set(all_litellm_callbacks_3) - set(all_litellm_callbacks_2), - ) - - assert abs(num_callbacks_3 - num_callbacks_2) <= 4 + leaks, samples = await _detect_leaks_confirmed(session, samples) + report = _format_report(samples, leaks) + print(report) + assert not leaks, f"Callback leak detected.\n{report}" @pytest.mark.asyncio @pytest.mark.order2 +@pytest.mark.flaky(reruns=2, reruns_delay=5) async def test_check_num_callbacks_on_lowest_latency(): """ - Test 1: num callbacks should NOT increase over time - -> Update to lowest latency - -> check current callbacks - -> sleep for 30s - -> check current callbacks - -> sleep for 30s - -> check current callbacks - -> update back to original routing-strategy + Same PROD invariant as test_check_num_callbacks, but after switching the + router to latency-based-routing. That switch is a *known, bounded* one-time + registration (it adds the latency strategy handler + Slack alerting); we + settle past it before baselining so only post-switch growth counts as a + leak. Also asserts the alerting count is stable. """ - from litellm._uuid import uuid - async with aiohttp.ClientSession() as session: await asyncio.sleep(30) original_routing_strategy = await get_current_routing_strategy(session=session) await config_update(session=session, routing_strategy="latency-based-routing") - await asyncio.sleep(30) + try: + # Absorb the deliberate one-time config/update registration step. + await asyncio.sleep(SETTLE_SECONDS) - num_callbacks_1, num_alerts_1, all_litellm_callbacks_1 = ( - await get_active_callbacks(session=session) - ) + samples, alerts = await _sample_callbacks( + session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS + ) - await asyncio.sleep(30) - - num_callbacks_2, num_alerts_2, all_litellm_callbacks_2 = ( - await get_active_callbacks(session=session) - ) - - print( - "diff in callbacks all_litellm_callbacks_2 - all_litellm_callbacks_1 =", - set(all_litellm_callbacks_2) - set(all_litellm_callbacks_1), - ) - - assert abs(num_callbacks_1 - num_callbacks_2) <= 4 - - await asyncio.sleep(30) - - num_callbacks_3, num_alerts_3, all_litellm_callbacks_3 = ( - await get_active_callbacks(session=session) - ) - - print( - "diff in callbacks all_litellm_callbacks_3 - all_litellm_callbacks_2 =", - set(all_litellm_callbacks_3) - set(all_litellm_callbacks_2), - ) - - assert abs(num_callbacks_2 - num_callbacks_3) <= 4 - - assert num_alerts_1 == num_alerts_2 == num_alerts_3 - - await config_update(session=session, routing_strategy=original_routing_strategy) + leaks, samples = await _detect_leaks_confirmed(session, samples) + report = _format_report(samples, leaks) + print(report) + assert not leaks, f"Callback leak detected.\n{report}" + assert ( + len(set(alerts)) == 1 + ), f"alerting count changed across samples: {alerts}" + finally: + await config_update( + session=session, routing_strategy=original_routing_strategy + ) diff --git a/tests/test_litellm/caching/test_caching_handler.py b/tests/test_litellm/caching/test_caching_handler.py index 742a4f410d4..3eb949d7f29 100644 --- a/tests/test_litellm/caching/test_caching_handler.py +++ b/tests/test_litellm/caching/test_caching_handler.py @@ -232,3 +232,207 @@ def test_combine_usage_handles_none_details(): combined = llm_caching_handler.combine_usage(usage_a, usage_c) assert combined.prompt_tokens_details is not None assert combined.prompt_tokens_details.image_count == 1 + + +def test_is_chat_completion_cached_dict(): + from litellm.caching.caching_handler import _is_chat_completion_cached_dict + + assert _is_chat_completion_cached_dict( + {"id": "chatcmpl-abc", "object": "chat.completion", "choices": []} + ) + assert _is_chat_completion_cached_dict( + {"id": "other", "object": "chat.completion.chunk", "choices": []} + ) + assert _is_chat_completion_cached_dict( + {"id": "no-object", "choices": [{"index": 0}]} + ) + assert not _is_chat_completion_cached_dict( + {"id": "resp_abc", "object": "response", "output": []} + ) + + +def _build_logging_obj(call_type: str, stream: bool): + import uuid as _uuid + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + return LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=call_type, + model="gpt-5.4", + messages=[], + function_id=str(_uuid.uuid4()), + stream=stream, + start_time=datetime.now(), + ) + + +def test_convert_cached_aresponses_bridge_chat_completion_stream(): + """openai/responses chat-completions bridge: streaming cache hit replays as chat stream.""" + from litellm import aresponses + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + from litellm.types.utils import CallTypes + + caching_handler = LLMCachingHandler( + original_function=aresponses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "chatcmpl-bridge-cache-test", + "object": "chat.completion", + "created": int(time.time()), + "model": "gpt-5.4", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hi!"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}, + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.aresponses.value, + kwargs={ + "model": "gpt-5.4", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + logging_obj=_build_logging_obj(CallTypes.aresponses.value, stream=True), + model="gpt-5.4", + args=(), + ) + + assert isinstance(result, CustomStreamWrapper) + + +def test_convert_cached_responses_bridge_chat_completion_nonstream(): + """openai/responses chat-completions bridge: non-streaming cache hit replays as ModelResponse.""" + from litellm import responses + from litellm.types.utils import CallTypes, ModelResponse + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "chatcmpl-bridge-nonstream", + "object": "chat.completion", + "created": int(time.time()), + "model": "gpt-5.4", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hi!"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}, + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={ + "model": "gpt-5.4", + "stream": False, + "messages": [{"role": "user", "content": "hi"}], + }, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), + model="gpt-5.4", + args=(), + ) + + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "Hi!" + + +def test_convert_cached_responses_legacy_nonstream_path(): + """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) falls through legacy path.""" + from litellm import responses + from litellm.types.llms.openai import ResponsesAPIResponse + from litellm.types.utils import CallTypes + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "resp_legacy_nonstream", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_legacy", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "legacy response", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "hi", "stream": False}, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), + model="gpt-4o", + args=(), + ) + + assert isinstance(result, ResponsesAPIResponse) + assert result.id == "resp_legacy_nonstream" + + +def test_convert_cached_responses_legacy_stream_path(): + """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) on stream falls through legacy path.""" + from litellm import responses + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + from litellm.types.utils import CallTypes + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "resp_legacy_stream", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_legacy_stream", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "legacy stream", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "hi", "stream": True}, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=True), + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py new file mode 100644 index 00000000000..734033ed6be --- /dev/null +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py @@ -0,0 +1,150 @@ +import os +import sys +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.completion_extras.litellm_responses_transformation.handler import ( + ResponsesToCompletionBridgeHandler, +) +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.types.utils import ModelResponse + + +def test_is_preformatted_cached_chat_stream_true(): + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "cached_response" + assert ( + ResponsesToCompletionBridgeHandler._is_preformatted_cached_chat_stream(stream) + is True + ) + + +def test_is_preformatted_cached_chat_stream_false_wrong_provider(): + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "openai" + assert ( + ResponsesToCompletionBridgeHandler._is_preformatted_cached_chat_stream(stream) + is False + ) + + +def test_is_preformatted_cached_chat_stream_false_wrong_type(): + assert ( + ResponsesToCompletionBridgeHandler._is_preformatted_cached_chat_stream( + {"object": "chat.completion.chunk"} + ) + is False + ) + + +def _bridge_kwargs(stream: bool): + logging_obj = LiteLLMLogging( + litellm_call_id="test-call", + call_type="completion", + model="gpt-5.4", + messages=[{"role": "user", "content": "hi"}], + function_id="fn-id", + stream=stream, + start_time=datetime.now(), + ) + return { + "model": "gpt-5.4", + "custom_llm_provider": "openai", + "messages": [{"role": "user", "content": "hi"}], + "optional_params": {"stream": stream}, + "litellm_params": {}, + "headers": {}, + "model_response": ModelResponse(), + "logging_obj": logging_obj, + } + + +def test_completion_returns_cached_model_response_directly(): + """Non-streaming bridge cache hit: responses() returns a ModelResponse -> bridge returns it as-is.""" + cached = ModelResponse(id="chatcmpl-cached-nonstream", model="gpt-5.4") + bridge = ResponsesToCompletionBridgeHandler() + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ), + patch("litellm.responses", return_value=cached), + ): + result = bridge.completion(**_bridge_kwargs(stream=False)) + + assert result is cached + + +@pytest.mark.asyncio +async def test_acompletion_returns_cached_model_response_directly(): + cached = ModelResponse(id="chatcmpl-cached-nonstream-async", model="gpt-5.4") + bridge = ResponsesToCompletionBridgeHandler() + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ), + patch("litellm.aresponses", new=AsyncMock(return_value=cached)), + ): + result = await bridge.acompletion(**_bridge_kwargs(stream=False)) + + assert result is cached + + +def test_completion_skips_rewrapping_preformatted_cached_chat_stream(): + """Streaming bridge cache hit returning CustomStreamWrapper(cached_response) -> bridge skips re-wrapping.""" + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "cached_response" + bridge = ResponsesToCompletionBridgeHandler() + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ), + patch("litellm.responses", return_value=stream), + patch.object( + bridge, + "_apply_post_stream_processing", + side_effect=lambda s, *a, **kw: s, + ) as post, + ): + result = bridge.completion(**_bridge_kwargs(stream=True)) + + post.assert_called_once() + assert result is stream + + +@pytest.mark.asyncio +async def test_acompletion_skips_rewrapping_preformatted_cached_chat_stream(): + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "cached_response" + bridge = ResponsesToCompletionBridgeHandler() + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ), + patch("litellm.aresponses", new=AsyncMock(return_value=stream)), + patch.object( + bridge, + "_apply_post_stream_processing", + side_effect=lambda s, *a, **kw: s, + ) as post, + ): + result = await bridge.acompletion(**_bridge_kwargs(stream=True)) + + post.assert_called_once() + assert result is stream diff --git a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py index 009f432fca1..05bdc40112c 100644 --- a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py @@ -230,3 +230,30 @@ def test_transform_request_drops_user_metadata_with_additional_drop_params(): assert "metadata" not in result assert result["litellm_metadata"]["internal_key"] == "secret" + + +def test_translate_responses_chunk_passthrough_chat_completion_chunk(): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + chat_chunk = { + "id": "chatcmpl-cache-passthrough", + "object": "chat.completion.chunk", + "created": 1779104834, + "model": "gpt-5.4", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "Hi! How can I help?"}, + "finish_reason": None, + } + ], + } + + result = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( + chat_chunk + ) + + assert result.choices[0].delta.content == "Hi! How can I help?" + assert result.choices[0].finish_reason is None diff --git a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py index 19ae819c85a..12f30ab6024 100644 --- a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py +++ b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py @@ -460,6 +460,104 @@ async def test_assemble_user_object_does_not_override_metadata_max_budget( ), "max_budget from metadata must not be replaced by the DB value" +async def test_assemble_user_object_populates_user_email_and_alias_from_db( + prometheus_logger, +): + db_user = MagicMock() + db_user.max_budget = None + db_user.budget_reset_at = None + db_user.user_email = "alice@example.com" + db_user.user_alias = "Alice" + + with patch("litellm.proxy.auth.auth_checks.get_user_object") as mock_get_user: + mock_get_user.return_value = db_user + user_object = await prometheus_logger._assemble_user_object( + user_id="user-abc-123", + spend=10.0, + max_budget=None, + response_cost=0.5, + ) + + assert user_object.user_email == "alice@example.com" + assert user_object.user_alias == "Alice" + + +def test_set_user_budget_metrics_default_no_email_alias_labels( + prometheus_logger, +): + """By default (flag off), only user label is emitted.""" + import litellm + from litellm.proxy._types import LiteLLM_UserTable + + litellm.prometheus_user_budget_label_include_email_alias = False + + user = LiteLLM_UserTable( + user_id="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + spend=25.0, + max_budget=100.0, + budget_reset_at=datetime(2026, 3, 1, tzinfo=timezone.utc), + ) + + prometheus_logger.litellm_remaining_user_budget_metric = MagicMock() + prometheus_logger.litellm_user_max_budget_metric = MagicMock() + prometheus_logger.litellm_user_budget_remaining_hours_metric = MagicMock() + + prometheus_logger._set_user_budget_metrics(user) + + prometheus_logger.litellm_remaining_user_budget_metric.labels.assert_called_once_with( + user="user-abc-123", + ) + + +def test_set_user_budget_metrics_includes_user_email_and_alias_labels_when_opted_in( + prometheus_logger, +): + """When prometheus_user_budget_label_include_email_alias=True, email+alias labels appear.""" + import litellm + from litellm.proxy._types import LiteLLM_UserTable + + litellm.prometheus_user_budget_label_include_email_alias = True + + user = LiteLLM_UserTable( + user_id="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + spend=25.0, + max_budget=100.0, + budget_reset_at=datetime(2026, 3, 1, tzinfo=timezone.utc), + ) + + prometheus_logger.litellm_remaining_user_budget_metric = MagicMock() + prometheus_logger.litellm_user_max_budget_metric = MagicMock() + prometheus_logger.litellm_user_budget_remaining_hours_metric = MagicMock() + + try: + prometheus_logger._set_user_budget_metrics(user) + + prometheus_logger.litellm_remaining_user_budget_metric.labels.assert_called_once_with( + user="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + ) + prometheus_logger.litellm_remaining_user_budget_metric.labels().set.assert_called_once_with( + 75.0 + ) + prometheus_logger.litellm_user_max_budget_metric.labels.assert_called_once_with( + user="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + ) + prometheus_logger.litellm_user_budget_remaining_hours_metric.labels.assert_called_once_with( + user="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + ) + finally: + litellm.prometheus_user_budget_label_include_email_alias = False + + async def test_set_user_budget_metrics_after_api_request_no_inf_when_metadata_budget_none( prometheus_logger, ): diff --git a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py index 758ff3ea38e..0ef97e25e44 100644 --- a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py +++ b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py @@ -9,15 +9,24 @@ Covers credential leak prevention changes: import os import sys -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest sys.path.insert(0, os.path.abspath("../../..")) +from litellm.interactions.litellm_responses_transformation.streaming_iterator import ( + LiteLLMResponsesInteractionsStreamingIterator, +) from litellm.llms.gemini.interactions.transformation import ( GoogleAIStudioInteractionsConfig, ) +from litellm.types.llms.openai import ( + ContentPartAddedEvent, + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponseCreatedEvent, +) from litellm.types.router import GenericLiteLLMParams _PATCH_GET_API_KEY = "litellm.llms.gemini.common_utils.GeminiModelInfo.get_api_key" @@ -113,6 +122,186 @@ class TestGetCompleteUrl: ) +class TestStreamingIterator: + def _make_iterator(self) -> LiteLLMResponsesInteractionsStreamingIterator: + return LiteLLMResponsesInteractionsStreamingIterator( + model="gpt-5.4", + litellm_custom_stream_wrapper=MagicMock(), + request_input="hi", + optional_params={}, + ) + + def _make_text_delta( + self, text: str, item_id: str = "item_1" + ) -> OutputTextDeltaEvent: + event = MagicMock(spec=OutputTextDeltaEvent) + event.delta = text + event.item_id = item_id + return event + + def _make_part_added(self, item_id: str = "item_1") -> ContentPartAddedEvent: + event = MagicMock(spec=ContentPartAddedEvent) + event.item_id = item_id + return event + + def _make_response_created(self) -> ResponseCreatedEvent: + event = MagicMock(spec=ResponseCreatedEvent) + event.response = MagicMock(id="resp_123") + return event + + def test_content_delta_includes_type_field(self): + """content.delta events must carry delta.type='text' so the UI can display them.""" + it = self._make_iterator() + it.sent_interaction_start = True + it.sent_content_start = True + + chunk = it._transform_responses_chunk_to_interactions_chunk( + self._make_text_delta("Hello") + ) + + assert chunk is not None + assert chunk.event_type == "content.delta" + assert chunk.delta == {"type": "text", "text": "Hello"} + + def test_response_part_added_emits_content_start(self): + """ContentPartAddedEvent (arrives before text deltas) should emit content.start + so the first OutputTextDeltaEvent immediately emits content.delta without dropping text. + """ + it = self._make_iterator() + it.sent_interaction_start = True + + chunk = it._transform_responses_chunk_to_interactions_chunk( + self._make_part_added() + ) + + assert chunk is not None + assert chunk.event_type == "content.start" + assert it.sent_content_start is True + + def test_first_text_delta_not_dropped_when_part_added_seen(self): + """After ContentPartAddedEvent, the first text delta must yield content.delta + (not content.start), preserving the token text.""" + it = self._make_iterator() + it.sent_interaction_start = True + it._transform_responses_chunk_to_interactions_chunk(self._make_part_added()) + + chunk = it._transform_responses_chunk_to_interactions_chunk( + self._make_text_delta("Hello") + ) + + assert chunk is not None + assert chunk.event_type == "content.delta" + assert chunk.delta is not None + assert chunk.delta.get("text") == "Hello" + + def test_part_added_emits_interaction_start_fallback_when_not_sent(self): + """If ContentPartAddedEvent arrives before any ResponseCreatedEvent, + the iterator must emit interaction.start before content.start to honor + the documented event ordering contract.""" + it = self._make_iterator() + + chunk = it._transform_responses_chunk_to_interactions_chunk( + self._make_part_added(item_id="item_42") + ) + + assert chunk is not None + assert chunk.event_type == "interaction.start" + assert chunk.id == "item_42" + assert chunk.status == "in_progress" + assert chunk.model == "gpt-5.4" + assert it.sent_interaction_start is True + assert it.sent_content_start is False + + def test_part_added_returns_none_when_already_started(self): + """A second ContentPartAddedEvent (after content.start was already emitted) + should be a no-op so we don't re-emit content.start.""" + it = self._make_iterator() + it.sent_interaction_start = True + it.sent_content_start = True + + chunk = it._transform_responses_chunk_to_interactions_chunk( + self._make_part_added() + ) + + assert chunk is None + + def test_part_added_without_item_id_falls_back_to_self_id(self): + """When ContentPartAddedEvent has no item_id and we emit the interaction.start + fallback, the id must default to an interaction_ string.""" + it = self._make_iterator() + event = MagicMock(spec=ContentPartAddedEvent) + event.item_id = None + + chunk = it._transform_responses_chunk_to_interactions_chunk(event) + + assert chunk is not None + assert chunk.event_type == "interaction.start" + assert chunk.id == f"interaction_{id(it)}" + + def test_first_text_delta_not_dropped_when_no_prior_start_events(self): + """When OutputTextDeltaEvent arrives before any ResponseCreatedEvent or + ContentPartAddedEvent, the iterator must emit interaction.start *and* + immediately follow with a content.start that carries this delta's text, + so the first token is never silently dropped from the stream.""" + events = [ + self._make_text_delta("Hello"), + self._make_text_delta(" World"), + ] + wrapper = MagicMock() + wrapper.__iter__ = lambda self: iter(events) + wrapper.__next__ = lambda self, _it=iter(events): next(_it) + it = LiteLLMResponsesInteractionsStreamingIterator( + model="gpt-5.4", + litellm_custom_stream_wrapper=wrapper, + request_input="hi", + optional_params={}, + ) + + first = it._transform_responses_chunk_to_interactions_chunk(events[0]) + assert first is not None + assert first.event_type == "interaction.start" + assert it.sent_interaction_start is True + assert it.sent_content_start is True + assert len(it._pending_events) == 1 + pending = it._pending_events[0] + assert pending.event_type == "content.start" + assert pending.delta == {"type": "text", "text": "Hello"} + + second = it._transform_responses_chunk_to_interactions_chunk(events[1]) + assert second is not None + assert second.event_type == "content.delta" + assert second.delta == {"type": "text", "text": " World"} + + +class TestTransformRequest: + def test_stream_param_included_in_request_body(self, config): + """When stream=True is in optional_params, the request body must include it + so the proxy forwards the SSE streaming flag to Google's backend.""" + body = config.transform_request( + model="gemini-2.5-flash", + agent=None, + input="Hello", + optional_params={"stream": True}, + litellm_params=GenericLiteLLMParams(api_key="test-key"), + headers={}, + ) + + assert body.get("stream") is True + assert body.get("input") == "Hello" + + def test_stream_false_not_included_when_absent(self, config): + body = config.transform_request( + model="gemini-2.5-flash", + agent=None, + input="Hello", + optional_params={}, + litellm_params=GenericLiteLLMParams(api_key="test-key"), + headers={}, + ) + + assert "stream" not in body + + class TestInteractionOperationUrls: """Test that get/delete/cancel interaction URLs exclude API key.""" diff --git a/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py b/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py new file mode 100644 index 00000000000..8de47331614 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py @@ -0,0 +1,119 @@ +""" +Test that BedrockBatchesConfig._get_openai_compatible_batch_metadata +sanitizes non-string metadata values injected by proxy guardrail hooks. + +The OpenAI Batch Pydantic model requires metadata: Dict[str, str]. +Proxy hooks (Model Armor, OpenAI Moderations, queue time tracking) inject +dicts, floats, and other non-string values that cause a ValidationError +when constructing LiteLLMBatch. This test suite verifies the sanitization +layer prevents that. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig + + +class TestGetOpenaiCompatibleBatchMetadata: + """Tests for _get_openai_compatible_batch_metadata.""" + + def test_string_values_pass_through_unchanged(self): + metadata = {"user_key": "user_value", "run_id": "abc123"} + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + assert result == {"user_key": "user_value", "run_id": "abc123"} + + def test_dict_values_serialized_to_json_string(self): + metadata = { + "_model_armor_response": { + "sanitizationResult": {"filterMatchState": "MATCH_FOUND"} + } + } + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + assert "_model_armor_response" in result + assert isinstance(result["_model_armor_response"], str) + assert "MATCH_FOUND" in result["_model_armor_response"] + + def test_float_values_serialized_to_string(self): + metadata = {"queue_time_seconds": 0.5} + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + assert result == {"queue_time_seconds": "0.5"} + + def test_none_values_excluded(self): + metadata = {"key": "value", "empty": None} + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + assert "empty" not in result + assert result == {"key": "value"} + + def test_standard_logging_guardrail_information_excluded(self): + metadata = { + "standard_logging_guardrail_information": {"some": "logging_data"}, + "user_key": "keep_me", + } + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + assert "standard_logging_guardrail_information" not in result + assert result == {"user_key": "keep_me"} + + def test_non_dict_input_returns_empty_dict(self): + assert BedrockBatchesConfig._get_openai_compatible_batch_metadata(None) == {} + assert BedrockBatchesConfig._get_openai_compatible_batch_metadata("string") == {} + assert BedrockBatchesConfig._get_openai_compatible_batch_metadata(123) == {} + + def test_empty_dict_returns_empty_dict(self): + assert BedrockBatchesConfig._get_openai_compatible_batch_metadata({}) == {} + + def test_mixed_metadata_from_guardrails(self): + """Simulate real metadata contaminated by proxy guardrails.""" + metadata = { + "_model_armor_response": {"sanitizationResult": {"key": "val"}}, + "_model_armor_status": "success", + "_openai_moderation_response": {"id": "mod-123", "flagged": False}, + "queue_time_seconds": 1.23, + "headers": {"Authorization": "Bearer sk-xxx"}, + "standard_logging_guardrail_information": {"internal": True}, + "user_metadata_key": "user_value", + "none_field": None, + } + result = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + + # All values must be strings + for key, value in result.items(): + assert isinstance(value, str), f"metadata[{key!r}] is {type(value)}, not str" + + # Excluded keys + assert "standard_logging_guardrail_information" not in result + assert "none_field" not in result + + # Preserved keys + assert result["_model_armor_status"] == "success" + assert result["user_metadata_key"] == "user_value" + + def test_result_compatible_with_litellm_batch(self): + """Verify sanitized metadata can construct a LiteLLMBatch without error.""" + import time + + from litellm.types.utils import LiteLLMBatch + + metadata = { + "_model_armor_response": {"blocked": True}, + "queue_time_seconds": 0.05, + "user_key": "value", + } + sanitized = BedrockBatchesConfig._get_openai_compatible_batch_metadata(metadata) + + # This would raise ValidationError before the fix + batch = LiteLLMBatch( + id="arn:aws:bedrock:us-east-1:123:model-invocation-job/test", + object="batch", + endpoint="/v1/chat/completions", + input_file_id="file-123", + completion_window="24h", + status="validating", + created_at=int(time.time()), + metadata=sanitized, + ) + assert batch.metadata == sanitized diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py index c67a8712340..9955851132c 100644 --- a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py +++ b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py @@ -957,3 +957,50 @@ def test_titan_image_embedding_cost_uses_per_image_rate(): assert response.usage is not None assert response.usage.prompt_tokens_details is not None assert response.usage.prompt_tokens_details.image_count == 1 + + +@pytest.mark.parametrize( + "encoding_format,expected_embedding_types", + [ + ("float", ["float"]), + ("base64", ["base64"]), + (["float", "int8"], ["float", "int8"]), + ], +) +def test_bedrock_cohere_embedding_types_wrapped_as_list( + encoding_format, expected_embedding_types +): + """ + Bedrock Cohere expects `embedding_types` as a JSON array, not a raw string. + + Regression test for: Bedrock returns + Malformed input request: #/embedding_types: expected type: JSONArray, found: String + when `encoding_format` is passed as a string. + """ + litellm.set_verbose = True + client = HTTPHandler() + model = "bedrock/cohere.embed-multilingual-v3" + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(cohere_embedding_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + response = litellm.embedding( + model=model, + input=test_input, + encoding_format=encoding_format, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key="test-bearer-token-12345", + ) + + assert isinstance(response, litellm.EmbeddingResponse) + + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + assert "embedding_types" in request_body + assert request_body["embedding_types"] == expected_embedding_types + assert isinstance(request_body["embedding_types"], list) diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py b/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py index 8fa9290d3de..c39fb427a01 100644 --- a/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py @@ -14,18 +14,46 @@ from litellm.llms.bedrock.common_utils import BedrockModelInfo # --------------------------------------------------------------------------- # -# BEDROCK_RESPONSE_STREAM_SHAPE eager-load tests # +# get_bedrock_response_stream_shape lazy-load tests # # --------------------------------------------------------------------------- # -def test_bedrock_response_stream_shape_loaded_at_import(): +@pytest.fixture(autouse=True) +def _reset_bedrock_response_stream_shape_cache(): + """Prevent lru_cache leakage between tests in this module.""" + import litellm.llms.bedrock.common_utils as mod + + mod.get_bedrock_response_stream_shape.cache_clear() + yield + mod.get_bedrock_response_stream_shape.cache_clear() + + +def test_bedrock_response_stream_shape_lazy_loads_once(): """ - BEDROCK_RESPONSE_STREAM_SHAPE is resolved at module import time. + get_bedrock_response_stream_shape() loads from botocore at most once per process. + """ + from unittest.mock import MagicMock, patch + + import litellm.llms.bedrock.common_utils as mod + + sentinel = MagicMock() + with patch.object( + mod, "_load_bedrock_response_stream_shape", return_value=sentinel + ) as mock_load: + assert mod.get_bedrock_response_stream_shape() is sentinel + assert mod.get_bedrock_response_stream_shape() is sentinel + mock_load.assert_called_once() + + +def test_bedrock_response_stream_shape_loaded_on_first_access(): + """ + get_bedrock_response_stream_shape() loads once on first use. In a standard environment with botocore installed it must be non-None. """ - from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE + pytest.importorskip("botocore") + from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape - assert BEDROCK_RESPONSE_STREAM_SHAPE is not None + assert get_bedrock_response_stream_shape() is not None def test_bedrock_response_stream_shape_load_failure_returns_none(): @@ -38,6 +66,7 @@ def test_bedrock_response_stream_shape_load_failure_returns_none(): import litellm.llms.bedrock.common_utils as mod + pytest.importorskip("botocore") with patch( "botocore.loaders.Loader.load_service_model", side_effect=Exception("no data"), @@ -51,31 +80,29 @@ def test_bedrock_response_stream_shape_is_structure_shape(): The loaded shape should be the botocore StructureShape for ResponseStream, not a plain dict or any other type. """ + pytest.importorskip("botocore") from botocore.model import StructureShape - from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE + from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape - assert BEDROCK_RESPONSE_STREAM_SHAPE is not None, ( - "BEDROCK_RESPONSE_STREAM_SHAPE is None — botocore may not be installed" - ) - shape: StructureShape = BEDROCK_RESPONSE_STREAM_SHAPE # remove Optional + loaded_shape = get_bedrock_response_stream_shape() + assert ( + loaded_shape is not None + ), "get_bedrock_response_stream_shape() is None — botocore may not be installed" + shape: StructureShape = loaded_shape assert isinstance(shape, StructureShape) assert shape.name == "ResponseStream" -def test_bedrock_response_stream_shape_same_object_across_imports(): +def test_bedrock_response_stream_shape_same_object_across_calls(): """ - Both bedrock modules that use the shape must reference the identical object — - confirming the constant is not re-loaded per import. + Repeated calls must return the identical cached object. """ - from litellm.llms.bedrock.chat.invoke_handler import ( - BEDROCK_RESPONSE_STREAM_SHAPE as invoke_shape, - ) - from litellm.llms.bedrock.common_utils import ( - BEDROCK_RESPONSE_STREAM_SHAPE as common_shape, - ) + from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape - assert common_shape is invoke_shape + first = get_bedrock_response_stream_shape() + second = get_bedrock_response_stream_shape() + assert first is second def test_bedrock_event_stream_decoder_base_uses_module_shape(): @@ -95,19 +122,23 @@ def test_bedrock_event_stream_decoder_base_uses_module_shape(): def test_bedrock_parse_message_from_event_raises_on_none_shape(): """ - When BEDROCK_RESPONSE_STREAM_SHAPE is None (botocore unavailable), + When get_bedrock_response_stream_shape() returns None (botocore unavailable), _parse_message_from_event must raise BedrockError before touching the botocore parser — not an opaque AttributeError from inside botocore. """ from unittest.mock import MagicMock, patch import litellm.llms.bedrock.common_utils as mod - from litellm.llms.bedrock.common_utils import BedrockError, BedrockEventStreamDecoderBase + from litellm.llms.bedrock.common_utils import ( + BedrockError, + BedrockEventStreamDecoderBase, + ) - decoder = BedrockEventStreamDecoderBase() + decoder = BedrockEventStreamDecoderBase.__new__(BedrockEventStreamDecoderBase) + decoder.parser = MagicMock() mock_event = MagicMock() - with patch.object(mod, "BEDROCK_RESPONSE_STREAM_SHAPE", None): + with patch.object(mod, "get_bedrock_response_stream_shape", return_value=None): with pytest.raises(BedrockError) as exc_info: decoder._parse_message_from_event(mock_event) diff --git a/tests/test_litellm/llms/deepseek/__init__.py b/tests/test_litellm/llms/deepseek/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/deepseek/messages/__init__.py b/tests/test_litellm/llms/deepseek/messages/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/deepseek/messages/test_deepseek_anthropic_messages_transformation.py b/tests/test_litellm/llms/deepseek/messages/test_deepseek_anthropic_messages_transformation.py new file mode 100644 index 00000000000..7c5f0483ded --- /dev/null +++ b/tests/test_litellm/llms/deepseek/messages/test_deepseek_anthropic_messages_transformation.py @@ -0,0 +1,189 @@ +import litellm +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, +) +from litellm.llms.deepseek.messages.transformation import ( + DeepSeekAnthropicMessagesConfig, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ProviderConfigManager + + +def test_deepseek_provider_uses_anthropic_messages_config(): + config = ProviderConfigManager.get_provider_anthropic_messages_config( + model="deepseek-v4-pro", + provider=litellm.LlmProviders.DEEPSEEK, + ) + + assert isinstance(config, DeepSeekAnthropicMessagesConfig) + assert config.custom_llm_provider == "deepseek" + + +def test_deepseek_anthropic_messages_config_defaults(): + config = DeepSeekAnthropicMessagesConfig() + + assert config.custom_llm_provider == "deepseek" + assert config.get_api_base() == "https://api.deepseek.com/anthropic" + + +def test_anthropic_provider_keeps_default_config_for_deepseek_named_model(): + config = ProviderConfigManager.get_provider_anthropic_messages_config( + model="deepseek-v4-pro", + provider=litellm.LlmProviders.ANTHROPIC, + ) + + assert isinstance(config, AnthropicMessagesConfig) + assert not isinstance(config, DeepSeekAnthropicMessagesConfig) + + +def test_deepseek_anthropic_messages_url_defaults_to_anthropic_endpoint(): + config = DeepSeekAnthropicMessagesConfig() + + assert ( + config.get_complete_url( + api_base=None, + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + assert ( + config.get_complete_url( + api_base="https://api.deepseek.com/anthropic/v1", + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + assert ( + config.get_complete_url( + api_base="https://api.deepseek.com/anthropic", + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + assert ( + config.get_complete_url( + api_base="https://api.deepseek.com", + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + assert ( + config.get_complete_url( + api_base="https://api.deepseek.com/v1", + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + assert ( + config.get_complete_url( + api_base="https://api.deepseek.com/v1/messages", + api_key=None, + model="deepseek-v4-pro", + optional_params={}, + litellm_params={}, + ) + == "https://api.deepseek.com/anthropic/v1/messages" + ) + + +def test_deepseek_anthropic_messages_headers_use_deepseek_key(): + config = DeepSeekAnthropicMessagesConfig() + + headers, api_base = config.validate_anthropic_messages_environment( + headers={}, + model="deepseek-v4-pro", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-deepseek", + api_base="https://example.test/anthropic", + ) + + assert api_base == "https://example.test/anthropic" + assert headers["x-api-key"] == "sk-deepseek" + assert headers["anthropic-version"] == "2023-06-01" + assert headers["content-type"] == "application/json" + + +def test_deepseek_anthropic_messages_preserves_thinking_and_sanitizes_custom_tools(): + config = DeepSeekAnthropicMessagesConfig() + messages = [ + { + "role": "user", + "content": "Use the tool.", + }, + { + "role": "assistant", + "content": [ + { + "type": "thinking", + "thinking": "I should call the tool.", + "signature": "sig", + }, + { + "type": "tool_use", + "id": "toolu_123", + "name": "get_weather", + "input": {"city": "Sao Paulo"}, + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_123", + "content": "Sunny", + } + ], + }, + ] + + request = config.transform_anthropic_messages_request( + model="deepseek-v4-pro", + messages=messages, + anthropic_messages_optional_request_params={ + "max_tokens": 100, + "thinking": {"type": "enabled", "budget_tokens": 1024}, + "tools": [ + { + "type": "custom", + "name": "get_weather", + "description": "Get weather", + "input_schema": {"type": "object"}, + }, + { + "type": "web_search_20260209", + "name": "web_search", + "max_uses": 1, + }, + ], + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert request["messages"] == messages + assert request["thinking"] == {"type": "enabled", "budget_tokens": 1024} + assert request["tools"][0] == { + "name": "get_weather", + "description": "Get weather", + "input_schema": {"type": "object"}, + } + assert request["tools"][1]["type"] == "web_search_20260209" diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py index 9d7706557b5..7e13459bca1 100644 --- a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py @@ -12,18 +12,46 @@ from litellm.llms.sagemaker.completion.transformation import SagemakerConfig # --------------------------------------------------------------------------- # -# SAGEMAKER_RESPONSE_STREAM_SHAPE eager-load tests # +# get_sagemaker_response_stream_shape lazy-load tests # # --------------------------------------------------------------------------- # -def test_sagemaker_response_stream_shape_loaded_at_import(): +@pytest.fixture(autouse=True) +def _reset_sagemaker_response_stream_shape_cache(): + """Prevent lru_cache leakage between tests in this module.""" + import litellm.llms.sagemaker.common_utils as mod + + mod.get_sagemaker_response_stream_shape.cache_clear() + yield + mod.get_sagemaker_response_stream_shape.cache_clear() + + +def test_sagemaker_response_stream_shape_lazy_loads_once(): """ - SAGEMAKER_RESPONSE_STREAM_SHAPE is resolved at module import time. + get_sagemaker_response_stream_shape() loads from botocore at most once per process. + """ + from unittest.mock import MagicMock, patch + + import litellm.llms.sagemaker.common_utils as mod + + sentinel = MagicMock() + with patch.object( + mod, "_load_sagemaker_response_stream_shape", return_value=sentinel + ) as mock_load: + assert mod.get_sagemaker_response_stream_shape() is sentinel + assert mod.get_sagemaker_response_stream_shape() is sentinel + mock_load.assert_called_once() + + +def test_sagemaker_response_stream_shape_loaded_on_first_access(): + """ + get_sagemaker_response_stream_shape() loads once on first use. In a standard environment with botocore installed it must be non-None. """ - from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + pytest.importorskip("botocore") + from litellm.llms.sagemaker.common_utils import get_sagemaker_response_stream_shape - assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None + assert get_sagemaker_response_stream_shape() is not None def test_sagemaker_response_stream_shape_load_failure_returns_none(): @@ -36,6 +64,7 @@ def test_sagemaker_response_stream_shape_load_failure_returns_none(): import litellm.llms.sagemaker.common_utils as mod + pytest.importorskip("botocore") with patch( "botocore.loaders.Loader.load_service_model", side_effect=Exception("no data"), @@ -49,14 +78,16 @@ def test_sagemaker_response_stream_shape_is_structure_shape(): The loaded shape should be the botocore StructureShape for InvokeEndpointWithResponseStreamOutput, not a plain dict or any other type. """ + pytest.importorskip("botocore") from botocore.model import StructureShape - from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + from litellm.llms.sagemaker.common_utils import get_sagemaker_response_stream_shape - assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None, ( - "SAGEMAKER_RESPONSE_STREAM_SHAPE is None — botocore may not be installed" - ) - shape: StructureShape = SAGEMAKER_RESPONSE_STREAM_SHAPE # remove Optional + shape = get_sagemaker_response_stream_shape() + assert ( + shape is not None + ), "get_sagemaker_response_stream_shape() is None — botocore may not be installed" + shape: StructureShape = shape # remove Optional assert isinstance(shape, StructureShape) assert shape.name == "InvokeEndpointWithResponseStreamOutput" @@ -64,29 +95,25 @@ def test_sagemaker_response_stream_shape_is_structure_shape(): def test_sagemaker_response_stream_shape_not_reloaded_on_new_decoder(): """ Creating multiple AWSEventStreamDecoder instances must not trigger - additional botocore Loader calls — the shape is resolved once at import - time and reused. + additional botocore Loader calls — the shape is cached after first access. """ - from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + from litellm.llms.sagemaker.common_utils import get_sagemaker_response_stream_shape - decoder_a = AWSEventStreamDecoder(model="test-model-a") - decoder_b = AWSEventStreamDecoder(model="test-model-b") + decoder_a = AWSEventStreamDecoder.__new__(AWSEventStreamDecoder) + decoder_b = AWSEventStreamDecoder.__new__(AWSEventStreamDecoder) - # Both decoders should use the same pre-loaded shape object (identity check) assert "_response_stream_shape_cache" not in decoder_a.__dict__ assert "_response_stream_shape_cache" not in decoder_b.__dict__ - # The module constant is still the same object - from litellm.llms.sagemaker.common_utils import ( - SAGEMAKER_RESPONSE_STREAM_SHAPE as shape_after, - ) - assert SAGEMAKER_RESPONSE_STREAM_SHAPE is shape_after + first = get_sagemaker_response_stream_shape() + second = get_sagemaker_response_stream_shape() + assert first is second def test_sagemaker_parse_message_from_event_raises_on_none_shape(): """ - When SAGEMAKER_RESPONSE_STREAM_SHAPE is None (botocore unavailable), - _parse_message_from_event must raise ValueError before touching the + When get_sagemaker_response_stream_shape() returns None (botocore unavailable), + _parse_message_from_event must raise SagemakerError before touching the botocore parser — not an opaque AttributeError from inside botocore. """ from unittest.mock import MagicMock, patch @@ -94,10 +121,14 @@ def test_sagemaker_parse_message_from_event_raises_on_none_shape(): import litellm.llms.sagemaker.common_utils as mod from litellm.llms.sagemaker.common_utils import SagemakerError - decoder = AWSEventStreamDecoder(model="test-model") + decoder = AWSEventStreamDecoder.__new__(AWSEventStreamDecoder) + decoder.model = "test-model" + decoder.parser = MagicMock() + decoder.content_blocks = [] + decoder.is_messages_api = None mock_event = MagicMock() - with patch.object(mod, "SAGEMAKER_RESPONSE_STREAM_SHAPE", None): + with patch.object(mod, "get_sagemaker_response_stream_shape", return_value=None): with pytest.raises(SagemakerError) as exc_info: decoder._parse_message_from_event(mock_event) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5c7bbc46c95..b450a262907 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -7943,3 +7943,103 @@ async def test_team_member_me_returns_404_for_unknown_team(mock_db_client): user_api_key_dict=caller_auth, ) assert exc_info.value.status_code == 404 + + +def _non_admin_auth(): + return UserAPIKeyAuth( + user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER + ) + + +def test_check_passthrough_routes_caller_permission_team(): + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.common_utils import ( + _check_passthrough_routes_caller_permission, + ) + + admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + non_admin = _non_admin_auth() + + _check_passthrough_routes_caller_permission( + NewTeamRequest(allowed_passthrough_routes=["/foo/*"]), admin, entity="team" + ) + + _check_passthrough_routes_caller_permission( + NewTeamRequest(), non_admin, entity="team" + ) + _check_passthrough_routes_caller_permission( + NewTeamRequest(allowed_passthrough_routes=[]), non_admin, entity="team" + ) + + with pytest.raises(HTTPException) as exc: + _check_passthrough_routes_caller_permission( + NewTeamRequest(allowed_passthrough_routes=["/admin/*"]), + non_admin, + entity="team", + ) + assert exc.value.status_code == 403 + assert "allowed_passthrough_routes" in str(exc.value.detail) + assert "team" in str(exc.value.detail) + + with pytest.raises(HTTPException) as exc: + _check_passthrough_routes_caller_permission( + NewTeamRequest(metadata={"allowed_passthrough_routes": ["/admin/*"]}), + non_admin, + entity="team", + ) + assert exc.value.status_code == 403 + assert "metadata.allowed_passthrough_routes" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_new_team_blocks_non_admin_passthrough_routes(mock_db_client): + """A non-proxy-admin cannot self-grant pass-through routes via /team/new.""" + mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest, ProxyException + from litellm.proxy.management_endpoints.team_endpoints import new_team + + with patch( + "litellm.proxy.management_endpoints.team_endpoints._check_user_team_limits", + AsyncMock(return_value=None), + ): + with pytest.raises(ProxyException) as exc: + await new_team( + data=NewTeamRequest( + team_alias="t", allowed_passthrough_routes=["/admin/*"] + ), + http_request=MagicMock(spec=Request), + user_api_key_dict=_non_admin_auth(), + ) + assert str(exc.value.code) == "403" + assert "allowed_passthrough_routes" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): + """Even a team manager (non-proxy-admin) cannot set pass-through routes via + /team/update — the gate runs after _verify_team_access.""" + from fastapi import Request + + from litellm.proxy._types import ProxyException, UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + existing = MagicMock() + existing.model_dump.return_value = {"team_id": "t1"} + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing) + + with patch( + "litellm.proxy.management_endpoints.team_endpoints._verify_team_access", + AsyncMock(return_value=None), + ): + with pytest.raises(ProxyException) as exc: + await update_team( + data=UpdateTeamRequest( + team_id="t1", allowed_passthrough_routes=["/admin/*"] + ), + http_request=MagicMock(spec=Request), + user_api_key_dict=_non_admin_auth(), + ) + assert str(exc.value.code) == "403" + assert "allowed_passthrough_routes" in str(exc.value.message) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 73d53631622..6d10d2a6353 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5065,6 +5065,66 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac mock_response.aclose.assert_awaited_once() +@pytest.mark.asyncio +async def test_async_data_generator_passes_through_google_native_sse_bytes(): + """ + Google-native streamGenerateContent yields raw SSE bytes; they must not be + re-wrapped as data: b'data: {...}'. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gemini-2.0-flash", + "messages": [{"role": "user", "content": "test"}], + } + gemini_event = b'data: {"candidates": [{"content": "hi"}]}\n\n' + gemini_event_without_terminator = b'data: {"candidates": [{"content": "there"}]}' + raw_payload = b'{"partial": true}' + + class MockStream: + def __aiter__(self): + return self._stream() + + async def _stream(self): + yield gemini_event + yield gemini_event_without_terminator + yield raw_payload + + async def aclose(self): + pass + + mock_response = MockStream() + mock_response.aclose = AsyncMock() + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.has_streaming_callbacks.return_value = False + mock_proxy_logging_obj.needs_iterator_wrap.return_value = False + mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock() + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + yielded_data = [] + async for data in async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ): + yielded_data.append(data) + + yielded_text = [ + chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk + for chunk in yielded_data + ] + assert yielded_text[0] == gemini_event.decode("utf-8") + assert yielded_text[1] == gemini_event_without_terminator.decode("utf-8") + "\n\n" + assert yielded_text[2] == f'data: {raw_payload.decode("utf-8")}\n\n' + assert "b'data:" not in "".join(yielded_text) + assert yielded_text[-1] == "data: [DONE]\n\n" + + @pytest.mark.asyncio async def test_async_data_generator_cleanup_on_normal_completion(): """ diff --git a/tests/unified_google_tests/conftest.py b/tests/unified_google_tests/conftest.py index d28f89a77b0..5b4f57b8036 100644 --- a/tests/unified_google_tests/conftest.py +++ b/tests/unified_google_tests/conftest.py @@ -12,14 +12,17 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm # noqa: E402,F401 -from tests._vcr_conftest_common import ( # noqa: E402 +from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, + _pin_multipart_boundary, apply_vcr_auto_marker_to_items, emit_cassette_cache_session_banner, emit_vcr_classification_summary, + emit_vcr_diagnostic_log, install_live_call_probe, record_vcr_outcome, register_persister_if_enabled, + reset_vcr_diag_dir, vcr_config_dict, ) @@ -84,6 +87,7 @@ def _vcr_outcome_gate(request, vcr): def pytest_configure(config): _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() def pytest_runtest_logreport(report): @@ -110,3 +114,4 @@ def pytest_collection_modifyitems(config, items): def pytest_terminal_summary(terminalreporter, exitstatus, config): emit_cassette_cache_session_banner(terminalreporter) emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx index 06a09ca2c37..a80cf935c4b 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx @@ -49,6 +49,7 @@ import { fetchAvailableModels, ModelGroup } from "../llm_calls/fetch_models"; import { makeOpenAIImageEditsRequest } from "../llm_calls/image_edits"; import { makeOpenAIImageGenerationRequest } from "../llm_calls/image_generation"; import { makeOpenAIResponsesRequest } from "../llm_calls/responses_api"; +import { makeInteractionsRequest } from "../llm_calls/interactions_api"; import A2AMetrics from "./A2AMetrics"; import AdditionalModelSettings from "./AdditionalModelSettings"; import AudioRenderer from "./AudioRenderer"; @@ -649,6 +650,7 @@ const ChatUI: React.FC = ({ EndpointType.ANTHROPIC_MESSAGES, EndpointType.EMBEDDINGS, EndpointType.TRANSCRIPTION, + EndpointType.INTERACTIONS, ]; if (modelRequiredEndpoints.includes(endpointType as EndpointType) && !selectedModel) { @@ -914,6 +916,16 @@ const ChatUI: React.FC = ({ customProxyBaseUrl || undefined, ); } + } else if (endpointType === EndpointType.INTERACTIONS) { + await makeInteractionsRequest( + inputMessage, + (text, model) => updateTextUI("assistant", text, model), + selectedModel, + effectiveApiKey, + selectedTags, + signal, + customProxyBaseUrl || undefined, + ); } } @@ -1241,10 +1253,11 @@ const ChatUI: React.FC = ({ return true; } const optionEndpoint = getEndpointType(option.mode); - // Show chat models for responses/anthropic_messages endpoints as they are compatible + // Show chat models for responses/anthropic_messages/interactions endpoints as they are compatible if ( endpointType === EndpointType.RESPONSES || - endpointType === EndpointType.ANTHROPIC_MESSAGES + endpointType === EndpointType.ANTHROPIC_MESSAGES || + endpointType === EndpointType.INTERACTIONS ) { return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT; } @@ -2089,7 +2102,8 @@ const ChatUI: React.FC = ({ endpointType === EndpointType.CHAT || endpointType === EndpointType.EMBEDDINGS || endpointType === EndpointType.RESPONSES || - endpointType === EndpointType.ANTHROPIC_MESSAGES + endpointType === EndpointType.ANTHROPIC_MESSAGES || + endpointType === EndpointType.INTERACTIONS ? "Type your message... (Shift+Enter for new line)" : endpointType === EndpointType.A2A_AGENTS ? "Send a message to the A2A agent..." diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/chatConstants.ts b/ui/litellm-dashboard/src/components/playground/chat_ui/chatConstants.ts index 919e3bc1c65..9592a521250 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/chatConstants.ts +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/chatConstants.ts @@ -45,4 +45,5 @@ export const ENDPOINT_OPTIONS = [ { value: EndpointType.A2A_AGENTS, label: "/v1/a2a/message/send" }, { value: EndpointType.MCP, label: "/mcp-rest/tools/call" }, { value: EndpointType.REALTIME, label: "/v1/realtime" }, + { value: EndpointType.INTERACTIONS, label: "/v1beta/interactions" }, ]; diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/mode_endpoint_mapping.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/mode_endpoint_mapping.tsx index f354efe641e..3b2a5e0fd60 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/mode_endpoint_mapping.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/mode_endpoint_mapping.tsx @@ -28,6 +28,7 @@ export enum EndpointType { A2A_AGENTS = "a2a_agents", MCP = "mcp", REALTIME = "realtime", + INTERACTIONS = "interactions", } // Create a mapping between the model mode and the corresponding endpoint type diff --git a/ui/litellm-dashboard/src/components/playground/llm_calls/interactions_api.tsx b/ui/litellm-dashboard/src/components/playground/llm_calls/interactions_api.tsx new file mode 100644 index 00000000000..6a6e4a2cfa0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/playground/llm_calls/interactions_api.tsx @@ -0,0 +1,124 @@ +import NotificationManager from "@/components/molecules/notifications_manager"; +import { getGlobalLitellmHeaderName, getProxyBaseUrl } from "@/components/networking"; + +export async function makeInteractionsRequest( + input: string, + updateUI: (text: string, model?: string) => void, + selectedModel: string, + accessToken: string, + tags?: string[], + signal?: AbortSignal, + customBaseUrl?: string, + previousInteractionId?: string, +): Promise { + if (!accessToken) { + throw new Error("Virtual Key is required"); + } + + const isLocal = process.env.NODE_ENV === "development"; + if (isLocal !== true) { + console.log = function () {}; + } + + const proxyBaseUrl = customBaseUrl || getProxyBaseUrl(); + const normalizedBaseUrl = proxyBaseUrl.endsWith("/") ? proxyBaseUrl.slice(0, -1) : proxyBaseUrl; + const requestUrl = `${normalizedBaseUrl}/v1beta/interactions`; + + const headers: Record = { + "Content-Type": "application/json", + [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`, + }; + if (tags && tags.length > 0) { + headers["x-litellm-tags"] = tags.join(","); + } + + const body: Record = { + model: selectedModel, + input, + stream: true, + }; + if (previousInteractionId) { + body.previous_interaction_id = previousInteractionId; + } + + try { + const response = await fetch(requestUrl, { + method: "POST", + headers, + body: JSON.stringify(body), + signal, + }); + + if (!response.ok) { + const errorText = await response.text(); + throw new Error(errorText || `Request failed with status ${response.status}`); + } + + if (!response.body) { + throw new Error("No response body received"); + } + + const reader = response.body.getReader(); + const decoder = new TextDecoder(); + let responseModel: string | undefined; + let buffer = ""; + + while (true) { + const { done, value } = await reader.read(); + if (done) break; + + buffer += decoder.decode(value, { stream: true }); + + // SSE lines are separated by double newlines; split on single newlines and + // look for "data: " prefixed lines. + const lines = buffer.split("\n"); + // Keep the last (potentially incomplete) line in the buffer + buffer = lines.pop() ?? ""; + + for (const line of lines) { + const trimmed = line.trim(); + if (!trimmed.startsWith("data:")) continue; + + const jsonStr = trimmed.slice("data:".length).trim(); + if (!jsonStr || jsonStr === "[DONE]") continue; + + let event: Record; + try { + event = JSON.parse(jsonStr); + } catch { + continue; + } + + const eventType = event.event_type as string | undefined; + + if (eventType === "interaction.start" || eventType === "interaction.complete") { + // Capture model from either the native Gemini shape (nested under + // `interaction`) or the bridge shape (top-level `model` field). + const interaction = event.interaction as Record | undefined; + if (typeof interaction?.model === "string" && interaction.model) { + responseModel = interaction.model; + } else if (typeof event.model === "string" && event.model) { + responseModel = event.model; + } + } else if (eventType === "content.delta" || eventType === "content.start") { + const delta = event.delta as Record | undefined; + // Accept both native Gemini format {"type":"text","text":"..."} and bridge + // format {"text":"..."} (no type discriminator) + if (typeof delta?.text === "string" && delta.text) { + updateUI(delta.text, responseModel ?? selectedModel); + } + } + // content.start, content.stop, interaction.status_update — no UI action needed + } + } + } catch (error: unknown) { + if (signal?.aborted) { + console.log("Interactions request was cancelled"); + throw error; + } + NotificationManager.fromBackend( + `Error occurred while making Interactions API request. Error: ${error}`, + ); + throw error; + } +}