Managed batches - Address PR bot comments from #22464

Made-with: Cursor
This commit is contained in:
Ephrim Stanley 2026-03-03 11:01:49 -05:00 committed by shivam
parent 0435375b12
commit b064ec8d7e
5 changed files with 380 additions and 7 deletions

View file

@ -336,7 +336,7 @@ async def afile_retrieve(
@client
def file_retrieve(
file_id: str,
custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "manus"] = "openai",
custom_llm_provider: Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,

View file

@ -108,11 +108,19 @@ class VertexAIBatchPrediction(VertexLLM):
client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.VERTEX_AI,
)
response = await client.post(
url=api_base,
headers=headers,
data=json.dumps(vertex_batch_request),
)
try:
response = await client.post(
url=api_base,
headers=headers,
data=json.dumps(vertex_batch_request),
)
except httpx.HTTPStatusError as e:
error_body = e.response.text
litellm.verbose_logger.error(
"Vertex AI batch create failed: status=%s, body=%s",
e.response.status_code, error_body[:1000],
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")

View file

@ -1,6 +1,7 @@
import json
import os
import time
import urllib.parse
from typing import Any, Dict, List, Optional, Tuple, Union
from httpx import Headers, Response
@ -365,7 +366,14 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> FileDeleted:
raise NotImplementedError("VertexAIFilesConfig does not support file deletion")
file_id = "deleted"
if hasattr(raw_response, "request") and raw_response.request:
url = str(raw_response.request.url)
if "/b/" in url and "/o/" in url:
bucket_part = url.split("/b/")[-1].split("/o/")[0]
encoded_name = url.split("/o/")[-1].split("?")[0]
file_id = f"gs://{bucket_part}/{urllib.parse.unquote(encoded_name)}"
return FileDeleted(id=file_id, deleted=True, object="file")
def transform_list_files_request(
self,

View file

@ -0,0 +1,127 @@
"""
Tests for Fix 1: file_retrieve Literal type was missing 'vertex_ai' and 'gemini',
causing a type mismatch when afile_retrieve delegated to the sync function.
"""
import pytest
from unittest.mock import MagicMock, patch
from litellm.files.main import file_retrieve
class TestFileRetrieveProviderRouting:
"""
Verify that file_retrieve accepts 'vertex_ai' and 'gemini' providers and
routes them through ProviderConfigManager / base_llm_http_handler.
"""
def _make_mock_file_object(self):
mock = MagicMock()
mock.model_dump.return_value = {
"id": "gs://my-bucket/file.jsonl",
"object": "file",
"bytes": 1024,
"created_at": 0,
"filename": "file.jsonl",
"purpose": "batch",
"status": "processed",
}
return mock
def test_should_route_vertex_ai_through_provider_config(self):
"""
Regression: file_retrieve Literal type was missing 'vertex_ai',
so passing custom_llm_provider='vertex_ai' would fail type-checking
and potentially cause a routing failure at runtime.
"""
mock_file = self._make_mock_file_object()
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file",
return_value=mock_file,
) as mock_retrieve:
result = file_retrieve(
file_id="gs://my-bucket/file.jsonl",
custom_llm_provider="vertex_ai",
)
mock_retrieve.assert_called_once()
assert result is not None
def test_should_route_gemini_through_provider_config(self):
"""
Regression: file_retrieve Literal type was also missing 'gemini'.
"""
mock_file = self._make_mock_file_object()
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file",
return_value=mock_file,
) as mock_retrieve:
result = file_retrieve(
file_id="some-gemini-file-id",
custom_llm_provider="gemini",
)
mock_retrieve.assert_called_once()
assert result is not None
def test_should_pass_file_id_to_handler_for_vertex_ai(self):
"""Verify the file_id is forwarded correctly to the underlying handler."""
mock_file = self._make_mock_file_object()
expected_file_id = "gs://my-bucket/path/to/file.jsonl"
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file",
return_value=mock_file,
) as mock_retrieve:
file_retrieve(
file_id=expected_file_id,
custom_llm_provider="vertex_ai",
)
call_kwargs = mock_retrieve.call_args.kwargs
assert call_kwargs.get("file_id") == expected_file_id
def test_should_not_raise_bad_request_for_vertex_ai(self):
"""
Before the fix, vertex_ai fell through to the else-branch which raised
BadRequestError. Verify it no longer does.
"""
import litellm
mock_file = self._make_mock_file_object()
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file",
return_value=mock_file,
):
try:
file_retrieve(
file_id="gs://my-bucket/file.jsonl",
custom_llm_provider="vertex_ai",
)
except litellm.exceptions.BadRequestError as e:
pytest.fail(
f"file_retrieve raised BadRequestError for vertex_ai: {e}"
)
def test_should_not_raise_bad_request_for_gemini(self):
"""Same as above but for 'gemini'."""
import litellm
mock_file = self._make_mock_file_object()
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file",
return_value=mock_file,
):
try:
file_retrieve(
file_id="some-file-id",
custom_llm_provider="gemini",
)
except litellm.exceptions.BadRequestError as e:
pytest.fail(
f"file_retrieve raised BadRequestError for gemini: {e}"
)

View file

@ -0,0 +1,230 @@
"""
Tests for VertexAIFilesConfig transformation methods (Issues 5-7).
"""
import json
import urllib.parse
import httpx
import pytest
from unittest.mock import MagicMock
from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig
from litellm.types.llms.openai import OpenAIFileObject, HttpxBinaryResponseContent
from openai.types.file_deleted import FileDeleted
@pytest.fixture
def config():
return VertexAIFilesConfig()
class TestParseGcsUri:
"""Tests for the _parse_gcs_uri helper used by retrieve / content / delete."""
def test_should_parse_standard_gs_uri(self, config):
bucket, encoded = config._parse_gcs_uri(
"gs://my-bucket/path/to/object.jsonl"
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("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)
assert bucket == "litellm-local"
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/some/path", safe="")
bucket, encoded = config._parse_gcs_uri(encoded_uri)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("some/path", safe="")
def test_should_handle_bucket_only(self, config):
bucket, encoded = config._parse_gcs_uri("gs://my-bucket")
assert bucket == "my-bucket"
assert encoded == ""
def test_should_handle_no_gs_prefix(self, config):
bucket, encoded = config._parse_gcs_uri("my-bucket/object.txt")
assert bucket == "my-bucket"
assert encoded == "object.txt"
class TestTransformRetrieveFile:
def test_should_build_correct_gcs_metadata_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
url, params = config.transform_retrieve_file_request(
file_id=file_id, optional_params={}, litellm_params={}
)
expected_encoded = urllib.parse.quote("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):
gcs_json = {
"id": "my-bucket/path/to/file.jsonl/123456",
"name": "path/to/file.jsonl",
"size": "4096",
"timeCreated": "2025-02-15T10:00:00.000Z",
"metadata": {"purpose": "batch"},
}
raw_response = MagicMock(spec=httpx.Response)
raw_response.json.return_value = gcs_json
result = config.transform_retrieve_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert isinstance(result, OpenAIFileObject)
assert result.id == "gs://my-bucket/path/to/file.jsonl"
assert result.filename == "path/to/file.jsonl"
assert result.bytes == 4096
assert result.object == "file"
assert result.status == "processed"
assert result.purpose == "batch"
def test_should_default_purpose_to_batch_when_metadata_missing(self, config):
gcs_json = {
"id": "bucket/obj/999",
"name": "obj",
"size": "0",
"timeCreated": "2025-01-01T00:00:00.000Z",
}
raw_response = MagicMock(spec=httpx.Response)
raw_response.json.return_value = gcs_json
result = config.transform_retrieve_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert result.purpose == "batch"
class TestTransformFileContent:
def test_should_build_gcs_media_download_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
url, params = config.transform_file_content_request(
file_content_request={"file_id": file_id},
optional_params={},
litellm_params={},
)
encoded = urllib.parse.quote("path/to/file.jsonl", safe="")
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):
raw_response = httpx.Response(
status_code=200,
content=b'{"line": 1}\n{"line": 2}\n',
headers={"content-type": "application/octet-stream"},
request=httpx.Request("GET", "https://example.com"),
)
result = config.transform_file_content_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert isinstance(result, HttpxBinaryResponseContent)
assert result.response.content == b'{"line": 1}\n{"line": 2}\n'
class TestTransformDeleteFile:
def test_should_build_correct_gcs_delete_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
url, params = config.transform_delete_file_request(
file_id=file_id, optional_params={}, litellm_params={}
)
encoded = urllib.parse.quote("path/to/file.jsonl", safe="")
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):
raw_response = MagicMock(spec=httpx.Response)
mock_request = MagicMock()
encoded_name = urllib.parse.quote(
"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}"
)
raw_response.request = mock_request
result = config.transform_delete_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
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"
def test_should_fallback_to_deleted_id_when_no_request(self, config):
raw_response = MagicMock(spec=httpx.Response)
raw_response.request = None
result = config.transform_delete_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert isinstance(result, FileDeleted)
assert result.id == "deleted"
assert result.deleted is True
def test_should_include_bucket_name_in_reconstructed_delete_id(self, config):
"""
Regression: the old code split on /o/ only, dropping the bucket from
the reconstructed gs:// URI. e.g. gs://path/to/file instead of
gs://my-bucket/path/to/file.
"""
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}"
)
raw_response.request = mock_request
result = config.transform_delete_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert result.id == "gs://my-bucket/path/to/file.jsonl"
def test_should_include_bucket_in_nested_object_path(self, config):
"""Verify bucket extraction works with deeply nested GCS object paths."""
raw_response = MagicMock(spec=httpx.Response)
mock_request = MagicMock()
encoded_object = urllib.parse.quote(
"litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123",
safe="",
)
mock_request.url = (
f"https://storage.googleapis.com/storage/v1/b/prod-bucket/o/{encoded_object}"
)
raw_response.request = mock_request
result = config.transform_delete_file_response(
raw_response=raw_response,
logging_obj=MagicMock(),
litellm_params={},
)
assert result.id == (
"gs://prod-bucket/litellm-vertex-files/publishers/google/"
"models/gemini-2.0-flash-001/abc-123"
)