Merge pull request #24340 from BerriAI/litellm_staging_03_21_2026

Litellm staging 03 21 2026
This commit is contained in:
yuneng-jiang 2026-05-01 11:57:44 -07:00 committed by GitHub
commit 02582466c4
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
8 changed files with 495 additions and 71 deletions

View File

@ -68,7 +68,7 @@ Managing LLM calls across providers gets complicated fast — different SDKs, au
<td><img height="60" alt="Stripe" src="https://github.com/user-attachments/assets/f7296d4f-9fbd-460d-9d05-e4df31697c4b" /></td>
<td><img height="60" alt="image" src="https://github.com/user-attachments/assets/436fca71-988b-40bb-b5fe-8450c80fdbd0" /></td>
<td><img height="60" alt="Google ADK" src="https://github.com/user-attachments/assets/caf270a2-5aee-45c4-8222-41a2070c4f19" /></td>
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/0be4bd8a-7cfa-48d3-9090-f415fe948280" /></td>
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/3db0ae72-0843-4005-a56d-bba1dde2193d" /></td>
<td><img height="60" alt="OpenHands" src="https://github.com/user-attachments/assets/a6150c4c-149e-4cae-888b-8b92be6e003f" /></td>
<td><h2>Netflix</h2></td>
<td><img height="60" alt="OpenAI Agents SDK" src="https://github.com/user-attachments/assets/c02f7be0-8c2e-4d27-aea7-7c024bfaebc0" /></td>

View File

@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint
"""
import json
from typing import Any, Dict, Literal, Optional, Union
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
import httpx
@ -13,8 +13,8 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
)
from litellm.types.llms.openai import EmbeddingInput
from litellm.types.llms.vertex_ai import (
GeminiEmbeddingInput,
VertexAIBatchEmbeddingsRequestBody,
VertexAIBatchEmbeddingsResponseObject,
)
@ -23,7 +23,6 @@ from litellm.types.utils import EmbeddingResponse
from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from .batch_embed_content_transformation import (
_is_file_reference,
_is_multimodal_input,
process_embed_content_response,
process_response,
transform_openai_input_gemini_content,
@ -32,9 +31,24 @@ from .batch_embed_content_transformation import (
class GoogleBatchEmbeddings(VertexLLM):
@staticmethod
def _flatten_and_detect_file_refs(
input: GeminiEmbeddingInput,
) -> Tuple[List[str], bool]:
"""Flatten nested input lists and detect file references."""
input_list = [input] if isinstance(input, str) else input
flat_elements = [
e
for item in input_list
for e in (item if isinstance(item, list) else [item])
if isinstance(e, str)
]
has_file_refs = any(_is_file_reference(e) for e in flat_elements)
return flat_elements, has_file_refs
def _resolve_file_references(
self,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
api_key: str,
sync_handler: HTTPHandler,
) -> Dict[str, Dict[str, str]]:
@ -42,7 +56,7 @@ class GoogleBatchEmbeddings(VertexLLM):
Resolve Gemini file references (files/...) to get mime_type and uri.
Args:
input: EmbeddingInput that may contain file references
input: GeminiEmbeddingInput that may contain file references
api_key: Gemini API key
sync_handler: HTTP client
@ -73,7 +87,7 @@ class GoogleBatchEmbeddings(VertexLLM):
async def _async_resolve_file_references(
self,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
api_key: str,
async_handler: AsyncHTTPHandler,
) -> Dict[str, Dict[str, str]]:
@ -81,7 +95,7 @@ class GoogleBatchEmbeddings(VertexLLM):
Async version of _resolve_file_references.
Args:
input: EmbeddingInput that may contain file references
input: GeminiEmbeddingInput that may contain file references
api_key: Gemini API key
async_handler: Async HTTP client
@ -110,10 +124,10 @@ class GoogleBatchEmbeddings(VertexLLM):
return resolved_files
def batch_embeddings(
def batch_embeddings( # noqa: PLR0915
self,
model: str,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
print_verbose,
model_response: EmbeddingResponse,
custom_llm_provider: Literal["gemini", "vertex_ai"],
@ -151,8 +165,7 @@ class GoogleBatchEmbeddings(VertexLLM):
optional_params = optional_params or {}
is_multimodal = _is_multimodal_input(input)
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
use_embed_content = custom_llm_provider == "vertex_ai"
mode: Literal["embedding", "batch_embedding"]
if use_embed_content:
mode = "embedding"
@ -215,8 +228,22 @@ class GoogleBatchEmbeddings(VertexLLM):
resolved_files=resolved_files,
)
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/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and has_file_refs:
resolved_files = self._resolve_file_references(
input=flat_elements, api_key=api_key, sync_handler=sync_handler
)
request_data = transform_openai_input_gemini_content(
input=input, model=model, optional_params=optional_params
input=input,
model=model,
optional_params=optional_params,
resolved_files=resolved_files,
)
## LOGGING
@ -264,7 +291,7 @@ class GoogleBatchEmbeddings(VertexLLM):
url: str,
data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]],
model_response: EmbeddingResponse,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
timeout: Optional[Union[float, httpx.Timeout]],
headers={},
client: Optional[AsyncHTTPHandler] = None,
@ -303,8 +330,22 @@ class GoogleBatchEmbeddings(VertexLLM):
resolved_files=resolved_files,
)
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/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and has_file_refs:
resolved_files = await self._async_resolve_file_references(
input=flat_elements, api_key=api_key, async_handler=async_handler
)
data = transform_openai_input_gemini_content(
input=input, model=model, optional_params=optional_params or {}
input=input,
model=model,
optional_params=optional_params or {},
resolved_files=resolved_files,
)
## LOGGING

View File

@ -6,12 +6,12 @@ Why separate file? Make it easy to see how transformation works
from typing import Dict, List, Optional, Tuple
from litellm.types.llms.openai import EmbeddingInput
from litellm.types.llms.vertex_ai import (
BlobType,
ContentType,
EmbedContentRequest,
FileDataType,
GeminiEmbeddingInput,
PartType,
VertexAIBatchEmbeddingsRequestBody,
VertexAIBatchEmbeddingsResponseObject,
@ -114,33 +114,77 @@ def _parse_data_url(data_url: str) -> Tuple[str, str]:
return media_type, base64_data
def _is_multimodal_input(input: EmbeddingInput) -> bool:
def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool:
"""
Check if the input contains multimodal data (data URIs, file references, or GCS URLs).
Check if the input contains multimodal data (data URIs, file references,
GCS URLs, or nested lists for combined embeddings).
Args:
input: EmbeddingInput (str or List[str])
input: GeminiEmbeddingInput str, List[str], or List[List[str]] for combined embeddings
Returns:
bool: True if any element is a data URI, file reference, or GCS URL
bool: True if any element is multimodal or a nested list
"""
if isinstance(input, str):
input_list = [input]
else:
input_list = input
return _is_multimodal_element(input)
for element in input_list:
if isinstance(element, str):
if element.startswith("data:") and ";base64," in element:
return True
if _is_file_reference(element):
return True
if _is_gcs_url(element):
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
def _is_multimodal_element(element: str) -> bool:
"""Check if a single string element is multimodal."""
if element.startswith("data:") and ";base64," in element:
return True
if _is_file_reference(element):
return True
if _is_gcs_url(element):
return True
return False
def _build_part_for_input(
element: str,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
) -> PartType:
"""
Build a single PartType for an input element, handling text, data URIs,
file references, and GCS URLs.
"""
resolved_files = resolved_files or {}
if element.startswith("data:") and ";base64," in element:
mime_type, base64_data = _parse_data_url(element)
blob: 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: 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 = resolved_files[element]
file_data_ref: FileDataType = {
"mime_type": file_info["mime_type"],
"file_uri": file_info["uri"],
}
return PartType(file_data=file_data_ref)
else:
return PartType(text=element)
_SUPPORTED_EMBED_PARAMS = {"outputDimensionality", "taskType", "title"}
@ -155,37 +199,60 @@ def _filter_embed_params(optional_params: dict) -> dict:
def transform_openai_input_gemini_content(
input: EmbeddingInput, model: str, optional_params: dict
input: GeminiEmbeddingInput,
model: str,
optional_params: dict,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
) -> VertexAIBatchEmbeddingsRequestBody:
"""
The content to embed. Only the parts.text fields will be counted.
Transform OpenAI embedding input to Gemini batchEmbedContents format.
Each input element becomes a separate EmbedContentRequest, supporting
text, data URIs, file references, and GCS URLs.
If an element is a list (nested input), all sub-elements are combined
into a single content with multiple parts, producing one combined
embedding for the group.
Examples:
input=["text", "image"] 2 separate embeddings
input=[["text", "image"]] 1 combined embedding
input=[["text", "image"], "x"] 2 embeddings (1 combined + 1 separate)
"""
gemini_model_name = "models/{}".format(model)
gemini_params = _filter_embed_params(optional_params)
input_list = [input] if isinstance(input, str) else input
requests: List[EmbedContentRequest] = []
if isinstance(input, str):
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
]
else:
parts = [_build_part_for_input(element, resolved_files=resolved_files)]
request = EmbedContentRequest(
model=gemini_model_name,
content=ContentType(parts=[PartType(text=input)]),
content=ContentType(parts=parts),
**gemini_params,
)
requests.append(request)
else:
for i in input:
request = EmbedContentRequest(
model=gemini_model_name,
content=ContentType(parts=[PartType(text=i)]),
**gemini_params,
)
requests.append(request)
return VertexAIBatchEmbeddingsRequestBody(requests=requests)
def transform_openai_input_gemini_embed_content(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model: str,
optional_params: dict,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
@ -194,7 +261,7 @@ def transform_openai_input_gemini_embed_content(
Transform OpenAI embedding input to Gemini embedContent format (multimodal).
Args:
input: EmbeddingInput (str or List[str]) with text, data URIs, or file references
input: GeminiEmbeddingInput with text, data URIs, or file references
model: Model name
optional_params: Additional parameters (taskType, outputDimensionality, etc.)
resolved_files: Dict mapping file names (files/abc) to {mime_type, uri}
@ -210,31 +277,14 @@ def transform_openai_input_gemini_embed_content(
parts: List[PartType] = []
for element in input_list:
if isinstance(element, list):
raise ValueError(
"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)}")
if element.startswith("data:") and ";base64," in element:
mime_type, base64_data = _parse_data_url(element)
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
parts.append(PartType(inline_data=blob))
elif _is_gcs_url(element):
mime_type = _infer_mime_type_from_gcs_url(element)
file_data: FileDataType = {
"mime_type": mime_type,
"file_uri": element,
}
parts.append(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 = resolved_files[element]
file_data_ref: FileDataType = {
"mime_type": file_info["mime_type"],
"file_uri": file_info["uri"],
}
parts.append(PartType(file_data=file_data_ref))
else:
parts.append(PartType(text=element))
parts.append(_build_part_for_input(element, resolved_files=resolved_files))
request_body: dict = {
"content": ContentType(parts=parts),
@ -245,7 +295,7 @@ def transform_openai_input_gemini_embed_content(
def process_embed_content_response(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model_response: EmbeddingResponse,
model: str,
response_json: dict,
@ -291,7 +341,7 @@ def process_embed_content_response(
def process_response(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model_response: EmbeddingResponse,
model: str,
_predictions: VertexAIBatchEmbeddingsResponseObject,
@ -308,8 +358,29 @@ def process_response(
model_response.data = openai_embeddings
model_response.model = model
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
prompt_tokens = token_counter(model=model, text=input_text)
has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input)
if _is_multimodal_input(input) or has_nested:
input_list = input if isinstance(input, list) else [input]
text_elements: List[str] = []
for e in input_list:
if isinstance(e, list):
text_elements.extend(
sub
for sub in e
if isinstance(sub, str) and not _is_multimodal_element(sub)
)
elif isinstance(e, str) and not _is_multimodal_element(e):
text_elements.append(e)
if text_elements:
input_text = get_formatted_prompt(
data={"input": text_elements}, call_type="embedding"
)
prompt_tokens = token_counter(model=model, text=input_text)
else:
prompt_tokens = 0
else:
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
prompt_tokens = token_counter(model=model, text=input_text)
model_response.usage = Usage(
prompt_tokens=prompt_tokens, total_tokens=prompt_tokens
)

View File

@ -6,6 +6,13 @@ from typing_extensions import (
TypedDict,
)
from litellm.types.llms.openai import 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 = Union[EmbeddingInput, List[List[str]]]
class FunctionResponse(TypedDict):
name: str

View File

@ -34894,6 +34894,20 @@
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-5": {
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3.2e-06,
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-4.7-flash": {
"input_cost_per_token": 7e-08,
"litellm_provider": "bedrock_converse",

View File

@ -477,3 +477,4 @@ def test_get_llm_provider_use_proxy_arg_true_with_direct_args():
assert provider == "litellm_proxy"
assert key == arg_api_key # Should use the argument key
assert base == arg_api_base # Should use the argument base

View File

@ -0,0 +1,290 @@
"""
Tests for Gemini batchEmbedContents transformation logic.
Covers:
- Text-only inputs (single and batch)
- Multimodal inputs (data URIs, GCS URLs, file references)
- Mixed text + multimodal inputs
- Response processing with correct indices
"""
import pytest
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
_build_part_for_input,
_is_multimodal_input,
process_response,
transform_openai_input_gemini_content,
transform_openai_input_gemini_embed_content,
)
from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject
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"
class TestIsMultimodalInput:
def test_text_only_string(self):
assert _is_multimodal_input("hello world") is False
def test_text_only_list(self):
assert _is_multimodal_input(["hello", "world"]) is False
def test_data_uri(self):
assert _is_multimodal_input([IMAGE_DATA_URI]) is True
def test_gcs_url(self):
assert _is_multimodal_input([GCS_URL]) is True
def test_file_reference(self):
assert _is_multimodal_input(["files/abc123"]) is True
def test_mixed_text_and_image(self):
assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True
def test_nested_text_is_not_multimodal(self):
"""Nested list with text is not multimodal."""
assert _is_multimodal_input([["text_a", "text_b"]]) is False
def test_nested_list_with_image_is_multimodal(self):
assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True
class TestBuildPartForInput:
def test_text_input(self):
part = _build_part_for_input("hello")
assert part["text"] == "hello"
assert part.get("inline_data") is None
def test_data_uri_input(self):
part = _build_part_for_input(IMAGE_DATA_URI)
assert part.get("text") is None
assert part["inline_data"] is not None
assert part["inline_data"]["mime_type"] == "image/png"
def test_gcs_url_input(self):
part = _build_part_for_input(GCS_URL)
assert part.get("text") is None
assert part["file_data"] is not None
assert part["file_data"]["mime_type"] == "image/png"
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"}}
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"):
_build_part_for_input("files/abc")
class TestTransformOpenaiInputGeminiContent:
"""Test that transform_openai_input_gemini_content creates separate requests per input."""
def test_single_text(self):
result = transform_openai_input_gemini_content(
input="hello", model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 1
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
def test_multiple_texts(self):
result = transform_openai_input_gemini_content(
input=["hello", "world"], model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 2
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
assert result["requests"][1]["content"]["parts"][0]["text"] == "world"
def test_multimodal_inputs_are_separate_requests(self):
"""Key regression test for #24209: each input becomes its own request."""
result = transform_openai_input_gemini_content(
input=["The food was delicious", IMAGE_DATA_URI],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 2
# First request is text
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
def test_dimensions_mapped_to_output_dimensionality(self):
result = transform_openai_input_gemini_content(
input="hello",
model="gemini-embedding-2-preview",
optional_params={"dimensions": 256},
)
assert result["requests"][0]["outputDimensionality"] == 256
def test_model_name_prefixed(self):
result = transform_openai_input_gemini_content(
input="hello", model="gemini-embedding-2-preview", optional_params={}
)
assert result["requests"][0]["model"] == "models/gemini-embedding-2-preview"
def test_gcs_url_input(self):
result = transform_openai_input_gemini_content(
input=[GCS_URL], model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 1
assert result["requests"][0]["content"]["parts"][0]["file_data"] is not None
def test_mixed_text_image_gcs(self):
result = transform_openai_input_gemini_content(
input=["hello", IMAGE_DATA_URI, GCS_URL],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 3
def test_nested_input_combined_embedding(self):
"""Nested list produces one request with multiple parts (combined embedding)."""
result = transform_openai_input_gemini_content(
input=[["a red shoe", IMAGE_DATA_URI]],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 1
parts = result["requests"][0]["content"]["parts"]
assert len(parts) == 2
assert parts[0]["text"] == "a red shoe"
assert parts[1]["inline_data"] is not None
def test_mixed_nested_and_flat(self):
"""Mixed nested + flat produces correct number of requests."""
result = transform_openai_input_gemini_content(
input=[["text", IMAGE_DATA_URI], "standalone"],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 2
# First: combined (2 parts)
assert len(result["requests"][0]["content"]["parts"]) == 2
# Second: standalone (1 part)
assert len(result["requests"][1]["content"]["parts"]) == 1
assert result["requests"][1]["content"]["parts"][0]["text"] == "standalone"
class TestTransformOpenaiInputGeminiEmbedContent:
"""Test transform_openai_input_gemini_embed_content (vertex_ai / embedContent path)."""
def test_text_and_image_combined(self):
result = transform_openai_input_gemini_embed_content(
input=["hello", IMAGE_DATA_URI],
model="gemini-embedding-2-preview",
optional_params={},
)
assert "content" in result
parts = result["content"]["parts"]
assert len(parts) == 2
assert parts[0]["text"] == "hello"
assert parts[1]["inline_data"] is not None
def test_gcs_url(self):
result = transform_openai_input_gemini_embed_content(
input=[GCS_URL],
model="gemini-embedding-2-preview",
optional_params={},
)
parts = result["content"]["parts"]
assert len(parts) == 1
assert parts[0]["file_data"]["file_uri"] == GCS_URL
def test_dimensions_mapped(self):
result = transform_openai_input_gemini_embed_content(
input="hello",
model="gemini-embedding-2-preview",
optional_params={"dimensions": 256},
)
assert result["outputDimensionality"] == 256
class TestProcessResponse:
"""Test that process_response sets correct indices."""
def test_single_embedding_index(self):
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}]
}
model_response = EmbeddingResponse()
result = process_response(
input="hello",
model_response=model_response,
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 1
assert result.data[0]["index"] == 0
def test_multiple_embeddings_have_correct_indices(self):
"""Regression test: indices should be 0, 1, 2... not all 0."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [
{"values": [0.1, 0.2]},
{"values": [0.3, 0.4]},
{"values": [0.5, 0.6]},
]
}
model_response = EmbeddingResponse()
result = process_response(
input=["a", "b", "c"],
model_response=model_response,
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 3
assert result.data[0]["index"] == 0
assert result.data[1]["index"] == 1
assert result.data[2]["index"] == 2
def test_multimodal_mixed_input(self):
"""process_response works with mixed text + multimodal inputs."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
}
result = process_response(
input=["hello", IMAGE_DATA_URI],
model_response=EmbeddingResponse(),
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 2
assert result.data[0]["index"] == 0
assert result.data[1]["index"] == 1
# Should count tokens only for the text element, not the image
assert result.usage.prompt_tokens > 0
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]}]
}
result = process_response(
input=[["a red shoe", IMAGE_DATA_URI]],
model_response=EmbeddingResponse(),
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 1
assert result.usage.prompt_tokens > 0
def test_nested_empty_list_raises(self):
with pytest.raises(ValueError, match="must not be empty"):
transform_openai_input_gemini_content(
input=[[]],
model="gemini-embedding-2-preview",
optional_params={},
)
def test_nested_non_string_element_raises(self):
with pytest.raises(ValueError, match="must be strings"):
transform_openai_input_gemini_content(
input=[[["doubly", "nested"]]],
model="gemini-embedding-2-preview",
optional_params={},
)