mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
Managed batches - Address PR bot comments from #22464
Made-with: Cursor
This commit is contained in:
parent
0435375b12
commit
b064ec8d7e
5 changed files with 380 additions and 7 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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"
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue