From 56bafc29cf7601304c5f2e7dbb8e18d614d91fc2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 13:20:03 -0700 Subject: [PATCH] feat(gemini): accept file content blocks with video_metadata on multimodal embeddings (#45305) * feat(gemini): accept file content blocks with video_metadata on multimodal embeddings * test(gemini): annotate the embedding mock helper and keep the type comment short * fix(gemini-embeddings): keep file blocks through a partial embedding cache hit and validate video_metadata strictly * test(caching): type the partial-hit cache test and pin its clock * fix(gemini-embeddings): accept the Files API URI /v1/files returns as a file block's file_id * fix(gemini-embeddings): 400 on bad input shapes, drop unknown block keys under drop_params, cache multi-block inputs A bare object `input` answered 500 from both the caching handler and the transformation; both now answer 400. An unresolved `files/` reference on vertex_ai/ answered a ValueError 500; it now answers 400 naming the gemini/ provider. An empty `format` passed through to the provider; it now answers 400 naming file.format. Unknown block keys (`detail`, an unknown video_metadata key, a top-level block key) are dropped under drop_params, global or per-request, and still answer 400 without it. The embedding cache counted file blocks against `max_messages`, so a request with 5 or more blocks was never cached; file blocks no longer count. * fix(caching): answer 400 for an object embedding input on the cache lookup * test(gemini-embeddings): annotate the new embedding tests with return types and Final * test(gemini-embeddings): annotate the transformation and files batch tests with return types and Final * test(gemini-embeddings): audit file content blocks on the wire, in the cache, in batch uploads and under chaos Integration cells for gemini/ and vertex_ai/ embedding file blocks with video_metadata: the batchEmbedContents and embedContent wire shapes, format overrides, nested and repeated inputs, every malformed block answering 400 before any provider call, drop_params at the request, deployment and YAML levels, the OpenAI SDK sync and async clients, Redis cache fills, hits, partial hits and the metadata in the key with zero-spend hit rows, chat, responses and messages controls, the vertex_ai/ batch JSONL upload, a mixed fast/slow/malformed/dropped burst and a worker kill. The wire peer now reads chunked request bodies, which the streamed GCS media upload sends * test(integration): treat an upload the client abandons before its terminating chunk as a disconnect, never a stored request --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/caching/caching.py | 14 +- litellm/caching/caching_handler.py | 21 +- litellm/llms/vertex_ai/common_utils.py | 16 + .../llms/vertex_ai/files/transformation.py | 21 +- .../llms/vertex_ai/gemini/transformation.py | 11 +- .../batch_embed_content_handler.py | 30 +- .../batch_embed_content_transformation.py | 342 +++++++++++++----- litellm/main.py | 6 + litellm/types/llms/vertex_ai.py | 15 +- tests/integration/_support/wire.py | 49 ++- .../test_embedding_file_block_cache.py | 216 +++++++++++ .../test_gemini_chat_video_metadata_wire.py | 144 ++++++++ .../test_gemini_embedding_file_block_chaos.py | 223 ++++++++++++ .../test_gemini_embedding_file_block_wire.py | 294 +++++++++++++++ ..._vertex_embedding_batch_file_block_wire.py | 109 ++++++ .../test_vertex_embedding_file_block_wire.py | 118 ++++++ tests/unit/caching/test_caching.py | 2 + tests/unit/caching/test_caching_handler.py | 86 +++++ .../test_vertex_ai_files_transformation.py | 27 ++ .../test_batch_embed_content_handler.py | 39 ++ ...test_batch_embed_content_transformation.py | 251 ++++++++++++- .../vertex_ai/test_gemini_batch_embeddings.py | 164 +++++++++ 22 files changed, 2048 insertions(+), 150 deletions(-) create mode 100644 tests/integration/caching/test_embedding_file_block_cache.py create mode 100644 tests/integration/providers/test_gemini_chat_video_metadata_wire.py create mode 100644 tests/integration/providers/test_gemini_embedding_file_block_chaos.py create mode 100644 tests/integration/providers/test_gemini_embedding_file_block_wire.py create mode 100644 tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py create mode 100644 tests/integration/providers/test_vertex_embedding_file_block_wire.py create mode 100644 tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 85ef5a93937..b928f33e961 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -71,15 +71,25 @@ class CacheMode(str, Enum): #### LiteLLM.Completion / Embedding Cache #### +def _is_conversation_item(item: object) -> bool: + if isinstance(item, BaseModel): + return True + if not isinstance(item, Mapping): + return False + block: Final = cast(Mapping[str, object], item) # cast-ok: isinstance leaves the key and value types unknown + return block.get("type") != "file" + + def _request_message_count(kwargs: Mapping[str, object]) -> int: - """Chat and Messages API `messages`, else Responses API `input` items; embedding `input` strings count as none""" + """Chat and Messages API `messages`, else Responses API `input` items; embedding strings and file blocks count as none""" messages: Final = kwargs.get("messages") if isinstance(messages, list): return len(messages) input_items: Final = kwargs.get("input") if not isinstance(input_items, list): return 0 - return sum(1 for item in input_items if isinstance(item, (Mapping, BaseModel))) + items: Final = cast(list[object], input_items) # cast-ok: isinstance leaves the element type unknown + return sum(1 for item in items if _is_conversation_item(item)) class Cache: diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index f9215864d1e..9ffc95c6260 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -21,7 +21,7 @@ import time from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Generator, Mapping from typing import TYPE_CHECKING, Any, Final, Optional, TypeVar -from pydantic import ConfigDict, ValidationError +from pydantic import ConfigDict, SkipValidation, ValidationError import litellm from litellm._internal_context import post_response_phase @@ -39,7 +39,7 @@ from litellm.litellm_core_utils.logging_utils import ( from litellm.types.caching import CACHED_STREAM_EVENTS_KEY, EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding from litellm.types.integrations.custom_logger import converted_stream_requested from litellm.types.llms.base import LiteLLMBaseModel -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import ChatCompletionFileObject, ResponsesAPIResponse from litellm.types.rerank import RerankResponse from litellm.types.utils import ( CachingDetails, @@ -71,6 +71,8 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +EmbeddingCacheInputElement = str | list[int] | ChatCompletionFileObject + class CachingHandlerResponse(LiteLLMBaseModel): """ @@ -82,7 +84,7 @@ class CachingHandlerResponse(LiteLLMBaseModel): cached_result: object | None = None final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call - embedding_uncached_input: list[str | list[int]] | None = None + embedding_uncached_input: SkipValidation[list[EmbeddingCacheInputElement]] | None = None in_memory_cache_obj: Final = InMemoryCache() @@ -508,7 +510,7 @@ class LLMCachingHandler: _sync_get_cache = sync_get_cache - def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[str]: + def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[EmbeddingCacheInputElement]: """ Handles the input of kwargs['input'] being a list or a string """ @@ -517,7 +519,11 @@ class LLMCachingHandler: elif isinstance(kwargs["input"], list): return kwargs["input"] else: - raise ValueError("input must be a string or a list") + raise litellm.BadRequestError( + message="input must be a string or a list of strings and content blocks", + model=str(kwargs.get("model")), + llm_provider=str(kwargs.get("custom_llm_provider")), + ) def _extract_model_from_cached_results(self, non_null_list: list[tuple[int, CachedEmbedding]]) -> str | None: """ @@ -851,10 +857,7 @@ class LLMCachingHandler: self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs) cached_result: object | None = None if call_type == CallTypes.aembedding.value: - if isinstance(new_kwargs["input"], str): - new_kwargs["input"] = [new_kwargs["input"]] - elif not isinstance(new_kwargs["input"], list): - raise ValueError("input must be a string or a list") + new_kwargs["input"] = self.handle_kwargs_input_list_or_str(new_kwargs) tasks: Final[list[Awaitable[object]]] = [] for idx, i in enumerate(new_kwargs["input"]): preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i}) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 8c77413bace..6b759065706 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -25,6 +25,22 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages +GEMINI_VIDEO_METADATA_KEYS: Final = MappingProxyType( + {"fps": "fps", "start_offset": "startOffset", "end_offset": "endOffset"} +) + + +GEMINI_FILES_API_URI_PREFIX: Final = "https://generativelanguage.googleapis.com/v1beta/files/" + + +def gemini_video_metadata_from_openai(video_metadata: Mapping[str, object]) -> dict[str, object]: + return { + gemini_key: video_metadata[openai_key] + for openai_key, gemini_key in GEMINI_VIDEO_METADATA_KEYS.items() + if openai_key in video_metadata + } + + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset( { "audio", diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index da5b974f186..09ec2ba6ee4 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -67,7 +67,7 @@ from litellm.types.llms.openai import ( OpenAIFilesPurpose, PathLike, ) -from litellm.types.llms.vertex_ai import GcsBucketResponse, GeminiEmbeddingInput +from litellm.types.llms.vertex_ai import GcsBucketResponse, GeminiEmbeddingElement, GeminiEmbeddingInput from litellm.types.utils import ( Embedding, EmbeddingResponse, @@ -555,18 +555,27 @@ def _is_responses_batch_entry(openai_entry: Mapping[str, object]) -> bool: return path == "responses" or path.endswith("/responses") +def _own_embedding_input( + element: GeminiEmbeddingElement | list[str] | list[GeminiEmbeddingElement], +) -> GeminiEmbeddingInput: + if isinstance(element, (str, list)): + return element + file_block_alone: Final[list[GeminiEmbeddingElement]] = [element] + return file_block_alone + + def _openai_embedding_input_elements( embedding_input: GeminiEmbeddingInput, -) -> tuple[str | list[str], ...]: +) -> tuple[GeminiEmbeddingInput, ...]: """ Split an OpenAI `input` into the elements that each get their own embedding. - A string is one embedding, a flat array is one embedding per element, and a nested - array is one combined embedding per inner array, matching the online - `batchEmbedContents` path. + A string or a file content block is one embedding, a flat array is one embedding + per element, and a nested array is one combined embedding per inner array, + matching the online `batchEmbedContents` path. """ if isinstance(embedding_input, list): - return tuple(embedding_input) + return tuple(_own_embedding_input(element) for element in embedding_input) return (embedding_input,) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 8749b81d7f1..15e5e440b91 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -57,7 +57,9 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import GenericImageParsingChunk, LlmProviders from ..common_utils import ( + GEMINI_FILES_API_URI_PREFIX, check_text_in_content, + gemini_video_metadata_from_openai, get_supports_response_schema, get_supports_system_message, ) @@ -69,7 +71,6 @@ _GCS_METADATA_VERTEX_BASE: object | None = None # Shared sync client for GCS JSON API metadata reads so proxy/SSL settings # from litellm's HTTP stack apply (see Greptile review on PR #27278). _GCS_METADATA_HTTP_HANDLER: HTTPHandler | None = None -GEMINI_FILES_API_URI_PREFIX: Final = "https://generativelanguage.googleapis.com/v1beta/files/" _GEMINI_MIME_TYPE_ALIASES: Final[dict[str, str]] = { "image/jpg": "image/jpeg", } @@ -197,13 +198,7 @@ def _apply_gemini_metadata( part_dict["media_resolution"] = media_resolution_enum if video_metadata is not None: - gemini_video_metadata: Final = {} - if "fps" in video_metadata: - gemini_video_metadata["fps"] = video_metadata["fps"] - if "start_offset" in video_metadata: - gemini_video_metadata["startOffset"] = video_metadata["start_offset"] - if "end_offset" in video_metadata: - gemini_video_metadata["endOffset"] = video_metadata["end_offset"] + gemini_video_metadata: Final = gemini_video_metadata_from_openai(video_metadata) if gemini_video_metadata: part_dict["video_metadata"] = gemini_video_metadata diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index d15cf89cf47..0a800497de1 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -22,6 +22,8 @@ from litellm.types.utils import EmbeddingResponse from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM from .batch_embed_content_transformation import ( + file_reference_name, + flatten_media_sources, is_file_reference, process_embed_content_response, process_response, @@ -38,11 +40,8 @@ class GoogleBatchEmbeddings(VertexLLM): def _flatten_and_detect_file_refs( input: GeminiEmbeddingInput, ) -> tuple[list[str], bool]: - """Flatten nested input lists and detect file references.""" - input_list: Final = [input] if isinstance(input, str) else input - flat_elements: Final = [ - e for item in input_list for e in (item if isinstance(item, list) else [item]) if isinstance(e, str) - ] + """Flatten nested input lists and file content blocks into their sources and detect file references.""" + flat_elements: Final = list(flatten_media_sources(input)) has_file_refs: Final = any(is_file_reference(e) for e in flat_elements) return flat_elements, has_file_refs @@ -68,7 +67,7 @@ class GoogleBatchEmbeddings(VertexLLM): for element in input_list: if isinstance(element, str) and is_file_reference(element): - url = f"https://generativelanguage.googleapis.com/v1beta/{element}" + url = f"https://generativelanguage.googleapis.com/v1beta/{file_reference_name(element)}" headers = {"x-goog-api-key": api_key} response = sync_handler.get(url=url, headers=headers) @@ -105,7 +104,7 @@ class GoogleBatchEmbeddings(VertexLLM): for element in input_list: if isinstance(element, str) and is_file_reference(element): - url = f"https://generativelanguage.googleapis.com/v1beta/{element}" + url = f"https://generativelanguage.googleapis.com/v1beta/{file_reference_name(element)}" headers = {"x-goog-api-key": api_key} response = await async_handler.get(url=url, headers=headers) @@ -139,6 +138,7 @@ class GoogleBatchEmbeddings(VertexLLM): timeout=300, client=None, extra_headers: dict | None = None, + drop_params: bool = False, ) -> EmbeddingResponse: _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -207,22 +207,26 @@ class GoogleBatchEmbeddings(VertexLLM): api_key=api_key, optional_params=optional_params, logging_obj=logging_obj, + drop_params=drop_params, ) ### TRANSFORMATION (sync path) ### request_data: VertexAIBatchEmbeddingsRequestBody | dict[str, object] + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if use_embed_content: resolved_files = {} if api_key: - resolved_files = self._resolve_file_references(input=input, api_key=api_key, sync_handler=sync_handler) + resolved_files = self._resolve_file_references( + input=flat_elements, api_key=api_key, sync_handler=sync_handler + ) request_data = transform_openai_input_gemini_embed_content( input=input, model=model, optional_params=optional_params, resolved_files=resolved_files, + drop_params=drop_params, ) else: - flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if has_file_refs and not api_key: raise ValueError( "An API key is required to resolve Gemini file references (files/...). " @@ -238,6 +242,7 @@ class GoogleBatchEmbeddings(VertexLLM): model=model, optional_params=optional_params, resolved_files=resolved_files, + drop_params=drop_params, ) ## LOGGING @@ -294,6 +299,7 @@ class GoogleBatchEmbeddings(VertexLLM): api_key: str | None = None, optional_params: dict | None = None, logging_obj: "LiteLLMLoggingObj | None" = None, + drop_params: bool = False, ) -> EmbeddingResponse: if client is None: _params: Final = {} @@ -312,20 +318,21 @@ class GoogleBatchEmbeddings(VertexLLM): async_handler = client ### TRANSFORMATION (async path) ### + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if use_embed_content: resolved_files = {} if api_key: resolved_files = await self._async_resolve_file_references( - input=input, api_key=api_key, async_handler=async_handler + input=flat_elements, api_key=api_key, async_handler=async_handler ) data = transform_openai_input_gemini_embed_content( input=input, model=model, optional_params=optional_params or {}, resolved_files=resolved_files, + drop_params=drop_params, ) else: - flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if has_file_refs and not api_key: raise ValueError( "An API key is required to resolve Gemini file references (files/...). " @@ -341,6 +348,7 @@ class GoogleBatchEmbeddings(VertexLLM): model=model, optional_params=optional_params or {}, resolved_files=resolved_files, + drop_params=drop_params, ) ## LOGGING diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 9ee2b71d30c..96a785c06c9 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -4,22 +4,27 @@ Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batc Why separate file? Make it easy to see how transformation works """ -from collections.abc import Mapping, Sequence -from typing import Final +from collections.abc import Iterator, Mapping, Sequence +from typing import Annotated, Final, Literal, cast -from pydantic import TypeAdapter, ValidationError +from pydantic import ConfigDict, Field, TypeAdapter, ValidationError +from litellm.exceptions import BadRequestError +from litellm.llms.vertex_ai.common_utils import GEMINI_FILES_API_URI_PREFIX, gemini_video_metadata_from_openai +from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.llms.vertex_ai import ( BlobType, ContentType, EmbedContentRequest, FileDataType, + GeminiEmbeddingElement, GeminiEmbeddingInput, PartType, PromptTokensDetails, UsageMetadata, VertexAIBatchEmbeddingsRequestBody, VertexAIBatchEmbeddingsResponseObject, + VideoMetadataType, ) from litellm.types.utils import ( Embedding, @@ -40,9 +45,16 @@ SUPPORTED_EMBEDDING_MIME_TYPES: Final = { } +_GEMINI_API_V1BETA: Final = GEMINI_FILES_API_URI_PREFIX.removesuffix("files/") + + def is_file_reference(s: str) -> bool: - """Check if string is a Gemini file reference (files/...).""" - return isinstance(s, str) and s.startswith("files/") + """A `files/...` name or the Files API URI that `/v1/files` returns as the file id.""" + return isinstance(s, str) and (s.startswith("files/") or s.startswith(GEMINI_FILES_API_URI_PREFIX)) + + +def file_reference_name(reference: str) -> str: + return reference.removeprefix(_GEMINI_API_V1BETA) _is_file_reference = is_file_reference @@ -88,12 +100,13 @@ def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: ) -def _parse_data_url(data_url: str) -> tuple[str, str]: +def _parse_data_url(data_url: str, mime_type_override: str | None = None) -> tuple[str, str]: """ Parse a data URL to extract the media type and base64 data. Args: data_url: Data URL in format: data:image/jpeg;base64,/9j/4AAQ... + mime_type_override: An explicit media type that replaces the declared one, skipping the allowlist Returns: tuple: (media_type, base64_data) @@ -110,50 +123,225 @@ def _parse_data_url(data_url: str) -> tuple[str, str]: raise ValueError(f"Invalid data URL format (missing comma): {data_url[:50]}...") metadata, base64_data = data_url.split(",", 1) + declared_media_type: Final = metadata[5:].split(";")[0] - metadata = metadata[5:] + if mime_type_override is not None: + return mime_type_override, base64_data - if ";" in metadata: - media_type = metadata.split(";")[0] - else: - media_type = metadata - - if media_type not in SUPPORTED_EMBEDDING_MIME_TYPES: + if declared_media_type not in SUPPORTED_EMBEDDING_MIME_TYPES: raise ValueError( - f"Unsupported MIME type for embedding: {media_type}. " + f"Unsupported MIME type for embedding: {declared_media_type}. " f"Supported types: {', '.join(sorted(SUPPORTED_EMBEDDING_MIME_TYPES))}" ) - return media_type, base64_data + return declared_media_type, base64_data + + +def _is_data_url(s: str) -> bool: + return s.startswith("data:") and ";base64," in s + + +class _EmbeddingVideoMetadata(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + + fps: float | None = None + start_offset: str | None = None + end_offset: str | None = None + + +class _EmbeddingFile(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid") + + file_id: str | None = None + file_data: str | None = None + filename: str | None = None + format: Annotated[str, Field(min_length=1)] | None = None + video_metadata: _EmbeddingVideoMetadata | None = None + + +class _EmbeddingFileBlock(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid") + + type: Literal["file"] + file: _EmbeddingFile + + +_file_block_adapter: Final = TypeAdapter(_EmbeddingFileBlock) +_video_metadata_adapter: Final = TypeAdapter(VideoMetadataType) +_input_shape_adapter: Final[TypeAdapter[str | list[object]]] = TypeAdapter(str | list[object]) +_mapping_adapter: Final = TypeAdapter(dict[str, object]) +_BLOCK_FIELDS: Final = frozenset(_EmbeddingFileBlock.model_fields) +_FILE_FIELDS: Final = frozenset(_EmbeddingFile.model_fields) +_VIDEO_METADATA_FIELDS: Final = frozenset(_EmbeddingVideoMetadata.model_fields) +_FILE_SOURCE_FORMS: Final = "a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI" + + +def _invalid_input(message: str) -> BadRequestError: + return BadRequestError(message=message, model=None, llm_provider="gemini") + + +def _validation_error_summary(error: ValidationError) -> str: + return "; ".join(f"{'.'.join(str(loc) for loc in detail['loc'])}: {detail['msg']}" for detail in error.errors()) + + +def _as_mapping(value: object) -> dict[str, object] | None: + try: + return _mapping_adapter.validate_python(value) + except ValidationError: + return None + + +def _only_fields(mapping: Mapping[str, object], fields: frozenset[str]) -> dict[str, object]: + return {key: value for key, value in mapping.items() if key in fields} + + +def _dropping_unsupported_keys(block: Mapping[str, object]) -> dict[str, object]: + """What `drop_params` keeps of a file content block: the keys this surface understands, at every level.""" + kept_block: Final = _only_fields(block, _BLOCK_FIELDS) + file: Final = _as_mapping(block.get("file")) + if file is None: + return kept_block + kept_file: Final = _only_fields(file, _FILE_FIELDS) + video_metadata: Final = _as_mapping(file.get("video_metadata")) + if video_metadata is None: + return {**kept_block, "file": kept_file} + return { + **kept_block, + "file": {**kept_file, "video_metadata": _only_fields(video_metadata, _VIDEO_METADATA_FIELDS)}, + } + + +def _parse_file_block(element: object, drop_params: bool) -> _EmbeddingFileBlock: + if not isinstance(element, Mapping): + raise _invalid_input( + f"Embedding input elements must be strings or file content blocks, got {type(element).__name__}" + ) + block: Final = cast(Mapping[str, object], element) # cast-ok: isinstance leaves the key and value types unknown + try: + return _file_block_adapter.validate_python(_dropping_unsupported_keys(block) if drop_params else block) + except ValidationError as error: + raise _invalid_input( + f"Invalid file content block in embedding input: {_validation_error_summary(error)}" + ) from error + + +def _file_block_source(block: _EmbeddingFileBlock) -> str: + match (block.file.file_id, block.file.file_data): + case (str() as file_id, None): + return file_id + case (None, str() as file_data): + return file_data + case (None, None): + raise _invalid_input("A file content block in embedding input needs file.file_id or file.file_data") + case _: + raise _invalid_input( + "A file content block in embedding input takes file.file_id or file.file_data, not both" + ) + + +def _gemini_video_metadata(video_metadata: _EmbeddingVideoMetadata) -> VideoMetadataType: + return _video_metadata_adapter.validate_python( + gemini_video_metadata_from_openai(video_metadata.model_dump(exclude_none=True)) + ) + + +def _source_mime_type( + source: str, + mime_type_override: str | None, + resolved_files: Mapping[str, Mapping[str, str]], +) -> str | None: + if mime_type_override is not None: + return mime_type_override + if _is_data_url(source): + try: + return _parse_data_url(source)[0] + except ValueError: + return None + if _is_gcs_url(source): + try: + return _infer_mime_type_from_gcs_url(source) + except ValueError: + return None + if is_file_reference(source): + file_info: Final = resolved_files.get(source) + return None if file_info is None else file_info.get("mime_type") + return None + + +def _media_part( + source: str, + mime_type_override: str | None, + resolved_files: Mapping[str, Mapping[str, str]], +) -> PartType: + if _is_data_url(source): + mime_type, base64_data = _parse_data_url(source, mime_type_override) + return PartType(inline_data=BlobType(mime_type=mime_type, data=base64_data)) + if _is_gcs_url(source): + gcs_mime_type: Final = mime_type_override or _infer_mime_type_from_gcs_url(source) + return PartType(file_data=FileDataType(mime_type=gcs_mime_type, file_uri=source)) + if is_file_reference(source): + file_info: Final = resolved_files.get(source) + if file_info is None: + raise _invalid_input( + f"File reference {source!r} could not be resolved: " + "Gemini Files API references are only supported through the gemini/ provider" + ) + return PartType( + file_data=FileDataType(mime_type=mime_type_override or file_info["mime_type"], file_uri=file_info["uri"]) + ) + raise _invalid_input(f"A file content block source must be {_FILE_SOURCE_FORMS}, got {source[:50]!r}") + + +def _top_level_elements( + input: GeminiEmbeddingInput, +) -> Sequence[GeminiEmbeddingElement | list[str] | list[GeminiEmbeddingElement]]: + try: + _input_shape_adapter.validate_python(input) + except ValidationError as error: + raise _invalid_input( + f"Embedding input must be a string or a list of strings and file content blocks, got {type(input).__name__}" + ) from error + return [input] if isinstance(input, str) else input + + +def _elements(input: GeminiEmbeddingInput) -> Iterator[GeminiEmbeddingElement]: + for element in _top_level_elements(input): + if isinstance(element, list): + yield from element + else: + yield element + + +def _element_source(element: GeminiEmbeddingElement) -> str: + if isinstance(element, str): + return element + return _file_block_source(_parse_file_block(element, drop_params=True)) + + +def flatten_media_sources(input: GeminiEmbeddingInput) -> tuple[str, ...]: + """Every string element plus every file content block's source, in input order.""" + return tuple(_element_source(element) for element in _elements(input)) def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool: """ Check if the input contains multimodal data (data URIs, file references, - GCS URLs, or nested lists for combined embeddings). + GCS URLs, file content blocks, or nested lists for combined embeddings). Args: - input: GeminiEmbeddingInput — str, List[str], or List[List[str]] for combined embeddings + input: GeminiEmbeddingInput — str, List[element], or List[List[element]] for combined embeddings Returns: - bool: True if any element is multimodal or a nested list + bool: True if any element is multimodal """ - if isinstance(input, str): - return _is_multimodal_element(input) - - for element in input: - if isinstance(element, list): - if any(_is_multimodal_element(sub) for sub in element if isinstance(sub, str)): - return True - elif isinstance(element, str) and _is_multimodal_element(element): - return True - - return False + return any(_is_multimodal_element(element) for element in _elements(input)) -def _is_multimodal_element(element: str) -> bool: - """Check if a single string element is multimodal.""" - if element.startswith("data:") and ";base64," in element: +def _is_multimodal_element(element: GeminiEmbeddingElement) -> bool: + """Check if a single element is multimodal.""" + if not isinstance(element, str): + return True + if _is_data_url(element): return True if is_file_reference(element): return True @@ -163,37 +351,29 @@ def _is_multimodal_element(element: str) -> bool: def _build_part_for_input( - element: str, - resolved_files: dict[str, dict[str, str]] | None = None, + element: GeminiEmbeddingElement, + resolved_files: Mapping[str, Mapping[str, str]] | None = None, + drop_params: bool = False, ) -> PartType: """ Build a single PartType for an input element, handling text, data URIs, - file references, and GCS URLs. + file references, GCS URLs, and file content blocks carrying a mime type + and video_metadata. """ - resolved_files = resolved_files or {} + files: Final = resolved_files or {} - if element.startswith("data:") and ";base64," in element: - mime_type, base64_data = _parse_data_url(element) - blob: Final[BlobType] = {"mime_type": mime_type, "data": base64_data} - return PartType(inline_data=blob) - elif _is_gcs_url(element): - mime_type = _infer_mime_type_from_gcs_url(element) - file_data: Final[FileDataType] = { - "mime_type": mime_type, - "file_uri": element, - } - return PartType(file_data=file_data) - elif is_file_reference(element): - if element not in resolved_files: - raise ValueError(f"File reference {element} not resolved") - file_info: Final = resolved_files[element] - file_data_ref: Final[FileDataType] = { - "mime_type": file_info["mime_type"], - "file_uri": file_info["uri"], - } - return PartType(file_data=file_data_ref) - else: - return PartType(text=element) + if isinstance(element, str): + return _media_part(element, None, files) if _is_multimodal_element(element) else PartType(text=element) + + block: Final = _parse_file_block(element, drop_params) + part: Final = _media_part(_file_block_source(block), block.file.format, files) + if block.file.video_metadata is None: + return part + video_metadata: Final = _gemini_video_metadata(block.file.video_metadata) + if not video_metadata: + return part + part_with_metadata: Final[PartType] = {**part, "video_metadata": video_metadata} + return part_with_metadata _SUPPORTED_EMBED_PARAMS: Final = {"outputDimensionality", "taskType", "title"} @@ -214,6 +394,7 @@ def transform_openai_input_gemini_content( model: str, optional_params: dict, resolved_files: dict[str, dict[str, str]] | None = None, + drop_params: bool = False, ) -> VertexAIBatchEmbeddingsRequestBody: """ Transform OpenAI embedding input to Gemini batchEmbedContents format. @@ -234,19 +415,18 @@ def transform_openai_input_gemini_content( gemini_params: Final = _filter_embed_params(optional_params) - input_list: Final = [input] if isinstance(input, str) else input + input_list: Final = _top_level_elements(input) requests: Final[list[EmbedContentRequest]] = [] for element in input_list: if isinstance(element, list): if not element: raise ValueError("Nested input list must not be empty") - for sub in element: - if not isinstance(sub, str): - raise ValueError(f"Elements inside a nested input list must be strings, got {type(sub)}") - parts = [_build_part_for_input(sub, resolved_files=resolved_files) for sub in element] + parts = [ + _build_part_for_input(sub, resolved_files=resolved_files, drop_params=drop_params) for sub in element + ] else: - parts = [_build_part_for_input(element, resolved_files=resolved_files)] + parts = [_build_part_for_input(element, resolved_files=resolved_files, drop_params=drop_params)] request = EmbedContentRequest( model=gemini_model_name, content=ContentType(parts=parts), @@ -262,6 +442,7 @@ def transform_openai_input_gemini_embed_content( model: str, optional_params: dict, resolved_files: dict[str, dict[str, str]] | None = None, + drop_params: bool = False, ) -> dict: """ Transform OpenAI embedding input to Gemini embedContent format (multimodal). @@ -279,7 +460,7 @@ def transform_openai_input_gemini_embed_content( gemini_params: Final = _filter_embed_params(optional_params) - input_list: Final = [input] if isinstance(input, str) else input + input_list: Final = _top_level_elements(input) parts: Final[list[PartType]] = [] for element in input_list: @@ -288,9 +469,7 @@ def transform_openai_input_gemini_embed_content( "Nested (combined) embeddings are not supported on the embedContent path. " "Use the batchEmbedContents path or pass a flat list instead." ) - if not isinstance(element, str): - raise ValueError(f"Unsupported input type: {type(element)}") - parts.append(_build_part_for_input(element, resolved_files=resolved_files)) + parts.append(_build_part_for_input(element, resolved_files=resolved_files, drop_params=drop_params)) request_body: Final[dict] = { "content": ContentType(parts=parts), @@ -313,31 +492,18 @@ def _parse_usage_metadata(raw_usage_metadata: object) -> UsageMetadata | None: return None -def _flatten_input(input: GeminiEmbeddingInput) -> tuple[str, ...]: - if isinstance(input, str): - return (input,) - return tuple(sub for element in input for sub in (element if isinstance(element, list) else [element])) +def _flatten_input(input: GeminiEmbeddingInput) -> tuple[GeminiEmbeddingElement, ...]: + return tuple(_elements(input)) def _is_image_element( - element: str, + element: GeminiEmbeddingElement, resolved_files: Mapping[str, Mapping[str, str]], ) -> bool: - if element.startswith("data:") and ";base64," in element: - try: - mime_type, _ = _parse_data_url(element) - except ValueError: - return False - return mime_type in _IMAGE_MIME_TYPES - if _is_gcs_url(element): - try: - return _infer_mime_type_from_gcs_url(element) in _IMAGE_MIME_TYPES - except ValueError: - return False - if is_file_reference(element): - file_info: Final = resolved_files.get(element) - return file_info is not None and file_info.get("mime_type") in _IMAGE_MIME_TYPES - return False + if isinstance(element, str): + return _source_mime_type(element, None, resolved_files) in _IMAGE_MIME_TYPES + block: Final = _parse_file_block(element, drop_params=True) + return _source_mime_type(_file_block_source(block), block.file.format, resolved_files) in _IMAGE_MIME_TYPES def _is_image_only_input( diff --git a/litellm/main.py b/litellm/main.py index 25c73ea281e..b02ccca9ecb 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -6379,6 +6379,10 @@ def embedding( """ azure: Final = kwargs.get("azure", None) client: Final = kwargs.pop("client", None) + drop_params_kwarg: Final = ( + cast(object, kwargs["drop_params"]) if "drop_params" in kwargs else None # cast-ok: untyped request kwargs + ) + drop_unsupported_params: Final = litellm.drop_params is True or normalize_drop_params(drop_params_kwarg) is True shared_session: Final = kwargs.get("shared_session", None) max_retries: Final = kwargs.get("max_retries", None) litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") @@ -6859,6 +6863,7 @@ def embedding( api_base=api_base, client=client, extra_headers=headers, + drop_params=drop_unsupported_params, ) elif custom_llm_provider == "vertex_ai": @@ -6911,6 +6916,7 @@ def embedding( api_base=api_base, client=client, extra_headers=headers, + drop_params=drop_unsupported_params, ) elif ( "image" in optional_params diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index ce51e46ef15..116fdda0b5d 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -2,16 +2,20 @@ from enum import Enum from typing import Any, Final, Literal, Protocol from typing_extensions import ( + ReadOnly, Required, TypedDict, ) -from litellm.types.llms.openai import EmbeddingInput +from litellm.types.llms.openai import ChatCompletionFileObject, EmbeddingInput # Gemini supports nested-list inputs (e.g. [["text", "image"]]) as an explicit # opt-in for combined embeddings — a provider-specific extension of the # OpenAI-faithful EmbeddingInput shape. -GeminiEmbeddingInput = EmbeddingInput | list[list[str]] +GeminiEmbeddingElement = str | ChatCompletionFileObject +GeminiEmbeddingInput = ( + EmbeddingInput | list[GeminiEmbeddingElement] | list[list[str]] | list[list[GeminiEmbeddingElement]] +) class FunctionResponse(TypedDict, total=False): @@ -46,6 +50,12 @@ class FunctionResponsePartType(TypedDict, total=False): file_data: FileDataType +class VideoMetadataType(TypedDict, total=False): + fps: ReadOnly[float] + startOffset: ReadOnly[str] + endOffset: ReadOnly[str] + + class PartType(TypedDict, total=False): text: str inline_data: BlobType @@ -55,6 +65,7 @@ class PartType(TypedDict, total=False): thought: bool thoughtSignature: str media_resolution: Literal["low", "medium", "high"] + video_metadata: ReadOnly[VideoMetadataType] class HttpxFunctionCall(TypedDict, total=False): diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index d118da03daf..e6a166e8457 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -3,13 +3,13 @@ from __future__ import annotations import ssl import threading import time -from collections.abc import Callable, Generator, Mapping +from collections.abc import Callable, Generator, Iterator, Mapping from contextlib import contextmanager from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from queue import SimpleQueue from types import MappingProxyType -from typing import Final +from typing import BinaryIO, Final @dataclass(frozen=True, slots=True) @@ -20,6 +20,37 @@ class Request: body: bytes +class AbortedBody(Exception): + """The client closed the connection before the body it announced was complete.""" + + +def exactly(stream: BinaryIO, size: int) -> bytes: + data: Final = stream.read(size) + if len(data) < size: + raise AbortedBody + return data + + +def chunked_body(stream: BinaryIO) -> Iterator[bytes]: + while True: + size_line: Final = stream.readline() + if not size_line: + raise AbortedBody + size: Final = int(size_line.split(b";")[0].strip(), 16) + if size == 0: + while stream.readline().strip(): + pass + return + yield exactly(stream, size) + stream.readline() + + +def read_body(headers: Mapping[str, str], stream: BinaryIO) -> bytes: + if headers.get("transfer-encoding", "").lower() == "chunked": + return b"".join(chunked_body(stream)) + return exactly(stream, int(headers.get("content-length", "0"))) + + @dataclass(frozen=True, slots=True) class Reply: status: int = 200 @@ -78,12 +109,14 @@ def wire_server( connected.put(f"{self.client_address[0]}:{self.client_address[1]}") def respond(self) -> None: - request: Final = Request( - self.command, - self.path, - {name.lower(): value for name, value in self.headers.items()}, - self.rfile.read(int(self.headers.get("content-length", "0"))), - ) + headers: Final = {name.lower(): value for name, value in self.headers.items()} + try: + body: Final = read_body(headers, self.rfile) + except AbortedBody: + self.close_connection = True + disconnected.put(self.path) + return + request: Final = Request(self.command, self.path, headers, body) received.put(request) try: reply = respond(request) diff --git a/tests/integration/caching/test_embedding_file_block_cache.py b/tests/integration/caching/test_embedding_file_block_cache.py new file mode 100644 index 00000000000..9f753e996e8 --- /dev/null +++ b/tests/integration/caching/test_embedding_file_block_cache.py @@ -0,0 +1,216 @@ +from __future__ import annotations + +import asyncio +import json +import threading +import zlib +from collections.abc import Mapping +from hashlib import sha256 +from typing import Final + +import pytest +from openai import AsyncOpenAI, OpenAI +from openai.types import CreateEmbeddingResponse +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +WIRE_METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +SETTLE_SECONDS: Final = 15 +SPEND_SQL: Final = ( + 'SELECT request_id, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE api_key = %s ORDER BY "startTime", request_id' +) + + +def clip(name: str) -> str: + return f"gs://scripted-bucket/cache/{name}.mp4" + + +def block(name: str, metadata: Mapping[str, JsonValue] | None = METADATA) -> dict[str, JsonValue]: + file: Final[dict[str, JsonValue]] = {"file_id": clip(name)} + return {"type": "file", "file": file if metadata is None else {**file, "video_metadata": dict(metadata)}} + + +def gcs_part(name: str, metadata: Mapping[str, JsonValue] | None = WIRE_METADATA) -> dict[str, JsonValue]: + part: Final[dict[str, JsonValue]] = {"file_data": {"mime_type": "video/mp4", "file_uri": clip(name)}} + return part if metadata is None else {**part, "video_metadata": dict(metadata)} + + +def list_value(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def source_of(part: JsonValue) -> str: + item: Final = object_value(part) + if "text" in item: + return string_value(item["text"]) + return string_value(object_value(item["file_data"])["file_uri"]) + + +def vector(parts: list[JsonValue]) -> list[float]: + return [zlib.crc32("|".join(source_of(part) for part in parts).encode()) / 2**32, 0.5] + + +def request_parts(request: Request) -> list[list[JsonValue]]: + body: Final = object_value(json.loads(request.body)) + return [list_value(object_value(object_value(item)["content"])["parts"]) for item in list_value(body["requests"])] + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + embeddings: Final = [{"values": vector(parts)} for parts in request_parts(request)] + return Reply(body=json.dumps({"embeddings": embeddings}).encode()) + + +def wire_parts(wire: Wire) -> list[list[JsonValue]]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + return request_parts(received[0]) + + +def embeddings_of(answer: Mapping[str, JsonValue]) -> list[JsonValue]: + data: Final = [object_value(item) for item in list_value(answer["data"])] + assert [item["index"] for item in data] == list(range(len(data))), answer + return [item["embedding"] for item in data] + + +def gemini_model(scenario: Scenario, url: str) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url) + + +def embed_body(model: str, elements: JsonValue) -> dict[str, JsonValue]: + return {"model": model, "input": elements} + + +def spend_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)) + + +def served_from_cache( + gateway: Gateway, wire: Wire, body: Mapping[str, JsonValue], key: str +) -> dict[str, JsonValue] | None: + answer: Final = gateway.post("/v1/embeddings", body, key=key) + return None if wire.drain() else answer + + +def await_hit(gateway: Gateway, wire: Wire, body: Mapping[str, JsonValue], key: str) -> dict[str, JsonValue]: + served: Final = eventually( + lambda: served_from_cache(gateway, wire, body, key), lambda answer: answer is not None, SETTLE_SECONDS + ) + assert served is not None + return served + + +def proxy_root(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def test_block_embedding_is_served_from_the_cache_with_a_zero_spend_row(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("one")]) + first: Final = gateway.request("POST", "/v1/embeddings", body, key=key) + assert first.status_code == 200, first.text + first_id: Final = first.headers["x-litellm-call-id"] + assert wire_parts(wire) == [[gcs_part("one")]] + served: Final = await_hit(gateway, wire, body, key) + assert embeddings_of(served) == embeddings_of(object_value(first.json())) == [vector([gcs_part("one")])] + rows: Final = eventually( + lambda: spend_rows(key), lambda found: any(row["cache_hit"] == "True" for row in found), seconds=70 + ) + by_id: Final = {string_value(row["request_id"]): row for row in rows} + assert by_id[first_id]["cache_hit"] != "True", rows + hits: Final = [row for row in rows if row["cache_hit"] == "True"] + assert hits and all(float(str(row["spend"])) == 0 for row in hits), rows + + +def test_cached_block_is_not_resent_when_a_new_text_joins_it(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + cached: Final = embed_body(model, [block("one")]) + gateway.post("/v1/embeddings", cached, key=key) + assert wire_parts(wire) == [[gcs_part("one")]] + await_hit(gateway, wire, cached, key) + mixed: Final = gateway.post("/v1/embeddings", embed_body(model, [block("one"), "a red bus"]), key=key) + assert wire_parts(wire) == [[TEXT_PART]] + assert embeddings_of(mixed) == [vector([gcs_part("one")]), vector([TEXT_PART])] + + +@pytest.mark.parametrize("names", (("a", "b", "c", "d", "e"), ("same", "same")), ids=("five-distinct", "two-identical")) +def test_block_lists_are_cached_whole(gateway: Gateway, names: tuple[str, ...]) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block(name) for name in names]) + first: Final = gateway.post("/v1/embeddings", body, key=key) + assert wire_parts(wire) == [[gcs_part(name)] for name in names] + assert embeddings_of(first) == [vector([gcs_part(name)]) for name in names] + assert embeddings_of(await_hit(gateway, wire, body, key)) == embeddings_of(first) + + +def test_provider_failure_is_not_cached(gateway: Gateway) -> None: + healed: Final = threading.Event() + + def respond(request: Request) -> Reply: + if healed.is_set(): + return embed_peer(request) + return Reply(status=500, body=b'{"error": {"message": "scripted outage"}}') + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("flaky")]) + failed: Final = gateway.request("POST", "/v1/embeddings", body, key=key) + assert failed.status_code >= 500, failed.text + assert wire.drain(), "the failing call never reached the provider" + healed.set() + recovered: Final = gateway.post("/v1/embeddings", body, key=key) + assert wire_parts(wire) == [[gcs_part("flaky")]] + assert embeddings_of(await_hit(gateway, wire, body, key)) == embeddings_of(recovered) + + +def test_video_metadata_is_part_of_the_cache_key(gateway: Gateway) -> None: + two_seconds: Final[dict[str, JsonValue]] = {**METADATA, "end_offset": "2s"} + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + settled: Final = embed_body(model, [block("keyed")]) + gateway.post("/v1/embeddings", settled, key=key) + assert wire_parts(wire) == [[gcs_part("keyed")]] + await_hit(gateway, wire, settled, key) + variants: Final = ( + (block("keyed", two_seconds), gcs_part("keyed", {**WIRE_METADATA, "endOffset": "2s"})), + (block("keyed", None), gcs_part("keyed", None)), + ) + for variant, part in variants: + gateway.post("/v1/embeddings", embed_body(model, [variant]), key=key) + assert wire_parts(wire) == [[part]] + + +def test_openai_sdk_clients_fill_and_hit_the_cache(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("sdk")]) + with OpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=key, max_retries=0) as client: + filled: Final = client.post("/embeddings", body=body, cast_to=CreateEmbeddingResponse) + assert wire_parts(wire) == [[gcs_part("sdk")]] + await_hit(gateway, wire, body, key) + + async def drive() -> CreateEmbeddingResponse: + async with AsyncOpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=key, max_retries=0) as client: + return await client.post("/embeddings", body=body, cast_to=CreateEmbeddingResponse) + + served: Final = asyncio.run(drive()) + assert served.data[0].embedding == filled.data[0].embedding == vector([gcs_part("sdk")]) + assert wire.drain() == (), "the cached answer reached the provider again" diff --git a/tests/integration/providers/test_gemini_chat_video_metadata_wire.py b/tests/integration/providers/test_gemini_chat_video_metadata_wire.py new file mode 100644 index 00000000000..32914341483 --- /dev/null +++ b/tests/integration/providers/test_gemini_chat_video_metadata_wire.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import base64 +import json +from collections.abc import Callable +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, object_value +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-3.7-flash" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +MODEL_PATH: Final = f"/v1/projects/{PROJECT}/locations/{LOCATION}/publishers/google/models/{BACKEND}" +CLIP: Final = base64.b64encode(b"\x00\x00\x00\x18ftypmp42" + bytes(24)).decode() +DATA_URI: Final = f"data:video/mp4;base64,{CLIP}" +PROMPT: Final = "describe the clip" +FILE_BLOCK: Final[dict[str, JsonValue]] = { + "type": "file", + "file": {"file_data": DATA_URI, "video_metadata": {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"}}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": PROMPT} +BARE_VIDEO_PART: Final[dict[str, JsonValue]] = {"inline_data": {"mime_type": "video/mp4", "data": CLIP}} +VIDEO_PART: Final[dict[str, JsonValue]] = { + **BARE_VIDEO_PART, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +REPLY: Final[dict[str, JsonValue]] = { + "candidates": [{"content": {"role": "model", "parts": [{"text": "a cat"}]}, "finishReason": "STOP"}], + "usageMetadata": {"promptTokenCount": 5, "candidatesTokenCount": 2, "totalTokenCount": 7}, +} + + +def chat_peer(path: str, authorization: str | None) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + assert authorization is None or request.headers["authorization"] == authorization + target, _, query = request.target.partition("?") + if target == f"{path}:streamGenerateContent": + assert "alt=sse" in query, request.target + return Reply(content_type="text/event-stream", chunks=(f"data: {json.dumps(REPLY)}\n\n".encode(),)) + assert target == f"{path}:generateContent", request.target + return Reply(body=json.dumps(REPLY).encode()) + + return respond + + +def wire_parts(wire: Wire) -> list[JsonValue]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + body: Final = object_value(json.loads(received[0].body)) + contents: Final = body["contents"] + assert isinstance(contents, list) and len(contents) == 1, body + parts: Final = object_value(contents[0])["parts"] + assert isinstance(parts, list), body + return parts + + +def gemini_model(scenario: Scenario, url: str) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url) + + +def vertex_model(gateway: Gateway, scenario: Scenario, url: str) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + ) + + +def chat(gateway: Gateway, model: str, stream: bool) -> httpx.Response: + body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": [{"type": "text", "text": PROMPT}, FILE_BLOCK]}], + "stream": stream, + } + return gateway.request("POST", "/v1/chat/completions", body) + + +def answered(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + assert "a cat" in response.text, response.text + + +@pytest.mark.parametrize("stream", (False, True), ids=("non-stream", "stream")) +def test_gemini_chat_file_block_video_metadata_reaches_the_wire(gateway: Gateway, stream: bool) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + answered(chat(gateway, model, stream)) + assert wire_parts(wire) == [TEXT_PART, VIDEO_PART] + + +@pytest.mark.parametrize("stream", (False, True), ids=("non-stream", "stream")) +def test_vertex_chat_file_block_video_metadata_reaches_the_wire(gateway: Gateway, stream: bool) -> None: + with wire_server(chat_peer(MODEL_PATH, "Bearer scripted-token")) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + answered(chat(gateway, model, stream)) + assert wire_parts(wire) == [TEXT_PART, VIDEO_PART] + + +def test_responses_input_file_video_reaches_the_wire_as_inline_data(gateway: Gateway) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": PROMPT}, + {"type": "input_file", "file_data": DATA_URI, "filename": "clip.mp4"}, + ], + } + ], + } + answered(gateway.request("POST", "/v1/responses", body)) + assert wire_parts(wire) == [TEXT_PART, BARE_VIDEO_PART] + + +def test_messages_document_video_reaches_the_wire_as_inline_data(gateway: Gateway) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "max_tokens": 32, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": PROMPT}, + {"type": "document", "source": {"type": "base64", "media_type": "video/mp4", "data": CLIP}}, + ], + } + ], + } + answered(gateway.request("POST", "/v1/messages", body)) + assert wire_parts(wire) == [TEXT_PART, BARE_VIDEO_PART] diff --git a/tests/integration/providers/test_gemini_embedding_file_block_chaos.py b/tests/integration/providers/test_gemini_embedding_file_block_chaos.py new file mode 100644 index 00000000000..cee0758b433 --- /dev/null +++ b/tests/integration/providers/test_gemini_embedding_file_block_chaos.py @@ -0,0 +1,223 @@ +from __future__ import annotations + +import json +import os +import signal +import threading +import uuid +import zlib +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import group_members, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +END_OFFSETS: Final = {"fast": "3s", "slow": "5s", "malformed": "3s", "dropped": "9s"} +PLAN: Final = tuple(enumerate(("fast",) * 16 + ("slow",) * 8 + ("malformed",) * 8 + ("dropped",) * 8)) +BURST: Final = 24 +SPEND_SQL: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s' + + +@dataclass(frozen=True, slots=True) +class Attempt: + kind: str + index: int + status: int + call_id: str + + +def source(kind: str, index: int) -> str: + return f"gs://scripted-bucket/chaos/{kind}-{index}.mp4" + + +def body(model: str, kind: str, index: int) -> dict[str, JsonValue]: + fps: Final[JsonValue] = "x" if kind == "malformed" else 1.0 + metadata: Final[dict[str, JsonValue]] = {"fps": fps, "start_offset": "0s", "end_offset": END_OFFSETS[kind]} + return { + "model": model, + "input": [{"type": "file", "file": {"file_id": source(kind, index), "video_metadata": metadata}}], + } + + +def first_part(request: Request) -> dict[str, JsonValue]: + requests: Final = object_value(json.loads(request.body))["requests"] + assert isinstance(requests, list) and len(requests) == 1, requests + parts: Final = object_value(object_value(requests[0])["content"])["parts"] + assert isinstance(parts, list) and len(parts) == 1, parts + return object_value(parts[0]) + + +def source_on_the_wire(request: Request) -> str: + return string_value(object_value(first_part(request)["file_data"])["file_uri"]) + + +def embed_reply(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + values: Final = [zlib.crc32(source_on_the_wire(request).encode()) / 2**32, 0.5] + return Reply(body=json.dumps({"embeddings": [{"values": values}]}).encode()) + + +def chaos_peer(healed: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + end_offset: Final = object_value(first_part(request)["video_metadata"])["endOffset"] + if end_offset == END_OFFSETS["dropped"] and not healed.is_set(): + return Reply(drop_connection=True) + reply: Final = embed_reply(request) + if end_offset == END_OFFSETS["slow"]: + return Reply(chunks=(reply.body[:8], reply.body[8:]), pause_between_chunks=1.5) + return reply + + return respond + + +def wire_sources(wire: Wire) -> tuple[str, ...]: + return tuple(source_on_the_wire(request) for request in wire.drain()) + + +def spend_row_count(call_id: str) -> int: + return len(read_rows(SPEND_SQL, (call_id,))) + + +def spend_row_counts(call_ids: tuple[str, ...]) -> tuple[int, ...]: + return tuple(spend_row_count(call_id) for call_id in call_ids) + + +def attempt(gateway: Gateway, model: str, index: int, kind: str) -> Attempt: + response: Final = gateway.request("POST", "/v1/embeddings", body(model, kind, index)) + return Attempt(kind, index, response.status_code, response.headers.get("x-litellm-call-id", "")) + + +def of_kind(attempts: tuple[Attempt, ...], *kinds: str) -> tuple[Attempt, ...]: + return tuple(item for item in attempts if item.kind in kinds) + + +@pytest.mark.timeout(240) +def test_mixed_burst_keeps_every_answer_and_spend_row_honest(gateway: Gateway) -> None: + healed: Final = threading.Event() + with wire_server(chaos_peer(healed)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=wire.url) + with ThreadPoolExecutor(max_workers=len(PLAN)) as pool: + attempts: Final = tuple(pool.map(lambda plan: attempt(gateway, model, plan[0], plan[1]), PLAN)) + served: Final = of_kind(attempts, "fast", "slow") + assert all(item.status == 200 for item in served), attempts + assert all(item.status == 400 for item in of_kind(attempts, "malformed")), attempts + assert all(item.status >= 500 for item in of_kind(attempts, "dropped")), attempts + reached: Final = wire_sources(wire) + assert sorted(item for item in reached if "/malformed-" not in item and "/dropped-" not in item) == sorted( + source(item.kind, item.index) for item in served + ), reached + assert not any("/malformed-" in item for item in reached), reached + assert {item for item in reached if "/dropped-" in item} == { + source(item.kind, item.index) for item in of_kind(attempts, "dropped") + }, reached + call_ids: Final = tuple(item.call_id for item in served) + assert all(call_ids), attempts + counts: Final = eventually( + lambda: spend_row_counts(call_ids), lambda found: all(count >= 1 for count in found), seconds=90 + ) + assert counts == (1,) * len(call_ids), counts + healed.set() + with ThreadPoolExecutor(max_workers=len(PLAN)) as pool: + resent: Final = tuple( + pool.map(lambda item: attempt(gateway, model, item.index, item.kind), of_kind(attempts, "dropped")) + ) + assert all(item.status == 200 for item in resent), resent + assert sorted(wire_sources(wire)) == sorted(source(item.kind, item.index) for item in resent) + + +def write_config(directory: Path, url: str, name: str) -> Path: + config: Final = directory / f"gemini_embeddings_{uuid.uuid4().hex}.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": f"gemini/{BACKEND}", + "api_key": "scripted-gemini-key", + "api_base": url, + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + }, + } + ) + ) + return config + + +def worker_pids(root_pid: int) -> frozenset[int]: + def is_worker(process: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(process.cmdline()) + except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess): + return False + + return frozenset(process.pid for process in group_members(root_pid) if is_worker(process)) + + +def outcome(candidate: Gateway, model: str, index: int) -> str: + try: + response: Final = candidate.request("POST", "/v1/embeddings", body(model, "fast", index)) + except httpx.TransportError as error: + return f"transport:{type(error).__name__}" + assert response.status_code == 200, response.text + return f"ok:{response.headers['x-litellm-call-id']}" + + +@pytest.mark.timeout(300) +def test_block_embeddings_survive_a_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + name: Final = f"gemini-embeddings-{uuid.uuid4().hex}" + with wire_server(embed_reply) as wire: + config: Final = write_config(tmp_path, wire.url, name) + with owned_proxy_process(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually(lambda: worker_pids(owned.process.pid), lambda pids: len(pids) == 2, seconds=30) + victim: Final = min(workers) + assert outcome(candidate, name, 0).startswith("ok:") + + def attempt_around_the_kill(index: int) -> str: + if index == 2: + os.kill(victim, signal.SIGKILL) + return outcome(candidate, name, index) + + with ThreadPoolExecutor(max_workers=BURST) as pool: + outcomes: Final = tuple(pool.map(attempt_around_the_kill, range(1, BURST + 1))) + assert outcomes.count("ok") == 0 and any(item.startswith("ok:") for item in outcomes), outcomes + assert all(item.startswith("ok:") or item.startswith("transport:") for item in outcomes), outcomes + respawned: Final = eventually( + lambda: worker_pids(owned.process.pid), + lambda pids: len(pids) == 2 and victim not in pids, + seconds=60, + ) + assert victim not in respawned, respawned + + def settled_burst() -> tuple[str, ...]: + with ThreadPoolExecutor(max_workers=BURST) as pool: + return tuple(pool.map(lambda index: outcome(candidate, name, BURST + 1 + index), range(BURST))) + + final: Final = eventually( + settled_burst, lambda values: all(v.startswith("ok:") for v in values), seconds=40 + ) + call_ids: Final = tuple(item.removeprefix("ok:") for item in (*outcomes, *final) if item.startswith("ok:")) + counts: Final = eventually( + lambda: spend_row_counts(call_ids), lambda found: all(count >= 1 for count in found), seconds=90 + ) + assert counts == (1,) * len(call_ids), counts + assert len(wire.drain()) >= len(call_ids) diff --git a/tests/integration/providers/test_gemini_embedding_file_block_wire.py b/tests/integration/providers/test_gemini_embedding_file_block_wire.py new file mode 100644 index 00000000000..d6439d27cd5 --- /dev/null +++ b/tests/integration/providers/test_gemini_embedding_file_block_wire.py @@ -0,0 +1,294 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import uuid +import zlib +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, object_value, string_value +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types import CreateEmbeddingResponse +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +CLIP: Final = base64.b64encode(b"\x00\x00\x00\x18ftypmp42" + bytes(24)).decode() +DATA_URI: Final = f"data:video/mp4;base64,{CLIP}" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +WIRE_METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"} +INLINE_DATA: Final[dict[str, JsonValue]] = {"mime_type": "video/mp4", "data": CLIP} +INLINE_PART: Final[dict[str, JsonValue]] = {"inline_data": INLINE_DATA, "video_metadata": WIRE_METADATA} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +ACCEPTED_FORMS: Final = "must be a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI" +LONG_OFFSET: Final = "9" * 5000 + "s" + + +def block(**file: JsonValue) -> dict[str, JsonValue]: + return {"type": "file", "file": file} + + +def gcs_part(uri: str = GCS_URI, metadata: dict[str, JsonValue] | None = WIRE_METADATA) -> dict[str, JsonValue]: + file_data: Final[dict[str, JsonValue]] = {"file_data": {"mime_type": "video/mp4", "file_uri": uri}} + return file_data if metadata is None else {**file_data, "video_metadata": metadata} + + +GCS_PART: Final = gcs_part() +NOISY_BLOCK: Final[dict[str, JsonValue]] = { + **block(file_id=GCS_URI, mime_type="video/mp4", video_metadata={**METADATA, "frame_rate": 2}), + "caption": "unknown block key", +} + + +def list_value(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def source_of(part: JsonValue) -> str: + item: Final = object_value(part) + if "text" in item: + return string_value(item["text"]) + if "file_data" in item: + return string_value(object_value(item["file_data"])["file_uri"]) + return string_value(object_value(item["inline_data"])["data"]) + + +def vector(parts: list[JsonValue]) -> list[float]: + return [zlib.crc32("|".join(source_of(part) for part in parts).encode()) / 2**32, 0.5] + + +def request_parts(request: Request) -> list[list[JsonValue]]: + body: Final = object_value(json.loads(request.body)) + return [list_value(object_value(object_value(item)["content"])["parts"]) for item in list_value(body["requests"])] + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + embeddings: Final = [{"values": vector(parts)} for parts in request_parts(request)] + return Reply(body=json.dumps({"embeddings": embeddings}).encode()) + + +def wire_parts(wire: Wire) -> list[list[JsonValue]]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + return request_parts(received[0]) + + +def embeddings(response: httpx.Response) -> list[JsonValue]: + assert response.status_code == 200, response.text + data: Final = [object_value(item) for item in list_value(object_value(response.json())["data"])] + assert [item["index"] for item in data] == list(range(len(data))), response.text + assert all(item["object"] == "embedding" for item in data), response.text + return [item["embedding"] for item in data] + + +def gemini_model(scenario: Scenario, url: str, **litellm_params: JsonValue) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url, **litellm_params) + + +def embed(gateway: Gateway, model: str, elements: JsonValue, **extra: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/embeddings", {"model": model, "input": elements, **extra}) + + +def proxy_root(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def write_config(directory: Path, url: str, name: str) -> Path: + config: Final = directory / f"gemini_embeddings_{uuid.uuid4().hex}.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": f"gemini/{BACKEND}", + "api_key": "scripted-gemini-key", + "api_base": url, + }, + } + ], + "litellm_settings": {"drop_params": True}, + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + }, + } + ) + ) + return config + + +def test_string_input_reaches_the_wire_as_a_text_part(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, ["a red bus"])) == [vector([TEXT_PART])] + assert wire_parts(wire) == [[TEXT_PART]] + + +@pytest.mark.parametrize( + ("file", "part"), + ( + pytest.param({"file_data": DATA_URI, "video_metadata": METADATA}, INLINE_PART, id="data-uri"), + pytest.param({"file_id": GCS_URI, "video_metadata": METADATA}, GCS_PART, id="gcs"), + pytest.param( + {"file_id": GCS_URI, "video_metadata": {"fps": 1, "start_offset": "", "end_offset": LONG_OFFSET}}, + gcs_part(metadata={"fps": 1, "startOffset": "", "endOffset": LONG_OFFSET}), + id="verbatim-strings-and-int-fps", + ), + pytest.param( + {"file_data": DATA_URI, "format": "video/quicktime"}, + {"inline_data": {"mime_type": "video/quicktime", "data": CLIP}}, + id="format-overrides-the-data-uri-mime", + ), + pytest.param( + {"file_id": "gs://scripted-bucket/clips/clip.bin", "format": "video/mp4"}, + gcs_part("gs://scripted-bucket/clips/clip.bin", None), + id="format-names-an-unlisted-extension", + ), + pytest.param({"file_id": GCS_URI}, gcs_part(metadata=None), id="no-metadata"), + pytest.param({"file_id": GCS_URI, "video_metadata": {}}, gcs_part(metadata=None), id="empty-metadata"), + ), +) +def test_file_block_reaches_the_wire_as_one_part( + gateway: Gateway, file: dict[str, JsonValue], part: dict[str, JsonValue] +) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, [block(**file)])) == [vector([part])] + assert wire_parts(wire) == [[part]] + + +def test_nested_text_and_block_become_one_request_with_two_parts(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [["a red bus", block(file_id=GCS_URI, video_metadata=METADATA)]]) + assert embeddings(response) == [vector([TEXT_PART, GCS_PART])] + assert wire_parts(wire) == [[TEXT_PART, GCS_PART]] + + +def test_repeated_blocks_become_one_request_each(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [block(file_id=GCS_URI, video_metadata=METADATA)] * 2) + assert embeddings(response) == [vector([GCS_PART])] * 2 + assert wire_parts(wire) == [[GCS_PART], [GCS_PART]] + + +@pytest.mark.parametrize( + ("element", "fragment"), + ( + pytest.param(block(file_id=GCS_URI, video_metadata={"fps": "1"}), "file.video_metadata.fps", id="fps-string"), + pytest.param( + block(file_id=GCS_URI, video_metadata={"start_offset": 0}), + "file.video_metadata.start_offset", + id="offset-int", + ), + pytest.param( + block(file_id=GCS_URI, video_metadata={"end_offset": ["3s"]}), + "file.video_metadata.end_offset", + id="offset-list", + ), + pytest.param(block(file_id=GCS_URI, video_metadata="1fps"), "file.video_metadata", id="metadata-string"), + pytest.param( + block(file_id=GCS_URI, video_metadata={**METADATA, "frame_rate": 2}), + "frame_rate", + id="unknown-metadata-key", + ), + pytest.param(block(file_id=GCS_URI, mime_type="video/mp4"), "mime_type", id="unknown-file-key"), + pytest.param({**block(file_id=GCS_URI), "caption": "x"}, "caption", id="unknown-block-key"), + pytest.param({"type": "video", "file": {"file_id": GCS_URI}}, "type", id="wrong-type"), + pytest.param( + block(file_id=GCS_URI, file_data=DATA_URI), + "takes file.file_id or file.file_data, not both", + id="both-sources", + ), + pytest.param(block(video_metadata=METADATA), "needs file.file_id or file.file_data", id="no-source"), + pytest.param(block(file_data=CLIP), ACCEPTED_FORMS, id="bare-base64"), + pytest.param(block(file_id="https://example.com/clip.mp4"), ACCEPTED_FORMS, id="http-url"), + pytest.param(block(file_id=""), ACCEPTED_FORMS, id="empty-file-id"), + pytest.param(block(file_id=GCS_URI, format=""), "file.format", id="empty-format"), + pytest.param(block(file_data=5), "file.file_data", id="file-data-int"), + pytest.param(block(file_data=[DATA_URI]), "file.file_data", id="file-data-list"), + pytest.param(1, "got int", id="int-element"), + ), +) +def test_malformed_blocks_answer_400_before_any_provider_call( + gateway: Gateway, element: JsonValue, fragment: str +) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [element]) + assert response.status_code == 400, response.text + assert fragment in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_bare_block_input_answers_400_before_any_provider_call(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, block(file_id=GCS_URI, video_metadata=METADATA)) + assert response.status_code == 400, response.text + assert "input must be a string or a list" in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_request_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, [NOISY_BLOCK], drop_params=True)) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +def test_deployment_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url, drop_params=True) + assert embeddings(embed(gateway, model, [NOISY_BLOCK])) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +@pytest.mark.timeout(300) +def test_yaml_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway, tmp_path: Path) -> None: + name: Final = f"gemini-embeddings-{uuid.uuid4().hex}" + with wire_server(embed_peer) as wire: + config: Final = write_config(tmp_path, wire.url, name) + with owned_proxy_process(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as owned: + assert embeddings(embed(owned.gateway, name, [NOISY_BLOCK])) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +def test_openai_sdk_clients_send_blocks_through_the_proxy(gateway: Gateway) -> None: + sync_uri: Final = "gs://scripted-bucket/clips/sync.mp4" + async_uri: Final = "gs://scripted-bucket/clips/async.mp4" + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + with OpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=gateway.key, max_retries=0) as client: + served: Final = client.post( + "/embeddings", + body={"model": model, "input": [block(file_id=sync_uri, video_metadata=METADATA)]}, + cast_to=CreateEmbeddingResponse, + ) + assert served.data[0].embedding == vector([gcs_part(sync_uri)]), served.model_dump_json() + assert wire_parts(wire) == [[gcs_part(sync_uri)]] + + async def drive() -> CreateEmbeddingResponse: + async with AsyncOpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=gateway.key, max_retries=0) as client: + return await client.post( + "/embeddings", + body={"model": model, "input": [block(file_id=async_uri, video_metadata=METADATA)]}, + cast_to=CreateEmbeddingResponse, + ) + + served_async: Final = asyncio.run(drive()) + assert served_async.data[0].embedding == vector([gcs_part(async_uri)]), served_async.model_dump_json() + assert wire_parts(wire) == [[gcs_part(async_uri)]] diff --git a/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py b/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py new file mode 100644 index 00000000000..042ca8b0889 --- /dev/null +++ b/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py @@ -0,0 +1,109 @@ +from __future__ import annotations + +import json +from typing import Final +from urllib.parse import unquote + +import httpx +from integration._support.client import Gateway, Scenario +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-2-preview" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +BUCKET: Final = "scripted-bucket" +UPLOAD_PREFIX: Final = f"/upload/storage/v1/b/{BUCKET}/o?uploadType=media&name=" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +BLOCK: Final[dict[str, JsonValue]] = { + "type": "file", + "file": {"file_id": GCS_URI, "video_metadata": {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"}}, +} +GCS_PART: Final[dict[str, JsonValue]] = { + "file_data": {"mime_type": "video/mp4", "file_uri": GCS_URI}, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} + + +def row(custom_id: str, model: str, elements: JsonValue) -> dict[str, JsonValue]: + return { + "custom_id": custom_id, + "method": "POST", + "url": "/v1/embeddings", + "body": {"model": model, "input": elements}, + } + + +def stored_row(key: str, parts: list[JsonValue]) -> dict[str, JsonValue]: + return {"key": key, "request": {"content": {"parts": parts}}} + + +def gcs_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.startswith(UPLOAD_PREFIX), request.target + assert request.headers["authorization"] == "Bearer scripted-token" + name: Final = unquote(request.target.removeprefix(UPLOAD_PREFIX)) + stored: Final = { + "kind": "storage#object", + "id": f"{BUCKET}/{name}/1759950000000000", + "name": name, + "bucket": BUCKET, + "size": str(len(request.body)), + "timeCreated": "2026-10-08T19:00:00.000Z", + "contentType": "application/octet-stream", + } + return Reply(body=json.dumps(stored).encode()) + + +def vertex_batch_model(gateway: Gateway, scenario: Scenario, url: str) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + gcs_bucket_name=BUCKET, + ) + + +def upload(gateway: Gateway, model: str, rows: tuple[dict[str, JsonValue], ...]) -> httpx.Response: + jsonl: Final = "\n".join(json.dumps(line) for line in rows) + "\n" + return gateway.request_multipart( + "/v1/files", + {"purpose": "batch", "target_model_names": model}, + {"file": ("in.jsonl", jsonl.encode(), "application/jsonl")}, + ) + + +def test_block_rows_upload_as_embed_content_requests(gateway: Gateway) -> None: + with wire_server(gcs_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_batch_model(gateway, scenario, wire.url) + rows: Final = ( + row("mixed", model, [BLOCK, "a red bus"]), + row("bare", model, "a red bus"), + row("nested", model, [[BLOCK, "a red bus"]]), + ) + response: Final = upload(gateway, model, rows) + assert response.status_code == 200, response.text + assert response.json()["object"] == "file" and response.json()["purpose"] == "batch", response.text + uploads: Final = wire.drain() + assert len(uploads) == 1, [upload.target for upload in uploads] + stored: Final = tuple(json.loads(line) for line in uploads[0].body.decode().splitlines() if line.strip()) + assert stored == ( + stored_row("mixed#0/2", [GCS_PART]), + stored_row("mixed#1/2", [TEXT_PART]), + stored_row("bare", [TEXT_PART]), + stored_row("nested", [GCS_PART, TEXT_PART]), + ) + + +def test_bare_block_row_fails_the_upload_before_any_storage_write(gateway: Gateway) -> None: + with wire_server(gcs_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_batch_model(gateway, scenario, wire.url) + response: Final = upload(gateway, model, (row("mixed", model, [BLOCK, "a red bus"]), row("bare", model, BLOCK))) + assert response.status_code >= 400, response.text + assert "got dict" in response.text, response.text + assert wire.drain() == (), "a rejected batch file reached the bucket" diff --git a/tests/integration/providers/test_vertex_embedding_file_block_wire.py b/tests/integration/providers/test_vertex_embedding_file_block_wire.py new file mode 100644 index 00000000000..ad66412a27a --- /dev/null +++ b/tests/integration/providers/test_vertex_embedding_file_block_wire.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +import json +from typing import Final + +import httpx +from integration._support.client import Gateway, Scenario, object_value +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-2-preview" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +TARGET: Final = f"/v1/projects/{PROJECT}/locations/{LOCATION}/publishers/google/models/{BACKEND}:embedContent" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +GCS_PART: Final[dict[str, JsonValue]] = { + "file_data": {"mime_type": "video/mp4", "file_uri": GCS_URI}, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +VALUES: Final = [0.75, 0.25] + + +def block(**file: JsonValue) -> dict[str, JsonValue]: + return {"type": "file", "file": file} + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target == TARGET, request.target + assert request.headers["authorization"] == "Bearer scripted-token" + return Reply(body=json.dumps({"embedding": {"values": VALUES}}).encode()) + + +def wire_parts(wire: Wire) -> list[JsonValue]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + body: Final = object_value(json.loads(received[0].body)) + parts: Final = object_value(body["content"])["parts"] + assert isinstance(parts, list), body + return parts + + +def vertex_model(gateway: Gateway, scenario: Scenario, url: str, **litellm_params: JsonValue) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + **litellm_params, + ) + + +def embed(gateway: Gateway, model: str, elements: JsonValue, **extra: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/embeddings", {"model": model, "input": elements, **extra}) + + +def single_embedding(response: httpx.Response) -> JsonValue: + assert response.status_code == 200, response.text + data: Final = object_value(response.json())["data"] + assert isinstance(data, list) and len(data) == 1, response.text + item: Final = object_value(data[0]) + assert item["index"] == 0 and item["object"] == "embedding", response.text + return item["embedding"] + + +def rejected(gateway: Gateway, wire: Wire, model: str, elements: JsonValue, fragment: str) -> None: + response: Final = embed(gateway, model, elements) + assert response.status_code == 400, response.text + assert fragment in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_string_input_reaches_embed_content_as_a_text_part(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, ["a red bus"])) == VALUES + assert wire_parts(wire) == [TEXT_PART] + + +def test_file_block_with_video_metadata_reaches_embed_content(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, [block(file_id=GCS_URI, video_metadata=METADATA)])) == VALUES + assert wire_parts(wire) == [GCS_PART] + + +def test_request_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + noisy: Final[dict[str, JsonValue]] = { + **block(file_id=GCS_URI, mime_type="video/mp4", video_metadata={**METADATA, "frame_rate": 2}), + "caption": "unknown block key", + } + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, [noisy], drop_params=True)) == VALUES + assert wire_parts(wire) == [GCS_PART] + + +def test_bad_fps_answers_400_before_any_provider_call(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + rejected(gateway, wire, model, [block(file_id=GCS_URI, video_metadata={"fps": "1"})], "file.video_metadata.fps") + + +def test_files_api_reference_answers_400_naming_the_gemini_provider(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + rejected( + gateway, + wire, + model, + [block(file_id="files/abc", video_metadata=METADATA)], + "Gemini Files API references are only supported through the gemini/ provider", + ) diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index a799c45e4c7..c94ee836f91 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -953,6 +953,7 @@ def test_redis_caching_multiple_namespaces(): _TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} +_FILE_BLOCK_ITEM: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA", "format": "video/mp4"}} @pytest.mark.parametrize( @@ -964,6 +965,7 @@ _TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} pytest.param({"input": [_TOOL_TURN_ITEM] * 5}, False, id="five-responses-items-skip-the-cache"), pytest.param({"input": "one prompt"}, True, id="string-input-is-one-message"), pytest.param({"input": ["a", "b", "c", "d", "e"]}, True, id="embedding-strings-are-not-messages"), + pytest.param({"input": [_FILE_BLOCK_ITEM] * 5}, True, id="embedding-file-blocks-are-not-messages"), ], ) def test_should_use_cache_stops_past_the_default_max_messages(kwargs: dict[str, object], expected: bool) -> None: diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index 22eea9935f6..d98c02667ee 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -41,6 +41,7 @@ from litellm.types.utils import ( from litellm.types.llms.openai import ResponsesAPIResponse from collections.abc import Awaitable, Callable from datetime import timedelta, datetime +from typing import Final 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 @@ -2714,3 +2715,88 @@ async def test_async_get_cache_forgets_the_worker_copy_of_a_stored_response_with assert lookup.cached_result is None assert await handler.dual_cache.async_get_cache(key) is None + + +@pytest.mark.asyncio +async def test_async_get_cache_partial_hit_keeps_file_block_items_uncached() -> None: + setup_cache() + fixed_start: Final = datetime(2026, 1, 1) + caching_handler: Final = LLMCachingHandler(original_function=aembedding, request_kwargs={}, start_time=fixed_start) + model: Final = "gemini/gemini-embedding-2-preview" + logging_obj: Final = LiteLLMLogging( + litellm_call_id=str(uuid.uuid4()), + call_type=CallTypes.aembedding.value, + model=model, + messages=[], + function_id=str(uuid.uuid4()), + stream=False, + start_time=fixed_start, + ) + await caching_handler.async_set_cache( + result=EmbeddingResponse(model=model, data=[Embedding(embedding=[0.1, 0.2], index=0, object="embedding")]), + original_function=aembedding, + kwargs={"model": model, "input": ["a red bus"], "caching": True}, + ) + clip_block: Final = { + "type": "file", + "file": { + "file_data": "data:video/mp4;base64,AAAA", + "format": "video/mp4", + "video_metadata": {"fps": 1, "start_offset": "0s", "end_offset": "1s"}, + }, + "detail": "left for the provider transformation to judge", + } + + cached_response: Final = await caching_handler.async_get_cache( + model=model, + original_function=aembedding, + logging_obj=logging_obj, + start_time=fixed_start, + call_type=CallTypes.aembedding.value, + kwargs={"model": model, "input": [clip_block, "a red bus"], "caching": True}, + ) + + assert cached_response.embedding_all_elements_cache_hit is False + assert cached_response.embedding_uncached_input == [clip_block] + assert cached_response.final_embedding_cached_response is not None + assert cached_response.final_embedding_cached_response.data[1].embedding == [0.1, 0.2] + assert cached_response.final_embedding_cached_response.data[0] is None + + +def test_handle_kwargs_input_answers_400_for_a_single_object_input() -> None: + caching_handler: Final = LLMCachingHandler( + original_function=aembedding, request_kwargs={}, start_time=datetime(2026, 1, 1) + ) + clip_block: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA"}} + with pytest.raises(litellm.BadRequestError, match="string or a list"): + caching_handler.handle_kwargs_input_list_or_str( + {"model": "gemini/gemini-embedding-2-preview", "custom_llm_provider": "gemini", "input": clip_block} + ) + + +@pytest.mark.asyncio +async def test_async_get_cache_answers_400_for_a_single_object_embedding_input() -> None: + setup_cache() + fixed_start: Final = datetime(2026, 1, 1) + caching_handler: Final = LLMCachingHandler(original_function=aembedding, request_kwargs={}, start_time=fixed_start) + model: Final = "gemini/gemini-embedding-2-preview" + logging_obj: Final = LiteLLMLogging( + litellm_call_id=str(uuid.uuid4()), + call_type=CallTypes.aembedding.value, + model=model, + messages=[], + function_id=str(uuid.uuid4()), + stream=False, + start_time=fixed_start, + ) + clip_block: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA"}} + + with pytest.raises(litellm.BadRequestError, match="string or a list"): + await caching_handler.async_get_cache( + model=model, + original_function=aembedding, + logging_obj=logging_obj, + start_time=fixed_start, + call_type=CallTypes.aembedding.value, + kwargs={"model": model, "custom_llm_provider": "gemini", "input": clip_block, "caching": True}, + ) diff --git a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 6f18a391f7b..dde3df6fcc0 100644 --- a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -1289,6 +1289,33 @@ class TestVertexEmbeddingsBatchInputTranslation: assert set(row["request"]) == {"content"} + def test_should_translate_a_file_block_with_video_metadata(self) -> None: + (row,) = _wrap_entries( + [ + _embeddings_entry( + body={ + "model": "gemini-embedding-2", + "input": [ + { + "type": "file", + "file": { + "file_id": "gs://my-bucket/clip.mp4", + "video_metadata": {"start_offset": "3s", "end_offset": "6s"}, + }, + } + ], + } + ) + ] + ) + + assert row["request"]["content"]["parts"] == [ + { + "file_data": {"mime_type": "video/mp4", "file_uri": "gs://my-bucket/clip.mp4"}, + "video_metadata": {"startOffset": "3s", "endOffset": "6s"}, + } + ] + def test_should_translate_multimodal_gcs_input(self): (row,) = _wrap_entries( [ diff --git a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py new file mode 100644 index 00000000000..483b350fd43 --- /dev/null +++ b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py @@ -0,0 +1,39 @@ +from typing import Final + +import httpx +import pytest + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler import GoogleBatchEmbeddings + +FILES_URI: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" +FILE_METADATA_URL: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + + +def _files_api_metadata(request: httpx.Request) -> httpx.Response: + if str(request.url) != FILE_METADATA_URL or request.headers.get("x-goog-api-key") != "gemini-key": + return httpx.Response(404, json={"error": {"message": "no such file"}}) + return httpx.Response(200, json={"mimeType": "video/mp4", "uri": FILES_URI}) + + +@pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) +def test_resolve_file_references_fetches_the_file_metadata_for_both_reference_forms(reference: str) -> None: + resolved: Final = GoogleBatchEmbeddings()._resolve_file_references( + input=[reference, "a red bus"], + api_key="gemini-key", + sync_handler=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_files_api_metadata))), + ) + assert resolved == {reference: {"mime_type": "video/mp4", "uri": FILES_URI}} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) +async def test_async_resolve_file_references_fetches_the_file_metadata_for_both_reference_forms( + reference: str, +) -> None: + resolved: Final = await GoogleBatchEmbeddings()._async_resolve_file_references( + input=[reference, "a red bus"], + api_key="gemini-key", + async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(_files_api_metadata)), + ) + assert resolved == {reference: {"mime_type": "video/mp4", "uri": FILES_URI}} diff --git a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py index ba2b26bf0a2..6584c27b895 100644 --- a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -8,12 +8,17 @@ Covers: - Response processing with correct indices """ +from collections.abc import Callable +from typing import Final + import pytest import litellm +from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( _build_part_for_input, + file_reference_name, _is_multimodal_input, process_embed_content_response, process_response, @@ -25,6 +30,12 @@ from litellm.types.utils import EmbeddingResponse IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII" GCS_URL = "gs://my-bucket/image.png" +VIDEO_DATA_URI = "data:video/mp4;base64,AAAAIGZ0eXBpc29tAAACAGlzb21pc28yYXZjMW1wNDEAAAAIZnJlZQAA" +FILES_URI = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + + +def _file_block(**file: object) -> dict[str, object]: + return {"type": "file", "file": file} @pytest.fixture(autouse=True) @@ -51,6 +62,7 @@ class TestIsMultimodalInput: def test_file_reference(self): assert _is_multimodal_input(["files/abc123"]) is True + assert _is_multimodal_input([FILES_URI]) is True def test_mixed_text_and_image(self): assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True @@ -59,9 +71,16 @@ class TestIsMultimodalInput: """Nested list with text is not multimodal.""" assert _is_multimodal_input([["text_a", "text_b"]]) is False + def test_file_content_block_is_multimodal(self) -> None: + assert _is_multimodal_input([_file_block(file_data=VIDEO_DATA_URI)]) is True + def test_nested_list_with_image_is_multimodal(self): assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True + def test_single_object_input_answers_400(self) -> None: + with pytest.raises(BadRequestError, match="string or a list"): + _is_multimodal_input(_file_block(file_data=VIDEO_DATA_URI)) + class TestBuildPartForInput: def test_text_input(self): @@ -83,15 +102,13 @@ class TestBuildPartForInput: assert part["file_data"]["file_uri"] == GCS_URL def test_file_reference_resolved(self): - resolved = { - "files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"} - } + resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}} part = _build_part_for_input("files/abc", resolved_files=resolved) assert part["file_data"] is not None assert part["file_data"]["mime_type"] == "image/jpeg" - def test_file_reference_unresolved_raises(self): - with pytest.raises(ValueError, match="not resolved"): + def test_file_reference_unresolved_answers_400_naming_the_gemini_provider(self) -> None: + with pytest.raises(BadRequestError, match="gemini/ provider"): _build_part_for_input("files/abc") @@ -124,10 +141,7 @@ class TestTransformOpenaiInputGeminiContent: ) assert len(result["requests"]) == 2 # First request is text - assert ( - result["requests"][0]["content"]["parts"][0]["text"] - == "The food was delicious" - ) + assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious" # Second request is image assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None @@ -226,9 +240,7 @@ class TestProcessResponse: """Test that process_response sets correct indices.""" def test_single_embedding_index(self): - predictions: VertexAIBatchEmbeddingsResponseObject = { - "embeddings": [{"values": [0.1, 0.2]}] - } + predictions: VertexAIBatchEmbeddingsResponseObject = {"embeddings": [{"values": [0.1, 0.2]}]} model_response = EmbeddingResponse() result = process_response( input="hello", @@ -279,9 +291,7 @@ class TestProcessResponse: def test_nested_input_token_counting(self): """Nested list: only plain-text sub-elements should be counted.""" - predictions: VertexAIBatchEmbeddingsResponseObject = { - "embeddings": [{"values": [0.1, 0.2]}] - } + predictions: VertexAIBatchEmbeddingsResponseObject = {"embeddings": [{"values": [0.1, 0.2]}]} result = process_response( input=[["a red shoe", IMAGE_DATA_URI]], model_response=EmbeddingResponse(), @@ -300,7 +310,7 @@ class TestProcessResponse: ) def test_nested_non_string_element_raises(self): - with pytest.raises(ValueError, match="must be strings"): + with pytest.raises(BadRequestError, match="must be strings or file content blocks, got list"): transform_openai_input_gemini_content( input=[[["doubly", "nested"]]], model="gemini-embedding-2-preview", @@ -408,3 +418,212 @@ class TestProcessEmbedContentResponseUsage: assert result.usage.prompt_tokens > 0 +class TestFileContentBlocks: + MODEL: Final = "gemini-embedding-2-preview" + CLIP_METADATA: Final = {"fps": 2, "start_offset": "3s", "end_offset": "6s"} + CLIP_PART: Final = {"fps": 2.0, "startOffset": "3s", "endOffset": "6s"} + + def test_data_uri_block_forwards_format_and_video_metadata(self) -> None: + part: Final = _build_part_for_input( + _file_block( + file_data="data:application/octet-stream;base64,QUJD", + format="video/mp4", + video_metadata=self.CLIP_METADATA, + ) + ) + assert part == { + "inline_data": {"mime_type": "video/mp4", "data": "QUJD"}, + "video_metadata": self.CLIP_PART, + } + + def test_gcs_block_infers_mime_type_and_forwards_video_metadata(self) -> None: + part: Final = _build_part_for_input( + _file_block(file_id="gs://my-bucket/clip.mp4", video_metadata={"start_offset": "0s", "end_offset": "3s"}) + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": "gs://my-bucket/clip.mp4"}, + "video_metadata": {"startOffset": "0s", "endOffset": "3s"}, + } + + def test_file_reference_block_uses_the_resolved_file(self) -> None: + resolved_files: Final = {"files/clip123": {"mime_type": "video/mp4", "uri": FILES_URI}} + part: Final = _build_part_for_input( + _file_block(file_id="files/clip123", video_metadata={"fps": 1}), + resolved_files=resolved_files, + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": FILES_URI}, + "video_metadata": {"fps": 1.0}, + } + + def test_files_api_uri_block_uses_the_resolved_file(self) -> None: + resolved_files: Final = {FILES_URI: {"mime_type": "video/mp4", "uri": FILES_URI}} + part: Final = _build_part_for_input( + _file_block(file_id=FILES_URI, video_metadata={"start_offset": "0s", "end_offset": "3s"}), + resolved_files=resolved_files, + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": FILES_URI}, + "video_metadata": {"startOffset": "0s", "endOffset": "3s"}, + } + + @pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) + def test_file_reference_name_is_the_files_name_for_both_forms(self, reference: str) -> None: + assert file_reference_name(reference) == "files/clip123" + + def test_block_without_video_metadata_sends_no_video_metadata_key(self) -> None: + part: Final = _build_part_for_input(_file_block(file_data=IMAGE_DATA_URI, filename="dot.png")) + assert part == {"inline_data": {"mime_type": "image/png", "data": IMAGE_DATA_URI.split(",", 1)[1]}} + + def test_block_with_empty_video_metadata_sends_no_video_metadata_key(self) -> None: + part: Final = _build_part_for_input(_file_block(file_data=VIDEO_DATA_URI, video_metadata={})) + assert part == {"inline_data": {"mime_type": "video/mp4", "data": VIDEO_DATA_URI.split(",", 1)[1]}} + + def test_batch_path_nested_block_and_text_share_one_request(self) -> None: + result: Final = transform_openai_input_gemini_content( + input=[[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"]], + model=self.MODEL, + optional_params={"dimensions": 768}, + ) + [request] = result["requests"] + assert request["outputDimensionality"] == 768 + assert request["content"]["parts"] == [ + { + "inline_data": {"mime_type": "video/mp4", "data": VIDEO_DATA_URI.split(",", 1)[1]}, + "video_metadata": self.CLIP_PART, + }, + {"text": "a solid color clip"}, + ] + + def test_batch_path_flat_block_and_text_are_separate_requests(self) -> None: + result: Final = transform_openai_input_gemini_content( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"], + model=self.MODEL, + optional_params={}, + ) + assert len(result["requests"]) == 2 + assert result["requests"][0]["content"]["parts"][0]["video_metadata"] == self.CLIP_PART + assert result["requests"][1]["content"]["parts"] == [{"text": "a solid color clip"}] + + def test_embed_content_path_accepts_flat_block_and_text(self) -> None: + result: Final = transform_openai_input_gemini_embed_content( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"], + model=self.MODEL, + optional_params={}, + ) + parts: Final = result["content"]["parts"] + assert parts[0]["video_metadata"] == self.CLIP_PART + assert parts[1] == {"text": "a solid color clip"} + + def test_embed_content_path_still_rejects_nested_lists(self) -> None: + with pytest.raises(ValueError, match="Nested"): + transform_openai_input_gemini_embed_content( + input=[[_file_block(file_data=VIDEO_DATA_URI), "a solid color clip"]], + model=self.MODEL, + optional_params={}, + ) + + @pytest.mark.parametrize( + "block, named_in_error", + [ + ( + _file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": 1, "startOffset": "1s"}), + "video_metadata.startOffset", + ), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": "fast"}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": "1"}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": True}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"start_offset": 5}), "video_metadata.start_offset"), + (_file_block(file_data=VIDEO_DATA_URI, detail="high"), "file.detail"), + (_file_block(file_data=VIDEO_DATA_URI, format=""), "file.format"), + (_file_block(file_id="gs://my-bucket/clip.mp4", file_data=VIDEO_DATA_URI), "not both"), + (_file_block(), "needs file.file_id or file.file_data"), + ( + _file_block(file_id="https://example.com/clip.mp4"), + "a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI", + ), + ({"type": "image_url", "image_url": {"url": IMAGE_DATA_URI}}, "Input should be 'file'"), + ], + ) + def test_malformed_block_answers_400_naming_the_field( + self, block: dict[str, object], named_in_error: str + ) -> None: + with pytest.raises(BadRequestError, match=named_in_error): + _build_part_for_input(block) + + def test_drop_params_drops_the_block_keys_this_surface_does_not_take(self) -> None: + block: Final = _file_block( + file_data=VIDEO_DATA_URI, + detail="high", + video_metadata={"fps": 1, "start_offset": "1s", "resolution": "low"}, + ) + part: Final = _build_part_for_input({**block, "cache_control": {"type": "ephemeral"}}, drop_params=True) + assert part["inline_data"]["mime_type"] == "video/mp4" + assert part["video_metadata"] == {"fps": 1.0, "startOffset": "1s"} + + def test_drop_params_still_answers_400_for_a_malformed_value(self) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high", video_metadata={"fps": "fast"}) + with pytest.raises(BadRequestError, match=r"video_metadata\.fps"): + _build_part_for_input(block, drop_params=True) + + @pytest.mark.parametrize( + "transform", [transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content] + ) + def test_transforms_forward_drop_params_to_every_block(self, transform: Callable[..., object]) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high") + with pytest.raises(BadRequestError, match=r"file\.detail"): + transform(input=[block], model="gemini-embedding-2-preview", optional_params={}) + transform(input=[block], model="gemini-embedding-2-preview", optional_params={}, drop_params=True) + + def test_batch_path_forwards_drop_params_into_nested_lists(self) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high") + body: Final = transform_openai_input_gemini_content( + input=[[block, "a caption"]], model="gemini-embedding-2-preview", optional_params={}, drop_params=True + ) + assert len(body["requests"][0]["content"]["parts"]) == 2 + + @pytest.mark.parametrize( + "transform", [transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content] + ) + def test_single_object_input_answers_400(self, transform: Callable[..., object]) -> None: + with pytest.raises(BadRequestError, match="string or a list"): + transform( + input=_file_block(file_data=VIDEO_DATA_URI), model="gemini-embedding-2-preview", optional_params={} + ) + + def test_process_response_counts_only_the_text_tokens_next_to_a_block(self) -> None: + text: Final = "a solid color clip" + with_block: Final = process_response( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), text], + model_response=EmbeddingResponse(), + model=self.MODEL, + _predictions={"embeddings": [{"values": [0.1]}, {"values": [0.2]}]}, + ) + text_only: Final = process_response( + input=[text], + model_response=EmbeddingResponse(), + model=self.MODEL, + _predictions={"embeddings": [{"values": [0.2]}]}, + ) + assert with_block.usage.prompt_tokens == text_only.usage.prompt_tokens > 0 + + def test_embed_content_usage_fallback_with_a_block_does_not_estimate(self) -> None: + result: Final = process_embed_content_response( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA)], + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json={"embedding": {"values": [0.1, 0.2]}}, + ) + assert result.usage.prompt_tokens == 0 + + def test_image_block_counts_as_image_only_input(self) -> None: + result: Final = process_embed_content_response( + input=[_file_block(file_data=IMAGE_DATA_URI)], + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json={ + "embedding": {"values": [0.1, 0.2]}, + "usageMetadata": {"promptTokenCount": 258, "totalTokenCount": 258}, + }, + ) + assert result.usage.prompt_tokens_details.image_tokens == 258 diff --git a/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py index 98abf5459df..7271fb48861 100644 --- a/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -8,11 +8,15 @@ This test ensures that: """ import json +from contextlib import AbstractContextManager +from typing import Final from unittest.mock import MagicMock, patch import pytest import litellm +import httpx + from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( _filter_embed_params, @@ -655,3 +659,163 @@ def test_batch_embeddings_response_has_correct_indices_and_order(): assert ( embedding.embedding == expected_values[i] ), f"embedding {i} has wrong values: {embedding.embedding}" + + +CLIP_BLOCK = { + "type": "file", + "file": { + "file_data": "data:video/mp4;base64,AAAAIGZ0eXBpc29tAAACAGlzb21pc28yYXZjMW1wNDEAAAAIZnJlZQAA", + "video_metadata": {"fps": 2, "start_offset": "3s", "end_offset": "6s"}, + }, +} +CLIP_PART_METADATA = {"fps": 2.0, "startOffset": "3s", "endOffset": "6s"} + + +def _mock_embedding_call( + response_json: dict[str, object], +) -> tuple[AbstractContextManager[MagicMock], AbstractContextManager[MagicMock], MagicMock]: + mock_get_token: Final = patch( + "litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._get_token_and_url" + ) + mock_auth: Final = patch( + "litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._ensure_access_token", + side_effect=lambda *args, **kwargs: (None, "test-project"), + ) + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = response_json + return mock_get_token, mock_auth, mock_response + + +def test_gemini_batch_path_sends_file_block_video_metadata() -> None: + client: Final = HTTPHandler() + mock_get_token, mock_auth, mock_response = _mock_embedding_call( + {"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]} + ) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ( + {"x-goog-api-key": "test-key"}, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:batchEmbedContents", + ) + response: Final = litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[CLIP_BLOCK, "a solid color clip"], + api_key="test-key", + client=client, + ) + + request_body: Final = json.loads(mock_post.call_args.kwargs["data"]) + clip_part: Final = request_body["requests"][0]["content"]["parts"][0] + assert clip_part["inline_data"]["mime_type"] == "video/mp4" + assert clip_part["video_metadata"] == CLIP_PART_METADATA + assert request_body["requests"][1]["content"]["parts"] == [{"text": "a solid color clip"}] + assert [row["index"] for row in response.data] == [0, 1] + + +def test_vertex_embed_content_path_sends_file_block_video_metadata() -> None: + client: Final = HTTPHandler() + url: Final = "https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/publishers/google/models/gemini-embedding-2-preview:embedContent" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embedding": {"values": [0.1, 0.2]}}) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ({"Authorization": "Bearer test-token"}, url) + response: Final = litellm.embedding( + model="vertex_ai/gemini-embedding-2-preview", + input=[CLIP_BLOCK, "a solid color clip"], + vertex_project="test-project", + vertex_location="us-central1", + client=client, + ) + + data: Final = json.loads(mock_post.call_args.kwargs["data"]) + assert data["content"]["parts"][0]["video_metadata"] == CLIP_PART_METADATA + assert data["content"]["parts"][1] == {"text": "a solid color clip"} + assert len(response.data) == 1 + + +def test_file_block_with_files_reference_is_resolved_through_the_files_api() -> None: + client: Final = HTTPHandler() + files_uri: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embeddings": [{"values": [0.1, 0.2]}]}) + file_lookup: Final = MagicMock() + file_lookup.status_code = 200 + file_lookup.json.return_value = {"mimeType": "video/mp4", "uri": files_uri} + with ( + patch.object(client, "post", return_value=mock_response) as mock_post, + patch.object(client, "get", return_value=file_lookup) as mock_get, + mock_auth, + mock_get_token as token, + ): + token.return_value = ( + {"x-goog-api-key": "test-key"}, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:batchEmbedContents", + ) + litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[{"type": "file", "file": {"file_id": "files/clip123", "video_metadata": {"fps": 1}}}], + api_key="test-key", + client=client, + ) + + assert mock_get.call_args.kwargs["url"] == "https://generativelanguage.googleapis.com/v1beta/files/clip123" + clip_part: Final = json.loads(mock_post.call_args.kwargs["data"])["requests"][0]["content"]["parts"][0] + assert clip_part == { + "file_data": {"mime_type": "video/mp4", "file_uri": files_uri}, + "video_metadata": {"fps": 1.0}, + } + + +CLIP_BLOCK_WITH_DETAIL: Final = {"type": "file", "file": {**CLIP_BLOCK["file"], "detail": "high"}} + + +def _recording_client(calls: list[httpx.Request], response_json: dict[str, object]) -> HTTPHandler: + def route(request: httpx.Request) -> httpx.Response: + calls.append(request) + return httpx.Response(200, json=response_json) + + return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(route))) + + +def test_gemini_drop_params_strips_the_block_keys_embeddings_do_not_take() -> None: + calls: Final[list[httpx.Request]] = [] + client: Final = _recording_client(calls, {"embeddings": [{"values": [0.1, 0.2]}]}) + response: Final = litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[CLIP_BLOCK_WITH_DETAIL], + api_key="test-key", + client=client, + drop_params=True, + ) + sent_part: Final = json.loads(calls[0].content)["requests"][0]["content"]["parts"][0] + assert sent_part["video_metadata"] == CLIP_PART_METADATA + assert "detail" not in calls[0].content.decode() + assert response.data[0].embedding == [0.1, 0.2] + + +def test_gemini_block_detail_answers_400_without_drop_params() -> None: + calls: Final[list[httpx.Request]] = [] + client: Final = _recording_client(calls, {"embeddings": [{"values": [0.1, 0.2]}]}) + with pytest.raises(litellm.BadRequestError, match=r"file\.detail"): + litellm.embedding( + model="gemini/gemini-embedding-2-preview", input=[CLIP_BLOCK_WITH_DETAIL], api_key="test-key", client=client + ) + assert calls == [] + + +def test_vertex_drop_params_strips_the_block_keys_embeddings_do_not_take() -> None: + client: Final = HTTPHandler() + url: Final = "https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/publishers/google/models/gemini-embedding-2-preview:embedContent" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embedding": {"values": [0.1, 0.2]}}) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ({"Authorization": "Bearer test-token"}, url) + litellm.embedding( + model="vertex_ai/gemini-embedding-2-preview", + input=[CLIP_BLOCK_WITH_DETAIL], + vertex_project="test-project", + vertex_location="us-central1", + client=client, + drop_params=True, + ) + + data: Final = json.loads(mock_post.call_args.kwargs["data"]) + assert data["content"]["parts"][0]["video_metadata"] == CLIP_PART_METADATA + assert "detail" not in mock_post.call_args.kwargs["data"]