fix(vertex_ai): percent-encode the custom_id in fanned-out vertex batch keys

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-07-29 20:50:27 +00:00
parent 3c979f0b47
commit bf723fa9c1
2 changed files with 92 additions and 12 deletions

View file

@ -17,7 +17,7 @@ from typing import (
Tuple,
Union,
)
from urllib.parse import unquote
from urllib.parse import quote, unquote
import httpx
from httpx import Headers, Response
@ -87,7 +87,7 @@ _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM = {
"taskType": "task_type",
"title": "title",
}
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN = re.compile(r"(?P<custom_id>.*)#(?P<index>\d+)/(?P<total>\d+)")
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
def _sanitize_gcp_label_value(value: str) -> str:
@ -160,7 +160,7 @@ def _get_litellm_batch_custom_id(vertex_output_row: Mapping[str, Any]) -> str:
"""
key = vertex_output_row.get(_VERTEX_BATCH_KEY_FIELD)
if key is not None:
return str(key)
return unquote(str(key))
request_data = vertex_output_row.get("request") or {}
return _get_litellm_batch_custom_id_from_labels(request_data.get("labels") or {})
@ -228,14 +228,17 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str,
Resolve `(custom_id, index within that custom_id)` for a Vertex batch output row.
A `/v1/embeddings` entry whose `input` is an array fans out into one Vertex row per
element, tagged `<custom_id>#<index>/<total>` (see `_vertex_batch_embeddings_key`),
so the rows can be reassembled into a single OpenAI response.
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 = _get_litellm_batch_custom_id(vertex_output_row)
match = _VERTEX_BATCH_FANNED_OUT_KEY_PATTERN.fullmatch(key)
if match is None or int(match["total"]) < 2:
return key, 0
return match["custom_id"], int(match["index"])
key = vertex_output_row.get(_VERTEX_BATCH_KEY_FIELD)
if key is None:
return _get_litellm_batch_custom_id(vertex_output_row), 0
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"])
def _vertex_embeddings_rows_to_openai_batch_output_row(
@ -359,9 +362,11 @@ def _vertex_batch_embeddings_key(custom_id: str, index: int, total: int) -> str:
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. Entries asking for a single embedding keep their bare `custom_id`.
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.
"""
return custom_id if total < 2 else f"{custom_id}#{index}/{total}"
encoded_custom_id = quote(custom_id, safe="")
return encoded_custom_id if total < 2 else f"{encoded_custom_id}#{index}/{total}"
def _openai_batch_jsonl_entry_to_vertex_embeddings_rows(

View file

@ -1476,6 +1476,19 @@ class TestVertexEmbeddingsBatchInputTranslation:
assert row["key"] == "request-1"
def test_should_encode_a_custom_id_that_looks_like_a_fan_out_tag(self):
"""A customer custom_id ending in `#<i>/<n>` must not read back as fan-out metadata."""
(row,) = _wrap_entries(
[
_embeddings_entry(
custom_id="request-1#0/2",
body={"model": "gemini-embedding-2", "input": "hello world"},
)
]
)
assert row["key"] == "request-1%230%2F2"
def test_should_combine_a_nested_input_into_one_multipart_row(self):
"""Nested arrays are the combined-embedding shape, as on the online path."""
(row,) = _wrap_entries(
@ -1711,6 +1724,68 @@ class TestVertexEmbeddingsBatchOutputTranslation:
assert len(results[0]["response"]["body"]["data"]) == 2
assert len(results[1]["response"]["body"]["data"]) == 1
def test_should_not_merge_an_entry_whose_custom_id_looks_like_a_fan_out_tag(self, config):
"""`request-1#0/2` is a legal custom_id, and a distinct entry from `request-1`."""
lookalike_row, plain_row = _wrap_entries(
[
_embeddings_entry(
custom_id="request-1#0/2",
body={"model": "gemini-embedding-2", "input": "lookalike"},
),
_embeddings_entry(
custom_id="request-1",
body={"model": "gemini-embedding-2", "input": "plain"},
),
]
)
results = self._transform(
config,
[
{**row, "status": "", "response": {"embedding": {"values": values}}}
for row, values in ((lookalike_row, [0.1]), (plain_row, [0.2]))
],
)
assert [result["custom_id"] for result in results] == [
"request-1#0/2",
"request-1",
]
assert [
result["response"]["body"]["data"][0]["embedding"] for result in results
] == [[0.1], [0.2]]
def test_should_round_trip_a_fan_out_of_a_custom_id_holding_the_separator(self, config):
rows = _wrap_entries(
[
_embeddings_entry(
custom_id="request#1/1",
body={
"model": "gemini-embedding-2",
"input": ["first", "second"],
},
)
]
)
assert [row["key"] for row in rows] == [
"request%231%2F1#0/2",
"request%231%2F1#1/2",
]
(result,) = self._transform(
config,
[
{**row, "status": "", "response": {"embedding": {"values": values}}}
for row, values in zip(reversed(rows), ([0.3], [0.1]))
],
)
assert result["custom_id"] == "request#1/1"
assert [
embedding["embedding"] for embedding in result["response"]["body"]["data"]
] == [[0.1], [0.3]]
def test_should_fail_the_whole_entry_when_one_of_its_rows_failed(self, config):
"""An OpenAI batch row is either a response or an error, never both."""
(result,) = self._transform(