mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
f663794fcd
commit
56bafc29cf
22 changed files with 2048 additions and 150 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
216
tests/integration/caching/test_embedding_file_block_cache.py
Normal file
216
tests/integration/caching/test_embedding_file_block_cache.py
Normal 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"
|
||||
|
|
@ -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]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)]]
|
||||
|
|
@ -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"
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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}}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue