mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
3c979f0b47
commit
bf723fa9c1
2 changed files with 92 additions and 12 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue