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>
This commit is contained in:
milan 2026-08-15 00:29:35 +00:00
parent f2a10f6331
commit 0d0c712df7
2 changed files with 42 additions and 11 deletions

View file

@ -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 `<percent-encoded custom_id>#<index>/<total>` (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)
)

View file

@ -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(
[