- Hello, world!
- This is a test of the text-to-speech API.
-
- """
-
- # Set up the mock for asynchronous calls
- with patch(
- "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
- new_callable=AsyncMock,
- ) as mock_async_post:
- mock_async_post.return_value = mock_response
- model = "vertex_ai/test"
-
- try:
- response = await litellm.aspeech(
- input=ssml,
- model=model,
- voice={
- "languageCode": "en-UK",
- "name": "en-UK-Studio-O",
- },
- audioConfig={
- "audioEncoding": "LINEAR22",
- "speakingRate": "10",
- },
- )
- except litellm.APIConnectionError as e:
- if "Your default credentials were not found" in str(e):
- pytest.skip("skipping test, credentials not found")
-
- # Assert asynchronous call
- mock_async_post.assert_called_once()
- _, kwargs = mock_async_post.call_args
- print("call args", kwargs)
-
- assert kwargs["url"] == "https://texttospeech.googleapis.com/v1/text:synthesize"
-
- assert "x-goog-user-project" in kwargs["headers"]
- assert kwargs["headers"]["Authorization"] is not None
-
- assert kwargs["json"] == {
- "input": {"ssml": ssml},
- "voice": {"languageCode": "en-UK", "name": "en-UK-Studio-O"},
- "audioConfig": {"audioEncoding": "LINEAR22", "speakingRate": "10"},
- }
-
-
-
-
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
async def test_azure_ava_tts_async():
diff --git a/tests/audio_tests/test_whisper.py b/tests/audio_tests/test_whisper.py
index 5380a57c871..99b084be1d3 100644
--- a/tests/audio_tests/test_whisper.py
+++ b/tests/audio_tests/test_whisper.py
@@ -15,29 +15,21 @@ from dotenv import load_dotenv
from openai import AsyncOpenAI
import litellm
-from litellm.integrations.custom_logger import CustomLogger
# Get the current directory of the file being run
pwd = os.path.dirname(os.path.realpath(__file__))
print(pwd)
file_path = os.path.join(pwd, "gettysburg.wav")
-file2_path = os.path.join(pwd, "eagle.wav")
with open(file_path, "rb") as _f:
_GETTYSBURG_BYTES = _f.read()
-with open(file2_path, "rb") as _f:
- _EAGLE_BYTES = _f.read()
def _audio_file():
return ("gettysburg.wav", _GETTYSBURG_BYTES, "audio/wav")
-def _audio_file2():
- return ("eagle.wav", _EAGLE_BYTES, "audio/wav")
-
-
load_dotenv()
from litellm import Router
@@ -75,104 +67,3 @@ async def test_transcription_azure_whisper(response_format, timestamp_granularit
response_format=response_format,
timestamp_granularities=timestamp_granularities,
)
-
-
-@pytest.mark.asyncio()
-async def test_transcription_caching():
- import litellm
- from litellm.caching.caching import Cache
-
- litellm.set_verbose = True
- litellm.cache = Cache()
-
- # make raw llm api call
-
- response_1 = await litellm.atranscription(
- model="whisper-1",
- file=_audio_file(),
- )
-
- await asyncio.sleep(5)
-
- # cache hit
-
- response_2 = await litellm.atranscription(
- model="whisper-1",
- file=_audio_file(),
- )
-
- print("response_1", response_1)
- print("response_2", response_2)
- print("response2 hidden params", response_2._hidden_params)
- assert response_2._hidden_params["cache_hit"] is True
-
- # cache miss
-
- response_3 = await litellm.atranscription(
- model="whisper-1",
- file=_audio_file2(),
- )
- print("response_3", response_3)
- print("response3 hidden params", response_3._hidden_params)
- assert response_3._hidden_params.get("cache_hit") is not True
- assert response_3.text != response_2.text
-
- litellm.cache = None
-
-
-@pytest.mark.asyncio
-async def test_whisper_log_pre_call():
- from litellm.litellm_core_utils.litellm_logging import Logging
- from datetime import datetime
- from unittest.mock import patch, MagicMock
-
- custom_logger = CustomLogger()
-
- litellm.callbacks = [custom_logger]
-
- with patch.object(custom_logger, "log_pre_api_call") as mock_log_pre_call:
- await litellm.atranscription(
- model="whisper-1",
- file=_audio_file(),
- )
- mock_log_pre_call.assert_called_once()
-
-
-@pytest.mark.asyncio
-async def test_gpt_4o_transcribe_model_mapping():
- """Test that GPT-4o transcription models are correctly mapped and not hardcoded to whisper-1"""
-
- # Test GPT-4o mini transcribe
- response = await litellm.atranscription(
- model="openai/gpt-4o-mini-transcribe",
- file=_audio_file(),
- response_format="json",
- )
-
- # Check that the response contains the correct model in hidden params
- assert response._hidden_params is not None
- assert response._hidden_params["model"] == "gpt-4o-mini-transcribe"
- assert response._hidden_params["custom_llm_provider"] == "openai"
- assert response.text is not None
-
- # Test GPT-4o transcribe
- response2 = await litellm.atranscription(
- model="openai/gpt-4o-transcribe", file=_audio_file(), response_format="json"
- )
-
- # Check that the response contains the correct model in hidden params
- assert response2._hidden_params is not None
- assert response2._hidden_params["model"] == "gpt-4o-transcribe"
- assert response2._hidden_params["custom_llm_provider"] == "openai"
- assert response2.text is not None
-
- # Test traditional whisper-1 still works
- response3 = await litellm.atranscription(
- model="openai/whisper-1", file=_audio_file(), response_format="json"
- )
-
- # Check that the response contains the correct model in hidden params
- assert response3._hidden_params is not None
- assert response3._hidden_params["model"] == "whisper-1"
- assert response3._hidden_params["custom_llm_provider"] == "openai"
- assert response3.text is not None
diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py
deleted file mode 100644
index 87d35e11c74..00000000000
--- a/tests/batches_tests/test_batch_rate_limits.py
+++ /dev/null
@@ -1,896 +0,0 @@
-"""
-Integration Tests for Batch Rate Limits
-"""
-
-import asyncio
-import json
-import os
-
-import pytest
-from fastapi import HTTPException
-
-
-import litellm
-from litellm import DualCache
-from litellm.proxy._types import UserAPIKeyAuth
-from litellm.proxy.hooks.batch_rate_limiter import (
- BatchFileUsage,
- PROXY_BatchRateLimiter,
-)
-from litellm.proxy.hooks.parallel_request_limiter_v3 import (
- PROXY_MaxParallelRequestsHandler_v3,
-)
-from litellm.proxy.utils import InternalUsageCache
-
-
-def _build_batch_limiter() -> PROXY_BatchRateLimiter:
- internal_usage_cache = InternalUsageCache(dual_cache=DualCache())
- return PROXY_BatchRateLimiter(
- internal_usage_cache=internal_usage_cache,
- parallel_request_limiter=PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- ),
- )
-
-
-def get_expected_batch_file_usage(file_path: str) -> tuple[int, int]:
- """
- Helper function to calculate expected request count and token count from a batch JSONL file.
-
- Returns:
- tuple[int, int]: (expected_request_count, expected_total_tokens)
- """
- with open(file_path, "r") as f:
- file_contents = [json.loads(line) for line in f if line.strip()]
-
- expected_request_count = len(file_contents)
- expected_total_tokens = 0
-
- for item in file_contents:
- body = item.get("body", {})
- model = body.get("model", "")
- messages = body.get("messages", [])
- if messages:
- item_tokens = litellm.token_counter(model=model, messages=messages)
- expected_total_tokens += item_tokens
-
- return expected_request_count, expected_total_tokens
-
-
-def _write_batch_file(tmp_path, file_name: str, content: str) -> str:
- path = tmp_path / file_name
- path.write_text(content)
- return str(path)
-
-
-@pytest.mark.asyncio()
-@pytest.mark.skipif(
- os.environ.get("OPENAI_API_KEY") is None,
- reason="OPENAI_API_KEY not set - skipping integration test",
-)
-async def test_batch_rate_limits():
- """
- Integration test for batch rate limits with real OpenAI API calls.
- Tests the full flow: file creation -> token counting -> cleanup
- """
- litellm.turn_on_debug()
- CUSTOM_LLM_PROVIDER = "openai"
- BATCH_LIMITER = _build_batch_limiter()
-
- file_name = "openai_batch_completions.jsonl"
- _current_dir = os.path.dirname(os.path.abspath(__file__))
- file_path = os.path.join(_current_dir, file_name)
-
- # Create file on OpenAI
- print(f"Creating file from {file_path}")
- file_obj = await litellm.acreate_file(
- file=open(file_path, "rb"),
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f"Response from creating file: {file_obj}")
-
- assert file_obj.id is not None, "File ID should not be None"
-
- # Give API a moment to process the file
- await asyncio.sleep(1)
-
- # Count requests and token usage in input file
- tracked_batch_file_usage: BatchFileUsage = (
- await BATCH_LIMITER.count_input_file_usage(
- file_id=file_obj.id,
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- )
- print(f"Actual total tokens: {tracked_batch_file_usage.total_tokens}")
- print(f"Actual request count: {tracked_batch_file_usage.request_count}")
-
- # Calculate expected values by reading the JSONL file
- expected_request_count, expected_total_tokens = get_expected_batch_file_usage(
- file_path=file_path
- )
-
- print(f"Expected request count: {expected_request_count}")
- print(f"Expected total tokens: {expected_total_tokens}")
-
- # Verify token counting results
- assert (
- tracked_batch_file_usage.request_count == expected_request_count
- ), f"Expected {expected_request_count} requests, got {tracked_batch_file_usage.request_count}"
- assert (
- tracked_batch_file_usage.total_tokens == expected_total_tokens
- ), f"Expected {expected_total_tokens} total_tokens, got {tracked_batch_file_usage.total_tokens}"
-
-
-@pytest.mark.asyncio()
-async def test_batch_rate_limit_single_file(tmp_path):
- """
- Test batch rate limiting with a single file.
-
- Key has TPM = 200
- - File with < 200 tokens: should go through
- - File with > 200 tokens: should hit rate limit
- """
- CUSTOM_LLM_PROVIDER = "openai"
-
- # Setup: Create internal usage cache and rate limiter
- dual_cache = DualCache()
- internal_usage_cache = InternalUsageCache(dual_cache=dual_cache)
- rate_limiter = PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- )
-
- # Setup: Get batch rate limiter
- batch_limiter = rate_limiter._get_batch_rate_limiter()
- assert batch_limiter is not None, "Batch rate limiter should be available"
-
- # Setup: Create user API key with TPM = 200
- user_api_key_dict = UserAPIKeyAuth(
- api_key="test-key-123",
- tpm_limit=200,
- rpm_limit=10,
- )
-
- # Test 1: File with < 200 tokens should go through
- print("\n=== Test 1: File under 200 tokens ===")
-
- # Create a small batch file with ~150 tokens
- small_batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}
-{"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hi"}]}}
-{"custom_id": "request-3", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey"}]}}"""
-
- small_file_path = _write_batch_file(
- tmp_path, "small-batch-rate-limit.jsonl", small_batch_content
- )
-
- try:
- # Upload file to OpenAI
- with open(small_file_path, "rb") as batch_file:
- file_obj_small = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f"Created small file: {file_obj_small.id}")
- await asyncio.sleep(1) # Give API time to process
-
- data_under_limit = {
- "model": "gpt-3.5-turbo",
- "input_file_id": file_obj_small.id,
- "custom_llm_provider": CUSTOM_LLM_PROVIDER,
- }
-
- # Should not raise an exception
- result = await batch_limiter.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=dual_cache,
- data=data_under_limit,
- call_type="acreate_batch",
- )
- print(f"✓ File with ~150 tokens passed (under limit of 200)")
- print(f" Actual tokens: {result.get('_batch_token_count')}")
- except HTTPException as e:
- pytest.fail(f"Should not have hit rate limit with small file: {e.detail}")
-
- # Test 2: File with > 200 tokens should hit rate limit
- print("\n=== Test 2: File over 200 tokens ===")
-
- # Reset cache for clean test
- dual_cache = DualCache()
- internal_usage_cache = InternalUsageCache(dual_cache=dual_cache)
- rate_limiter = PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- )
- batch_limiter = rate_limiter._get_batch_rate_limiter()
-
- # Create a larger batch file with ~10000+ tokens (100x larger to ensure it exceeds 200 token limit)
- base_message = (
- "This is a longer message that will consume more tokens from the rate limit. "
- * 100
- )
-
- # Build JSONL content with json.dumps to avoid f-string nesting issues
- import json as json_lib
-
- requests = []
- for i in range(1, 4):
- request_obj = {
- "custom_id": f"request-{i}",
- "method": "POST",
- "url": "/v1/chat/completions",
- "body": {
- "model": "gpt-3.5-turbo",
- "messages": [{"role": "user", "content": base_message}],
- },
- }
- requests.append(json_lib.dumps(request_obj))
-
- large_batch_content = "\n".join(requests)
-
- large_file_path = _write_batch_file(
- tmp_path, "large-batch-rate-limit.jsonl", large_batch_content
- )
-
- # Upload file to OpenAI
- with open(large_file_path, "rb") as batch_file:
- file_obj_large = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f"Created large file: {file_obj_large.id}")
- await asyncio.sleep(1) # Give API time to process
-
- data_over_limit = {
- "model": "gpt-3.5-turbo",
- "input_file_id": file_obj_large.id,
- "custom_llm_provider": CUSTOM_LLM_PROVIDER,
- }
-
- # Should raise HTTPException with 429 status
- with pytest.raises(HTTPException) as exc_info:
- await batch_limiter.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=dual_cache,
- data=data_over_limit,
- call_type="acreate_batch",
- )
-
- assert exc_info.value.status_code == 429, "Should return 429 status code"
- assert (
- "tokens" in exc_info.value.detail.lower()
- ), "Error message should mention tokens"
- print(f"✓ File with 250+ tokens correctly rejected (over limit of 200)")
- print(f" Error: {exc_info.value.detail}")
-
-
-@pytest.mark.asyncio()
-async def test_batch_rate_limit_multiple_requests(tmp_path):
- """
- Test batch rate limiting with multiple requests.
-
- Key has TPM = 200
- - Request 1: file with ~100 tokens (should go through, 100/200 used)
- - Request 2: file with ~105 tokens (should hit limit, 100+105=205 > 200)
- """
- CUSTOM_LLM_PROVIDER = "openai"
-
- # Setup: Create internal usage cache and rate limiter
- dual_cache = DualCache()
- internal_usage_cache = InternalUsageCache(dual_cache=dual_cache)
- rate_limiter = PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- )
-
- # Setup: Get batch rate limiter
- batch_limiter = rate_limiter._get_batch_rate_limiter()
- assert batch_limiter is not None, "Batch rate limiter should be available"
-
- # Setup: Create user API key with TPM = 200
- user_api_key_dict = UserAPIKeyAuth(
- api_key="test-key-456",
- tpm_limit=200,
- rpm_limit=10,
- )
-
- # Request 1: File with ~100 tokens
- print("\n=== Request 1: File with ~100 tokens ===")
-
- # Create file with ~100 tokens
- import json as json_lib
-
- message_1 = "This message has some content to reach about 100 tokens total. " * 4
- requests_1 = []
- for i in range(1, 3):
- request_obj = {
- "custom_id": f"request-{i}",
- "method": "POST",
- "url": "/v1/chat/completions",
- "body": {
- "model": "gpt-3.5-turbo",
- "messages": [{"role": "user", "content": message_1}],
- },
- }
- requests_1.append(json_lib.dumps(request_obj))
-
- batch_content_1 = "\n".join(requests_1)
-
- file_path_1 = _write_batch_file(
- tmp_path, "batch-rate-limit-request-1.jsonl", batch_content_1
- )
-
- try:
- # Upload file to OpenAI
- with open(file_path_1, "rb") as batch_file:
- file_obj_1 = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f"Created file 1: {file_obj_1.id}")
- await asyncio.sleep(1) # Give API time to process
-
- data_request1 = {
- "model": "gpt-3.5-turbo",
- "input_file_id": file_obj_1.id,
- "custom_llm_provider": CUSTOM_LLM_PROVIDER,
- }
-
- # Should not raise an exception
- result1 = await batch_limiter.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=dual_cache,
- data=data_request1,
- call_type="acreate_batch",
- )
- tokens_used_1 = result1.get("_batch_token_count", 0)
- print(
- f"✓ Request 1 with {tokens_used_1} tokens passed ({tokens_used_1}/200 used)"
- )
- except HTTPException as e:
- pytest.fail(f"Request 1 should not have hit rate limit: {e.detail}")
-
- # Request 2: File with ~105+ tokens (total would exceed 200)
- print("\n=== Request 2: File with ~105 tokens (should hit limit) ===")
-
- # Create file with ~105+ tokens
- message_2 = (
- "This is another message with more content to exceed the remaining limit. " * 11
- )
- requests_2 = []
- for i in range(1, 3):
- request_obj = {
- "custom_id": f"request-{i}",
- "method": "POST",
- "url": "/v1/chat/completions",
- "body": {
- "model": "gpt-3.5-turbo",
- "messages": [{"role": "user", "content": message_2}],
- },
- }
- requests_2.append(json_lib.dumps(request_obj))
-
- batch_content_2 = "\n".join(requests_2)
-
- file_path_2 = _write_batch_file(
- tmp_path, "batch-rate-limit-request-2.jsonl", batch_content_2
- )
-
- # Upload file to OpenAI
- with open(file_path_2, "rb") as batch_file:
- file_obj_2 = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f"Created file 2: {file_obj_2.id}")
- await asyncio.sleep(1) # Give API time to process
-
- data_request2 = {
- "model": "gpt-3.5-turbo",
- "input_file_id": file_obj_2.id,
- "custom_llm_provider": CUSTOM_LLM_PROVIDER,
- }
-
- # Should raise HTTPException with 429 status
- with pytest.raises(HTTPException) as exc_info:
- await batch_limiter.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=dual_cache,
- data=data_request2,
- call_type="acreate_batch",
- )
-
- assert exc_info.value.status_code == 429, "Should return 429 status code"
- assert (
- "tokens" in exc_info.value.detail.lower()
- ), "Error message should mention tokens"
- print(f"✓ Request 2 correctly rejected")
- print(f" Error: {exc_info.value.detail}")
-
-
-@pytest.mark.asyncio()
-@pytest.mark.skipif(
- os.environ.get("OPENAI_API_KEY") is None,
- reason="OPENAI_API_KEY not set - skipping integration test",
-)
-async def test_batch_rate_limiter_with_managed_files(tmp_path):
- """
- Test for GEN-2166: Verify batch rate limiter can read user files when managed files are enabled.
-
- This test ensures that:
- 1. The batch rate limiter passes user_api_key_dict to afile_content()
- 2. The managed files hook can verify file ownership correctly
- 3. Rate limiting is enforced (not silently bypassed)
- 4. No 403 Permission Denied errors occur for files owned by the user
- """
- from unittest.mock import AsyncMock, MagicMock, patch
-
- CUSTOM_LLM_PROVIDER = "openai"
-
- # Setup: Create internal usage cache and rate limiter
- dual_cache = DualCache()
- internal_usage_cache = InternalUsageCache(dual_cache=dual_cache)
- rate_limiter = PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- )
-
- # Setup: Get batch rate limiter
- batch_limiter = rate_limiter._get_batch_rate_limiter()
- assert batch_limiter is not None, "Batch rate limiter should be available"
-
- # Setup: Create user API key with TPM = 500, RPM = 10
- test_user_id = "test-user-abc123"
- user_api_key_dict = UserAPIKeyAuth(
- api_key="test-key-managed-files",
- user_id=test_user_id,
- tpm_limit=500,
- rpm_limit=10,
- )
-
- print(f"\n=== Testing Batch Rate Limiter with Managed Files ===")
- print(f"User ID: {test_user_id}")
-
- # Create a batch file with ~200 tokens
- import json as json_lib
-
- message = "This is a test message for batch rate limiting with managed files. " * 5
- requests = []
- for i in range(1, 4):
- request_obj = {
- "custom_id": f"request-{i}",
- "method": "POST",
- "url": "/v1/chat/completions",
- "body": {
- "model": "gpt-3.5-turbo",
- "messages": [{"role": "user", "content": message}],
- },
- }
- requests.append(json_lib.dumps(request_obj))
-
- batch_content = "\n".join(requests)
-
- file_path = _write_batch_file(
- tmp_path, "managed-files-batch-rate-limit.jsonl", batch_content
- )
-
- try:
- # Step 1: Upload file to OpenAI (simulating user upload)
- print("\n1. Uploading batch input file...")
- with open(file_path, "rb") as batch_file:
- file_obj = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- print(f" ✓ File uploaded: {file_obj.id}")
- await asyncio.sleep(1) # Give API time to process
-
- # Step 2: Mock managed files hook to simulate file ownership check
- # In a real scenario, the managed files hook would check if the user owns the file
- # For this test, we'll verify that user_api_key_dict is passed correctly
- print("\n2. Testing rate limiter file access with user context...")
-
- # Track if user_api_key_dict was passed to afile_content
- original_afile_content = litellm.afile_content
- user_context_passed = {"value": False}
-
- async def mock_afile_content(*args, **kwargs):
- # Check if user_api_key_dict was passed
- if (
- "user_api_key_dict" in kwargs
- and kwargs["user_api_key_dict"] is not None
- ):
- user_context_passed["value"] = True
- print(f" ✓ user_api_key_dict passed to afile_content")
- print(f" User ID: {kwargs['user_api_key_dict'].user_id}")
- else:
- print(f" ✗ user_api_key_dict NOT passed to afile_content (BUG!)")
-
- # Call original function
- return await original_afile_content(*args, **kwargs)
-
- # Patch afile_content to track the call
- with patch("litellm.afile_content", side_effect=mock_afile_content):
- data = {
- "model": "gpt-3.5-turbo",
- "input_file_id": file_obj.id,
- "custom_llm_provider": CUSTOM_LLM_PROVIDER,
- }
-
- # Step 3: Submit batch and verify rate limiting works
- print("\n3. Submitting batch with rate limiting...")
- result = await batch_limiter.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=dual_cache,
- data=data,
- call_type="acreate_batch",
- )
-
- tokens_used = result.get("_batch_token_count", 0)
- requests_count = result.get("_batch_request_count", 0)
- print(f" ✓ Batch submitted successfully")
- print(f" Tokens counted: {tokens_used}")
- print(f" Requests counted: {requests_count}")
- print(
- f" Rate limit usage: {tokens_used}/500 TPM, {requests_count}/10 RPM"
- )
-
- # Step 4: Verify user context was passed
- print("\n4. Verifying fix for GEN-2166...")
- assert user_context_passed["value"], (
- "FAILED: user_api_key_dict was not passed to afile_content(). "
- "This means the bug GEN-2166 is not fixed!"
- )
- print(" ✓ Fix verified: user_api_key_dict is correctly passed")
-
- # Step 5: Verify rate limiting is actually enforced (not bypassed)
- print("\n5. Verifying rate limiting is enforced...")
- assert tokens_used > 0, "Token count should be greater than 0"
- assert requests_count > 0, "Request count should be greater than 0"
- print(" ✓ Rate limiting is active (not silently bypassed)")
-
- print("\n=== Test Passed: GEN-2166 Fix Verified ===")
- print("✓ Batch rate limiter can access user files")
- print("✓ User context is correctly passed")
- print("✓ Rate limiting is enforced")
- print("✓ No silent failures")
-
- except HTTPException as e:
- if e.status_code == 403:
- pytest.fail(
- f"FAILED: Got 403 Permission Denied error. "
- f"This indicates the bug GEN-2166 is not fixed. "
- f"Error: {e.detail}"
- )
- else:
- raise
- except Exception as e:
- pytest.fail(f"Unexpected error: {str(e)}")
-
-
-@pytest.mark.asyncio()
-async def test_batch_rate_limiter_without_user_context(tmp_path):
- """
- Test that verifies the bug scenario from GEN-2166.
-
- When user_api_key_dict is NOT passed to count_input_file_usage(),
- the function should still work for non-managed files, but would fail
- for managed files (which is the bug we fixed).
-
- This test documents the expected behavior with and without user context.
- """
- CUSTOM_LLM_PROVIDER = "openai"
-
- # Setup
- BATCH_LIMITER = _build_batch_limiter()
-
- # Create a simple batch file
- batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}"""
-
- file_path = _write_batch_file(
- tmp_path, "without-user-context-batch-rate-limit.jsonl", batch_content
- )
-
- # Upload file
- with open(file_path, "rb") as batch_file:
- file_obj = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- )
- await asyncio.sleep(1)
-
- # Test 1: Without user context (old behavior - would fail with managed files)
- print("\n=== Test 1: count_input_file_usage WITHOUT user context ===")
- try:
- usage_without_context = await BATCH_LIMITER.count_input_file_usage(
- file_id=file_obj.id,
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- user_api_key_dict=None, # Explicitly passing None
- )
- print(
- f"✓ Works for non-managed files (tokens: {usage_without_context.total_tokens})"
- )
- print(" Note: Would fail with 403 for managed files (GEN-2166 bug)")
- except Exception as e:
- print(f"✗ Failed: {str(e)}")
-
- # Test 2: With user context (new behavior - works with managed files)
- print("\n=== Test 2: count_input_file_usage WITH user context ===")
- user_api_key_dict = UserAPIKeyAuth(
- api_key="test-key",
- user_id="test-user-123",
- )
-
- usage_with_context = await BATCH_LIMITER.count_input_file_usage(
- file_id=file_obj.id,
- custom_llm_provider=CUSTOM_LLM_PROVIDER,
- user_api_key_dict=user_api_key_dict, # Passing user context
- )
- print(f"✓ Works with user context (tokens: {usage_with_context.total_tokens})")
- print(" Note: This fixes GEN-2166 for managed files")
-
- # Verify both return the same results
- assert usage_with_context.total_tokens == usage_without_context.total_tokens
- assert usage_with_context.request_count == usage_without_context.request_count
- print("\n✓ Both methods return identical results for non-managed files")
-
-
-@pytest.mark.asyncio()
-async def test_batch_rate_limiter_managed_files_regression():
- """
- Regression test for GEN-2166: Batch Rate Limiter Cannot Access User Files
-
- This test ensures that the batch rate limiter can properly access managed files
- by verifying that:
- 1. Managed files are detected correctly (base64 encoded unified file IDs)
- 2. The _fetch_managed_file_content method uses the managed files hook
- 3. User context (user_api_key_dict) is properly passed through
- 4. No 403 errors occur when accessing files owned by the user
- 5. The fix doesn't break non-managed file access
-
- This is a unit test that doesn't require external API calls.
- """
- from unittest.mock import AsyncMock, MagicMock, patch
- from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
- from litellm.types.llms.openai import HttpxBinaryResponseContent
- import httpx
-
- print("\n=== Regression Test: GEN-2166 Batch Rate Limiter Managed Files ===")
-
- # Setup: Create batch rate limiter
- dual_cache = DualCache()
- internal_usage_cache = InternalUsageCache(dual_cache=dual_cache)
- rate_limiter = PROXY_MaxParallelRequestsHandler_v3(
- internal_usage_cache=internal_usage_cache
- )
- batch_limiter = rate_limiter._get_batch_rate_limiter()
- assert batch_limiter is not None
-
- # Setup: Create user API key dict
- user_api_key_dict = UserAPIKeyAuth(
- api_key="test-key-regression",
- user_id="test-user-regression",
- tpm_limit=1000,
- rpm_limit=10,
- )
-
- # Setup: Create mock file content (batch input file)
- batch_content = b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Test message for regression"}]}}'
-
- # Mock managed file ID (base64 encoded unified file ID format)
- managed_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxyZWdyZXNzaW9uLXRlc3QtZmlsZQ=="
-
- # Test 1: Verify managed file detection
- print("\n1. Verifying managed file detection...")
- from litellm.proxy.openai_files_endpoints.common_utils import (
- is_base64_encoded_unified_file_id,
- )
-
- is_managed = is_base64_encoded_unified_file_id(managed_file_id)
- assert is_managed, "Managed file should be detected correctly"
- print(" ✓ Managed file detected")
-
- # Test 2: Verify _fetch_managed_file_content uses managed files hook
- print("\n2. Verifying managed files hook integration...")
-
- # Create mock managed files hook
- class MockManagedFiles(BaseFileEndpoints):
- def __init__(self):
- self._afile_content_called = False
- self._last_call_args = None
-
- async def acreate_file(self, *args, **kwargs):
- pass
-
- async def afile_content(self, *args, **kwargs):
- self._afile_content_called = True
- self._last_call_args = kwargs
- # Return mock file content
- mock_response = httpx.Response(
- status_code=200,
- content=batch_content,
- headers={"content-type": "application/octet-stream"},
- )
- return HttpxBinaryResponseContent(response=mock_response)
-
- async def afile_delete(self, *args, **kwargs):
- pass
-
- async def afile_list(self, *args, **kwargs):
- pass
-
- async def afile_retrieve(self, *args, **kwargs):
- pass
-
- mock_managed_files = MockManagedFiles()
- mock_llm_router = MagicMock()
- mock_proxy_logging_obj = MagicMock()
- mock_proxy_logging_obj.get_proxy_hook.return_value = mock_managed_files
-
- # Patch proxy_server imports
- with patch.dict(
- "sys.modules",
- {
- "litellm.proxy.proxy_server": MagicMock(
- llm_router=mock_llm_router,
- proxy_logging_obj=mock_proxy_logging_obj,
- )
- },
- ):
- # Call _fetch_managed_file_content
- result = await batch_limiter._fetch_managed_file_content(
- file_id=managed_file_id,
- user_api_key_dict=user_api_key_dict,
- )
-
- # Verify managed files hook was called
- assert (
- mock_managed_files._afile_content_called
- ), "REGRESSION: managed_files_obj.afile_content was not called! Bug GEN-2166 has returned."
-
- # Verify user context was passed
- assert (
- mock_managed_files._last_call_args is not None
- ), "REGRESSION: No arguments passed to afile_content"
- assert (
- "file_id" in mock_managed_files._last_call_args
- ), "REGRESSION: file_id not passed to managed files hook"
- assert (
- mock_managed_files._last_call_args["file_id"] == managed_file_id
- ), "REGRESSION: Incorrect file_id passed"
- assert (
- "llm_router" in mock_managed_files._last_call_args
- ), "REGRESSION: llm_router not passed to managed files hook"
-
- print(" ✓ Managed files hook called correctly")
- print(" ✓ User context passed correctly")
-
- # Test 3: Verify count_input_file_usage uses managed files path
- print("\n3. Verifying count_input_file_usage integration...")
-
- with patch.object(batch_limiter, "_fetch_managed_file_content") as mock_fetch:
- mock_response = httpx.Response(
- status_code=200,
- content=batch_content,
- headers={"content-type": "application/octet-stream"},
- )
- mock_fetch.return_value = HttpxBinaryResponseContent(response=mock_response)
-
- # Call count_input_file_usage with managed file
- usage = await batch_limiter.count_input_file_usage(
- file_id=managed_file_id,
- custom_llm_provider="openai",
- user_api_key_dict=user_api_key_dict,
- )
-
- # Verify _fetch_managed_file_content was called
- assert (
- mock_fetch.called
- ), "REGRESSION: _fetch_managed_file_content not called for managed files! Bug GEN-2166 has returned."
-
- # Verify correct parameters were passed
- call_kwargs = mock_fetch.call_args.kwargs
- assert (
- call_kwargs["file_id"] == managed_file_id
- ), "REGRESSION: Incorrect file_id passed to _fetch_managed_file_content"
- assert (
- call_kwargs["user_api_key_dict"] == user_api_key_dict
- ), "REGRESSION: user_api_key_dict not passed! Bug GEN-2166 has returned."
-
- # Verify usage was calculated
- assert usage.total_tokens > 0, "Token count should be greater than 0"
- assert usage.request_count == 1, "Request count should be 1"
-
- print(" ✓ Managed file path used")
- print(f" ✓ Token count: {usage.total_tokens}")
- print(f" ✓ Request count: {usage.request_count}")
-
- # Test 4: Verify non-managed files still work
- print("\n4. Verifying non-managed files still work...")
-
- non_managed_file_id = "file-abc123" # Standard OpenAI file ID
-
- with patch("litellm.afile_content") as mock_afile_content:
- mock_response = httpx.Response(
- status_code=200,
- content=batch_content,
- headers={"content-type": "application/octet-stream"},
- )
- mock_afile_content.return_value = HttpxBinaryResponseContent(
- response=mock_response
- )
-
- # Call count_input_file_usage with non-managed file
- usage = await batch_limiter.count_input_file_usage(
- file_id=non_managed_file_id,
- custom_llm_provider="openai",
- user_api_key_dict=user_api_key_dict,
- )
-
- # Verify litellm.afile_content was called
- assert (
- mock_afile_content.called
- ), "REGRESSION: litellm.afile_content not called for non-managed files"
-
- print(" ✓ Standard file path used")
- print(f" ✓ Token count: {usage.total_tokens}")
-
- # Test 5: Verify the fix prevents 403 errors
- print("\n5. Verifying 403 error prevention...")
-
- # Simulate the bug scenario: managed files hook not being used
- with patch.object(batch_limiter, "_fetch_managed_file_content") as mock_fetch:
- # If this is NOT called for managed files, the bug has returned
- mock_fetch.side_effect = Exception("Should not be called if bug exists")
-
- # This should call _fetch_managed_file_content
- try:
- with patch("litellm.afile_content") as mock_afile_content:
- # If litellm.afile_content is called for managed files, bug exists
- mock_afile_content.side_effect = Exception(
- "Error code: 403 - User does not have access to the file"
- )
-
- # Reset mock_fetch to return valid content
- mock_response = httpx.Response(
- status_code=200,
- content=batch_content,
- headers={"content-type": "application/octet-stream"},
- )
- mock_fetch.side_effect = None
- mock_fetch.return_value = HttpxBinaryResponseContent(
- response=mock_response
- )
-
- # This should use _fetch_managed_file_content, not litellm.afile_content
- usage = await batch_limiter.count_input_file_usage(
- file_id=managed_file_id,
- custom_llm_provider="openai",
- user_api_key_dict=user_api_key_dict,
- )
-
- # Verify managed files path was used (not standard path that causes 403)
- assert (
- mock_fetch.called
- ), "REGRESSION: Managed files path not used! This would cause 403 errors."
- assert (
- not mock_afile_content.called
- ), "REGRESSION: Standard path used for managed files! This causes 403 errors."
-
- print(" ✓ 403 error prevention verified")
-
- except Exception as e:
- if "403" in str(e):
- pytest.fail(
- f"REGRESSION: 403 error occurred! Bug GEN-2166 has returned. Error: {str(e)}"
- )
- raise
-
- print("\n=== Regression Test Passed ===")
- print("✓ Bug GEN-2166 is fixed and protected against regression")
- print("✓ Managed files are properly accessed via managed files hook")
- print("✓ User context is correctly passed through")
- print("✓ No 403 errors occur")
- print("✓ Non-managed files still work correctly\n")
diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py
deleted file mode 100644
index 6341fe2f91c..00000000000
--- a/tests/batches_tests/test_openai_batches_and_files.py
+++ /dev/null
@@ -1,580 +0,0 @@
-# What is this?
-## Unit Tests for OpenAI Batches API
-import asyncio
-import json
-import os
-import tempfile
-from dotenv import load_dotenv
-
-load_dotenv()
-
-import logging
-import time
-
-import pytest
-from typing import Optional
-import litellm
-from litellm._logging import verbose_logger
-import openai
-
-verbose_logger.setLevel(logging.DEBUG)
-
-from litellm.integrations.custom_logger import CustomLogger
-from litellm.types.utils import StandardLoggingPayload
-import socket
-import httpx
-from unittest.mock import patch, MagicMock, AsyncMock
-
-
-def _can_resolve_openai():
- """Check if api.openai.com is reachable (DNS resolves)."""
- try:
- socket.getaddrinfo("api.openai.com", 443, socket.AF_UNSPEC, socket.SOCK_STREAM)
- return True
- except socket.gaierror:
- return False
-
-
-skip_if_no_openai_network = pytest.mark.skipif(
- not _can_resolve_openai(),
- reason="Cannot resolve api.openai.com - skipping integration test due to DNS issues",
-)
-
-
-async def _wait_for_standard_logging_object(
- custom_logger: "TestCustomLogger", timeout: float = 15.0
-) -> StandardLoggingPayload:
- from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
-
- deadline = time.monotonic() + timeout
- while time.monotonic() < deadline:
- await GLOBAL_LOGGING_WORKER.flush()
- if custom_logger.standard_logging_object is not None:
- return custom_logger.standard_logging_object
- await asyncio.sleep(0.25)
- assert custom_logger.standard_logging_object is not None
- return custom_logger.standard_logging_object
-
-
-def load_vertex_ai_credentials():
- # Define the path to the vertex_key.json file
- print("loading vertex ai credentials")
- os.environ["GCS_FLUSH_INTERVAL"] = "1"
- filepath = os.path.dirname(os.path.abspath(__file__))
- vertex_key_path = filepath + "/vertex_key.json"
-
- # Read the existing content of the file or create an empty dictionary
- try:
- with open(vertex_key_path, "r") as file:
- # Read the file content
- print("Read vertexai file path")
- content = file.read()
-
- # If the file is empty or not valid JSON, create an empty dictionary
- if not content or not content.strip():
- service_account_key_data = {}
- else:
- # Attempt to load the existing JSON content
- file.seek(0)
- service_account_key_data = json.load(file)
- except FileNotFoundError:
- # If the file doesn't exist, create an empty dictionary
- service_account_key_data = {}
-
- # Update the service_account_key_data with environment variables
- private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "")
- private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "")
- private_key = private_key.replace("\\n", "\n")
- service_account_key_data["private_key_id"] = private_key_id
- service_account_key_data["private_key"] = private_key
-
- # Create a temporary file
- with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file:
- # Write the updated content to the temporary files
- json.dump(service_account_key_data, temp_file, indent=2)
-
- # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS
- os.environ["GCS_PATH_SERVICE_ACCOUNT"] = os.path.abspath(temp_file.name)
- os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name)
- print("created gcs path service account=", os.environ["GCS_PATH_SERVICE_ACCOUNT"])
-
-
-async def cancel_batch_unless_already_terminal(batch_id: str, provider: str) -> None:
- try:
- cancel_batch_response = await litellm.acancel_batch(batch_id=batch_id, custom_llm_provider=provider)
- except openai.ConflictError as e:
- if "Cannot cancel a batch with status 'completed'" in str(e):
- print(f"Batch already completed, cannot cancel: {e}")
- return
- if "Cannot cancel a batch with status 'failed'" not in str(e):
- raise
- failed_batch = await litellm.aretrieve_batch(batch_id=batch_id, custom_llm_provider=provider)
- print(f"Batch failed before cancel, errors={failed_batch.errors}")
- failure_codes = {err.code for err in (failed_batch.errors.data if failed_batch.errors else None) or []}
- assert failure_codes == {"token_limit_exceeded"}, (
- f"batch failed for a reason other than the org's enqueued token limit: {failed_batch.errors}"
- )
- return
- print("cancel_batch_response=", cancel_batch_response)
-
-
-class TestCustomLogger(CustomLogger):
- def __init__(self):
- super().__init__()
- self.standard_logging_object: Optional[StandardLoggingPayload] = None
-
- async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
- print(
- "Success event logged with kwargs=",
- kwargs,
- "and response_obj=",
- response_obj,
- )
- self.standard_logging_object = kwargs["standard_logging_object"]
-
-
-def cleanup_azure_files():
- """
- Delete all files for Azure - helper for when we run out of Azure Files Quota
- """
- azure_files = litellm.file_list(
- custom_llm_provider="azure",
- api_key=os.getenv("AZURE_FT_API_KEY"),
- api_base=os.getenv("AZURE_FT_API_BASE"),
- )
- print("azure_files=", azure_files)
- for _file in azure_files:
- print("deleting file=", _file)
- delete_file_response = litellm.file_delete(
- file_id=_file.id,
- custom_llm_provider="azure",
- api_key=os.getenv("AZURE_FT_API_KEY"),
- api_base=os.getenv("AZURE_FT_API_BASE"),
- )
- print("delete_file_response=", delete_file_response)
- assert delete_file_response.id == _file.id
-
-
-def cleanup_azure_ft_models():
- """
- Test CLEANUP: Delete all existing fine tuning jobs for Azure
- """
- try:
- from openai import AzureOpenAI
- import requests
-
- client = AzureOpenAI(
- api_key=os.getenv("AZURE_AI_API_KEY"),
- azure_endpoint=os.getenv("AZURE_AI_API_BASE"),
- api_version=os.getenv("AZURE_AI_API_VERSION"),
- )
-
- _list_ft_jobs = client.fine_tuning.jobs.list()
- print("_list_ft_jobs=", _list_ft_jobs)
-
- # delete all ft jobs make post request to this
- # Delete all fine-tuning jobs
- for job in _list_ft_jobs:
- try:
- endpoint = os.getenv("AZURE_FT_API_BASE").rstrip("/")
- url = f"{endpoint}/openai/fine_tuning/jobs/{job.id}?api-version=2024-10-21"
- print("url=", url)
-
- headers = {
- "api-key": os.getenv("AZURE_FT_API_KEY"),
- "Content-Type": "application/json",
- }
-
- response = requests.delete(url, headers=headers)
- print(f"Deleting job {job.id}: Status {response.status_code}")
- if response.status_code != 204:
- print(f"Error deleting job {job.id}: {response.text}")
-
- except Exception as e:
- print(f"Error deleting job {job.id}: {str(e)}")
- except Exception as e:
- print(f"Error on cleanup_azure_ft_models: {str(e)}")
-
-
-@pytest.mark.parametrize("provider", ["openai"])
-@pytest.mark.asyncio()
-@skip_if_no_openai_network
-async def test_async_create_batch(provider, tmp_path):
- """
- 1. Create File for Batch completion
- 2. Create Batch Request
- 3. Retrieve the specific batch
- """
- litellm.turn_on_debug()
- print("Testing async create batch")
- litellm.logging_callback_manager._reset_all_callbacks()
-
- file_name = "openai_batch_completions.jsonl"
- _current_dir = os.path.dirname(os.path.abspath(__file__))
- file_path = os.path.join(_current_dir, file_name)
- with open(file_path, "rb") as batch_file:
- file_obj = await litellm.acreate_file(
- file=batch_file,
- purpose="batch",
- custom_llm_provider=provider,
- )
- print("Response from creating file=", file_obj)
-
- await asyncio.sleep(10)
- batch_input_file_id = file_obj.id
- assert (
- batch_input_file_id is not None
- ), "Failed to create file, expected a non null file_id but got {batch_input_file_id}"
-
- extra_metadata_field = {
- "user_api_key_alias": "special_api_key_alias",
- "user_api_key_team_alias": "special_team_alias",
- }
- custom_logger = TestCustomLogger()
- litellm.callbacks = [custom_logger, "datadog"]
- create_batch_response = await litellm.acreate_batch(
- completion_window="24h",
- endpoint="/v1/chat/completions",
- input_file_id=batch_input_file_id,
- custom_llm_provider=provider,
- metadata={"key1": "value1", "key2": "value2"},
- # litellm specific param - used for logging metadata on logging callback
- litellm_metadata=extra_metadata_field,
- )
-
- print("response from litellm.create_batch=", create_batch_response)
-
- assert (
- create_batch_response.id is not None
- ), f"Failed to create batch, expected a non null batch_id but got {create_batch_response.id}"
- assert (
- create_batch_response.endpoint == "/v1/chat/completions"
- or create_batch_response.endpoint == "/chat/completions"
- ), f"Failed to create batch, expected endpoint to be /v1/chat/completions but got {create_batch_response.endpoint}"
- assert (
- create_batch_response.input_file_id == batch_input_file_id
- ), f"Failed to create batch, expected input_file_id to be {batch_input_file_id} but got {create_batch_response.input_file_id}"
-
- # Assert that the create batch event is logged on CustomLogger
- standard_logging_object = await _wait_for_standard_logging_object(custom_logger)
- print(
- "standard_logging_object=",
- json.dumps(standard_logging_object, indent=4, default=str),
- )
- assert (
- standard_logging_object["metadata"]["user_api_key_alias"]
- == extra_metadata_field["user_api_key_alias"]
- )
- assert (
- standard_logging_object["metadata"]["user_api_key_team_alias"]
- == extra_metadata_field["user_api_key_team_alias"]
- )
-
- retrieved_batch = await litellm.aretrieve_batch(
- batch_id=create_batch_response.id, custom_llm_provider=provider
- )
- print("retrieved batch=", retrieved_batch)
- # just assert that we retrieved a non None batch
-
- assert retrieved_batch.id == create_batch_response.id
-
- # list all batches
- list_batches = await litellm.alist_batches(custom_llm_provider=provider, limit=2)
- print("list_batches=", list_batches)
-
- # try to get file content for our original file
-
- file_content = await litellm.afile_content(
- file_id=batch_input_file_id, custom_llm_provider=provider
- )
-
- print("file content = ", file_content)
-
- # file obj
- file_obj = await litellm.afile_retrieve(
- file_id=batch_input_file_id, custom_llm_provider=provider
- )
- print("file obj = ", file_obj)
- assert file_obj.id == batch_input_file_id
-
- # delete file
- delete_file_response = await litellm.afile_delete(
- file_id=batch_input_file_id, custom_llm_provider=provider
- )
-
- print("delete file response = ", delete_file_response)
-
- assert delete_file_response.id == batch_input_file_id
-
- all_files_list = await litellm.afile_list(
- custom_llm_provider=provider,
- )
-
- print("all_files_list = ", all_files_list)
-
- result_file_path = tmp_path / "batch_job_results_furniture.jsonl"
- result_file_path.write_bytes(file_content.content)
-
- await cancel_batch_unless_already_terminal(batch_id=create_batch_response.id, provider=provider)
-
-
-mock_file_response = {
- "kind": "storage#object",
- "id": "litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb/1739598666670574",
- "selfLink": "https://www.googleapis.com/storage/v1/b/litellm-local/o/litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2F5f7b99ad-9203-4430-98bf-3b45451af4cb",
- "mediaLink": "https://storage.googleapis.com/download/storage/v1/b/litellm-local/o/litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2F5f7b99ad-9203-4430-98bf-3b45451af4cb?generation=1739598666670574&alt=media",
- "name": "litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb",
- "bucket": "litellm-local",
- "generation": "1739598666670574",
- "metageneration": "1",
- "contentType": "application/json",
- "storageClass": "STANDARD",
- "size": "416",
- "md5Hash": "hbBNj7C8KJ7oVH+JmyRM6A==",
- "crc32c": "oDmiUA==",
- "etag": "CO7D0IT+xIsDEAE=",
- "timeCreated": "2025-02-15T05:51:06.741Z",
- "updated": "2025-02-15T05:51:06.741Z",
- "timeStorageClassUpdated": "2025-02-15T05:51:06.741Z",
- "timeFinalized": "2025-02-15T05:51:06.741Z",
-}
-
-mock_vertex_batch_response = {
- "name": "projects/123456789/locations/us-central1/batchPredictionJobs/test-batch-id-456",
- "displayName": "litellm_batch_job",
- "model": "projects/123456789/locations/us-central1/models/gemini-1.5-flash-001",
- "modelVersionId": "v1",
- "inputConfig": {
- "gcsSource": {
- "uris": [
- "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb"
- ]
- }
- },
- "outputConfig": {
- "gcsDestination": {"outputUriPrefix": "gs://litellm-local/batch-outputs/"}
- },
- "dedicatedResources": {
- "machineSpec": {
- "machineType": "n1-standard-4",
- "acceleratorType": "NVIDIA_TESLA_T4",
- "acceleratorCount": 1,
- },
- "startingReplicaCount": 1,
- "maxReplicaCount": 1,
- },
- "state": "JOB_STATE_RUNNING",
- "createTime": "2025-02-15T05:51:06.741Z",
- "startTime": "2025-02-15T05:51:07.741Z",
- "updateTime": "2025-02-15T05:51:08.741Z",
- "labels": {"key1": "value1", "key2": "value2"},
- "completionStats": {"successfulCount": 0, "failedCount": 0, "remainingCount": 100},
-}
-
-mock_vertex_list_response = {
- "batchPredictionJobs": [
- mock_vertex_batch_response,
- {
- **mock_vertex_batch_response,
- "name": "projects/123456789/locations/us-central1/batchPredictionJobs/test-batch-id-789",
- "state": "JOB_STATE_SUCCEEDED",
- },
- ],
- "nextPageToken": "",
-}
-
-
-@pytest.mark.asyncio
-async def test_avertex_batch_prediction(monkeypatch):
- monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local")
- monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project")
- monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1")
-
- # Mock Google auth so the test doesn't need real credentials
- mock_creds = MagicMock()
- mock_creds.token = "mock-token"
- mock_creds.valid = True
- mock_creds.expiry = None
- monkeypatch.setattr(
- "google.auth.default",
- lambda *args, **kwargs: (mock_creds, "mock-project"),
- )
-
- from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
-
- # Configure mock response object
- mock_response = MagicMock()
- mock_response.raise_for_status.return_value = None
-
- async def mock_side_effect(*args, **kwargs):
- print("args", args, "kwargs", kwargs)
- url = kwargs.get("url", "")
- if "files" in url:
- mock_response.json.return_value = mock_file_response
- elif "batch" in url:
- mock_response.json.return_value = mock_vertex_batch_response
- mock_response.status_code = 200
- return mock_response
-
- # Batch jsonl creation now stages the body to a temp file and issues a single
- # uploadType=media POST against the raw httpx.AsyncClient (client.client) inside
- # _astage_and_upload_media, not AsyncHTTPHandler.post. Patch that raw POST so the
- # real staging/upload + response transform run while the GCS object response is
- # mocked; AsyncHTTPHandler.post still handles the batch-prediction call.
- with (
- patch(
- "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
- side_effect=mock_side_effect,
- ),
- patch.object(
- httpx.AsyncClient,
- "post",
- new_callable=AsyncMock,
- return_value=httpx.Response(
- 200,
- json=mock_file_response,
- request=httpx.Request("POST", "https://storage.googleapis.com/upload"),
- ),
- ) as mock_gcs_upload,
- ):
- litellm.set_verbose = True
- litellm.turn_on_debug()
- file_name = "vertex_batch_completions.jsonl"
- _current_dir = os.path.dirname(os.path.abspath(__file__))
- file_path = os.path.join(_current_dir, file_name)
-
- # Create file
- file_obj = await litellm.acreate_file(
- file=open(file_path, "rb"),
- purpose="batch",
- custom_llm_provider="vertex_ai",
- )
- print("Response from creating file=", file_obj)
-
- assert (
- file_obj.id
- == "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb"
- )
-
- mock_gcs_upload.assert_awaited_once()
- upload_url = str(mock_gcs_upload.call_args.args[0])
- assert "uploadType=media" in upload_url
- assert "/b/litellm-local/o" in upload_url
- assert (
- mock_gcs_upload.call_args.kwargs["headers"]["Content-Type"]
- == "application/json"
- )
-
- # Create batch
- create_batch_response = await litellm.acreate_batch(
- completion_window="24h",
- endpoint="/v1/chat/completions",
- input_file_id=file_obj.id,
- custom_llm_provider="vertex_ai",
- metadata={"key1": "value1", "key2": "value2"},
- )
- print("create_batch_response=", create_batch_response)
-
- assert create_batch_response.id == "test-batch-id-456"
- assert (
- create_batch_response.input_file_id
- == "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb"
- )
-
- # Mock the retrieve batch response
- with patch(
- "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
- ) as mock_get:
- mock_get_response = MagicMock()
- mock_get_response.json.return_value = mock_vertex_batch_response
- mock_get_response.status_code = 200
- mock_get_response.is_redirect = False
- mock_get_response.raise_for_status.return_value = None
- mock_get_response.is_redirect = False
- mock_get.return_value = mock_get_response
-
- retrieved_batch = await litellm.aretrieve_batch(
- batch_id=create_batch_response.id,
- custom_llm_provider="vertex_ai",
- )
- print("retrieved_batch=", retrieved_batch)
-
- assert retrieved_batch.id == "test-batch-id-456"
-
-
-@pytest.mark.asyncio
-
-
-@pytest.mark.asyncio
-
-
-@pytest.mark.asyncio
-@skip_if_no_openai_network
-async def test_delete_batch_output_file():
- """
- Test that deleting a batch output file works correctly.
-
- This test verifies the fix for:
- - When a batch is retrieved and has an output_file_id, the file object is properly stored
- - The output file can be deleted without validation errors
- - The file_object is fetched and stored with proper metadata instead of None
- """
- litellm.turn_on_debug()
- print("Testing delete batch output file")
-
- file_name = "openai_batch_completions.jsonl"
- _current_dir = os.path.dirname(os.path.abspath(__file__))
- file_path = os.path.join(_current_dir, file_name)
-
- # Create file for batch
- file_obj = await litellm.acreate_file(
- file=open(file_path, "rb"),
- purpose="batch",
- custom_llm_provider="openai",
- )
- print("Response from creating file=", file_obj)
- batch_input_file_id = file_obj.id
-
- # Create batch
- create_batch_response = await litellm.acreate_batch(
- completion_window="24h",
- endpoint="/v1/chat/completions",
- input_file_id=batch_input_file_id,
- custom_llm_provider="openai",
- )
- print("Batch created with ID=", create_batch_response.id)
-
- # Retrieve batch to get output_file_id
- retrieved_batch = await litellm.aretrieve_batch(
- batch_id=create_batch_response.id, custom_llm_provider="openai"
- )
- print("Retrieved batch=", retrieved_batch)
-
- # If batch has completed and has output file, test deleting it
- if retrieved_batch.output_file_id:
- print(f"Testing deletion of output file: {retrieved_batch.output_file_id}")
-
- # This is the key test - deleting the output file should work
- # without validation errors (file_object should not be None)
- delete_output_file_response = await litellm.afile_delete(
- file_id=retrieved_batch.output_file_id, custom_llm_provider="openai"
- )
-
- print("Delete output file response=", delete_output_file_response)
- assert delete_output_file_response.id == retrieved_batch.output_file_id
- assert delete_output_file_response.deleted is True or hasattr(
- delete_output_file_response, "id"
- )
- print("✓ Successfully deleted batch output file")
- else:
- print(
- "⚠ Batch has not completed yet or no output file available, skipping output file deletion test"
- )
-
- # Clean up - delete the input file
- delete_input_file_response = await litellm.afile_delete(
- file_id=batch_input_file_id, custom_llm_provider="openai"
- )
- print("Delete input file response=", delete_input_file_response)
- assert delete_input_file_response.id == batch_input_file_id
- print("✓ Successfully deleted batch input file")
diff --git a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md
index d9380098891..e3f3bf009dc 100644
--- a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md
+++ b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md
@@ -43,7 +43,7 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
| Tag | `test_update_daily_tag_spend.py` | partial | yes (`test_tag_spend_matches_sum_of_tagged_logs`) |
| End-user | `test_proxy_update_spend.py` | covered | yes |
| Spend == sum(logs) consistency | none | gap | yes (key + tag aggregate == sum of rows) |
-| Concurrent increments (one key, parallel writers) | `tests/spend_tracking_tests/test_spend_accuracy_tests.py` (burst) | partial | yes (`test_burst_of_concurrent_calls_loses_no_spend`) |
+| Concurrent increments (one key, parallel writers) | `tests/integration/spend/test_spend_rollup_accuracy.py`, `tests/integration/spend/test_chaos_burst_spend_once.py` (burst) | partial | yes (`test_burst_of_concurrent_calls_loses_no_spend`) |
## Spend read endpoints (verification surface)
diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py
deleted file mode 100644
index aeeb34f4e8d..00000000000
--- a/tests/guardrails_tests/test_bedrock_guardrails.py
+++ /dev/null
@@ -1,100 +0,0 @@
-import io, asyncio
-import pytest
-
-import litellm
-from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
- BedrockGuardrail,
- _redact_pii_matches,
-)
-from litellm.proxy._types import UserAPIKeyAuth
-from unittest.mock import MagicMock, AsyncMock, patch
-
-
-@pytest.mark.asyncio
-async def test_bedrock_guardrails_pii_masking():
- # Create proper mock objects
- mock_user_api_key_dict = UserAPIKeyAuth()
-
- guardrail = BedrockGuardrail(
- guardrailIdentifier="wf0hkdb5x07f",
- guardrailVersion="DRAFT",
- )
-
- request_data = {
- "model": "gpt-5.5",
- "messages": [
- {"role": "user", "content": "Hello, my phone number is +1 412 555 1212"},
- {"role": "assistant", "content": "Hello, how can I help you today?"},
- {"role": "user", "content": "I need to cancel my order"},
- {
- "role": "user",
- "content": "ok, my credit card number is 1234-5678-9012-3456",
- },
- ],
- }
-
- response = await guardrail.async_moderation_hook(
- data=request_data,
- user_api_key_dict=mock_user_api_key_dict,
- call_type="completion",
- )
- print("response after moderation hook", response)
-
- if response: # Only assert if response is not None
- assert response["messages"][0]["content"] == "Hello, my phone number is {PHONE}"
- assert response["messages"][1]["content"] == "Hello, how can I help you today?"
- assert response["messages"][2]["content"] == "I need to cancel my order"
- assert (
- response["messages"][3]["content"]
- == "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}"
- )
-
-
-@pytest.mark.asyncio
-async def test_bedrock_guardrails_pii_masking_content_list():
- # Create proper mock objects
- mock_user_api_key_dict = UserAPIKeyAuth()
-
- guardrail = BedrockGuardrail(
- guardrailIdentifier="wf0hkdb5x07f",
- guardrailVersion="DRAFT",
- )
-
- request_data = {
- "model": "gpt-5.5",
- "messages": [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Hello, my phone number is +1 412 555 1212",
- },
- {"type": "text", "text": "what time is it?"},
- ],
- },
- {"role": "assistant", "content": "Hello, how can I help you today?"},
- {"role": "user", "content": "who is the president of the united states?"},
- ],
- }
-
- response = await guardrail.async_moderation_hook(
- data=request_data,
- user_api_key_dict=mock_user_api_key_dict,
- call_type="completion",
- )
- print(response)
-
- if response: # Only assert if response is not None
- # Verify that the list content is properly masked
- assert isinstance(response["messages"][0]["content"], list)
- assert (
- response["messages"][0]["content"][0]["text"]
- == "Hello, my phone number is {PHONE}"
- )
- assert response["messages"][0]["content"][1]["text"] == "what time is it?"
- assert response["messages"][1]["content"] == "Hello, how can I help you today?"
- assert (
- response["messages"][2]["content"]
- == "who is the president of the united states?"
- )
diff --git a/tests/guardrails_tests/test_presidio_pii.py b/tests/guardrails_tests/test_presidio_pii.py
deleted file mode 100644
index 595cc95fa9f..00000000000
--- a/tests/guardrails_tests/test_presidio_pii.py
+++ /dev/null
@@ -1,211 +0,0 @@
-import os
-import pytest
-from litellm import mock_completion
-from unittest.mock import patch
-
-import litellm
-from litellm.proxy.guardrails.guardrail_hooks.presidio import (
- OPTIONAL_PresidioPIIMasking,
- PresidioPerRequestConfig,
-)
-from litellm.types.guardrails import PiiEntityType, PiiAction
-from litellm.proxy._types import UserAPIKeyAuth
-from litellm.caching.caching import DualCache
-from litellm.exceptions import BlockedPiiEntityError
-
-
-@pytest.mark.asyncio
-async def test_presidio_with_blocked_entities():
- """Test for Presidio guardrail with blocked entities - requires actual Presidio API"""
- # Setup the guardrail with specific entities config - BLOCK for credit card
- litellm.turn_on_debug()
- pii_entities_config = {
- PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, # This entity should cause a block
- PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked
- }
-
- presidio_guardrail = OPTIONAL_PresidioPIIMasking(
- pii_entities_config=pii_entities_config,
- presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
- presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
- )
-
- # Test text with blocked PII type
- test_text = (
- "My credit card number is 4111-1111-1111-1111 and my email is test@example.com"
- )
-
- # Verify the analyze request configuration
- analyze_request = presidio_guardrail._get_presidio_analyze_request_payload(
- text=test_text, presidio_config=None, request_data={}
- )
-
- # Verify entities were passed correctly
- assert "entities" in analyze_request
- assert set(analyze_request["entities"]) == set(pii_entities_config.keys())
-
- # Test that BlockedPiiEntityError is raised when check_pii is called
- with pytest.raises(BlockedPiiEntityError) as excinfo:
- await presidio_guardrail.check_pii(
- text=test_text, output_parse_pii=True, presidio_config=None, request_data={}
- )
-
- # Verify the error contains the correct entity type
- assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD
- assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name
-
-
-@pytest.mark.asyncio
-async def test_presidio_pre_call_hook_with_blocked_entities():
- """Test for Presidio guardrail pre-call hook with blocked entities on a chat completion request"""
- # Setup the guardrail with specific entities config
- pii_entities_config = {
- PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, # This entity should cause a block
- PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked
- }
-
- presidio_guardrail = OPTIONAL_PresidioPIIMasking(
- pii_entities_config=pii_entities_config,
- presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
- presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
- )
-
- # Create a sample chat completion request with PII data
- data = {
- "messages": [
- {"role": "system", "content": "You are a helpful assistant."},
- {
- "role": "user",
- "content": "My credit card is 4111-1111-1111-1111 and my email is test@example.com.",
- },
- ],
- "model": "gpt-5-mini",
- }
-
- # Mock objects needed for the pre-call hook
- user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
- cache = DualCache()
-
- # Call the pre-call hook and expect BlockedPiiEntityError
- with pytest.raises(BlockedPiiEntityError) as excinfo:
- await presidio_guardrail.async_pre_call_hook(
- user_api_key_dict=user_api_key_dict,
- cache=cache,
- data=data,
- call_type="completion",
- )
-
- print(f"got error: {excinfo}")
-
- # Verify the error contains the correct entity type
- assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD
- assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name
-
-
-
-
-
-
-# asyncio.run(test_output_parsing())
-
-
-### UNIT TESTS FOR PRESIDIO PII MASKING ###
-
-input_a_anonymizer_results = {
- "text": "hello world, my name is