mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
d70cc14981
3 changed files with 999 additions and 230 deletions
|
|
@ -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}",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue