Managed batches fixes for Gemini/Vertex

This commit is contained in:
Ephrim Stanley 2026-02-28 20:45:16 -05:00 • committed by Sameer Kankute
parent 54cc4b7b35
commit 33d3c6022a
4 changed files with 405 additions and 145 deletions

View file

@ -29,6 +29,7 @@ verbose_logger.setLevel(logging.DEBUG)
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload
import random
import httpx
from unittest.mock import patch, MagicMock
@ -579,6 +580,48 @@ async def test_vertex_list_batches(monkeypatch):
assert list_response["data"][1].id == "test-batch-id-789"
@pytest.mark.asyncio
async def test_vertex_async_create_batch_logs_error_body_on_http_error():
"""
When Vertex AI returns an HTTP error (e.g. 400), _async_create_batch should
re-raise httpx.HTTPStatusError (not swallow it) and log the response body.
Before the fix the error body was lost because AsyncHTTPHandler.post()
calls raise_for_status() internally, raising before the handler's own
status-code check could log the body.
"""
from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction
handler = VertexAIBatchPrediction(gcs_bucket_name="test-bucket")
error_body = '{"error": {"code": 400, "message": "Do not support publisher model gemini-2.0-flash"}}'
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 400
mock_response.text = error_body
mock_response.headers = {}
http_error = httpx.HTTPStatusError(
message="Bad Request",
request=httpx.Request("POST", "https://fake-vertex-url"),
response=mock_response,
)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
side_effect=http_error,
):
with pytest.raises(httpx.HTTPStatusError) as exc_info:
await handler._async_create_batch(
vertex_batch_request={},
api_base="https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/batchPredictionJobs",
headers={"Authorization": "Bearer fake-token"},
)
assert exc_info.value.response.status_code == 400
assert "gemini-2.0-flash" in exc_info.value.response.text
@pytest.mark.asyncio
async def test_delete_batch_output_file():
"""

View file

@ -6,6 +6,7 @@ from fastapi import HTTPException
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm.caching import DualCache
from litellm.proxy._types import CallTypes
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
@ -61,6 +62,109 @@ async def test_async_pre_call_hook_batch_retrieve():
assert response["model"] == "my-general-azure-deployment"
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_resolves_model_id_from_litellm_metadata():
"""
For batch operations the router stores model_info under
kwargs["litellm_metadata"]["model_info"] (not top-level kwargs["model_info"]).
async_pre_call_deployment_hook must check both locations so the managed
file ID is resolved to the provider-specific file ID.
"""
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=MagicMock()
)
managed_file_id = "managed-file-abc"
model_id = "deployment-xyz"
provider_file_id = "gs://bucket/path/to/file.jsonl"
# model_info is nested under litellm_metadata (batch path)
kwargs = {
"input_file_id": managed_file_id,
"model_file_id_mapping": {
managed_file_id: {model_id: provider_file_id},
},
"litellm_metadata": {
"model_info": {"id": model_id},
},
}
result = await proxy_managed_files.async_pre_call_deployment_hook(
kwargs=kwargs, call_type=CallTypes.acreate_batch
)
assert result["input_file_id"] == provider_file_id, (
f"Expected provider file ID '{provider_file_id}', got '{result['input_file_id']}'"
)
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_prefers_top_level_model_info():
"""
When model_info exists at top-level kwargs, async_pre_call_deployment_hook
should use it without falling back to litellm_metadata.
"""
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=MagicMock()
)
managed_file_id = "managed-file-abc"
top_level_model_id = "deployment-top"
nested_model_id = "deployment-nested"
top_level_provider_file = "file-top-123"
nested_provider_file = "file-nested-456"
kwargs = {
"input_file_id": managed_file_id,
"model_file_id_mapping": {
managed_file_id: {
top_level_model_id: top_level_provider_file,
nested_model_id: nested_provider_file,
},
},
"model_info": {"id": top_level_model_id},
"litellm_metadata": {
"model_info": {"id": nested_model_id},
},
}
result = await proxy_managed_files.async_pre_call_deployment_hook(
kwargs=kwargs, call_type=CallTypes.acreate_batch
)
assert result["input_file_id"] == top_level_provider_file, (
"Should prefer top-level model_info over litellm_metadata"
)
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_unchanged():
"""
When model_info is absent from both top-level and litellm_metadata,
the managed file ID should remain unchanged.
"""
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=MagicMock()
)
managed_file_id = "managed-file-abc"
kwargs = {
"input_file_id": managed_file_id,
"model_file_id_mapping": {
managed_file_id: {"some-model": "provider-file-xyz"},
},
}
result = await proxy_managed_files.async_pre_call_deployment_hook(
kwargs=kwargs, call_type=CallTypes.acreate_batch
)
assert result["input_file_id"] == managed_file_id, (
"File ID should remain unchanged when model_info is not available"
)
# def test_list_managed_files():
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())

View file

@ -0,0 +1,184 @@
"""
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 "litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc" in result.id
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

View file

@ -227,52 +227,29 @@ class TestVertexAIBatchPassthroughHandler:
mock_managed_files_hook.store_unified_object_id.assert_called_once()
def test_batch_cost_calculation_integration(self):
"""Test integration with batch cost calculation"""
"""Single Vertex AI response → non-zero cost with correct token counts."""
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
# Mock Vertex AI batch responses
vertex_ai_batch_responses = [
{
"status": "JOB_STATE_SUCCEEDED",
"response": {
"candidates": [
{
"content": {
"parts": [
{"text": "Hello, world!"}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 5,
"totalTokenCount": 15
"totalTokenCount": 15,
}
}
}
]
with patch('litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexGeminiConfig') as mock_config:
with patch('litellm.completion_cost') as mock_completion_cost:
# Setup mocks
mock_config.return_value._transform_google_generate_content_to_openai_model_response.return_value = Mock(
usage=Mock(total_tokens=15, prompt_tokens=10, completion_tokens=5)
)
mock_completion_cost.return_value = 0.001
# Test the cost calculation
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses,
model_name="gemini-1.5-flash"
)
# Verify results
assert total_cost == 0.001
assert usage.total_tokens == 15
assert usage.prompt_tokens == 10
assert usage.completion_tokens == 5
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses, model_name="gemini-1.5-flash-001"
)
assert usage.total_tokens == 15
assert usage.prompt_tokens == 10
assert usage.completion_tokens == 5
assert total_cost > 0, "batch_cost_calculator should return a non-zero cost"
def test_batch_response_transformation(self):
"""Test transformation of Vertex AI batch responses to OpenAI format"""
@ -385,155 +362,107 @@ class TestVertexAIBatchPassthroughHandler:
class TestVertexAIBatchCostCalculation:
"""Test cases for Vertex AI batch cost calculation functionality"""
"""Test cases for Vertex AI batch cost calculation functionality.
def test_calculate_vertex_ai_batch_cost_and_usage_success(self):
"""Test successful batch cost and usage calculation"""
The function under test (calculate_vertex_ai_batch_cost_and_usage) extracts
usageMetadata directly from Vertex AI response dicts and calls
batch_cost_calculator — no VertexGeminiConfig transformation involved.
"""
def test_should_aggregate_cost_and_usage_across_responses(self):
"""Two successful responses → costs and token counts are summed."""
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
# Mock successful batch responses
vertex_ai_batch_responses = [
responses = [
{
"status": "JOB_STATE_SUCCEEDED",
"response": {
"candidates": [
{
"content": {
"parts": [
{"text": "Hello, world!"}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 5,
"totalTokenCount": 15
"totalTokenCount": 15,
}
}
},
{
"status": "JOB_STATE_SUCCEEDED",
"response": {
"candidates": [
{
"content": {
"parts": [
{"text": "How are you?"}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 8,
"candidatesTokenCount": 3,
"totalTokenCount": 11
"totalTokenCount": 11,
}
}
}
},
]
with patch('litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexGeminiConfig') as mock_config:
with patch('litellm.completion_cost') as mock_completion_cost:
# Setup mocks
mock_model_response = Mock()
mock_model_response.usage = Mock(total_tokens=15, prompt_tokens=10, completion_tokens=5)
mock_config.return_value._transform_google_generate_content_to_openai_model_response.return_value = mock_model_response
mock_completion_cost.return_value = 0.001
# Test the calculation
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses,
model_name="gemini-1.5-flash"
)
# Verify results
assert total_cost == 0.002 # 2 responses * 0.001 each
assert usage.total_tokens == 30 # 15 + 15
assert usage.prompt_tokens == 20 # 10 + 10
assert usage.completion_tokens == 10 # 5 + 5
def test_calculate_vertex_ai_batch_cost_and_usage_with_failed_responses(self):
"""Test batch cost calculation with some failed responses"""
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
responses, model_name="gemini-1.5-flash-001"
)
assert usage.prompt_tokens == 18
assert usage.completion_tokens == 8
assert usage.total_tokens == 26
assert total_cost > 0, "batch_cost_calculator should return a non-zero cost"
def test_should_skip_responses_with_null_response_body(self):
"""Failed lines (response: None) are skipped without error."""
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
# Mock batch responses with some failures
vertex_ai_batch_responses = [
responses = [
{
"status": "JOB_STATE_SUCCEEDED",
"response": {
"candidates": [
{
"content": {
"parts": [
{"text": "Hello, world!"}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 5,
"totalTokenCount": 15
"totalTokenCount": 15,
}
}
},
{"status": "JOB_STATE_FAILED", "response": None},
{
"status": "JOB_STATE_FAILED", # Failed response
"response": None
},
{
"status": "JOB_STATE_SUCCEEDED",
"response": {
"candidates": [
{
"content": {
"parts": [
{"text": "How are you?"}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 8,
"candidatesTokenCount": 3,
"totalTokenCount": 11
"totalTokenCount": 11,
}
}
}
},
]
with patch('litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexGeminiConfig') as mock_config:
with patch('litellm.completion_cost') as mock_completion_cost:
# Setup mocks
mock_model_response = Mock()
mock_model_response.usage = Mock(total_tokens=15, prompt_tokens=10, completion_tokens=5)
mock_config.return_value._transform_google_generate_content_to_openai_model_response.return_value = mock_model_response
mock_completion_cost.return_value = 0.001
# Test the calculation
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses,
model_name="gemini-1.5-flash"
)
# Verify results - should only process successful responses
assert total_cost == 0.002 # 2 successful responses * 0.001 each
assert usage.total_tokens == 30 # 15 + 15
assert usage.prompt_tokens == 20 # 10 + 10
assert usage.completion_tokens == 10 # 5 + 5
def test_calculate_vertex_ai_batch_cost_and_usage_empty_responses(self):
"""Test batch cost calculation with empty response list"""
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
responses, model_name="gemini-1.5-flash-001"
)
assert usage.prompt_tokens == 18
assert usage.completion_tokens == 8
assert usage.total_tokens == 26
assert total_cost > 0
def test_should_return_zeros_for_empty_response_list(self):
"""Empty input → zero cost and zero usage."""
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
# Test with empty list
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage([], model_name="gemini-1.5-flash")
# Verify results
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
[], model_name="gemini-1.5-flash-001"
)
assert total_cost == 0.0
assert usage.total_tokens == 0
assert usage.prompt_tokens == 0
assert usage.completion_tokens == 0
def test_should_handle_missing_usage_metadata_gracefully(self):
"""Response without usageMetadata → 0 tokens, 0 cost for that line."""
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
responses = [
{"response": {"candidates": [{"content": {"parts": [{"text": "hi"}]}}]}},
]
total_cost, usage = calculate_vertex_ai_batch_cost_and_usage(
responses, model_name="gemini-1.5-flash-001"
)
assert usage.prompt_tokens == 0
assert usage.completion_tokens == 0
assert usage.total_tokens == 0