test(vertex_ai): cover duplicated fan-out rows in embeddings batch reassembly
Some checks failed
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

Also ruff-formats the batch transformation test file, which the formatter
gate flags once the file is touched.
This commit is contained in:
mateo-berri 2026-08-14 17:37:07 -07:00
parent 0d0c712df7
commit de77711cf9

View file

@ -32,40 +32,26 @@ class TestParseGcsUri:
def test_should_parse_standard_gs_uri(self, config):
file_id = "gs://my-bucket/litellm-vertex-files/path/to/object.jsonl"
bucket, encoded = config._parse_gcs_uri(
file_id, litellm_params={"gcs_bucket_name": "my-bucket"}
)
bucket, encoded = config._parse_gcs_uri(file_id, litellm_params={"gcs_bucket_name": "my-bucket"})
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote(
"litellm-vertex-files/path/to/object.jsonl", safe=""
)
assert encoded == urllib.parse.quote("litellm-vertex-files/path/to/object.jsonl", safe="")
def test_should_parse_uri_with_nested_publisher_path(self, config):
uri = "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
bucket, encoded = config._parse_gcs_uri(
uri, litellm_params={"gcs_bucket_name": "litellm-local"}
)
bucket, encoded = config._parse_gcs_uri(uri, litellm_params={"gcs_bucket_name": "litellm-local"})
assert bucket == "litellm-local"
expected_path = (
"litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
)
expected_path = "litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
assert encoded == urllib.parse.quote(expected_path, safe="")
def test_should_handle_url_encoded_input(self, config):
encoded_uri = urllib.parse.quote(
"gs://my-bucket/litellm-vertex-files/some/path", safe=""
)
bucket, encoded = config._parse_gcs_uri(
encoded_uri, litellm_params={"gcs_bucket_name": "my-bucket"}
)
encoded_uri = urllib.parse.quote("gs://my-bucket/litellm-vertex-files/some/path", safe="")
bucket, encoded = config._parse_gcs_uri(encoded_uri, litellm_params={"gcs_bucket_name": "my-bucket"})
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("litellm-vertex-files/some/path", safe="")
def test_should_reject_bucket_only(self, config):
with pytest.raises(ValueError, match="object name"):
config._parse_gcs_uri(
"gs://my-bucket", litellm_params={"gcs_bucket_name": "my-bucket"}
)
config._parse_gcs_uri("gs://my-bucket", litellm_params={"gcs_bucket_name": "my-bucket"})
def test_should_reject_no_gs_prefix(self, config):
with pytest.raises(ValueError, match="gs://"):
@ -110,9 +96,7 @@ class TestParseGcsUri:
"gs://my-bucket/private/object.txt",
litellm_params={
"gcs_bucket_name": "my-bucket",
"_litellm_internal_model_credentials": {
"allow_legacy_cloud_file_ids": True
},
"_litellm_internal_model_credentials": {"allow_legacy_cloud_file_ids": True},
},
)
@ -176,7 +160,6 @@ class TestCreateFileUrl:
class TestTransformRetrieveFile:
def test_should_build_correct_gcs_metadata_url(self, config):
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"
url, params = config.transform_retrieve_file_request(
@ -184,13 +167,8 @@ class TestTransformRetrieveFile:
optional_params={},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
expected_encoded = urllib.parse.quote(
"litellm-vertex-files/path/to/file.jsonl", safe=""
)
assert (
url
== f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{expected_encoded}"
)
expected_encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{expected_encoded}"
assert params == {}
def test_should_return_openai_file_object_from_gcs_response(self, config):
@ -237,7 +215,6 @@ class TestTransformRetrieveFile:
class TestTransformFileContent:
def test_should_build_gcs_media_download_url(self, config):
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"
url, params = config.transform_file_content_request(
@ -246,10 +223,7 @@ class TestTransformFileContent:
litellm_params={"gcs_bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
url
== f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}?alt=media"
)
assert url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}?alt=media"
assert params == {}
def test_should_return_binary_response_content(self, config):
@ -269,9 +243,7 @@ class TestTransformFileContent:
assert isinstance(result, HttpxBinaryResponseContent)
assert result.response.content == b'{"line": 1}\n{"line": 2}\n'
def test_should_not_mutate_caller_logging_obj_for_batch_output_transform(
self, config, monkeypatch
):
def test_should_not_mutate_caller_logging_obj_for_batch_output_transform(self, config, monkeypatch):
original_model = "vertex_ai/original-model"
original_start_time = 123.456
original_optional_params = {"temperature": 0.1}
@ -283,9 +255,7 @@ class TestTransformFileContent:
"processed_time": "2024-11-01T18:13:16.826+00:00",
"request": {"labels": {"litellm_custom_id": "request-1"}},
"response": {
"candidates": [
{"content": {"parts": [{"text": "ok"}], "role": "model"}}
],
"candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
"modelVersion": "gemini-2.0-flash-001@default",
},
}
@ -308,9 +278,7 @@ class TestTransformFileContent:
captured["logging_obj"] = logging_obj
logging_obj.model = "gemini-2.0-flash-001"
logging_obj.start_time = 789.0
return {
"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]
}
return {"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]}
monkeypatch.setattr(
config,
@ -330,9 +298,7 @@ class TestTransformFileContent:
assert logging_obj.optional_params == original_optional_params
assert result.response is not raw_response
def test_should_skip_batch_output_transformation_when_opt_out_flag_set(
self, config, monkeypatch
):
def test_should_skip_batch_output_transformation_when_opt_out_flag_set(self, config, monkeypatch):
"""When `litellm.disable_vertex_batch_output_transformation` is True the
Vertex predictions.jsonl content must be returned untouched, so callers
that parse raw `candidates`/`modelVersion` keep working."""
@ -344,9 +310,7 @@ class TestTransformFileContent:
"processed_time": "2024-11-01T18:13:16.826+00:00",
"request": {"labels": {"litellm_custom_id": "request-1"}},
"response": {
"candidates": [
{"content": {"parts": [{"text": "ok"}], "role": "model"}}
],
"candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
"modelVersion": "gemini-2.0-flash-001@default",
},
}
@ -358,9 +322,7 @@ class TestTransformFileContent:
request=httpx.Request("GET", "https://example.com"),
)
monkeypatch.setattr(
litellm, "disable_vertex_batch_output_transformation", True, raising=False
)
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
result = config.transform_file_content_response(
raw_response=raw_response,
@ -381,9 +343,7 @@ class TestTransformDeleteFile:
litellm_params={"gcs_bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}"
)
assert url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}"
assert params == {}
def test_should_return_file_deleted_with_reconstructed_id(self, config):
@ -393,9 +353,7 @@ class TestTransformDeleteFile:
"litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc",
safe="",
)
mock_request.url = (
f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded_name}"
)
mock_request.url = f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded_name}"
raw_response.request = mock_request
result = config.transform_delete_file_response(
@ -407,10 +365,7 @@ class TestTransformDeleteFile:
assert isinstance(result, FileDeleted)
assert result.deleted is True
assert result.object == "file"
assert (
result.id
== "gs://my-bucket/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc"
)
assert result.id == "gs://my-bucket/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc"
def test_should_fallback_to_deleted_id_when_no_request(self, config):
raw_response = MagicMock(spec=httpx.Response)
@ -435,9 +390,7 @@ class TestTransformDeleteFile:
raw_response = MagicMock(spec=httpx.Response)
mock_request = MagicMock()
encoded_object = urllib.parse.quote("path/to/file.jsonl", safe="")
mock_request.url = (
f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded_object}"
)
mock_request.url = f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded_object}"
raw_response.request = mock_request
result = config.transform_delete_file_response(
@ -466,8 +419,7 @@ class TestTransformDeleteFile:
)
assert result.id == (
"gs://prod-bucket/litellm-vertex-files/publishers/google/"
"models/gemini-2.0-flash-001/abc-123"
"gs://prod-bucket/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
)
@ -504,9 +456,7 @@ class TestVertexBatchOutputTransformation:
}
content = json.dumps(vertex_output).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
result = json.loads(transformed_content.decode("utf-8"))
# Verify OpenAI format
@ -548,9 +498,7 @@ class TestVertexBatchOutputTransformation:
}
content = json.dumps(vertex_output).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
result = json.loads(transformed_content.decode("utf-8"))
# Per OpenAI Batch output spec, error entries set response to null
@ -584,9 +532,7 @@ class TestVertexBatchOutputTransformation:
}
class _RaisingGeminiConfig(VertexGeminiConfig):
def _transform_google_generate_content_to_openai_model_response(
self, *args, **kwargs
):
def _transform_google_generate_content_to_openai_model_response(self, *args, **kwargs):
raise ValueError("simulated transform failure")
mock_response = httpx.Response(
@ -637,9 +583,7 @@ class TestVertexBatchOutputTransformation:
}
content = json.dumps(vertex_output).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
result = json.loads(transformed_content.decode("utf-8"))
assert result["custom_id"] == "myrequest-1"
@ -651,9 +595,7 @@ class TestVertexBatchOutputTransformation:
"status": "",
"processed_time": "2024-11-01T18:13:16.826+00:00",
"request": {
"contents": [
{"role": "user", "parts": [{"text": "First request"}]}
],
"contents": [{"role": "user", "parts": [{"text": "First request"}]}],
"labels": {"litellm_custom_id": "request-1"},
},
"response": {
@ -678,9 +620,7 @@ class TestVertexBatchOutputTransformation:
"status": "",
"processed_time": "2024-11-01T18:13:17.826+00:00",
"request": {
"contents": [
{"role": "user", "parts": [{"text": "Second request"}]}
],
"contents": [{"role": "user", "parts": [{"text": "Second request"}]}],
"labels": {"litellm_custom_id": "request-2"},
},
"response": {
@ -703,12 +643,8 @@ class TestVertexBatchOutputTransformation:
},
]
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode(
"utf-8"
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
lines = transformed_content.decode("utf-8").strip().split("\n")
assert len(lines) == 2
@ -718,14 +654,12 @@ class TestVertexBatchOutputTransformation:
assert "id" in result
assert "response" in result
assert result["response"]["status_code"] == 200
assert result["custom_id"] == f"request-{i+1}"
assert result["custom_id"] == f"request-{i + 1}"
body = result["response"]["body"]
assert "choices" in body
assert len(body["choices"]) > 0
def test_transform_vertex_batch_output_with_first_line_prompt_feedback(
self, config, monkeypatch
):
def test_transform_vertex_batch_output_with_first_line_prompt_feedback(self, config, monkeypatch):
"""Test that promptFeedback-only first lines are detected as Vertex batch output."""
vertex_outputs = [
{
@ -751,9 +685,7 @@ class TestVertexBatchOutputTransformation:
logging_obj,
mock_httpx_response,
):
return {
"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]
}
return {"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]}
monkeypatch.setattr(
config,
@ -761,15 +693,9 @@ class TestVertexBatchOutputTransformation:
mock_transform_single,
)
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode(
"utf-8"
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
results = [
json.loads(line) for line in transformed_content.decode("utf-8").split("\n")
]
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
results = [json.loads(line) for line in transformed_content.decode("utf-8").split("\n")]
assert [result["custom_id"] for result in results] == [
"blocked-request",
@ -786,9 +712,7 @@ class TestVertexBatchOutputTransformation:
}
content = json.dumps(non_batch_output).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
assert transformed_content == content
@ -818,9 +742,7 @@ class TestVertexBatchOutputTransformation:
id(mock_httpx_response),
)
)
return {
"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]
}
return {"custom_id": vertex_output["request"]["labels"]["litellm_custom_id"]}
monkeypatch.setattr(
config,
@ -828,12 +750,8 @@ class TestVertexBatchOutputTransformation:
mock_transform_single,
)
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode(
"utf-8"
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
content = "\n".join(json.dumps(output) for output in vertex_outputs).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
assert len(transformed_content.decode("utf-8").strip().split("\n")) == 2
assert len(set(helper_ids)) == 1
@ -841,17 +759,13 @@ class TestVertexBatchOutputTransformation:
def test_non_batch_output_passthrough(self, config):
"""Test that non-batch output is returned as-is"""
regular_content = b"This is just a regular file content"
transformed_content = config._try_transform_vertex_batch_output_to_openai(
regular_content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(regular_content)
assert transformed_content == regular_content
def test_invalid_json_passthrough(self, config):
"""Test that invalid JSON is returned as-is"""
invalid_content = b'{"invalid": json content}'
transformed_content = config._try_transform_vertex_batch_output_to_openai(
invalid_content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(invalid_content)
assert transformed_content == invalid_content
def test_binary_content_passthrough(self, config):
@ -903,9 +817,7 @@ class TestVertexBatchOutputTransformation:
},
}
content = ("\n".join(json.dumps(vertex_row(i)) for i in range(4000))).encode(
"utf-8"
)
content = ("\n".join(json.dumps(vertex_row(i)) for i in range(4000))).encode("utf-8")
def list_pipeline() -> bytes:
gemini_config = VertexGeminiConfig()
@ -944,9 +856,7 @@ class TestVertexBatchOutputTransformation:
finally:
tracemalloc.stop()
streaming_peak = peak_of(
lambda: config._try_transform_vertex_batch_output_to_openai(content)
)
streaming_peak = peak_of(lambda: config._try_transform_vertex_batch_output_to_openai(content))
list_peak = peak_of(list_pipeline)
assert streaming_peak < list_peak * 0.75, (
@ -999,9 +909,9 @@ class TestTryTransformDoesNotMutateCallerLoggingObj:
logging_obj=logging_obj,
)
assert (
logging_obj.model == sentinel_model
), "logging_obj.model was mutated by _try_transform_vertex_batch_output_to_openai"
assert logging_obj.model == sentinel_model, (
"logging_obj.model was mutated by _try_transform_vertex_batch_output_to_openai"
)
def test_should_not_overwrite_start_time_on_caller_logging_obj(self, config):
sentinel_start = 1234567890.0
@ -1014,9 +924,9 @@ class TestTryTransformDoesNotMutateCallerLoggingObj:
logging_obj=logging_obj,
)
assert (
logging_obj.start_time == sentinel_start
), "logging_obj.start_time was mutated by _try_transform_vertex_batch_output_to_openai"
assert logging_obj.start_time == sentinel_start, (
"logging_obj.start_time was mutated by _try_transform_vertex_batch_output_to_openai"
)
def test_should_not_overwrite_optional_params_on_caller_logging_obj(self, config):
sentinel_params = {"temperature": 0.5, "top_p": 0.9}
@ -1028,9 +938,9 @@ class TestTryTransformDoesNotMutateCallerLoggingObj:
logging_obj=logging_obj,
)
assert (
logging_obj.optional_params is sentinel_params
), "logging_obj.optional_params was replaced by _try_transform_vertex_batch_output_to_openai"
assert logging_obj.optional_params is sentinel_params, (
"logging_obj.optional_params was replaced by _try_transform_vertex_batch_output_to_openai"
)
assert logging_obj.optional_params == {
"temperature": 0.5,
"top_p": 0.9,
@ -1060,9 +970,7 @@ def _wrap_entries(openai_jsonl_content):
return [
row
for entry in openai_jsonl_content
for row in _openai_batch_jsonl_entry_to_vertex_rows(
entry, cfg._map_openai_to_vertex_params
)
for row in _openai_batch_jsonl_entry_to_vertex_rows(entry, cfg._map_openai_to_vertex_params)
]
@ -1123,9 +1031,7 @@ class TestVertexBatchCustomIdLabels:
assert "litellm_custom_id_raw_1" in labels_a
assert "litellm_custom_id_raw_1" in labels_b
assert labels_a["litellm_custom_id_raw"] == labels_b["litellm_custom_id_raw"]
assert (
labels_a["litellm_custom_id_raw_1"] != labels_b["litellm_custom_id_raw_1"]
)
assert labels_a["litellm_custom_id_raw_1"] != labels_b["litellm_custom_id_raw_1"]
assert _get_litellm_batch_custom_id_from_labels(labels_a) == custom_id_a
assert _get_litellm_batch_custom_id_from_labels(labels_b) == custom_id_b
@ -1134,12 +1040,12 @@ class TestVertexBatchCustomIdLabels:
openai_jsonl_content = [
{
"custom_id": f"request-{i+1}",
"custom_id": f"request-{i + 1}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "gemini-1.5-flash-001",
"messages": [{"role": "user", "content": f"Question {i+1}"}],
"messages": [{"role": "user", "content": f"Question {i + 1}"}],
},
}
for i in range(3)
@ -1150,11 +1056,8 @@ class TestVertexBatchCustomIdLabels:
assert len(vertex_jsonl_content) == 3
for i, vertex_request in enumerate(vertex_jsonl_content):
expected_custom_id = f"request-{i+1}"
assert (
vertex_request["request"]["labels"]["litellm_custom_id"]
== expected_custom_id
)
expected_custom_id = f"request-{i + 1}"
assert vertex_request["request"]["labels"]["litellm_custom_id"] == expected_custom_id
raw_label = vertex_request["request"]["labels"]["litellm_custom_id_raw"]
assert raw_label != expected_custom_id
assert _sanitize_gcp_label_value(raw_label) == raw_label
@ -1201,9 +1104,7 @@ class TestVertexBatchCustomIdLabels:
vertex_input = _wrap_entries(openai_input)
# Verify both labels are GCP-safe and encoded raw preserves round-trip.
assert (
vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1"
)
assert vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1"
raw_label = vertex_input[0]["request"]["labels"]["litellm_custom_id_raw"]
assert raw_label != "MyRequest-1"
assert _sanitize_gcp_label_value(raw_label) == raw_label
@ -1231,9 +1132,7 @@ class TestVertexBatchCustomIdLabels:
# Step 3: Transform Vertex AI output back to OpenAI format
content = json.dumps(vertex_output).encode("utf-8")
transformed_content = config._try_transform_vertex_batch_output_to_openai(
content
)
transformed_content = config._try_transform_vertex_batch_output_to_openai(content)
openai_output = json.loads(transformed_content.decode("utf-8"))
# Step 4: Verify custom_id was preserved (original casing, not sanitized label)
@ -1269,9 +1168,7 @@ class TestVertexBatchCustomIdLabels:
vertex_input = _wrap_entries(openai_input)
# Verify both labels are safe for GCP labels.
assert (
vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1"
)
assert vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1"
raw_label = vertex_input[0]["request"]["labels"]["litellm_custom_id_raw"]
assert raw_label != "MyRequest-1"
assert _sanitize_gcp_label_value(raw_label) == raw_label
@ -1280,26 +1177,15 @@ class TestVertexBatchCustomIdLabels:
class TestConfiguredBucketNameResolution:
def test_should_resolve_new_gcs_bucket_name_key(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name({"gcs_bucket_name": "my-new-bucket"})
== "my-new-bucket"
)
assert config._get_configured_bucket_name({"gcs_bucket_name": "my-new-bucket"}) == "my-new-bucket"
def test_should_resolve_legacy_bucket_name_key(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name({"bucket_name": "my-legacy-bucket"})
== "my-legacy-bucket"
)
assert config._get_configured_bucket_name({"bucket_name": "my-legacy-bucket"}) == "my-legacy-bucket"
def test_should_prefer_new_key_over_legacy(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name(
{"gcs_bucket_name": "new", "bucket_name": "legacy"}
)
== "new"
)
assert config._get_configured_bucket_name({"gcs_bucket_name": "new", "bucket_name": "legacy"}) == "new"
def test_should_fall_back_to_env(self, config, monkeypatch):
monkeypatch.setenv("GCS_BUCKET_NAME", "env-bucket")
@ -1427,9 +1313,7 @@ class TestVertexEmbeddingsBatchInputTranslation:
def test_should_raise_when_input_empty(self):
with pytest.raises(ValueError, match="must not be empty"):
_wrap_entries(
[_embeddings_entry(body={"model": "gemini-embedding-2", "input": []})]
)
_wrap_entries([_embeddings_entry(body={"model": "gemini-embedding-2", "input": []})])
def test_should_fan_an_input_array_out_into_one_row_per_element(self):
"""
@ -1466,13 +1350,7 @@ class TestVertexEmbeddingsBatchInputTranslation:
]
def test_should_keep_the_bare_custom_id_for_single_element_arrays(self):
(row,) = _wrap_entries(
[
_embeddings_entry(
body={"model": "gemini-embedding-2", "input": ["only one"]}
)
]
)
(row,) = _wrap_entries([_embeddings_entry(body={"model": "gemini-embedding-2", "input": ["only one"]})])
assert row["key"] == "request-1"
@ -1533,9 +1411,7 @@ class TestVertexEmbeddingsBatchInputTranslation:
]
)
assert row["request"]["contents"] == [
{"role": "user", "parts": [{"text": "Hello"}]}
]
assert row["request"]["contents"] == [{"role": "user", "parts": [{"text": "Hello"}]}]
assert row["request"]["labels"]["litellm_custom_id"] == "request-1"
assert "key" not in row
@ -1553,9 +1429,7 @@ class TestVertexEmbeddingsBatchInputTranslation:
]
)
assert row["request"]["contents"] == [
{"role": "user", "parts": [{"text": "Hello"}]}
]
assert row["request"]["contents"] == [{"role": "user", "parts": [{"text": "Hello"}]}]
def test_should_translate_each_line_by_its_own_url(self):
chat_row, embeddings_row = _wrap_entries(
@ -1603,10 +1477,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
logging_obj=MagicMock(),
litellm_params={},
)
return [
json.loads(line)
for line in result.response.content.decode("utf-8").split("\n")
]
return [json.loads(line) for line in result.response.content.decode("utf-8").split("\n")]
def test_should_transform_embeddings_output_to_openai_batch_row(self, config):
(result,) = self._transform(config, [self._vertex_embeddings_output_row()])
@ -1616,9 +1487,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
assert result["response"]["status_code"] == 200
body = result["response"]["body"]
assert body["object"] == "list"
assert body["data"] == [
{"embedding": [-0.015, 0.024], "index": 0, "object": "embedding"}
]
assert body["data"] == [{"embedding": [-0.015, 0.024], "index": 0, "object": "embedding"}]
assert body["usage"]["prompt_tokens"] == 2
assert body["usage"]["total_tokens"] == 2
@ -1645,20 +1514,14 @@ class TestVertexEmbeddingsBatchOutputTranslation:
)
url = f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{object_path}?alt=media"
(result,) = self._transform(
config, [self._vertex_embeddings_output_row()], url=url
)
(result,) = self._transform(config, [self._vertex_embeddings_output_row()], url=url)
assert result["response"]["body"]["model"] == "gemini-embedding-2"
def test_should_surface_failed_embeddings_row_as_error(self, config):
(result,) = self._transform(
config,
[
self._vertex_embeddings_output_row(
status="Failed to parse JSON into proto", response={}
)
],
[self._vertex_embeddings_output_row(status="Failed to parse JSON into proto", response={})],
)
assert result["custom_id"] == "request-1"
@ -1669,10 +1532,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
def test_should_transform_every_row_of_a_multi_row_file(self, config):
results = self._transform(
config,
[
self._vertex_embeddings_output_row(key=f"request-{index}")
for index in range(3)
],
[self._vertex_embeddings_output_row(key=f"request-{index}") for index in range(3)],
)
assert [result["custom_id"] for result in results] == [
@ -1751,9 +1611,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
"request-1#0/2",
"request-1",
]
assert [
result["response"]["body"]["data"][0]["embedding"] for result in results
] == [[0.1], [0.2]]
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(
@ -1782,9 +1640,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
)
assert result["custom_id"] == "request#1/1"
assert [
embedding["embedding"] for embedding in result["response"]["body"]["data"]
] == [[0.1], [0.3]]
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."""
@ -1792,9 +1648,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
config,
[
self._vertex_embeddings_output_row(key="request-1#0/2"),
self._vertex_embeddings_output_row(
key="request-1#1/2", status="Quota exceeded", response={}
),
self._vertex_embeddings_output_row(key="request-1#1/2", status="Quota exceeded", response={}),
],
)
@ -1809,11 +1663,25 @@ class TestVertexEmbeddingsBatchOutputTranslation:
[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_fail_the_whole_entry_when_a_fanned_out_row_is_duplicated(self, config):
(result,) = self._transform(
config,
[
self._vertex_embeddings_output_row(key="request-1#0/2"),
self._vertex_embeddings_output_row(key="request-1#0/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"
"Vertex returned embeddings for input positions [0, 0] of the 2 requested"
)
def test_should_end_to_end_round_trip_a_fanned_out_embeddings_batch(self, config):
@ -1842,9 +1710,7 @@ class TestVertexEmbeddingsBatchOutputTranslation:
)
assert result["custom_id"] == "MyRequest-1"
assert [
embedding["embedding"] for embedding in result["response"]["body"]["data"]
] == [[0.1], [0.3]]
assert [embedding["embedding"] for embedding in result["response"]["body"]["data"]] == [[0.1], [0.3]]
def test_should_end_to_end_round_trip_openai_embeddings_batch(self, config):
(vertex_row,) = _wrap_entries(