mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
Merge branch 'litellm_internal_staging' into litellm_oci_genai_ci
This commit is contained in:
commit
1f064a76e2
61 changed files with 2579 additions and 328 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
133
litellm/llms/deepseek/messages/transformation.py
Normal file
133
litellm/llms/deepseek/messages/transformation.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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/",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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: <block message>" 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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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: "<bound method Cls.m of <... at 0x..>>" -> "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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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_<id(self)> 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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
0
tests/test_litellm/llms/deepseek/__init__.py
Normal file
0
tests/test_litellm/llms/deepseek/__init__.py
Normal file
0
tests/test_litellm/llms/deepseek/messages/__init__.py
Normal file
0
tests/test_litellm/llms/deepseek/messages/__init__.py
Normal file
|
|
@ -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"
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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<ChatUIProps> = ({
|
|||
EndpointType.ANTHROPIC_MESSAGES,
|
||||
EndpointType.EMBEDDINGS,
|
||||
EndpointType.TRANSCRIPTION,
|
||||
EndpointType.INTERACTIONS,
|
||||
];
|
||||
|
||||
if (modelRequiredEndpoints.includes(endpointType as EndpointType) && !selectedModel) {
|
||||
|
|
@ -914,6 +916,16 @@ const ChatUI: React.FC<ChatUIProps> = ({
|
|||
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<ChatUIProps> = ({
|
|||
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<ChatUIProps> = ({
|
|||
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..."
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
];
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<void> {
|
||||
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<string, string> = {
|
||||
"Content-Type": "application/json",
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
};
|
||||
if (tags && tags.length > 0) {
|
||||
headers["x-litellm-tags"] = tags.join(",");
|
||||
}
|
||||
|
||||
const body: Record<string, unknown> = {
|
||||
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<string, unknown>;
|
||||
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<string, unknown> | 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<string, unknown> | 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;
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue