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>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 13:20:03 -07:00 • committed by GitHub
parent f663794fcd
commit 56bafc29cf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 2048 additions and 150 deletions

View file

@ -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:

View file

@ -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})

View file

@ -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",

View file

@ -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,)

View file

@ -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

View file

@ -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

View file

@ -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(

View file

@ -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

View file

@ -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):

View file

@ -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)

View file

@ -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"

View file

@ -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]

View file

@ -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)

View file

@ -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)]]

View file

@ -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"

View file

@ -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",
)

View file

@ -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:

View file

@ -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},
)

View file

@ -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(
[

View file

@ -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}}

View file

@ -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

View file

@ -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"]