From 0d0c712df7872aff1fde85b78940dadd4296e50a Mon Sep 17 00:00:00 2001 From: milan Date: Sat, 15 Aug 2026 00:29:35 +0000 Subject: [PATCH] fix(vertex_ai): fail an embeddings batch entry whose fan-out came back incomplete Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llms/vertex_ai/files/transformation.py | 39 +++++++++++++------ .../test_vertex_ai_files_transformation.py | 14 +++++++ 2 files changed, 42 insertions(+), 11 deletions(-) diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index d5363a08a92..b7f91bfba0d 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -259,9 +259,10 @@ def _openai_batch_output_row( } -def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str, int]: +def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str, int, int]: """ - Resolve `(custom_id, index within that custom_id)` for a Vertex batch output row. + 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 `#/` (see @@ -270,11 +271,11 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str, """ key = vertex_output_row.get(_VERTEX_BATCH_KEY_FIELD) if key is None: - return _get_litellm_batch_custom_id(vertex_output_row), 0 + 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 - return unquote(match["custom_id"]), int(match["index"]) + 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: @@ -293,6 +294,8 @@ def _embedding_prompt_token_count(vertex_response: Mapping[str, Any]) -> int: 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: """ @@ -303,9 +306,11 @@ def _vertex_embeddings_rows_to_openai_batch_output_row( {"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 element fails the - whole entry, since an OpenAI batch row is either a response or an error. Rows carry - no `modelVersion`, so the model comes from the batch they belong to. + 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: @@ -315,6 +320,16 @@ def _vertex_embeddings_rows_to_openai_batch_output_row( 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( @@ -345,16 +360,18 @@ def _transform_vertex_embeddings_batch_output_to_openai( """ keyed_rows = tuple((_split_vertex_batch_key(row), row) for row in vertex_output_rows) grouped_rows = { - custom_id: tuple(row for _, row in group) + 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=grouped_rows[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) + for custom_id in dict.fromkeys(custom_id for (custom_id, _, _), _ in keyed_rows) ) diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index f95a63e4421..c280e44ff4d 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -1802,6 +1802,20 @@ class TestVertexEmbeddingsBatchOutputTranslation: assert result["response"] is None assert result["error"]["message"] == "Quota exceeded" + def test_should_fail_the_whole_entry_when_a_fanned_out_row_is_missing(self, config): + """A partial `data` array would shift embeddings onto the wrong input positions.""" + (result,) = self._transform( + config, + [self._vertex_embeddings_output_row(key="request-1#1/2")], + ) + + assert result["custom_id"] == "request-1" + assert result["response"] is None + assert result["error"]["code"] == "vertex_ai_error" + assert result["error"]["message"] == ( + "Vertex returned embeddings for input positions [1] of the 2 requested" + ) + def test_should_end_to_end_round_trip_a_fanned_out_embeddings_batch(self, config): first_row, second_row = _wrap_entries( [