Merge pull request #35092 from BerriAI/litellm_vertex_batch_embeddings_translation

fix(vertex_ai): translate /v1/embeddings batch rows to the Gemini embedding shape
This commit is contained in:
Mateo Wang 2026-08-14 21:52:32 -07:00 • committed by GitHub
commit d70cc14981
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 999 additions and 230 deletions

View file

@ -7,6 +7,7 @@ import re
import time
from collections.abc import Callable, Iterable, Iterator, Mapping
from typing import Any, Final, TypedDict
from urllib.parse import quote, unquote
import httpx
from httpx import Headers, Response
@ -43,6 +44,9 @@ from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
transform_openai_input_gemini_embed_content,
)
from litellm.types.files import StreamingMediaUploadConfig
from litellm.types.llms.openai import (
AllMessageValues,
@ -54,14 +58,28 @@ from litellm.types.llms.openai import (
OpenAIFilesPurpose,
PathLike,
)
from litellm.types.llms.vertex_ai import GcsBucketResponse
from litellm.types.utils import LlmProviders, ModelResponse
from litellm.types.llms.vertex_ai import GcsBucketResponse, GeminiEmbeddingInput
from litellm.types.utils import (
Embedding,
EmbeddingResponse,
LlmProviders,
ModelResponse,
Usage,
)
from ..common_utils import VertexAIError
from ..vertex_llm_base import VertexBase
_GCP_LABEL_VALUE_MAX_LEN: Final = 63
_CUSTOM_ID_RAW_LABEL_PREFIX: Final = "b32_"
_VERTEX_BATCH_KEY_FIELD: Final = "key"
_MANAGED_GCS_MODEL_PATH_PATTERN: Final = re.compile(r"publishers/[^/]+/models/([^/?]+)")
_EMBED_REQUEST_FIELD_BY_GEMINI_PARAM: Final = (
("outputDimensionality", "output_dimensionality"),
("taskType", "task_type"),
("title", "title"),
)
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN: Final = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
class _GcsObjectMetadataJson(TypedDict, total=False):
@ -164,8 +182,26 @@ def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: objec
labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk
def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object]) -> str:
def _get_litellm_batch_custom_id(vertex_output_row: Mapping[str, object]) -> str:
"""
Resolve the OpenAI `custom_id` for a Vertex batch output row.
Embedding rows carry it in the top-level `key` field that Vertex echoes back;
`generateContent` rows have no such field, so it is smuggled through request
labels instead (see `_set_litellm_batch_custom_id_labels`).
"""
key = vertex_output_row.get(_VERTEX_BATCH_KEY_FIELD)
if key is not None:
return unquote(str(key))
request_data = vertex_output_row.get("request")
labels = request_data.get("labels") if isinstance(request_data, Mapping) else None
return _get_litellm_batch_custom_id_from_labels(labels)
def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object] | None) -> str:
"""Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels)."""
if not labels:
return "unknown"
raw: Final = labels.get("litellm_custom_id_raw")
if raw:
raw_chunks: Final = [str(raw)]
@ -182,17 +218,311 @@ def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object]) -> st
return str(labels.get("litellm_custom_id", "unknown"))
def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
def _is_vertex_embeddings_batch_output_row(vertex_output_row: Mapping[str, Any]) -> bool:
"""
Whether a Vertex batch output row came from an `EmbedContentRequest`.
Successful rows hold the vector under `response.embedding.values`; failed rows only
carry `status`, so they are recognized from the singular `content` that the
embeddings request shape echoes back.
"""
if "request" not in vertex_output_row:
return False
response = vertex_output_row.get("response")
if isinstance(response, dict) and isinstance(response.get("embedding"), dict):
return True
request_data = vertex_output_row.get("request")
return bool(vertex_output_row.get("status")) and isinstance(request_data, dict) and "content" in request_data
def _openai_batch_output_row(
custom_id: str,
body: Mapping[str, Any] | None = None,
error_code: str | None = None,
error_message: str = "",
) -> _OpenAIBatchOutputRow:
"""
One row of an OpenAI batch output file. Per the OpenAI Batch spec, failed rows set
`response` to null and populate `error` instead.
"""
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": None
if body is None
else {
"status_code": 200,
"request_id": body.get("id", ""),
"body": body,
},
"error": None if error_code is None else {"code": error_code, "message": error_message},
}
def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str, int, int]:
"""
Resolve `(custom_id, index within that custom_id, group size)` for a Vertex batch
output row.
A `/v1/embeddings` entry whose `input` is an array fans out into one Vertex row per
element, tagged `<percent-encoded custom_id>#<index>/<total>` (see
`_vertex_batch_embeddings_key`), so the rows can be reassembled into a single OpenAI
response.
"""
key = vertex_output_row.get(_VERTEX_BATCH_KEY_FIELD)
if key is None:
return _get_litellm_batch_custom_id(vertex_output_row), 0, 1
match = _VERTEX_BATCH_FANNED_OUT_KEY_PATTERN.fullmatch(str(key))
if match is None:
return unquote(str(key)), 0, 1
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
def _embedding_prompt_token_count(vertex_response: Mapping[str, Any]) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
a fallback.
"""
usage_metadata = vertex_response.get("usageMetadata")
if isinstance(usage_metadata, Mapping):
return int(usage_metadata.get("promptTokenCount") or 0)
return int(vertex_response.get("tokenCount") or 0)
def _vertex_embeddings_rows_to_openai_batch_output_row(
custom_id: str,
vertex_output_rows: tuple[Mapping[str, Any], ...],
element_indices: tuple[int, ...],
element_count: int,
model: str | None,
) -> _OpenAIBatchOutputRow:
"""
Transforms the Vertex Gemini Embedding batch output rows belonging to one OpenAI
batch entry into an OpenAI batch output row holding an `/v1/embeddings` response.
Example Vertex jsonl
{"key": "id_1", "request": {...}, "response": {"embedding": {"values": [-0.015, 0.024]}, "usageMetadata": {"promptTokenCount": 2}}}
An entry that asked for several embeddings at once maps to several rows here, which
become the indexed elements of a single `data` array. One failed or missing element
fails the whole entry, since an OpenAI batch row is either a response or an error and
a partial `data` array would silently shift the remaining embeddings onto the wrong
input positions. Rows carry no `modelVersion`, so the model comes from the batch they
belong to.
"""
status = next((row["status"] for row in vertex_output_rows if row.get("status")), "")
if status:
return _openai_batch_output_row(
custom_id=custom_id,
error_code="vertex_ai_error",
error_message=status,
)
if element_indices != tuple(range(element_count)):
return _openai_batch_output_row(
custom_id=custom_id,
error_code="vertex_ai_error",
error_message=(
f"Vertex returned embeddings for input positions {list(element_indices)} "
f"of the {element_count} requested"
),
)
responses = tuple(row["response"] for row in vertex_output_rows)
token_count = sum(_embedding_prompt_token_count(response) for response in responses)
body = EmbeddingResponse(
model=model or "",
data=[
Embedding(
embedding=response["embedding"]["values"],
index=index,
object="embedding",
)
for index, response in enumerate(responses)
],
usage=Usage(prompt_tokens=token_count, total_tokens=token_count),
).model_dump()
return _openai_batch_output_row(custom_id=custom_id, body=body)
def _transform_vertex_embeddings_batch_output_to_openai(
vertex_output_rows: Iterable[Mapping[str, Any]],
model: str | None,
) -> tuple[_OpenAIBatchOutputRow, ...]:
"""
Transforms a whole Vertex Gemini Embedding batch output into OpenAI batch output
rows, one per OpenAI batch entry, in the order the entries first appear.
Rows are grouped rather than mapped one to one because a single entry can fan out
into several Vertex rows, and Vertex returns them in arbitrary order.
"""
keyed_rows = tuple((_split_vertex_batch_key(row), row) for row in vertex_output_rows)
grouped_rows = {
custom_id: tuple(group)
for custom_id, group in itertools.groupby(sorted(keyed_rows, key=lambda kr: kr[0]), key=lambda kr: kr[0][0])
}
return tuple(
_vertex_embeddings_rows_to_openai_batch_output_row(
custom_id=custom_id,
vertex_output_rows=tuple(row for _, row in grouped_rows[custom_id]),
element_indices=tuple(index for (_, index, _), _ in grouped_rows[custom_id]),
element_count=max(total for (_, _, total), _ in grouped_rows[custom_id]),
model=model,
)
for custom_id in dict.fromkeys(custom_id for (custom_id, _, _), _ in keyed_rows)
)
def _model_from_managed_gcs_url(url: str) -> str | None:
"""
Extracts the model from a LiteLLM-managed Vertex batch GCS url.
Batch inputs and their sibling outputs are stored under
`.../publishers/google/models/<model>/...`, which is the only place the model of an
embeddings batch output row can be recovered from; unlike `generateContent`
responses, embedding rows carry no `modelVersion`.
"""
match = _MANAGED_GCS_MODEL_PATH_PATTERN.search(unquote(url))
return match.group(1) if match else None
def _is_embeddings_batch_entry(openai_entry: Mapping[str, Any]) -> bool:
"""
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
OpenAI puts the target route on each line's `url` (e.g. `/v1/embeddings`); Vertex
has no equivalent per-line field, so the route decides which Vertex request shape
the line has to be translated into.
"""
url = openai_entry.get("url")
if not isinstance(url, str):
return False
path = url.split("?")[0].rstrip("/")
return path == "embeddings" or path.endswith("/embeddings")
def _openai_embedding_input_elements(
embedding_input: GeminiEmbeddingInput,
) -> tuple[str | list[str], ...]:
"""
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.
"""
if isinstance(embedding_input, list):
return tuple(embedding_input)
return (embedding_input,)
def _vertex_batch_embeddings_key(custom_id: str, index: int, total: int) -> str:
"""
The top-level `key` Vertex echoes back on an embeddings row.
An entry asking for several embeddings needs several Vertex rows, so its key also
carries the element index and the group size; `_split_vertex_batch_key` reads them
back out. The `custom_id` is percent-encoded so that a customer one ending in
`#<index>/<total>` cannot be mistaken for that tag, which would merge two entries.
"""
encoded_custom_id = quote(custom_id, safe="")
return encoded_custom_id if total < 2 else f"{encoded_custom_id}#{index}/{total}"
def _vertex_embeddings_row(key: str | None, embed_content_request: Mapping[str, Any]) -> Mapping[str, Any]:
"""
One Vertex Gemini Embedding batch input row.
The config fields live inside the `EmbedContentRequest` under their snake_case batch
names, and the OpenAI `custom_id` rides along in the top-level `key` that Vertex
echoes back.
"""
request = {
"content": embed_content_request["content"],
**{
request_field: embed_content_request[gemini_param]
for gemini_param, request_field in _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM
if gemini_param in embed_content_request
},
}
if key is None:
return {"request": request}
return {_VERTEX_BATCH_KEY_FIELD: key, "request": request}
def _openai_batch_jsonl_entry_to_vertex_embeddings_rows(
openai_entry: Mapping[str, Any],
) -> tuple[Mapping[str, Any], ...]:
"""
Transforms a single OpenAI `/v1/embeddings` batch entry into Vertex Gemini Embedding
batch rows, one per requested embedding.
Example Vertex jsonl
{"key": "id_1", "request": {"content": {"parts": [{"text": "Hello World"}]}, "output_dimensionality": 768, "task_type": "RETRIEVAL_DOCUMENT"}}
Note that `content` is singular (an `EmbedContentRequest`, not a
`GenerateContentRequest`) and that the `custom_id` round-trips through the top-level
`key`. An `EmbedContentRequest` returns exactly one vector, so an entry whose `input`
is an array fans out into one row per element and is reassembled on the way back.
The docs put the per-row config in an `embed_content_config` sibling of `request`,
but the API rejects that key outright and fails the whole batch job, so the config
fields go inside the `EmbedContentRequest` itself.
API Ref: https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/batch-prediction-genai-embeddings
"""
openai_request_body = openai_entry.get("body")
if not isinstance(openai_request_body, dict):
raise TypeError(
"`body` on /v1/embeddings batch requests must be a JSON object, but was missing or not an object"
)
embedding_input = openai_request_body.get("input")
if embedding_input is None:
raise ValueError("`input` is required on /v1/embeddings batch requests, but was not provided")
elements = _openai_embedding_input_elements(embedding_input)
if not elements:
raise ValueError("`input` on /v1/embeddings batch requests must not be empty")
embed_content_requests = tuple(
transform_openai_input_gemini_embed_content(
input=element,
model=openai_request_body.get("model", ""),
optional_params=openai_request_body,
)
for element in elements
)
custom_id = openai_entry.get("custom_id")
return tuple(
_vertex_embeddings_row(
key=None
if custom_id is None
else _vertex_batch_embeddings_key(
custom_id=str(custom_id),
index=index,
total=len(embed_content_requests),
),
embed_content_request=embed_content_request,
)
for index, embed_content_request in enumerate(embed_content_requests)
)
def _openai_batch_jsonl_entry_to_vertex_rows(
openai_entry: dict[str, Any],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]],
) -> dict[str, Any]:
) -> tuple[Mapping[str, Any], ...]:
"""
Transforms a single OpenAI JSONL batch entry into its Vertex wrapped request.
Transforms a single OpenAI JSONL batch entry into the Vertex rows it maps to.
jsonl body for vertex is {"request": <request_body>}
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
"""
if _is_embeddings_batch_entry(openai_entry):
return _openai_batch_jsonl_entry_to_vertex_embeddings_rows(openai_entry)
openai_request_body: Final = openai_entry.get("body") or {}
vertex_request_body: Final = _transform_request_body(
messages=openai_request_body.get("messages", []),
@ -209,7 +539,7 @@ def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
vertex_request_body["labels"] = {}
_set_litellm_batch_custom_id_labels(vertex_request_body["labels"], custom_id)
return {"request": vertex_request_body}
return ({"request": vertex_request_body},)
def _iter_stripped_lines(raw_lines: Iterable[str | bytes]) -> Iterator[str]:
@ -312,10 +642,10 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
def _iter_vertex_jsonl_chunks(self) -> Iterator[bytes]:
first = True
for entry in _iter_openai_jsonl_entries(self._openai_file_content):
wrapped = _openai_batch_jsonl_entry_to_vertex_wrapped_request(entry, self._map_openai_to_vertex_params)
prefix = b"" if first else b"\n"
first = False
yield prefix + json.dumps(wrapped).encode("utf-8")
for wrapped in _openai_batch_jsonl_entry_to_vertex_rows(entry, self._map_openai_to_vertex_params):
prefix = b"" if first else b"\n"
first = False
yield prefix + json.dumps(wrapped).encode("utf-8")
def iter_bytes(self) -> Iterator[bytes]:
return self._iter_vertex_jsonl_chunks()
@ -667,6 +997,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
transformed_content: Final = self._try_transform_vertex_batch_output_to_openai(
content=content,
logging_obj=logging_obj,
model=_model_from_managed_gcs_url(str(raw_response.request.url)),
)
if transformed_content != content:
# Create a new response with transformed content and updated Content-Length
@ -688,7 +1019,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
return HttpxBinaryResponseContent(response=raw_response)
def _try_transform_vertex_batch_output_to_openai(
self, content: bytes, logging_obj: LiteLLMLoggingObj | None = None
self,
content: bytes,
logging_obj: LiteLLMLoggingObj | None = None,
model: str | None = None,
) -> bytes:
"""
Try to transform Vertex AI batch output to OpenAI format.
@ -730,7 +1064,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# first line is not valid UTF-8/JSON) raises and falls through to the
# passthrough below, leaving the content untouched.
first_row: Final = _parse_vertex_batch_output_row(first_line)
is_vertex_batch_output: Final = (
is_vertex_batch_output: Final = _is_vertex_embeddings_batch_output_row(first_row) or (
"request" in first_row
and "response" in first_row
and "processed_time" in first_row
@ -763,11 +1097,23 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
request=httpx.Request(method="POST", url="https://example.com"),
)
all_lines = itertools.chain((first_line,), lines)
# Embedding rows are grouped by `custom_id` rather than transformed one at a
# time, since an entry that asked for several embeddings comes back as
# several rows, in arbitrary order.
if _is_vertex_embeddings_batch_output_row(first_row):
openai_outputs = _transform_vertex_embeddings_batch_output_to_openai(
vertex_output_rows=(json.loads(line) for line in all_lines),
model=model,
)
return b"\n".join(json.dumps(openai_output).encode("utf-8") for openai_output in openai_outputs)
# Transform each row straight into the output buffer, so peak memory
# stays at ~one row plus the output. If any row fails, return the
# original content unchanged.
output = bytearray()
for line in itertools.chain([first_line], lines):
for line in all_lines:
try:
openai_output = self._transform_single_vertex_batch_output_to_openai(
vertex_output=_parse_vertex_batch_output_row(line),
@ -798,25 +1144,18 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
Transform a single Vertex AI batch output line to OpenAI format.
Uses the existing VertexGeminiConfig transformation for the response.
"""
# Extract custom_id from request labels (prefer raw for OpenAI round-trip)
request_data: Final = vertex_output.get("request", {})
labels: Final[Mapping[str, object]] = request_data.get("labels", {}) or {}
custom_id: Final = _get_litellm_batch_custom_id_from_labels(labels)
custom_id: Final = _get_litellm_batch_custom_id(vertex_output)
# Check if there's an error
status: Final = vertex_output.get("status", "")
has_error: Final = bool(status)
if has_error:
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": None,
"error": {
"code": "vertex_ai_error",
"message": status,
},
}
return _openai_batch_output_row(
custom_id=custom_id,
error_code="vertex_ai_error",
error_message=status,
)
# Transform successful response using existing transformation
vertex_response: Final = vertex_output.get("response", {})
@ -842,24 +1181,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
response_dict: Final = transformed_response.model_dump()
# Return in OpenAI batch format
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": response_dict.get("id", ""),
"body": response_dict,
},
"error": None,
}
return _openai_batch_output_row(custom_id=custom_id, body=response_dict)
except Exception as e:
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": None,
"error": {
"code": "transformation_error",
"message": f"Failed to transform response: {e}",
},
}
return _openai_batch_output_row(
custom_id=custom_id,
error_code="transformation_error",
error_message=f"Failed to transform response: {e}",
)

View file

@ -37,7 +37,7 @@ from litellm.llms.vertex_ai.files.transformation import (
_get_litellm_batch_custom_id_from_labels,
_iter_openai_jsonl_entries,
_iter_openai_jsonl_lines,
_openai_batch_jsonl_entry_to_vertex_wrapped_request,
_openai_batch_jsonl_entry_to_vertex_rows,
)
from litellm.types.llms.openai import CreateFileRequest
@ -84,8 +84,9 @@ def _reference_vertex_jsonl_string(cfg: VertexAIFilesConfig, content: str) -> st
transform, so the streaming path can be checked against it for parity."""
entries = [json.loads(line) for line in content.splitlines() if line.strip()]
return "\n".join(
json.dumps(_openai_batch_jsonl_entry_to_vertex_wrapped_request(entry, cfg._map_openai_to_vertex_params))
json.dumps(row)
for entry in entries
for row in _openai_batch_jsonl_entry_to_vertex_rows(entry, cfg._map_openai_to_vertex_params)
)