diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 78a16b002e7..90fc5fae082 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -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.*)#(?P\d+)/(?P\d+)") +_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN = re.compile(r"(?P[^#]*)#(?P\d+)/(?P\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 `#/` (see `_vertex_batch_embeddings_key`), - so the rows can be reassembled into a single OpenAI response. + element, tagged `#/` (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 + `#/` 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( 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 aef7f8b4684..f95a63e4421 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 @@ -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 `#/` 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(