diff --git a/.circleci/config.yml b/.circleci/config.yml index d8dc40433dc..58860e3ed95 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -2074,98 +2074,6 @@ jobs: # Store test results - store_test_results: path: test-results - proxy_spend_accuracy_tests: - machine: - image: ubuntu-2204:2024.04.1 - resource_class: large - working_directory: ~/project - steps: - - checkout - - run: - name: Generate LiteLLM master key - command: | - key="$(openssl rand -hex 16)" - printf 'export LITELLM_MASTER_KEY=sk-%s\n' "$key" >> "$BASH_ENV" - - skip_if_unrelated_changes - - setup_google_dns - - install_uv - - install_rust - - run: - name: Install Dependencies - command: | - uv sync --frozen --all-groups --all-extras --python 3.12 - - start_postgres - - start_redis - - start_fake_openai_endpoint - - attach_workspace: - at: ~/project - - run: - name: Load Docker Database Image - command: | - zstd -d litellm-docker-database.tar.zst --stdout | docker load - docker images | grep litellm-docker-database - - run: - name: Run Docker container - # Point the proxy at the job-local Redis (start_redis) instead of the - # shared remote Redis. The Redis transaction buffer uses a single - # global pod-lock key (cronjob_lock:db_spend_update_job) and a single - # global buffer list (litellm_spend_update_buffer); sharing those - # across concurrent CI pipelines causes spend flushes to stall or - # land in the wrong DB, which is what makes this test flaky. - command: | - docker run -d \ - -p 4000:4000 \ - -e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \ - -e REDIS_HOST=host.docker.internal \ - -e REDIS_PORT=6379 \ - -e LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" \ - -e OPENAI_API_KEY=$OPENAI_API_KEY \ - -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ - -e LITELLM_LICENSE=$LITELLM_LICENSE \ - -e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \ - -e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \ - -e USE_DDTRACE=True \ - -e DD_API_KEY=$DD_API_KEY \ - -e DD_SITE=$DD_SITE \ - -e AWS_REGION_NAME=$AWS_REGION_NAME \ - -e PROXY_BATCH_WRITE_AT=2 \ - -e LITELLM_LOG=ERROR \ - --add-host host.docker.internal:host-gateway \ - --name my-app \ - -v $(pwd)/litellm/proxy/example_config_yaml/spend_tracking_config.yaml:/app/config.yaml \ - litellm-docker-database:ci \ - --config /app/config.yaml \ - --port 4000 - - run: - name: Start outputting logs - command: docker logs -f my-app - background: true - - wait_for_service: - url: http://localhost:4000 - timeout: "300" - - run: - name: Run tests - command: | - mkdir -p test-results - TEST_FILES=$(circleci tests glob "tests/spend_tracking_tests/**/test_*.py") - echo "$TEST_FILES" | circleci tests run \ - --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ - -vv \ - --junitxml=test-results/junit.xml \ - --durations=5" - no_output_timeout: 15m - - store_test_results: - path: test-results - - run: - name: Stop and remove first container - when: always - command: | - docker stop my-app - docker rm my-app - docker stop redis-cache - docker rm redis-cache - proxy_multi_instance_tests: machine: image: ubuntu-2204:2024.04.1 @@ -3591,9 +3499,6 @@ workflows: - proxy_logging_guardrails_model_info_tests: requires: - build_docker_database_image - - proxy_spend_accuracy_tests: - requires: - - build_docker_database_image - proxy_multi_instance_tests: requires: - build_docker_database_image diff --git a/tests/audio_tests/test_audio_speech.py b/tests/audio_tests/test_audio_speech.py index 998de5ecc3b..781dff654f7 100644 --- a/tests/audio_tests/test_audio_speech.py +++ b/tests/audio_tests/test_audio_speech.py @@ -8,248 +8,12 @@ from dotenv import load_dotenv load_dotenv() from pathlib import Path -from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm -async def _run_audio_speech_litellm(sync_mode, model, api_base, api_key): - litellm.turn_on_debug() - speech_file_path = Path(__file__).parent / "speech.mp3" - - if sync_mode: - response = litellm.speech( - model=model, - voice="alloy", - input="the quick brown fox jumped over the lazy dogs", - api_base=api_base, - api_key=api_key, - organization=None, - project=None, - max_retries=1, - timeout=600, - client=None, - optional_params={}, - ) - - from litellm.types.llms.openai import HttpxBinaryResponseContent - - assert isinstance(response, HttpxBinaryResponseContent) - else: - response = await litellm.aspeech( - model=model, - voice="alloy", - input="the quick brown fox jumped over the lazy dogs", - api_base=api_base, - api_key=api_key, - organization=None, - project=None, - max_retries=1, - timeout=600, - client=None, - optional_params={}, - ) - - from litellm.llms.openai.openai import HttpxBinaryResponseContent - - assert isinstance(response, HttpxBinaryResponseContent) - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_audio_speech_litellm_azure(sync_mode): - await _run_audio_speech_litellm( - sync_mode=sync_mode, - model="azure/tts", - api_base=os.getenv("AZURE_TTS_API_BASE"), - api_key=os.getenv("AZURE_TTS_API_KEY"), - ) - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_audio_speech_litellm_openai(sync_mode): - await _run_audio_speech_litellm( - sync_mode=sync_mode, - model="openai/tts-1", - api_base=None, - api_key=os.getenv("OPENAI_API_KEY"), - ) - - - - -@pytest.mark.flaky(retries=6, delay=2) -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - # 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( - model=model, - input="async hello what llm guardrail do you have", - ) - 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": {"text": "async hello what llm guardrail do you have"}, - "voice": {"languageCode": "en-US", "name": "en-US-Studio-O"}, - "audioConfig": {"audioEncoding": "LINEAR16", "speakingRate": "1"}, - } - - -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async_with_voice(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - # 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( - model=model, - input="async hello what llm guardrail do you have", - 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": {"text": "async hello what llm guardrail do you have"}, - "voice": {"languageCode": "en-UK", "name": "en-UK-Studio-O"}, - "audioConfig": {"audioEncoding": "LINEAR22", "speakingRate": "10"}, - } - - -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async_with_voice_ssml(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - ssml = """ - -

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 . My number is: ", - "items": [ - { - "start": 48, - "end": 62, - "entity_type": "PHONE_NUMBER", - "text": "", - "operator": "replace", - }, - { - "start": 24, - "end": 32, - "entity_type": "PERSON", - "text": "", - "operator": "replace", - }, - ], -} - -input_b_anonymizer_results = { - "text": "My name is , who are you? Say my name in your response", - "items": [ - { - "start": 11, - "end": 19, - "entity_type": "PERSON", - "text": "", - "operator": "replace", - } - ], -} - - -# Test if PII masking works with input A - - -# Test if PII masking works with input B (also test if the response != A's response) - - - - -@pytest.mark.asyncio -@patch.dict( - os.environ, - { - "PRESIDIO_ANALYZER_API_BASE": "http://localhost:5002", - "PRESIDIO_ANONYMIZER_API_BASE": "http://localhost:5001", - }, -) -async def test_presidio_pii_masking_logging_output_only_logged_response_guardrails_config(): - from typing import Dict, List, Optional - - import litellm - from litellm.proxy.guardrails.init_guardrails import initialize_guardrails - from litellm.types.guardrails import ( - GuardrailItemSpec, - GuardrailEventHooks, - ) - - litellm.set_verbose = True - # Environment variables are now patched via the decorator instead of setting them directly - - guardrails_config: List[Dict[str, GuardrailItemSpec]] = [ - { - "pii_masking": { - "callbacks": ["presidio"], - "default_on": True, - "logging_only": True, - } - } - ] - litellm_settings = {"guardrails": guardrails_config} - - assert len(litellm.guardrail_name_config_map) == 0 - initialize_guardrails( - guardrails_config=guardrails_config, - premium_user=True, - config_file_path="", - litellm_settings=litellm_settings, - ) - - assert len(litellm.guardrail_name_config_map) == 1 - - pii_masking_obj: Optional[OPTIONAL_PresidioPIIMasking] = None - for callback in litellm.callbacks: - print(f"CALLBACK: {callback}") - if isinstance(callback, OPTIONAL_PresidioPIIMasking): - pii_masking_obj = callback - - assert pii_masking_obj is not None - - assert hasattr(pii_masking_obj, "logging_only") - assert pii_masking_obj.event_hook == GuardrailEventHooks.logging_only - - assert pii_masking_obj.should_run_guardrail( - data={}, event_type=GuardrailEventHooks.logging_only - ) diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py deleted file mode 100644 index e7d526beeb8..00000000000 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ /dev/null @@ -1,112 +0,0 @@ -import logging -import traceback - -from dotenv import load_dotenv -from openai.types.image import Image - - -from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( - AmazonNovaCanvasConfig, -) - -logging.basicConfig(level=logging.DEBUG) -load_dotenv() -import asyncio - -import pytest -from litellm.llms.bedrock.image_generation.cost_calculator import cost_calculator -from litellm.types.utils import ImageResponse, ImageObject - -import litellm -from litellm.llms.bedrock.image_generation.amazon_stability3_transformation import ( - AmazonStability3Config, -) -from litellm.llms.bedrock.image_generation.amazon_stability1_transformation import ( - AmazonStabilityConfig, -) -from litellm.types.llms.bedrock import ( - AmazonStability3TextToImageRequest, - AmazonStability3TextToImageResponse, -) -from unittest.mock import MagicMock, patch -from litellm.llms.bedrock.image_generation.image_handler import ( - BedrockImageGeneration, - BedrockImagePreparedRequest, -) -from litellm.llms.bedrock.common_utils import BedrockError - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -# Test cases for issue #14373 - Bedrock Application Inference Profiles with Nova Canvas - - - - - - - - -def test_amazon_nova_canvas_image_gen(): - """Test Amazon Nova Canvas image generation with cost tracking.""" - from litellm import image_generation - - model_id = "bedrock/amazon.nova-canvas-v1:0" - - response = litellm.image_generation( - model=model_id, - prompt="A serene mountain landscape at sunset with a lake reflection", - aws_region_name="us-east-1", - ) - - print(f"response cost: {response._hidden_params['response_cost']}") - - assert response._hidden_params["response_cost"] > 0 diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index fdcbf6fcd8b..eca22196287 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -193,182 +193,3 @@ async def test_openai_image_edit_litellm_router(): f.write(image_bytes) except litellm.ContentPolicyViolationError as e: pass - - -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_openai_image_edit_with_bytesio(): - """Test image editing using BytesIO objects instead of file readers""" - from litellm import aimage_edit, image_edit - - litellm.turn_on_debug() - try: - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - - # Get images as BytesIO objects - bytesio_images = get_test_images_as_bytesio() - - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=bytesio_images, - ) - print("result from image edit with BytesIO", result) - - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) - - # Save the image to a file - with open("test_image_edit_bytesio.png", "wb") as f: - f.write(image_bytes) - except litellm.ContentPolicyViolationError as e: - pass - - - - - - -@pytest.mark.asyncio -async def test_azure_image_edit_cost_tracking(): - """Test Azure image edit cost tracking with custom logger""" - from litellm import aimage_edit, image_edit - - test_custom_logger = TestCustomLogger() - litellm.logging_callback_manager._reset_all_callbacks() - litellm.callbacks = [test_custom_logger] - - # Mock response for Azure image edit with usage data for cost tracking - mock_response = { - "created": 1589478378, - "data": [ - { - "b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg==" - } - ], - "usage": { - "total_tokens": 1100, - "input_tokens": 100, - "input_tokens_details": {"image_tokens": 50, "text_tokens": 50}, - "output_tokens": 1000, - }, - } - - class MockResponse: - def __init__(self, json_data, status_code): - self._json_data = json_data - self.status_code = status_code - self.text = json.dumps(json_data) - self.headers = {} - - def json(self): - return self._json_data - - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_post: - # Configure the mock to return our response - mock_post.return_value = MockResponse(mock_response, 200) - - litellm.turn_on_debug() - - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - - # Set up test environment variables - - result = await aimage_edit( - prompt=prompt, - model="azure/CUSTOM_AZURE_DEPLOYMENT_NAME", - base_model="azure/gpt-image-1", - image=_make_test_images(), - ) - - # Verify the request was made correctly - mock_post.assert_called_once() - - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) - - # Save the image to a file - with open("test_image_edit.png", "wb") as f: - f.write(image_bytes) - - await asyncio.sleep(5) - print( - "standard logging payload", - json.dumps( - test_custom_logger.standard_logging_payload, indent=4, default=str - ), - ) - - # check model - assert ( - test_custom_logger.standard_logging_payload["model"] - == "CUSTOM_AZURE_DEPLOYMENT_NAME" - ) - assert ( - test_custom_logger.standard_logging_payload["custom_llm_provider"] - == "azure" - ) - - # check response_cost - assert test_custom_logger.standard_logging_payload["response_cost"] is not None - assert test_custom_logger.standard_logging_payload["response_cost"] > 0 - - - - - - -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_multiple_image_edit_with_different_formats(): - """Test multiple images editing with different file formats and types""" - from litellm import aimage_edit - - litellm.turn_on_debug() - - try: - prompt = "Create a cohesive artistic style across all images" - - mixed_images = [ - _make_single_test_image(), - get_test_images_as_bytesio()[1], - ] - - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=mixed_images, - ) - - print("Mixed format images result:", result) - ImageResponse.model_validate(result) - - assert result is not None - assert result.data is not None - assert len(result.data) > 0 - - # Save result if available - if result.data and result.data[0].b64_json: - image_bytes = base64.b64decode(result.data[0].b64_json) - with open("test_multiple_image_edit_mixed.png", "wb") as f: - f.write(image_bytes) - - except litellm.ContentPolicyViolationError as e: - pytest.skip(f"Content policy violation: {e}") diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py new file mode 100644 index 00000000000..834e10e32b8 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py @@ -0,0 +1,584 @@ +import json +import uuid +from collections.abc import Callable, Generator, Mapping +from contextlib import contextmanager +from hashlib import sha256 +from pathlib import Path +from typing import Final + +from integration._support.client import Gateway, JsonValue, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server + +_MODEL: Final = "claude-sonnet-4-5-20250929" +_KEY: Final = "synthetic-anthropic-key" + +_PROXY_CONFIG: Final = ( + "model_list: []\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " disable_spend_logs: false\n" + " proxy_batch_write_at: 1\n" +) + + +def _message(message_id: str, model: str = _MODEL) -> dict[str, JsonValue]: + return { + "id": message_id, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": "hello test"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 11, "output_tokens": 7}, + } + + +def _thinking_message(message_id: str) -> dict[str, JsonValue]: + return { + **_message(message_id, "claude-haiku-4-5-20251001"), + "content": [ + {"type": "thinking", "thinking": "pondering the joke", "signature": "sig1"}, + {"type": "text", "text": "hello thinking"}, + ], + "usage": {"input_tokens": 11, "output_tokens": 30}, + } + + +def _sse(event: str, payload: Mapping[str, JsonValue]) -> bytes: + return f"event: {event}\ndata: {json.dumps(payload)}\n\n".encode() + + +def _stream_chunks(message_id: str) -> tuple[bytes, ...]: + return ( + _sse( + "message_start", + { + "type": "message_start", + "message": { + "id": message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [], + "usage": {"input_tokens": 11, "output_tokens": 1}, + }, + }, + ), + _sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello stream"}}, + ), + _sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _sse( + "message_delta", + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 7}}, + ), + _sse("message_stop", {"type": "message_stop"}), + ) + + +def _bad_request_reply() -> Reply: + return Reply( + status=400, + body=json.dumps( + {"type": "error", "error": {"type": "invalid_request_error", "message": "messages must be objects"}} + ).encode(), + ) + + +def _messages_are_objects(body: Mapping[str, JsonValue]) -> bool: + messages: Final = body.get("messages") + return isinstance(messages, list) and all(isinstance(message, dict) for message in messages) + + +_SPEND_COLUMNS: Final = ( + "SELECT request_id, status, call_type, prompt_tokens, completion_tokens, total_tokens, spend, request_tags, " + 'end_user, api_base, custom_llm_provider, model, cache_hit, ("startTime" <= "endTime") AS times_ordered ' + 'FROM "LiteLLM_SpendLogs" ' +) + + +def _spend_row(request_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows(_SPEND_COLUMNS + "WHERE request_id=%s", (request_id,)), + lambda values: len(values) == 1, + seconds=90, + ) + return rows[0] + + +def _key_spend_row(key: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + _SPEND_COLUMNS + "WHERE api_key=%s AND call_type=%s", + (sha256(key.encode()).hexdigest(), "pass_through_endpoint"), + ), + lambda values: len(values) == 1, + seconds=90, + ) + return rows[0] + + +_MODEL_LIST_PROBE: Final = ("GET", "/v1/models") + + +def _is_model_list_probe(request: Request) -> bool: + return (request.method, request.target) == _MODEL_LIST_PROBE + + +@contextmanager +def _upstream(respond: Callable[[Request], Reply]) -> Generator[Wire]: + with wire_server( + lambda request: ( + Reply(body=b'{"object":"list","data":[]}') if _is_model_list_probe(request) else respond(request) + ) + ) as wire: + yield wire + + +def _provider_calls(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if not _is_model_list_probe(request)) + + +def _tags(row: Mapping[str, JsonValue]) -> list[JsonValue]: + raw: Final = row["request_tags"] + tags: Final = json.loads(raw) if isinstance(raw, str) else raw + assert isinstance(tags, list), row + return [tag for tag in tags if not (isinstance(tag, str) and tag.startswith("User-Agent: "))] + + +def _assert_usage_row(row: Mapping[str, JsonValue], call_type: str, tags: list[str]) -> None: + assert row["status"] == "success", row + assert row["call_type"] == call_type, row + assert row["prompt_tokens"] == 11, row + assert row["completion_tokens"] == 7, row + assert row["total_tokens"] == 18, row + spend: Final = row["spend"] + assert isinstance(spend, (int, float)) and spend > 0, row + assert _tags(row) == tags, row + assert row["custom_llm_provider"] == "anthropic", row + assert str(row["cache_hit"]).lower() != "true", row + assert row["times_ordered"] is True, row + + +def _stream_text(gateway: Gateway, path: str, body: Mapping[str, JsonValue], key: str | None = None) -> str: + with gateway.client.stream( + "POST", + f"{gateway.client.base_url}{path}", + json=body, + headers={"Authorization": f"Bearer {key or gateway.key}"}, + ) as stream: + assert stream.status_code == 200, stream.read() + return "".join(stream.iter_text()) + + +def _owned_config(tmp_path: Path, text: str) -> Path: + config: Final = tmp_path / "proxy_config.yaml" + config.write_text(text) + return config + + +def test_passthrough_basic_completion_spend_row_v1_messages(gateway: Gateway) -> None: + marker: Final = "pt-basic-" + uuid.uuid4().hex + prompt: Final = f"say hello {marker}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + assert request.headers["x-api-key"] == _KEY + body: Final = json.loads(request.body) + assert body["model"] == _MODEL + assert body["messages"] == [{"role": "user", "content": prompt}] + assert "litellm_metadata" not in body + return Reply(body=json.dumps(_message(f"msg_{marker}")).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 100, + "messages": [{"role": "user", "content": prompt}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"], "user": f"end-user-{marker}"}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["id"] == f"msg_{marker}" + assert len(_provider_calls(wire)) == 1 + row: Final = _spend_row(f"msg_{marker}") + _assert_usage_row(row, "anthropic_messages", [f"{marker}-1", f"{marker}-2"]) + assert row["end_user"] == f"end-user-{marker}", row + + +def test_passthrough_streaming_spend_row_v1_messages(gateway: Gateway) -> None: + marker: Final = "pt-stream-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.target == "/v1/messages", request.target + body: Final = json.loads(request.body) + assert body["model"] == _MODEL + assert body["stream"] is True + return Reply(content_type="text/event-stream", chunks=_stream_chunks(f"msg_{marker}")) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + text: Final = _stream_text( + gateway, + "/v1/messages", + { + "model": model, + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": f"say hello {marker}"}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"], "user": f"end-user-{marker}"}, + }, + ) + assert "hello stream" in text + row: Final = _spend_row(f"msg_{marker}") + _assert_usage_row(row, "anthropic_messages", [f"{marker}-1", f"{marker}-2"]) + assert row["end_user"] == f"end-user-{marker}", row + + +def test_passthrough_wildcard_model_strips_provider_prefix(gateway: Gateway) -> None: + marker: Final = "pt-wildcard-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert body["model"] == "claude-haiku-4-5-20251001" + return Reply(body=json.dumps(_message(f"msg_{marker}", "claude-haiku-4-5-20251001")).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + created: Final = gateway.post( + "/model/new", + { + "model_name": "anthropic/*", + "litellm_params": {"model": "anthropic/*", "api_base": wire.url, "api_key": _KEY}, + }, + ) + model_info: Final = created["model_info"] + assert isinstance(model_info, dict), created + identity: Final = model_info["id"] + assert isinstance(identity, str), created + scenario.cleanups.callback(scenario.delete_model, identity) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": "anthropic/claude-haiku-4-5-20251001", + "max_tokens": 100, + "messages": [{"role": "user", "content": f"hello wildcard {marker}"}], + }, + ) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "hello test" + assert len(_provider_calls(wire)) == 1 + + +def test_passthrough_thinking_block_round_trips_v1_messages(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert body["model"] == "claude-haiku-4-5-20251001" + assert body["thinking"] == {"type": "enabled", "budget_tokens": 16000} + assert body["max_tokens"] == 20000 + return Reply(body=json.dumps(_thinking_message("msg_" + uuid.uuid4().hex)).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="anthropic/claude-haiku-4-5-20251001", api_base=wire.url, api_key=_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 20000, + "thinking": {"type": "enabled", "budget_tokens": 16000}, + "messages": [{"role": "user", "content": "Just pinging with thinking enabled"}], + }, + ) + assert response.status_code == 200, response.text + content: Final = response.json()["content"] + assert content[0]["type"] == "thinking" + assert content[0]["thinking"] == "pondering the joke" + + +def test_passthrough_bad_request_returns_400_v1_messages(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert not _messages_are_objects(body), body + return _bad_request_reply() + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + responses: Final = tuple( + gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 10, "stream": stream, "messages": ["hi"]}, + ) + for stream in (False, True) + ) + assert [response.status_code for response in responses] == [400, 400], [r.text for r in responses] + + +def test_native_anthropic_route_completion_stream_thinking_and_bad_request(gateway: Gateway, tmp_path: Path) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + assert request.headers["x-api-key"] == _KEY + body: Final = json.loads(request.body) + if not _messages_are_objects(body): + return _bad_request_reply() + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_stream_chunks("msg_" + uuid.uuid4().hex)) + if body.get("thinking") is not None: + return Reply(body=json.dumps(_thinking_message("msg_" + uuid.uuid4().hex)).encode()) + return Reply(body=json.dumps(_message("msg_" + uuid.uuid4().hex)).encode()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": _KEY}, + config=_owned_config(tmp_path, _PROXY_CONFIG), + ) as candidate, + ): + completion: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + {"model": _MODEL, "max_tokens": 100, "messages": [{"role": "user", "content": "say hello native"}]}, + ) + assert completion.status_code == 200, completion.text + assert completion.json()["content"][0]["text"] == "hello test" + thinking: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + { + "model": "claude-haiku-4-5-20251001", + "max_tokens": 20000, + "thinking": {"type": "enabled", "budget_tokens": 16000}, + "messages": [{"role": "user", "content": "ping"}], + }, + ) + assert thinking.status_code == 200, thinking.text + assert thinking.json()["content"][0]["type"] == "thinking" + assert thinking.json()["content"][0]["thinking"] == "pondering the joke" + bad: Final = tuple( + candidate.request( + "POST", + "/anthropic/v1/messages", + {"model": _MODEL, "max_tokens": 10, "stream": stream, "messages": ["hi"]}, + ) + for stream in (False, True) + ) + assert [response.status_code for response in bad] == [400, 400], [r.text for r in bad] + text: Final = _stream_text( + candidate, + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": "hello native stream"}], + }, + ) + assert "hello stream" in text + + +def test_native_passthrough_spend_rows_record_usage_tags_and_spend(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "pt-native-" + uuid.uuid4().hex + completion_id: Final = f"msg_{marker}_completion" + stream_id: Final = f"msg_{marker}_stream" + + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert "litellm_metadata" not in body + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_stream_chunks(stream_id)) + return Reply(body=json.dumps(_message(completion_id)).encode()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": _KEY}, + config=_owned_config(tmp_path, _PROXY_CONFIG), + ) as candidate, + candidate.scenario() as scenario, + ): + completion_key: Final = scenario.key() + stream_key: Final = scenario.key() + response: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"]}, + }, + key=completion_key, + ) + assert response.status_code == 200, response.text + assert response.json()["id"] == completion_id + text: Final = _stream_text( + candidate, + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 10, + "stream": True, + "messages": [{"role": "user", "content": "Say 'hello stream test' and nothing else"}], + "litellm_metadata": {"tags": [f"{marker}-s1", f"{marker}-s2"], "user": f"end-user-{marker}"}, + }, + key=stream_key, + ) + assert "hello stream" in text + completion_row: Final = _key_spend_row(completion_key) + stream_row: Final = _key_spend_row(stream_key) + assert completion_row["request_id"] == completion_id, completion_row + assert stream_row["request_id"] == stream_id, stream_row + _assert_usage_row(completion_row, "pass_through_endpoint", [f"{marker}-1", f"{marker}-2"]) + assert completion_row["api_base"] == f"{wire.url}/v1/messages", completion_row + assert "claude" in str(completion_row["model"]), completion_row + _assert_usage_row(stream_row, "pass_through_endpoint", [f"{marker}-s1", f"{marker}-s2"]) + assert stream_row["end_user"] == f"end-user-{marker}", stream_row + + +def _openai_responses_stream() -> tuple[bytes, ...]: + response: Final[dict[str, JsonValue]] = { + "id": "resp_pt1", + "object": "response", + "created_at": 1700000000, + "model": "gpt-4o-mini", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi from openai"}], + } + ], + "usage": {"input_tokens": 12, "output_tokens": 8, "total_tokens": 20}, + } + return ( + _sse( + "response.created", + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + ), + _sse( + "response.output_text.delta", + { + "type": "response.output_text.delta", + "item_id": "msg_pto", + "output_index": 0, + "content_index": 0, + "delta": "hi from openai", + }, + ), + _sse("response.completed", {"type": "response.completed", "response": response}), + ) + + +def _openai_chat_stream() -> tuple[bytes, ...]: + chunk: Final[dict[str, JsonValue]] = { + "id": "chatcmpl-pt1", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4o", + } + return ( + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': 'Hi'}}]})}\n\n".encode(), + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}]})}\n\n".encode(), + f"data: {json.dumps({**chunk, 'choices': [], 'usage': {'prompt_tokens': 12, 'completion_tokens': 8, 'total_tokens': 20}})}\n\n".encode(), + b"data: [DONE]\n\n", + ) + + +def _delta_usages(text: str) -> list[Mapping[str, JsonValue]]: + events: Final = [json.loads(line[len("data: ") :]) for line in text.splitlines() if line.startswith("data: ")] + return [event["usage"] for event in events if event.get("type") == "message_delta" and "usage" in event] + + +def _cost_config(wire_url: str) -> str: + return ( + "model_list:\n" + " - model_name: amsg\n" + " litellm_params:\n" + f" model: anthropic/{_MODEL}\n" + f" api_base: {wire_url}\n" + f" api_key: {_KEY}\n" + " - model_name: omsg\n" + " litellm_params:\n" + " model: openai/gpt-4o-mini\n" + f" api_base: {wire_url}\n" + " api_key: synthetic-openai-key\n" + "litellm_settings:\n" + " include_cost_in_streaming_usage: true\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " disable_spend_logs: false\n" + " proxy_batch_write_at: 1\n" + ) + + +def _assert_cost_in_every_delta(gateway: Gateway, model: str) -> None: + text: Final = _stream_text( + gateway, + "/v1/messages", + {"model": model, "max_tokens": 20, "stream": True, "messages": [{"role": "user", "content": "Say 'Hi'"}]}, + ) + usages: Final = _delta_usages(text) + assert usages, (model, text) + costs: Final = [usage.get("cost") for usage in usages] + assert all(isinstance(cost, (int, float)) and cost > 0 for cost in costs), (model, text) + + +def test_streaming_cost_injected_into_usage_for_anthropic_and_openai_responses( + gateway: Gateway, tmp_path: Path +) -> None: + def respond(request: Request) -> Reply: + if request.target.endswith("/responses"): + return Reply(content_type="text/event-stream", chunks=_openai_responses_stream()) + assert request.target == "/v1/messages", request.target + return Reply(content_type="text/event-stream", chunks=_stream_chunks("msg_" + uuid.uuid4().hex)) + + with ( + _upstream(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_owned_config(tmp_path, _cost_config(wire.url))) as candidate, + ): + _assert_cost_in_every_delta(candidate, "amsg") + _assert_cost_in_every_delta(candidate, "omsg") + targets: Final = [request.target for request in _provider_calls(wire)] + assert targets == ["/v1/messages", "/responses"], targets + + +def test_streaming_cost_injected_into_usage_for_openai_chat_completions_bridge( + gateway: Gateway, tmp_path: Path +) -> None: + def respond(request: Request) -> Reply: + assert request.target.endswith("/chat/completions"), request.target + assert json.loads(request.body)["stream"] is True + return Reply(content_type="text/event-stream", chunks=_openai_chat_stream()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES": "true"}, + config=_owned_config(tmp_path, _cost_config(wire.url)), + ) as candidate, + ): + _assert_cost_in_every_delta(candidate, "omsg") + assert len(_provider_calls(wire)) == 1 diff --git a/tests/integration/providers/test_ocr_router_wire.py b/tests/integration/providers/test_ocr_router_wire.py new file mode 100644 index 00000000000..cc6455aa5f1 --- /dev/null +++ b/tests/integration/providers/test_ocr_router_wire.py @@ -0,0 +1,69 @@ +import json +from typing import Final + +import pytest + +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server + +_DOCUMENT_URL: Final = "https://example.com/doc.pdf" +_OCR_COST_PER_PAGE: Final = 0.0125 +_MISTRAL_OCR_BODY: Final = json.dumps( + { + "model": "mistral-ocr-latest", + "pages": [{"index": 0, "markdown": "Test PDF File"}], + "usage_info": {"pages_processed": 1, "doc_size_bytes": 1024}, + } +).encode() + + +def _mistral_ocr_peer(request: Request) -> Reply: + return Reply(body=_MISTRAL_OCR_BODY) + + +def test_router_aocr_routes_to_mistral_and_logs_spend(gateway: Gateway) -> None: + with wire_server(_mistral_ocr_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="mistral/mistral-ocr-latest", + api_base=wire.url, + api_key="fake-mistral-key", + ocr_cost_per_page=_OCR_COST_PER_PAGE, + ) + response: Final = gateway.request( + "POST", "/v1/ocr", {"model": model, "document": {"type": "document_url", "document_url": _DOCUMENT_URL}} + ) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + assert len(upstream) == 1, upstream + assert (upstream[0].method, upstream[0].target) == ("POST", "/v1/ocr"), upstream[0] + sent: Final = json.loads(upstream[0].body) + assert sent["model"] == "mistral-ocr-latest", sent + assert sent["document"]["type"] == "document_url", sent + assert sent["document"]["document_url"] == _DOCUMENT_URL, sent + payload: Final = response.json() + assert payload["object"] == "ocr", payload + assert payload["model"] == model, payload + assert [page["index"] for page in payload["pages"]] == [0], payload + assert payload["pages"][0]["markdown"] == "Test PDF File", payload + assert payload["usage_info"]["pages_processed"] == len(payload["pages"]), payload + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(_OCR_COST_PER_PAGE), response.headers + + request_id: Final = string_value(response.headers["x-litellm-call-id"]) + rows: Final = eventually( + lambda: read_rows( + "SELECT status, call_type, custom_llm_provider, model, model_group, spend, prompt_tokens, " + 'completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + row: Final = rows[0] + assert row["status"] == "success", row + assert row["call_type"] == "aocr", row + assert row["custom_llm_provider"] == "mistral", row + assert row["model"] == "mistral/mistral-ocr-latest", row + assert row["model_group"] == model, row + assert float(row["spend"]) == pytest.approx(_OCR_COST_PER_PAGE), row + assert (row["prompt_tokens"], row["completion_tokens"], row["total_tokens"]) == (0, 0, 0), row diff --git a/tests/integration/providers/test_openai_passthrough_files_wire.py b/tests/integration/providers/test_openai_passthrough_files_wire.py new file mode 100644 index 00000000000..bf62db8aec6 --- /dev/null +++ b/tests/integration/providers/test_openai_passthrough_files_wire.py @@ -0,0 +1,49 @@ +import uuid +from pathlib import Path +from typing import Final + +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +_UPSTREAM_KEY: Final = "synthetic-openai-key" + + +def test_openai_passthrough_file_upload_and_delete(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "openai-file-" + uuid.uuid4().hex + file_id: Final = f"file-{marker}" + + def respond(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {_UPSTREAM_KEY}", request.headers + if request.method == "POST" and request.target == "/files": + assert request.headers["content-type"].startswith("multipart/form-data"), request.headers + assert b'name="purpose"\r\n\r\nassistants\r\n' in request.body, request.body[:400] + assert b'filename="notes.txt"' in request.body and marker.encode() in request.body, request.body[:400] + return Reply( + body=( + b'{"id": "' + file_id.encode() + b'", "object": "file", "bytes": 12, ' + b'"created_at": 1700000000, "purpose": "assistants", "filename": "notes.txt"}' + ), + ) + if request.method == "DELETE" and request.target == f"/files/{file_id}": + return Reply(body=b'{"id": "' + file_id.encode() + b'", "object": "file", "deleted": true}') + return Reply(status=404) + + with wire_server(respond) as wire: + with owned_proxy( + gateway, + tmp_path, + {"OPENAI_API_BASE": wire.url, "OPENAI_API_KEY": _UPSTREAM_KEY}, + ) as candidate: + upload: Final = candidate.request_multipart( + "/openai/files", + {"purpose": "assistants"}, + {"file": ("notes.txt", f"contents {marker}".encode(), "text/plain")}, + ) + assert upload.status_code == 200, upload.text + assert upload.json()["id"] == file_id + delete: Final = candidate.request("DELETE", f"/openai/files/{file_id}") + assert delete.status_code == 200, delete.text + assert delete.json()["deleted"] is True + forwarded: Final = tuple((request.method, request.target) for request in wire.drain()) + assert forwarded == (("POST", "/files"), ("DELETE", f"/files/{file_id}")), forwarded diff --git a/tests/integration/providers/test_responses_error_status_wire.py b/tests/integration/providers/test_responses_error_status_wire.py new file mode 100644 index 00000000000..01372c8ae7a --- /dev/null +++ b/tests/integration/providers/test_responses_error_status_wire.py @@ -0,0 +1,71 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + + +def test_unknown_model_provider_404_surfaces_to_client_as_404(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + body: Final = json.loads(request.body) + assert body["model"] == "non-existent-model" + return Reply( + status=404, + body=json.dumps( + {"error": {"message": "model not found", "type": "invalid_request_error", "code": "404"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/non-existent-model", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": "say hi"}) + assert response.status_code == 404, response.text + assert len(wire.drain()) == 1 + + +def test_provider_400_for_bad_temperature_surfaces_to_client_as_400(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + body: Final = json.loads(request.body) + assert body["model"] == "gpt-4o" + assert body["temperature"] == 2000 + return Reply( + status=400, + body=json.dumps( + {"error": {"message": "temperature out of range", "type": "invalid_request_error", "code": "400"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-4o", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "temperature": 2000} + ) + assert response.status_code == 400, response.text + assert len(wire.drain()) == 1 + + +def test_cancel_invalid_response_id_surfaces_error_status(gateway: Gateway) -> None: + response_id: Final = "invalid_response_id_12345" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == f"/responses/{response_id}/cancel", request.target + return Reply( + status=404, + body=json.dumps( + {"error": {"message": "No such response", "type": "invalid_request_error", "code": "404"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-4o", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request("POST", f"/v1/responses/{response_id}/cancel", {"model": model}) + assert response.status_code == 404, response.text + assert len(wire.drain()) == 1 diff --git a/tests/ocr_tests/test_ocr_matrix.py b/tests/ocr_tests/test_ocr_matrix.py index 13cffbbc9a1..3f8c9b402fc 100644 --- a/tests/ocr_tests/test_ocr_matrix.py +++ b/tests/ocr_tests/test_ocr_matrix.py @@ -24,7 +24,6 @@ from typing import Final, Literal import pytest import litellm -from litellm import Router from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse @@ -301,17 +300,3 @@ async def test_ocr(case: Case, monkeypatch: pytest.MonkeyPatch, logger: Recordin _assert_logged(await logger.wait_for_call(), response, case.provider.model, response.model, case.call) -async def test_router_aocr(monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None: - case: Final = Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "async") - router: Final = Router( - model_list=[ - { - "model_name": "ocr-alias", - "litellm_params": {"model": MISTRAL.model, **case.bind_credentials(monkeypatch)}, - } - ] - ) - response: Final = await router.aocr(model="ocr-alias", document=PDF_BY_URL.build()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Router.aocr is untyped - assert isinstance(response, OCRResponse) - _assert_ocr_response(response, MISTRAL.model, PDF_TEXT) - _assert_logged(await logger.wait_for_call(), response, MISTRAL.model, MISTRAL.model, case.call) diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py deleted file mode 100644 index 260656600ec..00000000000 --- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py +++ /dev/null @@ -1,115 +0,0 @@ -import os -import time -from collections.abc import Iterator -from typing import Final - -import httpx -import pytest -from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream -from openai.types.responses import ResponseStreamEvent - -BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90 - - -def generate_key(): - """Generate a key for testing""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - data = {} - - response = httpx.post(url, headers=headers, json=data) - if response.status_code != 200: - raise Exception(f"Key generation failed with status: {response.status_code}") - return response.json()["key"] - - -def get_test_client(): - """Create OpenAI client with generated key""" - key = generate_key() - return OpenAI(api_key=key, base_url="http://0.0.0.0:4000") - - -def validate_response(response): - """ - Validate basic response structure from OpenAI responses API - """ - assert response is not None - assert hasattr(response, "choices") - assert len(response.choices) > 0 - assert hasattr(response.choices[0], "message") - assert hasattr(response.choices[0].message, "content") - assert isinstance(response.choices[0].message.content, str) - assert hasattr(response, "id") - assert isinstance(response.id, str) - assert hasattr(response, "model") - assert isinstance(response.model, str) - assert hasattr(response, "created") - assert isinstance(response.created, int) - assert hasattr(response, "usage") - assert hasattr(response.usage, "prompt_tokens") - assert hasattr(response.usage, "completion_tokens") - assert hasattr(response.usage, "total_tokens") - - -def validate_stream_chunk(chunk): - """ - Validate streaming chunk structure from OpenAI responses API - """ - assert chunk is not None - assert hasattr(chunk, "choices") - assert len(chunk.choices) > 0 - assert hasattr(chunk.choices[0], "delta") - - # Some chunks might not have content in the delta - if ( - hasattr(chunk.choices[0].delta, "content") - and chunk.choices[0].delta.content is not None - ): - assert isinstance(chunk.choices[0].delta.content, str) - - assert hasattr(chunk, "id") - assert isinstance(chunk.id, str) - assert hasattr(chunk, "model") - assert isinstance(chunk.model, str) - assert hasattr(chunk, "created") - assert isinstance(chunk.created, int) - - -def test_model_not_found_error(): - client = get_test_client() - with pytest.raises(NotFoundError): - client.responses.create(model="non-existent-model", input="This should fail") - - -def test_bad_request_bad_param_error(): - client = get_test_client() - with pytest.raises(BadRequestError): - # Out-of-range temperature on a non-reasoning model, so drop_params forwards it - client.responses.create( - model="gpt-4.1", input="This should fail", temperature=2000 - ) - - -def admitted_response_id(chunk: ResponseStreamEvent) -> str | None: - response: Final = getattr(chunk, "response", None) - return None if response is None else response.id - - -def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]: - for chunk in stream: - print("stream chunk=", chunk) - yield chunk - if admitted_response_id(chunk) is not None: - return - if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: - return - - -def test_cancel_invalid_response_id(): - client = get_test_client() - with pytest.raises(APIStatusError): - # Try to cancel a non-existent response ID - client.responses.cancel("invalid_response_id_12345") diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py deleted file mode 100644 index f488095aa12..00000000000 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ /dev/null @@ -1,463 +0,0 @@ -# What this tests ? -## Tests /batches endpoints -import os -import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI -from typing import Optional, List, Union -from test_openai_files_endpoints import upload_file, delete_file -import sys -import time - - -BASE_URL = "http://localhost:4000" # Replace with your actual base URL -API_KEY = os.environ["LITELLM_MASTER_KEY"] # Replace with your actual API key - - -client = OpenAI(base_url=BASE_URL, api_key=API_KEY) - - -def create_batch_oai_sdk(filepath: str, custom_llm_provider: str) -> str: - batch_input_file = client.files.create( - file=open(filepath, "rb"), - purpose="batch", - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - batch_input_file_id = batch_input_file.id - - print("waiting for file to be processed......") - time.sleep(5) - rq = client.batches.create( - input_file_id=batch_input_file_id, - endpoint="/v1/chat/completions", - completion_window="24h", - metadata={ - "description": filepath, - }, - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - - print(f"Batch submitted. ID: {rq.id}") - return rq.id - - -def await_batch_completion(batch_id: str, custom_llm_provider: str): - max_tries = 3 - tries = 0 - - while tries < max_tries: - batch = client.batches.retrieve( - batch_id, extra_headers={"custom-llm-provider": custom_llm_provider} - ) - if batch.status == "completed": - print(f"Batch {batch_id} completed.") - return batch.id - - tries += 1 - print(f"waiting for batch to complete... (attempt {tries}/{max_tries})") - time.sleep(10) - - print( - f"Reached maximum number of attempts ({max_tries}). Batch may still be processing." - ) - - -def write_content_to_file( - batch_id: str, output_path: str, custom_llm_provider: str -) -> str: - batch = client.batches.retrieve( - batch_id=batch_id, extra_headers={"custom-llm-provider": custom_llm_provider} - ) - content = client.files.content( - file_id=batch.output_file_id, - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - print("content from files.content", content.content) - content.write_to_file(output_path) - - -def read_jsonl(filepath: str): - import json - - results = [] - with open(filepath, "r") as f: - for line in f: - if line.strip(): - results.append(json.loads(line)) - - for item in results: - print(item) - custom_id = item["custom_id"] - print(custom_id) - - -def get_any_completed_batch_id_azure(): - print("AZURE getting any completed batch id") - list_of_batches = client.batches.list( - extra_headers={"custom-llm-provider": "azure"} - ) - print("list of batches", list_of_batches) - for batch in list_of_batches: - if batch.status == "completed": - return batch.id - return None - - -@pytest.mark.skip(reason="Local only test to verify if things work well") -def test_vertex_batches_endpoint(): - """ - Test VertexAI Batches Endpoint - """ - import os - - oai_client = OpenAI(api_key=API_KEY, base_url=BASE_URL) - file_name = "local_testing/vertex_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - file_obj = oai_client.files.create( - file=open(file_path, "rb"), - purpose="batch", - extra_headers={"custom-llm-provider": "vertex_ai"}, - ) - print("Response from creating file=", file_obj) - - batch_input_file_id = file_obj.id - assert ( - batch_input_file_id is not None - ), f"Failed to create file, expected a non null file_id but got {batch_input_file_id}" - - create_batch_response = oai_client.batches.create( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=batch_input_file_id, - extra_headers={"custom-llm-provider": "vertex_ai"}, - metadata={"key1": "value1", "key2": "value2"}, - ) - print("response from create batch", create_batch_response) - pass - - -@pytest.mark.asyncio -async def test_batch_status_sync_from_provider_to_database(): - """ - Test that when batch status changes at the provider, - it gets synced to the ManagedObjectTable database. - - This tests the new refactored utility functions: - - get_batch_from_database() - - update_batch_in_database() - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - get_batch_from_database, - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - import json - - # Setup: Create mock objects - batch_id = "batch_test123" - unified_batch_id = "litellm_proxy:test_unified_batch" - - # Mock database batch object with "validating" status - mock_db_batch = MagicMock() - mock_db_batch.unified_object_id = batch_id - mock_db_batch.status = "validating" - mock_db_batch.file_object = json.dumps( - { - "id": batch_id, - "object": "batch", - "status": "validating", - "endpoint": "/v1/chat/completions", - "input_file_id": "file-test123", - "completion_window": "24h", - "created_at": 1234567890, - } - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=mock_db_batch - ) - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.debug = MagicMock() - mock_logger.info = MagicMock() - mock_logger.warning = MagicMock() - mock_logger.error = MagicMock() - - # Test 1: Retrieve batch from database (initial state) - db_batch_object, response_batch = await get_batch_from_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - ) - - # Verify database was queried - mock_prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with( - where={"unified_object_id": batch_id} - ) - - # Verify batch was retrieved correctly - assert db_batch_object is not None - assert response_batch is not None - assert response_batch.id == batch_id - assert response_batch.status == "validating" - - # Test 2: Simulate provider returning updated status - updated_batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="completed", # Status changed from "validating" to "completed" - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - output_file_id="file-output123", - ) - - # Test 3: Update database with new status from provider - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=updated_batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - db_batch_object=db_batch_object, - operation="retrieve", - ) - - # Verify database was updated - mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() - update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args - - # Verify the update call had correct parameters - assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id - assert ( - update_call_args.kwargs["data"]["status"] == "complete" - ) # "completed" normalized to "complete" - assert "file_object" in update_call_args.kwargs["data"] - assert "updated_at" in update_call_args.kwargs["data"] - # batch_processed must be set to True when batch transitions to complete - assert update_call_args.kwargs["data"]["batch_processed"] is True - - # Verify logger was called with status change message - mock_logger.info.assert_called() - log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:] - assert "validating" in log_message - assert "completed" in log_message - - print("✅ Test passed: Batch status synced from provider to database") - - -@pytest.mark.asyncio -async def test_batch_cancel_updates_database(): - """ - Test that canceling a batch updates the database status. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - - # Setup - batch_id = "batch_cancel_test" - unified_batch_id = "litellm_proxy:cancel_test" - - # Mock cancelled batch response from provider - cancelled_batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="cancelled", - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - cancelled_at=1234567999, - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=None - ) - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.info = MagicMock() - mock_logger.error = MagicMock() - - # Call update_batch_in_database for cancel operation - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=cancelled_batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - operation="cancel", - ) - - # Verify database was updated - mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() - update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args - - # Verify the update call had correct parameters - assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id - assert update_call_args.kwargs["data"]["status"] == "cancelled" - assert "file_object" in update_call_args.kwargs["data"] - - # Verify logger was called - mock_logger.info.assert_called() - log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:] - assert "cancel" in log_message.lower() - assert "cancelled" in log_message - - print("✅ Test passed: Batch cancel updates database") - - -@pytest.mark.asyncio -async def test_batch_terminal_state_skip_provider_call(): - """ - Test that when a batch is in a terminal state (completed, failed, cancelled, expired), - it returns immediately from database without calling the provider. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - get_batch_from_database, - ) - from litellm.types.utils import LiteLLMBatch - import json - - # Setup: Create mock objects for a completed batch - batch_id = "batch_completed_test" - unified_batch_id = "litellm_proxy:completed_test" - - # Mock database batch object with "completed" status - mock_db_batch = MagicMock() - mock_db_batch.unified_object_id = batch_id - mock_db_batch.status = "complete" - mock_db_batch.file_object = json.dumps( - { - "id": batch_id, - "object": "batch", - "status": "completed", - "endpoint": "/v1/chat/completions", - "input_file_id": "file-test123", - "output_file_id": "file-output123", - "completion_window": "24h", - "created_at": 1234567890, - "completed_at": 1234567999, - } - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=mock_db_batch - ) - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.debug = MagicMock() - - # Retrieve batch from database - db_batch_object, response_batch = await get_batch_from_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - ) - - # Verify batch was retrieved - assert db_batch_object is not None - assert response_batch is not None - assert response_batch.status == "completed" - - # In the actual endpoint, when status is in terminal states, - # it should return immediately without calling the provider - # This test verifies the database retrieval works correctly - assert response_batch.status in ["completed", "failed", "cancelled", "expired"] - - print("✅ Test passed: Terminal state batch retrieved from database") - - -@pytest.mark.asyncio -async def test_batch_no_status_change_skip_update(): - """ - Test that when batch status hasn't changed, database update is skipped. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - - # Setup - batch_id = "batch_no_change_test" - unified_batch_id = "litellm_proxy:no_change_test" - - # Mock database batch object with "validating" status - mock_db_batch = MagicMock() - mock_db_batch.status = "validating" - - # Mock batch response from provider with same status - batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="validating", # Same status as in database - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.info = MagicMock() - - # Call update_batch_in_database - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - db_batch_object=mock_db_batch, - operation="retrieve", - ) - - # Verify database update was NOT called (status hasn't changed) - mock_prisma_client.db.litellm_managedobjecttable.update.assert_not_called() - - # Verify logger info was NOT called (no status change to log) - mock_logger.info.assert_not_called() - - print("✅ Test passed: Database update skipped when status unchanged") diff --git a/tests/openai_endpoints_tests/test_openai_files_endpoints.py b/tests/openai_endpoints_tests/test_openai_files_endpoints.py deleted file mode 100644 index 9398a0d1c53..00000000000 --- a/tests/openai_endpoints_tests/test_openai_files_endpoints.py +++ /dev/null @@ -1,113 +0,0 @@ -import os -# What this tests ? -## Tests /chat/completions by generating a key and then making a chat completions request -import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI -from typing import Optional, List, Union - - -BASE_URL = "http://localhost:4000" # Replace with your actual base URL -API_KEY = os.environ["LITELLM_MASTER_KEY"] # Replace with your actual API key - - -@pytest.mark.asyncio -async def test_file_operations(): - openai_client = AsyncOpenAI(api_key=API_KEY, base_url=BASE_URL) - file_content = b'{"prompt": "Hello", "completion": "Hi"}' - uploaded_file = await openai_client.files.create( - purpose="fine-tune", - file=file_content, - ) - list_files = await openai_client.files.list() - print("list_files=", list_files) - - get_file = await openai_client.files.retrieve(file_id=uploaded_file.id) - print("get_file=", get_file) - - get_file_content = await openai_client.files.content(file_id=uploaded_file.id) - print("get_file_content=", get_file_content.content) - response = get_file_content.response - - assert get_file_content.content == file_content - assert response.status_code == 200 - assert response.headers.get("content-type") == "application/octet-stream" - assert response.headers.get("content-length") is not None - assert int(response.headers["content-length"]) == len(get_file_content.content) - assert response.headers.get("content-disposition") is not None - assert uploaded_file.filename in response.headers["content-disposition"] - assert response.headers.get("x-request-id") is not None - # try get_file_content.write_to_file - get_file_content.write_to_file("get_file_content.jsonl") - - delete_file = await openai_client.files.delete(file_id=uploaded_file.id) - print("delete_file=", delete_file) - - -async def upload_file(session, purpose="fine-tune"): - url = f"{BASE_URL}/v1/files" - headers = {"Authorization": f"Bearer {API_KEY}"} - data = aiohttp.FormData() - data.add_field("purpose", purpose) - data.add_field( - "file", b'{"prompt": "Hello", "completion": "Hi"}', filename="mydata.jsonl" - ) - - async with session.post(url, headers=headers, data=data) as response: - assert response.status == 200 - result = await response.json() - assert "id" in result - print(f"File upload successful. File ID: {result['id']}") - return result["id"] - - -async def list_files(session): - url = f"{BASE_URL}/v1/files" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert "data" in result - print("List files successful") - - -async def get_file(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert result["id"] == file_id - assert result["object"] == "file" - assert "bytes" in result - assert "created_at" in result - assert "filename" in result - assert result["purpose"] == "fine-tune" - print(f"Get file successful for file ID: {file_id}") - - -async def get_file_content(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}/content" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - content = await response.text() - print("content from /files/{file_id}/content=", content) - assert content # Check if content is not empty - print(f"Get file content successful for file ID: {file_id}") - - -async def delete_file(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.delete(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert "deleted" in result - assert result["id"] == file_id - print(f"Delete file successful for file ID: {file_id}") diff --git a/tests/pass_through_tests/base_anthropic_messages_test.py b/tests/pass_through_tests/base_anthropic_messages_test.py index 95f709e3880..710dfa453cf 100644 --- a/tests/pass_through_tests/base_anthropic_messages_test.py +++ b/tests/pass_through_tests/base_anthropic_messages_test.py @@ -13,61 +13,8 @@ class BaseAnthropicMessagesTest(ABC): def get_client(self): return anthropic.Anthropic() - def test_anthropic_basic_completion(self): - print("making basic completion request to anthropic passthrough") - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=1024, - messages=[{"role": "user", "content": "Say 'hello test' and nothing else"}], - extra_body={ - "litellm_metadata": { - "tags": ["test-tag-1", "test-tag-2"], - } - }, - ) - print(response) - def test_anthropic_streaming(self): - print("making streaming request to anthropic passthrough") - collected_output = [] - client = self.get_client() - with client.messages.stream( - max_tokens=10, - messages=[ - {"role": "user", "content": "Say 'hello stream test' and nothing else"} - ], - model="claude-sonnet-4-5-20250929", - extra_body={ - "litellm_metadata": { - "tags": ["test-tag-stream-1", "test-tag-stream-2"], - } - }, - ) as stream: - for text in stream.text_stream: - collected_output.append(text) - full_response = "".join(collected_output) - print(full_response) - - def test_anthropic_messages_with_thinking(self): - print("making request to anthropic passthrough with thinking") - client = self.get_client() - response = client.messages.create( - model="claude-haiku-4-5-20251001", - max_tokens=20000, - thinking={"type": "enabled", "budget_tokens": 16000}, - messages=[ - {"role": "user", "content": "Just pinging with thinking enabled"} - ], - ) - - print(response) - - # Verify the first content block is a thinking block - response_thinking = response.content[0].thinking - assert response_thinking is not None - assert len(response_thinking) > 0 def test_anthropic_streaming_with_thinking(self): print("making streaming request to anthropic passthrough with thinking enabled") @@ -105,41 +52,4 @@ class BaseAnthropicMessagesTest(ABC): assert len(collected_response) > 0 assert len(full_response) > 0 - def test_bad_request_error_handling_streaming(self): - print("making request to anthropic passthrough with bad request") - try: - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=10, - stream=True, - messages=["hi"], - ) - print(response) - assert pytest.fail("Expected BadRequestError") - except anthropic.BadRequestError as e: - print("Got BadRequestError from anthropic, e=", e) - print(e.__cause__) - print(e.status_code) - print(e.response) - except Exception as e: - pytest.fail(f"Got unexpected exception: {e}") - def test_bad_request_error_handling_non_streaming(self): - print("making request to anthropic passthrough with bad request") - try: - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=10, - messages=["hi"], - ) - print(response) - assert pytest.fail("Expected BadRequestError") - except anthropic.BadRequestError as e: - print("Got BadRequestError from anthropic, e=", e) - print(e.__cause__) - print(e.status_code) - print(e.response) - except Exception as e: - pytest.fail(f"Got unexpected exception: {e}") diff --git a/tests/pass_through_tests/test_anthropic_passthrough.py b/tests/pass_through_tests/test_anthropic_passthrough.py deleted file mode 100644 index 5d6ddb1fbd0..00000000000 --- a/tests/pass_through_tests/test_anthropic_passthrough.py +++ /dev/null @@ -1,472 +0,0 @@ -""" -This test ensures that the proxy can passthrough anthropic requests -""" - -import os -import pytest -import anthropic -import aiohttp -import asyncio -import json - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_basic_completion_with_headers(): - print("making basic completion request to anthropic passthrough with aiohttp") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "Anthropic-Version": "2023-06-01", - } - - payload = { - "model": "claude-sonnet-4-5-20250929", - "max_tokens": 10, - "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], - "litellm_metadata": { - "tags": ["test-tag-1", "test-tag-2"], - }, - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/anthropic/v1/messages", json=payload, headers=headers - ) as response: - response_text = await response.text() - print(f"Response text: {response_text}") - - response_json = await response.json() - response_headers = response.headers - print( - "non-streaming response", - json.dumps(response_json, indent=4, default=str), - ) - reported_usage = response_json.get("usage", None) - # fix null checks for reported_usage - anthropic_api_input_tokens = ( - reported_usage.get("input_tokens", None) if reported_usage else None - ) - anthropic_api_output_tokens = ( - reported_usage.get("output_tokens", None) if reported_usage else None - ) - anthropic_message_id = response_json.get("id") - - print(f"Anthropic message ID: {anthropic_message_id}") - - # Wait for spend to be logged - await asyncio.sleep(15) - - # Check spend logs for this specific request with retry logic - spend_data = None - max_retries = 2 - for attempt in range(max_retries): - print(f"Attempt {attempt + 1}/{max_retries} to check spend logs") - - async with session.get( - f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - ) as spend_response: - print("text spend response") - print(f"Spend response: {spend_response}") - spend_data = await spend_response.json() - print(f"Spend data: {spend_data}") - - # Check if spend data exists and has entries - if spend_data and len(spend_data) > 0: - print("Spend logs found!") - break - else: - print("Spend logs not found yet...") - if ( - attempt < max_retries - 1 - ): # Don't wait after the last attempt - print("Waiting 10 seconds before retry...") - await asyncio.sleep(10) - - if not isinstance(spend_data, list): - print(f"Spend endpoint answered with an error response: {spend_data}") - print("Skipping spend assertions (spend logs unreachable in CI)") - return - - assert spend_data, ( - f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id " - "the caller received" - ) - - log_entry = spend_data[0] - - # Basic existence checks - assert isinstance(log_entry, dict), "Log entry should be a dictionary" - - # Request metadata assertions - assert ( - log_entry["request_id"] == anthropic_message_id - ), "Request ID should be the message id the caller received" - assert ( - log_entry["call_type"] == "pass_through_endpoint" - ), "Call type should be pass_through_endpoint" - assert ( - log_entry["api_base"] == "https://api.anthropic.com/v1/messages" - ), "API base should be Anthropic's endpoint" - - # Token and spend assertions - assert log_entry["spend"] > 0, "Spend value should not be None" - assert isinstance( - log_entry["spend"], (int, float) - ), "Spend should be a number" - assert log_entry["total_tokens"] > 0, "Should have some tokens" - assert ( - log_entry["prompt_tokens"] == anthropic_api_input_tokens - ), f"Should have prompt tokens matching anthropic api. Expected {anthropic_api_input_tokens} but got {log_entry['prompt_tokens']}" - assert ( - log_entry["completion_tokens"] == anthropic_api_output_tokens - ), f"Should have completion tokens matching anthropic api. Expected {anthropic_api_output_tokens} but got {log_entry['completion_tokens']}" - assert ( - log_entry["total_tokens"] - == log_entry["prompt_tokens"] + log_entry["completion_tokens"] - ), "Total tokens should equal prompt + completion" - - # Time assertions - assert all( - key in log_entry - for key in ["startTime", "endTime", "completionStartTime"] - ), "Should have all time fields" - assert ( - log_entry["startTime"] < log_entry["endTime"] - ), "Start time should be before end time" - - # Metadata assertions - assert str(log_entry["cache_hit"]).lower() != "true", "Cache should be off" - assert log_entry["request_tags"] == [ - "test-tag-1", - "test-tag-2", - ], "Tags should match input" - assert ( - "user_api_key" in log_entry["metadata"] - ), "Should have user API key in metadata" - - assert "claude" in log_entry["model"] - assert log_entry["custom_llm_provider"] == "anthropic" - - -@pytest.mark.asyncio -async def test_anthropic_streaming_with_headers(): - print("making streaming request to anthropic passthrough with aiohttp") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "Anthropic-Version": "2023-06-01", - } - - payload = { - "model": "claude-sonnet-4-5-20250929", - "max_tokens": 10, - "messages": [ - {"role": "user", "content": "Say 'hello stream test' and nothing else"} - ], - "stream": True, - "litellm_metadata": { - "tags": ["test-tag-stream-1", "test-tag-stream-2"], - "user": "test-user-1", - }, - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/anthropic/v1/messages", json=payload, headers=headers - ) as response: - print("response status") - print(response.status) - assert response.status == 200, "Response should be successful" - response_headers = response.headers - print(f"Response headers: {response_headers}") - - collected_output = [] - async for line in response.content: - if line: - text = line.decode("utf-8").strip() - if text.startswith("data: "): - collected_output.append(text[6:]) # Remove 'data: ' prefix - - print("Collected output:", "".join(collected_output)) - anthropic_api_usage_chunks = [] - anthropic_message_id = None - for chunk in collected_output: - chunk_json = json.loads(chunk) - if chunk_json.get("type") == "message_start": - anthropic_message_id = chunk_json.get("message", {}).get("id") - if "usage" in chunk_json: - anthropic_api_usage_chunks.append(chunk_json["usage"]) - elif "message" in chunk_json and "usage" in chunk_json["message"]: - anthropic_api_usage_chunks.append(chunk_json["message"]["usage"]) - - print(f"Anthropic message ID: {anthropic_message_id}") - - print( - "anthropic_api_usage_chunks", - json.dumps(anthropic_api_usage_chunks, indent=4, default=str), - ) - - print("anthropic_api_usage_chunks: ", anthropic_api_usage_chunks) - # Get the most recent value of input tokens (iterate backwards to find last non-zero value) - anthropic_api_input_tokens = 0 - for usage in reversed(anthropic_api_usage_chunks): - if usage.get("input_tokens", 0) > 0: - anthropic_api_input_tokens = usage.get("input_tokens", 0) - break - anthropic_api_output_tokens = 0 - for usage in reversed(anthropic_api_usage_chunks): - if usage.get("output_tokens", 0) > 0: - anthropic_api_output_tokens = usage.get("output_tokens", 0) - break - - print("anthropic_api_input_tokens", anthropic_api_input_tokens) - print("anthropic_api_output_tokens", anthropic_api_output_tokens) - - # Wait for spend to be logged - await asyncio.sleep(20) - - # Check spend logs for this specific request with retry logic - spend_data = None - max_retries = 2 - for attempt in range(max_retries): - print(f"Attempt {attempt + 1}/{max_retries} to check spend logs") - - async with session.get( - f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - ) as spend_response: - spend_data = await spend_response.json() - print(f"Spend data: {spend_data}") - - # Check if spend data exists and has entries - if spend_data and len(spend_data) > 0: - print("Spend logs found!") - break - else: - print("Spend logs not found yet...") - if ( - attempt < max_retries - 1 - ): # Don't wait after the last attempt - print("Waiting 10 seconds before retry...") - await asyncio.sleep(10) - - if not isinstance(spend_data, list): - print(f"Spend endpoint answered with an error response: {spend_data}") - print("Skipping spend assertions (spend logs unreachable in CI)") - return - - assert spend_data, ( - f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id " - "the caller received" - ) - - log_entry = spend_data[0] - - # Basic existence checks - assert isinstance(log_entry, dict), "Log entry should be a dictionary" - - # Request metadata assertions - assert ( - log_entry["request_id"] == anthropic_message_id - ), "Request ID should be the message id the caller received" - assert ( - log_entry["call_type"] == "pass_through_endpoint" - ), "Call type should be pass_through_endpoint" - # assert ( - # log_entry["api_base"] == "https://api.anthropic.com/v1/messages" - # ), "API base should be Anthropic's endpoint" - - # Token and spend assertions - assert log_entry["spend"] > 0, "Spend value should not be None" - assert isinstance( - log_entry["spend"], (int, float) - ), "Spend should be a number" - assert log_entry["total_tokens"] > 0, "Should have some tokens" - assert ( - log_entry["prompt_tokens"] == anthropic_api_input_tokens - ), f"Should have prompt tokens matching anthropic api. Expected {anthropic_api_input_tokens} but got {log_entry['prompt_tokens']}" - assert ( - log_entry["completion_tokens"] == anthropic_api_output_tokens - ), f"Should have completion tokens matching anthropic api. Expected {anthropic_api_output_tokens} but got {log_entry['completion_tokens']}" - assert ( - log_entry["total_tokens"] - == log_entry["prompt_tokens"] + log_entry["completion_tokens"] - ), "Total tokens should equal prompt + completion" - - # Time assertions - assert all( - key in log_entry - for key in ["startTime", "endTime", "completionStartTime"] - ), "Should have all time fields" - assert ( - log_entry["startTime"] < log_entry["endTime"] - ), "Start time should be before end time" - - # Metadata assertions - assert str(log_entry["cache_hit"]).lower() != "true", "Cache should be off" - assert log_entry["request_tags"] == [ - "test-tag-stream-1", - "test-tag-stream-2", - ], "Tags should match input" - assert ( - "user_api_key" in log_entry["metadata"] - ), "Should have user API key in metadata" - - assert "claude" in log_entry["model"] - - assert log_entry["end_user"] == "test-user-1" - assert log_entry["custom_llm_provider"] == "anthropic" - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_messages_streaming_cost_injection(): - """ - Test that cost is injected into message_delta usage for Anthropic Messages API streaming - """ - print("Testing cost injection in Anthropic Messages API streaming response") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "anthropic-version": "2023-06-01", - } - - payload = { - "model": "claude-haiku-4-5-20251001", - "max_tokens": 10, - "stream": True, - "messages": [{"role": "user", "content": "Say 'Hi'"}], - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/v1/messages", - json=payload, - headers=headers, - ) as response: - assert response.status == 200 - - # Collect all SSE events. - # Split each chunk by newlines to handle both: - # - Anthropic direct path: chunks arrive as individual lines - # - OpenAI/Responses API path: chunks are full multi-line SSE events - events = [] - async for chunk in response.content: - chunk_str = chunk.decode("utf-8") - for line in chunk_str.split("\n"): - line = line.strip() - if line.startswith("data: "): - try: - data = json.loads(line[6:]) # Remove 'data: ' prefix - events.append(data) - except json.JSONDecodeError: - continue - - # Find message_delta event with usage - message_delta_events = [ - event - for event in events - if event.get("type") == "message_delta" and "usage" in event - ] - - assert ( - len(message_delta_events) > 0 - ), "No message_delta events with usage found" - - # Check that cost is included in usage - for event in message_delta_events: - usage = event.get("usage", {}) - assert "cost" in usage, f"Cost not found in usage: {usage}" - assert isinstance( - usage["cost"], (int, float) - ), f"Cost should be numeric: {usage['cost']}" - assert ( - usage["cost"] >= 0 - ), f"Cost should be non-negative: {usage['cost']}" - - print(f"Found message_delta with cost: {usage}") - - print( - f"Test passed: Found {len(message_delta_events)} message_delta events with cost" - ) - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_messages_openai_model_streaming_cost_injection(): - """ - Test that cost is injected into message_delta usage for OpenAI model via Anthropic Messages API - """ - print("Testing cost injection in Anthropic Messages API with OpenAI model") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "anthropic-version": "2023-06-01", - } - - payload = { - "model": "openai/gpt-4o", - "max_tokens": 20, - "stream": True, - "messages": [{"role": "user", "content": "Say 'Hi'"}], - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/v1/messages", - json=payload, - headers=headers, - ) as response: - assert response.status == 200 - - # Collect all SSE events. - # Split each chunk by newlines to handle both: - # - Direct API paths: chunks arrive as individual lines - # - OpenAI/Responses API path: AnthropicResponsesStreamWrapper yields - # full multi-line SSE events as single bytes objects, so a naive - # startswith('data: ') check on the whole chunk misses them. - events = [] - async for chunk in response.content: - chunk_str = chunk.decode("utf-8") - for line in chunk_str.split("\n"): - line = line.strip() - if line.startswith("data: "): - try: - data = json.loads(line[6:]) # Remove 'data: ' prefix - events.append(data) - except json.JSONDecodeError: - continue - - # Find message_delta event with usage - message_delta_events = [ - event - for event in events - if event.get("type") == "message_delta" and "usage" in event - ] - - assert ( - len(message_delta_events) > 0 - ), "No message_delta events with usage found" - - # Check that cost is included in usage - for event in message_delta_events: - usage = event.get("usage", {}) - assert "cost" in usage, f"Cost not found in usage: {usage}" - assert isinstance( - usage["cost"], (int, float) - ), f"Cost should be numeric: {usage['cost']}" - assert ( - usage["cost"] >= 0 - ), f"Cost should be non-negative: {usage['cost']}" - - print(f"Found message_delta with cost: {usage}") - - print( - f"Test passed: Found {len(message_delta_events)} message_delta events with cost" - ) diff --git a/tests/pass_through_tests/test_anthropic_passthrough_basic.py b/tests/pass_through_tests/test_anthropic_passthrough_basic.py index c7e9fea867c..4ef9887bc89 100644 --- a/tests/pass_through_tests/test_anthropic_passthrough_basic.py +++ b/tests/pass_through_tests/test_anthropic_passthrough_basic.py @@ -19,11 +19,3 @@ class TestAnthropicMessagesEndpoint(BaseAnthropicMessagesTest): api_key=os.environ["LITELLM_MASTER_KEY"], ) - def test_anthropic_messages_to_wildcard_model(self): - client = self.get_client() - response = client.messages.create( - model="anthropic/claude-haiku-4-5-20251001", - messages=[{"role": "user", "content": "Hello, world!"}], - max_tokens=100, - ) - print(response) diff --git a/tests/pass_through_tests/test_assembly_ai.py b/tests/pass_through_tests/test_assembly_ai.py deleted file mode 100644 index 09999bc2bed..00000000000 --- a/tests/pass_through_tests/test_assembly_ai.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -This test ensures that the proxy can passthrough requests to assemblyai -""" - -import os -import time - -import pytest -import httpx -import aiohttp -import asyncio - -TEST_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"] -TEST_BASE_URL = "http://0.0.0.0:4000/assemblyai" - - -def _transcribe_and_verify(virtual_key: str, base_url: str): - file_url = "https://assembly.ai/wildfires.mp3" - headers = { - "Authorization": f"Bearer {virtual_key}", - "Content-Type": "application/json", - } - create_payload = { - "audio_url": file_url, - "speech_models": ["universal-2"], - } - - create_response = httpx.post( - url=f"{base_url}/v2/transcript", - headers=headers, - json=create_payload, - timeout=60.0, - ) - if create_response.status_code != 200: - pytest.fail( - "Failed to create transcript request: " - f"status={create_response.status_code}, body={create_response.text}" - ) - - transcript = create_response.json() - transcript_id = transcript.get("id") - if not transcript_id: - pytest.fail("Failed to get transcript id") - - for _ in range(60): - poll_response = httpx.get( - url=f"{base_url}/v2/transcript/{transcript_id}", - headers=headers, - timeout=30.0, - ) - if poll_response.status_code != 200: - pytest.fail( - "Failed to poll transcript status: " - f"status={poll_response.status_code}, body={poll_response.text}" - ) - transcript = poll_response.json() - if transcript.get("status") in ("completed", "error"): - break - time.sleep(1) - - httpx.delete( - url=f"{base_url}/v2/transcript/{transcript_id}", - headers=headers, - timeout=30.0, - ) - - if transcript.get("status") == "error": - pytest.fail(f"Failed to transcribe file error: {transcript.get('error')}") - - print(transcript.get("text")) - - -def test_assemblyai_basic_transcribe(): - print("making basic transcribe request to assemblyai passthrough") - _transcribe_and_verify(TEST_MASTER_KEY, TEST_BASE_URL) - - -async def generate_key(calling_key: str) -> str: - """Helper function to generate a new API key""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {calling_key}", - "Content-Type": "application/json", - } - - async with aiohttp.ClientSession() as session: - async with session.post(url, headers=headers, json={}) as response: - if response.status == 200: - data = await response.json() - return data.get("key") - raise Exception(f"Failed to generate key: {response.status}") - - -@pytest.mark.asyncio -async def test_assemblyai_transcribe_with_non_admin_key(): - non_admin_key = await generate_key(TEST_MASTER_KEY) - print(f"Generated non-admin key: {non_admin_key}") - - request_start_time = time.time() - _transcribe_and_verify(non_admin_key, TEST_BASE_URL) - request_end_time = time.time() - print(f"Request took {request_end_time - request_start_time} seconds") diff --git a/tests/pass_through_tests/test_hosted_vllm_passthrough.py b/tests/pass_through_tests/test_hosted_vllm_passthrough.py deleted file mode 100644 index 272b4e1bb00..00000000000 --- a/tests/pass_through_tests/test_hosted_vllm_passthrough.py +++ /dev/null @@ -1,71 +0,0 @@ -import asyncio -from unittest.mock import AsyncMock, patch - -import httpx -import pytest - -from litellm.passthrough.main import allm_passthrough_route -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -from litellm.utils import ProviderConfigManager -from litellm.types.utils import LlmProviders -from litellm.llms.vllm.passthrough.transformation import ( - VLLMPassthroughConfig, -) - - -def test_get_provider_passthrough_config_for_hosted_vllm_returns_vllm_config(): - # When requesting passthrough config for HOSTED_VLLM - cfg = ProviderConfigManager.get_provider_passthrough_config( - model="hosted_vllm/my-deployment", - provider=LlmProviders.HOSTED_VLLM, - ) - - # Then we should get a VLLMPassthroughConfig instance - assert isinstance(cfg, VLLMPassthroughConfig) - - -@pytest.mark.asyncio -async def test_allm_passthrough_route_with_hosted_vllm_model_does_not_raise(): - # Given a hosted_vllm model and an async http client - client = AsyncHTTPHandler() - - # Mock the provider resolution to ensure we use hosted_vllm and provide api_base - with patch( - "litellm.passthrough.main.get_llm_provider", - return_value=( - "my-deployment", # normalized model name - "hosted_vllm", # provider - "fake-api-key", # api key (not required for vllm) - "http://localhost:8090", # api base - ), - ): - # Mock the underlying AsyncClient.send to avoid real network I/O - fake_request = httpx.Request( - method="POST", url="http://localhost:8090/v1/chat/completions" - ) - fake_response = httpx.Response( - status_code=200, - content=b'{\n "ok": true\n}', - request=fake_request, - headers={"content-type": "application/json"}, - ) - - with patch.object( - client.client, "send", new=AsyncMock(return_value=fake_response) - ): - # When calling the async passthrough route with a hosted_vllm/* model - response = await allm_passthrough_route( - method="POST", - endpoint="v1/chat/completions", - model="hosted_vllm/my-deployment", - api_base="http://localhost:8090", - json={ - "model": "anything", # will be replaced internally with normalized model - "messages": [{"role": "user", "content": "Hello"}], - }, - client=client, - ) - - # Then it should not raise and return a successful httpx.Response - assert isinstance(response, httpx.Response) - assert response.status_code == 200 diff --git a/tests/pass_through_tests/test_openai_assistants_passthrough.py b/tests/pass_through_tests/test_openai_assistants_passthrough.py deleted file mode 100644 index 4da84ce5ca3..00000000000 --- a/tests/pass_through_tests/test_openai_assistants_passthrough.py +++ /dev/null @@ -1,23 +0,0 @@ -import os -import openai -import tempfile - - -client = openai.OpenAI(base_url="http://0.0.0.0:4000/openai", api_key=os.environ["LITELLM_MASTER_KEY"]) - - -def test_pass_through_file_operations(): - with tempfile.NamedTemporaryFile( - mode="w+", suffix=".txt", delete=False - ) as temp_file: - temp_file.write("This is a test file for the OpenAI Assistants API.") - temp_file.flush() - - file = client.files.create( - file=open(temp_file.name, "rb"), - purpose="assistants", - ) - print("file created", file) - - delete_file = client.files.delete(file.id) - print("file deleted", delete_file) diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index c1de9ae777d..834373df650 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -82,93 +82,6 @@ def get_tracked_spend() -> float: return sum(float(row.get("spend") or 0.0) for row in rows) -VERTEX_PROJECT = "litellm-ci-cd" -VERTEX_MODEL = "gemini-3.1-flash-lite" -VERTEX_GENERATE_CONTENT_URL = ( - f"{LITE_LLM_ENDPOINT}/vertex_ai/v1/projects/{VERTEX_PROJECT}" - f"/locations/global/publishers/google/models/{VERTEX_MODEL}:generateContent" -) - - -def _vertex_access_token() -> str: - import google.auth - import google.auth.transport.requests - - credentials, _ = google.auth.default( - scopes=["https://www.googleapis.com/auth/cloud-platform"] - ) - credentials.refresh(google.auth.transport.requests.Request()) - return credentials.token - - -def _spend_log_for_request(call_id: str) -> dict | None: - response = requests.get( - f"{LITE_LLM_ENDPOINT}/spend/logs?request_id={call_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - timeout=30, - ) - if response.status_code != 200: - return None - rows = response.json() - return rows[0] if rows else None - - -def _is_vertex_quota_error(response: requests.Response) -> bool: - return response.status_code == 429 or "RESOURCE_EXHAUSTED" in response.text - - -@pytest.mark.asyncio() -async def test_basic_vertex_ai_pass_through_with_spendlog(): - load_vertex_ai_credentials() - access_token = _vertex_access_token() - - # Drive the pass-through over HTTP instead of the vertexai SDK: the SDK intermittently - # routes generateContent to the public Vertex endpoint rather than the proxy override, - # so the call never reaches LiteLLM and no spend is logged. A direct request always - # hits the proxy. Spend logging then runs on a best-effort background worker that can - # drop a single event, so retry a few billed calls and assert that one specific call's - # spend log lands. Failing every attempt still fails hard, which is the signal we want - # if cost tracking is broken. - max_attempts = 3 - poll_seconds = 60 - poll_interval = 5 - - for attempt in range(1, max_attempts + 1): - response = requests.post( - VERTEX_GENERATE_CONTENT_URL, - headers={ - "Authorization": f"Bearer {access_token}", - "Content-Type": "application/json", - }, - json={"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}, - timeout=60, - ) - if _is_vertex_quota_error(response): - pytest.skip("Vertex AI quota exhausted") - assert ( - response.status_code == 200 - ), f"vertex pass-through call failed: {response.status_code} {response.text}" - - call_id = response.headers.get("x-litellm-call-id") - assert call_id, "proxy response missing x-litellm-call-id header" - - for _ in range(poll_seconds // poll_interval): - await asyncio.sleep(poll_interval) - row = _spend_log_for_request(call_id) - if row is not None and float(row.get("spend") or 0) > 0: - assert "gemini" in row["model"], f"unexpected model in spend log: {row}" - assert ( - row["custom_llm_provider"] == "vertex_ai" - ), f"unexpected provider in spend log: {row}" - return - - print(f"attempt {attempt}: spend log for call {call_id} not found yet, re-billing") - - pytest.fail( - f"Vertex pass-through spend never recorded after {max_attempts} billed calls" - ) - - @pytest.mark.asyncio() @pytest.mark.skip(reason="skip flaky test - vertex pass through streaming is flaky") async def test_basic_vertex_ai_pass_through_streaming_with_spendlog(): diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py index a7f04466d14..41b351631f6 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py @@ -202,50 +202,6 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): f"but got {cache_read}. Full usage: {usage}" ) - @pytest.mark.asyncio - async def test_prompt_caching_with_system_message(self): - """ - E2E test: Prompt caching with system message should work. - """ - _skip_live_prompt_caching_test() - litellm.turn_on_debug() - - messages = [ - { - "role": "user", - "content": "What are the key terms?", - }, - ] - - system = [ - { - "type": "text", - "text": LARGE_DOCUMENT_FOR_CACHING, - "cache_control": {"type": "ephemeral"}, - }, - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - system=system, - max_tokens=100, - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - usage = response.get("usage", {}) - cache_creation = usage.get("cache_creation_input_tokens", 0) - cache_read = usage.get("cache_read_input_tokens", 0) - - print(f"cache_creation_input_tokens: {cache_creation}") - print(f"cache_read_input_tokens: {cache_read}") - - assert cache_creation > 0 or cache_read > 0, ( - f"Expected cache tokens > 0 for system message caching, " - f"but got cache_creation={cache_creation}, cache_read={cache_read}" - ) - def _parse_sse_chunks(self, chunk: bytes) -> list: """ Parse SSE format chunks and return list of JSON objects. @@ -432,94 +388,3 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): f"Expected cache_read_input_tokens > 0 on second streaming call, " f"but got {cache_read}" ) - - @pytest.mark.asyncio - async def test_prompt_caching_message_start_indicates_caching_support(self): - """ - E2E test: message_start event should contain cache fields to indicate caching support. - - This validates that the message_start event includes cache_creation_input_tokens - and cache_read_input_tokens fields (even if initialized to 0) so that clients - like Claude Code can detect that prompt caching is supported. - - This test specifically addresses the issue where Bedrock converse API streaming - didn't include cache fields in message_start, causing clients to think caching - wasn't supported. - """ - _skip_live_prompt_caching_test() - litellm.turn_on_debug() - - messages = self.get_messages_with_cache_control() - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - max_tokens=100, - stream=True, - ) - - # Look for message_start event and validate it has cache fields - message_start_found = False - message_start_has_cache_creation_field = False - message_start_has_cache_read_field = False - - async for chunk in response: - # Handle SSE format chunks (bytes) - if isinstance(chunk, bytes): - json_chunks = self._parse_sse_chunks(chunk) - for json_data in json_chunks: - if json_data.get("type") == "message_start": - message_start_found = True - message = json_data.get("message", {}) - usage = message.get("usage", {}) - - print( - f"message_start usage: {json.dumps(usage, indent=2, default=str)}" - ) - - # Check that cache fields are present (even if 0) - if "cache_creation_input_tokens" in usage: - message_start_has_cache_creation_field = True - if "cache_read_input_tokens" in usage: - message_start_has_cache_read_field = True - - # Break after first message_start - break - elif isinstance(chunk, dict): - if chunk.get("type") == "message_start": - message_start_found = True - message = chunk.get("message", {}) - usage = message.get("usage", {}) - - print( - f"message_start usage: {json.dumps(usage, indent=2, default=str)}" - ) - - # Check that cache fields are present (even if 0) - if "cache_creation_input_tokens" in usage: - message_start_has_cache_creation_field = True - if "cache_read_input_tokens" in usage: - message_start_has_cache_read_field = True - - # Break after first message_start - break - - # Break if we found message_start - if message_start_found: - break - - # Validate that message_start was found - assert ( - message_start_found - ), "Expected to find message_start event in streaming response" - - # Validate that cache fields are present in message_start - assert message_start_has_cache_creation_field, ( - "Expected cache_creation_input_tokens field in message_start event. " - "This field should be present (even if 0) to indicate caching support to clients." - ) - - assert message_start_has_cache_read_field, ( - "Expected cache_read_input_tokens field in message_start event. " - "This field should be present (even if 0) to indicate caching support to clients." - ) diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py index 9e706f99316..966048e609e 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py @@ -147,48 +147,6 @@ class BaseAnthropicMessagesToolSearchTest(ABC): content = response.get("content", []) assert len(content) > 0, "Response should have content" - @pytest.mark.asyncio - async def test_tool_search_discovers_tool(self): - """ - E2E test: Tool search should discover and use a deferred tool. - - This validates that when the user asks about weather, the model - discovers the get_weather tool via tool search and attempts to use it. - """ - litellm.turn_on_debug() - - tools = self.get_tools_with_tool_search() - messages = [ - { - "role": "user", - "content": "I need to know the current weather in New York City. Please use the appropriate tool.", - } - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - tools=tools, - max_tokens=1024, - extra_headers=self.get_extra_headers(), - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - content = response.get("content", []) - - # Check if the model used tool_use (either tool_search or get_weather) - tool_uses = [block for block in content if block.get("type") == "tool_use"] - - print(f"Tool uses: {json.dumps(tool_uses, indent=2, default=str)}") - - # The model should attempt to use tools when asked about weather - # It might use tool_search first, or directly use get_weather if discovered - if response.get("stop_reason") == "tool_use": - assert ( - len(tool_uses) > 0 - ), "Expected tool_use blocks when stop_reason is tool_use" - @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=5) async def test_tool_search_streaming(self): @@ -234,39 +192,3 @@ class BaseAnthropicMessagesToolSearchTest(ABC): # Should have message_start message_starts = [c for c in chunks if c.get("type") == "message_start"] assert len(message_starts) > 0, "Expected message_start in streaming response" - - @pytest.mark.asyncio - async def test_tool_search_with_multiple_deferred_tools(self): - """ - E2E test: Tool search should work with multiple deferred tools. - - This validates that the model can discover the appropriate tool - from a larger catalog of deferred tools. - """ - litellm.turn_on_debug() - - tools = self.get_tools_with_tool_search() - messages = [ - {"role": "user", "content": "What's the stock price of Apple (AAPL)?"} - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - tools=tools, - max_tokens=1024, - extra_headers=self.get_extra_headers(), - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - # Validate response - assert "content" in response, "Response should contain content" - - content = response.get("content", []) - tool_uses = [block for block in content if block.get("type") == "tool_use"] - - # If the model decides to use a tool, it should be related to stocks - if tool_uses: - tool_names = [t.get("name") for t in tool_uses] - print(f"Tools used: {tool_names}") diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index 16055b5a29b..858a6713f7b 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -62,234 +62,3 @@ class BaseAnthropicMessagesTest: assert "content" in response assert "model" in response assert response.get("role") == "assistant" - - @pytest.mark.asyncio - async def test_non_streaming_base(self): - """Base test for non-streaming requests""" - litellm.turn_on_debug() - - request_params = self.model_config - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Prepare call arguments - call_args = { - "messages": messages, - "max_tokens": 100, - } - - # Add any additional config from subclass - call_args.update(request_params) - - # Call the handler - response = await litellm.anthropic.messages.acreate(**call_args) - - print(f"Non-streaming {request_params['model']} response: ", response) - - # Verify response - self._validate_response(response) - - print(f"Non-streaming response: {json.dumps(response, indent=2, default=str)}") - return response - - @pytest.mark.asyncio - async def test_response_format_consistency(self): - """ - Test that response content blocks are consistently dicts (not Pydantic objects). - - This ensures that code like response["content"][0]["type"] works - regardless of the target provider. - - Issue: https://github.com/BerriAI/litellm/issues/20342 - """ - litellm.turn_on_debug() - - request_params = self.model_config - - # Set up test parameters - messages = [{"role": "user", "content": "Say hi"}] - - # Prepare call arguments - call_args = { - "messages": messages, - "max_tokens": 100, - } - - # Add any additional config from subclass - call_args.update(request_params) - - # Call the handler - response = await litellm.anthropic.messages.acreate(**call_args) - - print( - f"Response for {request_params['model']}: {json.dumps(response, indent=2, default=str)}" - ) - - # Verify response structure - assert "content" in response, "Response should have 'content' field" - assert len(response["content"]) > 0, "Response content should not be empty" - - # Get the first content block - block = response["content"][0] - - # Check that the block is a dict, not a Pydantic object - assert isinstance(block, dict), ( - f"Content block should be a dict, but got {type(block)}. " - f"This means response format is inconsistent across providers." - ) - - # Verify we can access fields using dict syntax (not object attributes) - try: - block_type = block["type"] - print(f"✓ Successfully accessed block['type']: {block_type}") - except TypeError as e: - pytest.fail( - f"Cannot access content block using dict syntax: {e}. " - f"Block type: {type(block)}" - ) - - # Verify the block has expected structure - assert "type" in block, "Content block should have 'type' field" - if block["type"] == "text": - assert "text" in block, "Text content block should have 'text' field" - - print( - f"✓ Response format consistency test passed for {request_params['model']}" - ) - - @pytest.mark.asyncio - async def test_anthropic_messages_litellm_router_streaming_with_logging(self): - """ - Test that logging and cost tracking works for anthropic_messages with streaming request - """ - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": {**self.model_config}, - } - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="claude-special-alias", - max_tokens=100, - stream=True, - ) - - response_prompt_tokens = 0 - response_completion_tokens = 0 - all_anthropic_usage_chunks = [] - buffer = "" - - async for chunk in response: - # Decode chunk if it's bytes - print("chunk=", chunk) - - # Handle SSE format chunks - if isinstance(chunk, bytes): - chunk_str = chunk.decode("utf-8") - buffer += chunk_str - # Extract the JSON data part from SSE format - for line in buffer.split("\n"): - if line.startswith("data: "): - try: - json_data = json.loads(line[6:]) # Skip the 'data: ' prefix - print( - "\n\nJSON data:", - json.dumps(json_data, indent=4, default=str), - ) - - # Extract usage information - if ( - json_data.get("type") == "message_start" - and "message" in json_data - ): - if "usage" in json_data["message"]: - usage = json_data["message"]["usage"] - all_anthropic_usage_chunks.append(usage) - print( - "USAGE BLOCK", - json.dumps(usage, indent=4, default=str), - ) - elif "usage" in json_data: - usage = json_data["usage"] - all_anthropic_usage_chunks.append(usage) - print( - "USAGE BLOCK", - json.dumps(usage, indent=4, default=str), - ) - except json.JSONDecodeError: - print(f"Failed to parse JSON from: {line[6:]}") - elif hasattr(chunk, "message"): - if chunk.message.usage: - print( - "USAGE BLOCK", - json.dumps(chunk.message.usage, indent=4, default=str), - ) - all_anthropic_usage_chunks.append(chunk.message.usage) - elif hasattr(chunk, "usage"): - print("USAGE BLOCK", json.dumps(chunk.usage, indent=4, default=str)) - all_anthropic_usage_chunks.append(chunk.usage) - - print( - "all_anthropic_usage_chunks", - json.dumps(all_anthropic_usage_chunks, indent=4, default=str), - ) - - # Extract token counts from usage data - if all_anthropic_usage_chunks: - response_prompt_tokens = max( - [usage.get("input_tokens", 0) for usage in all_anthropic_usage_chunks] - ) - response_completion_tokens = max( - [usage.get("output_tokens", 0) for usage in all_anthropic_usage_chunks] - ) - - print("input_tokens_anthropic_api", response_prompt_tokens) - print("output_tokens_anthropic_api", response_completion_tokens) - - await asyncio.sleep(4) - - print( - "logged_standard_logging_payload", - json.dumps( - test_custom_logger.logged_standard_logging_payload, - indent=4, - default=str, - ), - ) - - assert ( - test_custom_logger.logged_standard_logging_payload is not None - ), "Logging payload should not be None" - assert ( - test_custom_logger.logged_standard_logging_payload["messages"] == messages - ) - assert ( - test_custom_logger.logged_standard_logging_payload["response"] is not None - ) - assert ( - test_custom_logger.logged_standard_logging_payload["model"] - == self.expected_model_name_in_logging - ) - - # check logged usage + spend - assert test_custom_logger.logged_standard_logging_payload["response_cost"] > 0 - assert ( - test_custom_logger.logged_standard_logging_payload["prompt_tokens"] - == response_prompt_tokens - ) - assert ( - test_custom_logger.logged_standard_logging_payload["completion_tokens"] - == response_completion_tokens - ) diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 0f0bb8f091e..001de2f7d52 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -118,13 +118,6 @@ class TestAnthropicOpenAIAPI(BaseAnthropicMessagesTest): """ return "gpt-4.1-mini" - @pytest.mark.asyncio - async def test_anthropic_messages_litellm_router_streaming_with_logging(self): - """ - Test the anthropic_messages with streaming request - """ - pass - @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_non_streaming(): @@ -163,283 +156,3 @@ async def test_anthropic_messages_litellm_router_non_streaming(): print(f"Non-streaming response: {json.dumps(response, indent=2)}") return response - -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_routing_strategy(): - """ - Test the anthropic_messages with routing strategy + non-streaming request - """ - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ], - routing_strategy="latency-based-routing", - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="claude-special-alias", - max_tokens=100, - metadata={ - "user_id": "hello", - }, - ) - - # Verify response - assert "id" in response - assert "content" in response - assert "model" in response - assert response["role"] == "assistant" - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - return response - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallbacks(): - """ - E2E test the anthropic_messages fallbacks from Anthropic API to Bedrock - """ - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "anthropic/claude-opus-4-7", - "litellm_params": { - "model": "anthropic/claude-opus-4-7", - "api_key": "bad-key", - }, - }, - { - "model_name": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - ], - fallbacks=[ - { - "anthropic/claude-opus-4-7": [ - "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0" - ] - } - ], - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="anthropic/claude-opus-4-7", - max_tokens=100, - metadata={ - "user_id": "hello", - }, - ) - - # Verify response - assert "id" in response - assert "content" in response - assert "model" in response - assert response["role"] == "assistant" - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - return response - - -class TestCustomLogger(CustomLogger): - def __init__(self): - super().__init__() - self.logged_standard_logging_payload: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - print("inside async_log_success_event") - self.logged_standard_logging_payload = kwargs.get("standard_logging_object") - - pass - - -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_non_streaming_with_logging(): - """ - Test the anthropic_messages with non-streaming request - - - Ensure Cost + Usage is tracked - """ - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - litellm.turn_on_debug() - MODEL_GROUP = "claude-special-alias" - router = Router( - model_list=[ - { - "model_name": MODEL_GROUP, - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model=MODEL_GROUP, - max_tokens=100, - ) - - # Verify response - _validate_anthropic_response(response) - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - - await asyncio.sleep(1) - - assert ( - test_custom_logger.logged_standard_logging_payload is not None - ), "Logging payload should not be None" - print( - "tracked standard logging payload", - json.dumps( - test_custom_logger.logged_standard_logging_payload, indent=4, default=str - ), - ) - assert test_custom_logger.logged_standard_logging_payload["messages"] == messages - assert test_custom_logger.logged_standard_logging_payload["response"] is not None - assert ( - test_custom_logger.logged_standard_logging_payload["model"] - == "claude-haiku-4-5-20251001" - ) - - # check logged usage + spend - assert test_custom_logger.logged_standard_logging_payload["response_cost"] > 0 - assert ( - test_custom_logger.logged_standard_logging_payload["prompt_tokens"] - == response["usage"]["input_tokens"] - ) - assert ( - test_custom_logger.logged_standard_logging_payload["completion_tokens"] - == response["usage"]["output_tokens"] - ) - - # assert model_group - assert ( - test_custom_logger.logged_standard_logging_payload["model_group"] == MODEL_GROUP - ) - - -# @pytest.mark.asyncio -# async def test_bedrock_messages_api_header_forwarding(): -# """ -# Test that headers from kwargs (set by proxy's add_headers_to_llm_call_by_model_group) -# are correctly passed to validate_anthropic_messages_environment for Bedrock Invoke API. - -# This verifies that forward_client_headers_to_llm_api works for Bedrock Invoke API (Messages API). - -# Issue: When calling Anthropic models via the Messages API, LiteLLM makes a call to -# Bedrock's Invoke API, and custom headers were not being forwarded, even though -# they worked correctly for Chat Completions API with Bedrock's Converse API. -# """ -# from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler -# from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -# from litellm.types.router import GenericLiteLLMParams - -# handler = BaseLLMHTTPHandler() - -# # Headers that would be set by the proxy when forward_client_headers_to_llm_api is configured -# custom_headers = { -# "X-Custom-Header": "CustomValue", -# "X-Request-ID": "req-123", -# } - -# # Mock the provider config -# mock_provider_config = MagicMock() - -# # We'll check what headers are passed to this method -# mock_provider_config.validate_anthropic_messages_environment.return_value = ( -# {"Authorization": "Bearer test"}, -# "https://bedrock-runtime.us-east-1.amazonaws.com/invoke" -# ) -# mock_provider_config.transform_anthropic_messages_request.return_value = {"model": "test"} -# mock_provider_config.get_complete_url.return_value = "https://test.com" -# mock_provider_config.sign_request.return_value = ({}, None) -# mock_provider_config.transform_anthropic_messages_response.return_value = {"id": "test"} - -# # Mock HTTP client to prevent actual network calls -# with unittest.mock.patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client") as mock_get_client: -# mock_http_client = AsyncMock() -# mock_response = MagicMock() -# mock_response.status_code = 200 -# mock_response.json.return_value = {"id": "test", "content": []} -# mock_response.text = "{}" -# mock_http_client.post.return_value = mock_response -# mock_get_client.return_value = mock_http_client - -# # Mock logging object -# mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) -# mock_logging_obj.model_call_details = {} - -# # Call the handler with headers in kwargs -# try: -# await handler.async_anthropic_messages_handler( -# model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", -# messages=[{"role": "user", "content": "Hello"}], -# anthropic_messages_provider_config=mock_provider_config, -# anthropic_messages_optional_request_params={"max_tokens": 100}, -# custom_llm_provider="bedrock", -# litellm_params=GenericLiteLLMParams( -# api_key="test-key", -# aws_region_name="us-east-1" -# ), -# logging_obj=mock_logging_obj, -# api_key="test-key", -# stream=False, -# kwargs={"headers": custom_headers} # Headers set by proxy -# ) -# except Exception: -# pass # Ignore errors, we're only checking if headers were passed - -# # Verify that validate_anthropic_messages_environment was called -# assert mock_provider_config.validate_anthropic_messages_environment.called - -# # Get the headers that were passed -# call_args = mock_provider_config.validate_anthropic_messages_environment.call_args -# passed_headers = call_args[1]["headers"] - -# # The custom headers from kwargs should be in the passed headers -# assert "X-Custom-Header" in passed_headers or "x-custom-header" in passed_headers -# assert "X-Request-ID" in passed_headers or "x-request-id" in passed_headers - - -def test_sync_openai_messages(): - """ - Test the anthropic_messages with sync request - """ - litellm.turn_on_debug() - response = litellm.anthropic.messages.create( - messages=[{"role": "user", "content": "Hello, can you tell me a short joke?"}], - model="openai/gpt-4.1-mini", - max_tokens=100, - ) - print("ANT response", response) - - assert response is not None - assert isinstance(response, dict) - assert response["content"][0]["text"] is not None diff --git a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py index 451a09d30bb..aacf65ea4a3 100644 --- a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py +++ b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py @@ -14,54 +14,6 @@ from base_anthropic_unified_messages_test import BaseAnthropicMessagesTest INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST = BaseAnthropicMessagesTest() -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_bedrock(): - """ - Test the anthropic_messages with non-streaming request - """ - - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - { - "model_name": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call 1 using bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0 - response = await router.aanthropic_messages( - messages=messages, - model="bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - max_tokens=100, - ) - - # Verify response - INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST._validate_response(response) - - # Call 2 using bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0 - response = await router.aanthropic_messages( - messages=messages, - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - max_tokens=100, - ) - - # Verify response - INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST._validate_response(response) - - @pytest.mark.asyncio async def test_anthropic_messages_bedrock_converse_with_thinking(): """ diff --git a/tests/search_tests/base_search_unit_tests.py b/tests/search_tests/base_search_unit_tests.py index 7028f58a1a3..8c16e790607 100644 --- a/tests/search_tests/base_search_unit_tests.py +++ b/tests/search_tests/base_search_unit_tests.py @@ -113,62 +113,3 @@ class BaseSearchTest(ABC): except Exception as e: pytest.fail(f"Search call failed: {str(e)}") - - def test_search_response_structure(self): - """ - Test that the Search response has the correct structure. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="artificial intelligence recent news", - search_provider=search_provider, - ) - - # Validate response structure - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert hasattr(response, "object"), "Response should have 'object' attribute" - - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert response.object == "search", "object should be 'search'" - - # Validate first result structure - first_result = response.results[0] - assert hasattr(first_result, "title"), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - assert isinstance(first_result.title, str), "title should be a string" - assert isinstance(first_result.url, str), "url should be a string" - assert isinstance(first_result.snippet, str), "snippet should be a string" - - print(f"\nResponse structure validated:") - print(f" - object: {response.object}") - print(f" - results: {len(response.results)}") - print(f" - first result has all required fields") - - def test_search_with_optional_params(self): - """ - Test search with optional parameters. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="machine learning", - search_provider=search_provider, - max_results=5, - ) - - # Validate response - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert len(response.results) <= 5, "Should have at most 5 results as requested" - - print(f"\nSearch with optional params validated:") - print(f" - Requested max_results: 5") - print(f" - Received results: {len(response.results)}") diff --git a/tests/search_tests/test_duckduckgo_search.py b/tests/search_tests/test_duckduckgo_search.py deleted file mode 100644 index 682221326bb..00000000000 --- a/tests/search_tests/test_duckduckgo_search.py +++ /dev/null @@ -1,138 +0,0 @@ -""" -Tests for DuckDuckGo Search API integration. -""" - -import os - -import pytest - -import litellm -from tests.search_tests.base_search_unit_tests import BaseSearchTest - - -class TestDuckDuckGoSearch(BaseSearchTest): - """ - Tests for DuckDuckGo Search functionality. - """ - - def get_search_provider(self) -> str: - """ - Return search_provider for DuckDuckGo Search. - """ - return "duckduckgo" - - @pytest.mark.asyncio - async def test_basic_search(self): - """ - Test basic search functionality with a simple query. - """ - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - litellm.turn_on_debug() - search_provider = self.get_search_provider() - print("Search Provider=", search_provider) - - try: - response = await litellm.asearch( - query="india", - search_provider=search_provider, - ) - print("Search response=", response.model_dump_json(indent=4)) - - print(f"\n{'='*80}") - print(f"Response type: {type(response)}") - print( - f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}" - ) - - # Check if response has expected Search format - assert hasattr( - response, "results" - ), "Response should have 'results' attribute" - assert hasattr( - response, "object" - ), "Response should have 'object' attribute" - assert ( - response.object == "search" - ), f"Expected object='search', got '{response.object}'" - - # Validate results structure - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - - # Check first result structure - first_result = response.results[0] - assert hasattr( - first_result, "title" - ), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - - print(f"Total results: {len(response.results)}") - print(f"First result title: {first_result.title}") - print(f"First result URL: {first_result.url}") - print(f"First result snippet: {first_result.snippet[:100]}...") - print(f"{'='*80}\n") - - assert len(first_result.title) > 0, "Title should not be empty" - assert len(first_result.url) > 0, "URL should not be empty" - assert len(first_result.snippet) > 0, "Snippet should not be empty" - - # Validate cost tracking in _hidden_params - assert hasattr( - response, "_hidden_params" - ), "Response should have '_hidden_params' attribute" - hidden_params = response._hidden_params - assert ( - "response_cost" in hidden_params - ), "_hidden_params should contain 'response_cost'" - - response_cost = hidden_params["response_cost"] - assert response_cost is not None, "response_cost should not be None" - assert isinstance( - response_cost, (int, float) - ), "response_cost should be a number" - assert response_cost == 0, "response_cost should be 0" - - print(f"Cost tracking: ${response_cost:.6f}") - - except Exception as e: - pytest.fail(f"Search call failed: {str(e)}") - - def test_search_response_structure(self): - """ - Test that the Search response has the correct structure. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="india", - search_provider=search_provider, - ) - - # Validate response structure - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert hasattr(response, "object"), "Response should have 'object' attribute" - - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert response.object == "search", "object should be 'search'" - - # Validate first result structure - first_result = response.results[0] - assert hasattr(first_result, "title"), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - assert isinstance(first_result.title, str), "title should be a string" - assert isinstance(first_result.url, str), "url should be a string" - assert isinstance(first_result.snippet, str), "snippet should be a string" - - print(f"\nResponse structure validated:") - print(f" - object: {response.object}") - print(f" - results: {len(response.results)}") - print(f" - first result has all required fields") diff --git a/tests/search_tests/test_firecrawl_search.py b/tests/search_tests/test_firecrawl_search.py deleted file mode 100644 index eec74d48e26..00000000000 --- a/tests/search_tests/test_firecrawl_search.py +++ /dev/null @@ -1,42 +0,0 @@ -from unittest.mock import Mock, patch -import litellm - - -def test_firecrawl_search_request_body(): - """ - Test that validates the Firecrawl search request body is correctly formatted. - """ - mock_response = Mock() - mock_response.status_code = 200 - mock_response.json.return_value = { - "success": True, - "data": { - "web": [ - { - "title": "Test Title", - "url": "https://example.com", - "markdown": "Test content", - } - ] - }, - } - - with patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", - return_value=mock_response, - ) as mock_post: - litellm.search( - query="test query", - search_provider="firecrawl", - max_results=10, - country="US", - ) - - assert mock_post.called - call_kwargs = mock_post.call_args.kwargs - request_body = call_kwargs.get("json") - - assert request_body is not None - assert request_body["query"] == "test query" - assert request_body["limit"] == 10 - assert request_body["country"] == "US" diff --git a/tests/spend_tracking_tests/test_ocr_spend_tracking.py b/tests/spend_tracking_tests/test_ocr_spend_tracking.py deleted file mode 100644 index 3ce77c56361..00000000000 --- a/tests/spend_tracking_tests/test_ocr_spend_tracking.py +++ /dev/null @@ -1,296 +0,0 @@ -""" -Unit tests for OCR spend tracking in get_logging_payload. - -This test file verifies that OCR/AOCR calls correctly extract usage_info -and populate the spend logs payload with pages_processed instead of token counts. -""" - -import pytest -from datetime import datetime, timezone -from unittest.mock import Mock -from pydantic import BaseModel -from typing import Optional - -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - get_logging_payload, - _extract_usage_for_ocr_call, -) - - -class MockUsageInfo(BaseModel): - """Mock Pydantic model for OCR usage_info""" - - pages_processed: int - doc_size_bytes: Optional[int] = None - - -class MockOCRResponse(BaseModel): - """Mock Pydantic model for OCR response""" - - id: str - object: str - model: str - usage_info: MockUsageInfo - - -class TestExtractUsageForOCRCall: - """Test the _extract_usage_for_ocr_call helper method""" - - def test_extract_usage_from_dict(self): - """Test extracting usage from dict response""" - response_obj_dict = {"usage_info": {"pages_processed": 5}} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage["prompt_tokens"] == 0 - assert usage["completion_tokens"] == 0 - assert usage["total_tokens"] == 0 - assert usage["pages_processed"] == 5 - - def test_extract_usage_from_pydantic_model(self): - """Test extracting usage from Pydantic model response""" - usage_info = MockUsageInfo(pages_processed=10, doc_size_bytes=1024) - response_obj = MockOCRResponse( - id="ocr-123", object="ocr", model="test-ocr-model", usage_info=usage_info - ) - response_obj_dict = response_obj.model_dump() - - usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict) - - assert usage["prompt_tokens"] == 0 - assert usage["completion_tokens"] == 0 - assert usage["total_tokens"] == 0 - assert usage["pages_processed"] == 10 - - def test_extract_usage_with_object_attributes(self): - """Test extracting usage from object with __dict__""" - - class SimpleUsageInfo: - def __init__(self, pages_processed): - self.pages_processed = pages_processed - - class SimpleOCRResponse: - def __init__(self): - self.usage_info = SimpleUsageInfo(pages_processed=3) - - response_obj = SimpleOCRResponse() - response_obj_dict = {} - - usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict) - - assert usage.get("prompt_tokens") == 0 - assert usage.get("completion_tokens") == 0 - assert usage.get("total_tokens") == 0 - assert usage.get("pages_processed") == 3 - - def test_extract_usage_missing_usage_info(self): - """Test handling missing usage_info""" - response_obj_dict = {} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage == {} - - def test_extract_usage_empty_usage_info(self): - """Test handling empty usage_info""" - response_obj_dict = {"usage_info": {}} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage.get("prompt_tokens") == 0 - assert usage.get("completion_tokens") == 0 - assert usage.get("total_tokens") == 0 - assert usage.get("pages_processed") == 0 - - -class TestGetLoggingPayloadOCR: - """Test get_logging_payload with OCR call types""" - - @pytest.fixture - def mock_datetime(self): - """Fixture for consistent timestamps""" - return datetime.now(timezone.utc) - - @pytest.fixture - def base_kwargs(self): - """Fixture for base kwargs used in tests""" - return { - "model": "test-ocr-model", - "call_type": "ocr", - "litellm_params": {}, - "response_cost": 0.05, - } - - def test_ocr_call_with_dict_response(self, mock_datetime, base_kwargs): - """Test OCR call with dict response containing usage_info""" - response_obj = { - "id": "ocr-test-123", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 7, "doc_size_bytes": 2048}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - assert payload["spend"] == 0.05 - - # Verify pages_processed is in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert "additional_usage_values" in metadata - assert metadata["additional_usage_values"]["pages_processed"] == 7 - - def test_aocr_call_with_pydantic_response(self, mock_datetime, base_kwargs): - """Test AOCR (async OCR) call with Pydantic model response""" - base_kwargs["call_type"] = "aocr" - - usage_info = MockUsageInfo(pages_processed=12) - response_obj = MockOCRResponse( - id="aocr-test-456", - object="ocr", - model="test-ocr-model", - usage_info=usage_info, - ) - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "aocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - # Verify pages_processed is in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert "additional_usage_values" in metadata - assert metadata["additional_usage_values"]["pages_processed"] == 12 - - def test_ocr_call_missing_usage_info(self, mock_datetime, base_kwargs): - """Test OCR call with missing usage_info returns empty usage""" - response_obj = { - "id": "ocr-test-789", - "object": "ocr", - "model": "test-ocr-model", - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - def test_ocr_call_with_zero_pages(self, mock_datetime, base_kwargs): - """Test OCR call with zero pages processed""" - response_obj = { - "id": "ocr-test-000", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 0}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - # Verify pages_processed is 0 - import json - - metadata = json.loads(payload["metadata"]) - assert metadata["additional_usage_values"]["pages_processed"] == 0 - - def test_non_ocr_call_uses_token_based_usage(self, mock_datetime): - """Test that non-OCR calls still use token-based usage""" - kwargs = { - "model": "gpt-5.5", - "call_type": "completion", - "litellm_params": {}, - "response_cost": 0.02, - } - - response_obj = { - "id": "completion-test-123", - "object": "chat.completion", - "model": "gpt-5.5", - "usage": { - "prompt_tokens": 50, - "completion_tokens": 100, - "total_tokens": 150, - }, - } - - payload = get_logging_payload( - kwargs=kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "completion" - assert payload["prompt_tokens"] == 50 - assert payload["completion_tokens"] == 100 - assert payload["total_tokens"] == 150 - - def test_ocr_with_metadata(self, mock_datetime, base_kwargs): - """Test OCR call with additional metadata""" - base_kwargs["litellm_params"] = { - "metadata": { - "user_api_key_user_id": "test-user", - "user_api_key_team_id": "test-team", - } - } - - response_obj = { - "id": "ocr-metadata-test", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 5, "doc_size_bytes": 1024}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["user"] == "test-user" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - - # Verify pages_processed and doc_size_bytes are both in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert metadata["additional_usage_values"]["pages_processed"] == 5 - assert metadata["additional_usage_values"]["doc_size_bytes"] == 1024 diff --git a/tests/test_spend_logs.py b/tests/test_spend_logs.py index 57109ea01e5..f239fbbb53f 100644 --- a/tests/test_spend_logs.py +++ b/tests/test_spend_logs.py @@ -116,7 +116,7 @@ async def generate_team(session: aiohttp.ClientSession, org_id: str) -> dict: @pytest.mark.skip( - reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and the proxy_spend_accuracy_tests CircleCI job." + reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and tests/integration/spend/." ) @pytest.mark.asyncio async def test_spend_logs_with_org_id(): diff --git a/tests/unit/batches/test_main.py b/tests/unit/batches/test_main.py index 3303c13b6a3..fb97e0e4919 100644 --- a/tests/unit/batches/test_main.py +++ b/tests/unit/batches/test_main.py @@ -34,6 +34,16 @@ import pytest import litellm import litellm.batches.main as bm +import asyncio +import datetime +import json +from collections.abc import Mapping +from typing import Final +import httpx +import respx +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict +from litellm.integrations.custom_logger import CustomLogger # --------------------------------------------------------------------------- # @@ -986,3 +996,211 @@ async def test_batch_logging_azure_credentials_regression(): print("✓ Batch output files can be fetched with Azure credentials") print("✓ Cost and usage tracking works for Azure batches") print("✓ Backwards compatibility maintained\n") + + +_OPENAI_FILE_JSON: Final = MappingProxyType( + { + "id": "file-abc123", + "object": "file", + "purpose": "batch", + "filename": "batch.jsonl", + "bytes": 416, + "created_at": 1739598666, + "status": "processed", + } +) + + +_OPENAI_BATCH_JSON: Final = MappingProxyType( + { + "id": "batch_abc123", + "object": "batch", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc123", + "status": "validating", + "completion_window": "24h", + "created_at": 1739598666, + } +) + + +class _KeyAliasMetadata(TypedDict): + user_api_key_alias: ReadOnly[str | None] + user_api_key_team_alias: ReadOnly[str | None] + + +class _LoggedCall(TypedDict): + call_type: ReadOnly[str] + metadata: ReadOnly[_KeyAliasMetadata] + + +_LOGGED_CALL: Final = TypeAdapter(_LoggedCall) + + +class _SuccessPayloadRecorder(CustomLogger): + def __init__(self, call_type: str) -> None: + super().__init__() + self._call_type: Final = call_type + self.logged: Final = asyncio.Event() + self.payload: _LoggedCall | None = None + + async def async_log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + payload: Final = _LOGGED_CALL.validate_python(kwargs["standard_logging_object"]) + if payload["call_type"] != self._call_type: + return + self.payload = payload + self.logged.set() + + +@pytest.mark.asyncio +async def test_acreate_batch_full_crud_and_logging_metadata( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.logging_callback_manager._reset_all_callbacks() + recorder: Final = _SuccessPayloadRecorder("acreate_batch") + monkeypatch.setattr(litellm, "callbacks", [recorder]) + + upload_route: Final = respx_mock.post("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_FILE_JSON)) + ) + create_route: Final = respx_mock.post("https://api.openai.com/v1/batches").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_BATCH_JSON)) + ) + retrieve_route: Final = respx_mock.get("https://api.openai.com/v1/batches/batch_abc123").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_BATCH_JSON)) + ) + list_batches_route: Final = respx_mock.get("https://api.openai.com/v1/batches").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_OPENAI_BATCH_JSON)]}) + ) + respx_mock.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response(200, content=b'{"custom_id": "request-1"}\n') + ) + respx_mock.get("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_FILE_JSON)) + ) + respx_mock.delete("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json={"id": "file-abc123", "object": "file", "deleted": True}) + ) + list_files_route: Final = respx_mock.get("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_OPENAI_FILE_JSON)]}) + ) + cancel_route: Final = respx_mock.post("https://api.openai.com/v1/batches/batch_abc123/cancel").mock( + return_value=httpx.Response(200, json={**_OPENAI_BATCH_JSON, "status": "cancelling"}) + ) + + batch_file: Final = ( + "batch.jsonl", + b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", ' + b'"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}}\n', + "application/jsonl", + ) + file_obj: Final = await litellm.acreate_file( + file=batch_file, purpose="batch", custom_llm_provider="openai", api_key="fake-key" + ) + assert file_obj.id == "file-abc123" + upload_body: Final = upload_route.calls.last.request.content + assert b'name="purpose"\r\n\r\nbatch' in upload_body + assert batch_file[1] in upload_body + + extra_metadata_field: Final = { + "user_api_key_alias": "special_api_key_alias", + "user_api_key_team_alias": "special_team_alias", + } + create_batch_response: Final = await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id=file_obj.id, + custom_llm_provider="openai", + api_key="fake-key", + metadata={"key1": "value1", "key2": "value2"}, + litellm_metadata=extra_metadata_field, + ) + + assert json.loads(create_route.calls.last.request.content) == { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc123", + "metadata": {"key1": "value1", "key2": "value2"}, + } + assert create_batch_response.id == "batch_abc123" + assert create_batch_response.endpoint == "/v1/chat/completions" + assert create_batch_response.input_file_id == file_obj.id + + await asyncio.wait_for(recorder.logged.wait(), timeout=10) + assert recorder.payload is not None + standard_logging_object: Final = recorder.payload + 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: Final = await litellm.aretrieve_batch( + batch_id=create_batch_response.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieve_route.called + assert retrieved_batch.id == create_batch_response.id + + list_batches: Final = await litellm.alist_batches(custom_llm_provider="openai", limit=2, api_key="fake-key") + assert list_batches_route.calls.last.request.url.params["limit"] == "2" + assert [batch.id for batch in list_batches.data] == ["batch_abc123"] + + file_content: Final = await litellm.afile_content( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert file_content.content == b'{"custom_id": "request-1"}\n' + + retrieved_file: Final = await litellm.afile_retrieve( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieved_file.id == file_obj.id + + delete_file_response: Final = await litellm.afile_delete( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert delete_file_response.id == file_obj.id + + all_files_list: Final = await litellm.afile_list(custom_llm_provider="openai", api_key="fake-key") + assert list_files_route.called + assert [file.id for file in all_files_list.data] == ["file-abc123"] + + cancel_batch_response: Final = await litellm.acancel_batch( + batch_id=create_batch_response.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert cancel_route.called + assert cancel_batch_response.id == create_batch_response.id + + +@pytest.mark.asyncio +async def test_delete_batch_output_file(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + batch_with_output: Final = { + **_OPENAI_BATCH_JSON, + "status": "completed", + "output_file_id": "file-output123", + } + respx_mock.get("https://api.openai.com/v1/batches/batch_abc123").mock( + return_value=httpx.Response(200, json=batch_with_output) + ) + delete_route: Final = respx_mock.delete("https://api.openai.com/v1/files/file-output123").mock( + return_value=httpx.Response(200, json={"id": "file-output123", "object": "file", "deleted": True}) + ) + + batch: Final = await litellm.aretrieve_batch( + batch_id="batch_abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert batch.output_file_id == "file-output123" + + delete_response: Final = await litellm.afile_delete( + file_id=batch.output_file_id, custom_llm_provider="openai", api_key="fake-key" + ) + assert delete_route.call_count == 1 + assert delete_response.id == "file-output123" + assert delete_response.deleted is True diff --git a/tests/unit/files/test_main.py b/tests/unit/files/test_main.py index cb70b39f4d5..e91951fb5ba 100644 --- a/tests/unit/files/test_main.py +++ b/tests/unit/files/test_main.py @@ -1,3 +1,4 @@ +from types import MappingProxyType from typing import Final from urllib.parse import parse_qs, urlparse @@ -121,3 +122,68 @@ async def test_afile_retrieve_rejects_a_provider_file_without_its_size(): assert exc_info.value.title == "OpenAIFileObject" assert [error["loc"] for error in exc_info.value.errors()] == [("bytes",)] + + +_FILE_BODY: Final = b'{"prompt": "Hello", "completion": "Hi"}' +_FINE_TUNE_FILE_JSON: Final = MappingProxyType( + { + "id": "file-abc123", + "object": "file", + "bytes": len(_FILE_BODY), + "created_at": 1699000000, + "filename": "mydata.jsonl", + "purpose": "fine-tune", + } +) + + +@pytest.mark.asyncio +async def test_openai_file_operations_roundtrip(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + files_route: Final = respx_mock.post("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json=dict(_FINE_TUNE_FILE_JSON)) + ) + list_route: Final = respx_mock.get("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_FINE_TUNE_FILE_JSON)]}) + ) + retrieve_route: Final = respx_mock.get("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json=dict(_FINE_TUNE_FILE_JSON)) + ) + content_route: Final = respx_mock.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response(200, content=_FILE_BODY) + ) + delete_route: Final = respx_mock.delete("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json={"id": "file-abc123", "object": "file", "deleted": True}) + ) + + uploaded: Final = await litellm.acreate_file( + file=("mydata.jsonl", _FILE_BODY), purpose="fine-tune", custom_llm_provider="openai", api_key="fake-key" + ) + assert files_route.call_count == 1 + upload_body: Final = files_route.calls.last.request.content + assert b'name="purpose"\r\n\r\nfine-tune' in upload_body + assert b'filename="mydata.jsonl"' in upload_body + assert _FILE_BODY in upload_body + assert uploaded.id == "file-abc123" + + listed: Final = await litellm.afile_list(custom_llm_provider="openai", api_key="fake-key") + assert list_route.call_count == 1 + assert [file.id for file in listed.data] == ["file-abc123"] + + retrieved: Final = await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieve_route.call_count == 1 + assert retrieved.filename == "mydata.jsonl" + assert retrieved.purpose == "fine-tune" + + content: Final = await litellm.afile_content( + file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert content_route.call_count == 1 + assert content.content == _FILE_BODY + + deleted: Final = await litellm.afile_delete(file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key") + assert delete_route.call_count == 1 + assert deleted.id == "file-abc123" + assert deleted.deleted is True diff --git a/tests/unit/images/test_image_edit.py b/tests/unit/images/test_image_edit.py new file mode 100644 index 00000000000..583474cba90 --- /dev/null +++ b/tests/unit/images/test_image_edit.py @@ -0,0 +1,152 @@ +import asyncio +import io +from collections.abc import Iterator, Mapping +from datetime import datetime +from typing import Final + +import httpx +import pytest +import respx +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict, override + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.utils import ImageResponse + +_PNG_SIGNATURE: Final = b"\x89PNG\r\n\x1a\n" +_FIRST_IMAGE: Final = _PNG_SIGNATURE + b"first-reference-image" +_SECOND_IMAGE: Final = _PNG_SIGNATURE + b"second-reference-image" +_EDITED_IMAGE_B64: Final = ( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg==" +) +_TEXT_TOKENS: Final = 50 +_IMAGE_TOKENS: Final = 50 +_OUTPUT_TOKENS: Final = 1000 +_EDIT_RESPONSE: Final = { + "created": 1589478378, + "data": [{"b64_json": _EDITED_IMAGE_B64}], + "usage": { + "total_tokens": _TEXT_TOKENS + _IMAGE_TOKENS + _OUTPUT_TOKENS, + "input_tokens": _TEXT_TOKENS + _IMAGE_TOKENS, + "input_tokens_details": {"image_tokens": _IMAGE_TOKENS, "text_tokens": _TEXT_TOKENS}, + "output_tokens": _OUTPUT_TOKENS, + }, +} + + +class _LoggedImageEdit(TypedDict): + model: ReadOnly[str] + custom_llm_provider: ReadOnly[str] + response_cost: ReadOnly[float] + + +_LOGGED_IMAGE_EDIT: Final = TypeAdapter(_LoggedImageEdit) + + +class _SuccessLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payload: _LoggedImageEdit | None = None + self.logged: Final = asyncio.Event() + + @override + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self.payload = _LOGGED_IMAGE_EDIT.validate_python(kwargs.get("standard_logging_object")) + self.logged.set() + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +def _multipart_image_parts(request: httpx.Request) -> tuple[bytes, ...]: + body: Final = request.read() + boundary: Final = request.headers["content-type"].split("boundary=", 1)[1].encode() + parts: Final = body.split(b"--" + boundary) + return tuple(part.split(b"\r\n\r\n", 1)[1].removesuffix(b"\r\n") for part in parts if b'name="image[]"' in part) + + +@pytest.mark.asyncio +async def test_openai_image_edit_accepts_bytesio_images(respx_mock: respx.MockRouter, httpx_transport: None) -> None: + route: Final = respx_mock.post("https://api.openai.com/v1/images/edits").mock( + return_value=httpx.Response(200, json=_EDIT_RESPONSE) + ) + + result: Final = await litellm.aimage_edit( + prompt="combine the reference images", + model="gpt-image-1", + image=[io.BytesIO(_FIRST_IMAGE), io.BytesIO(_SECOND_IMAGE)], + api_key="fake-key", + ) + + assert isinstance(result, ImageResponse) + assert result.data is not None and result.data[0].b64_json == _EDITED_IMAGE_B64 + assert route.call_count == 1 + assert _multipart_image_parts(route.calls[0].request) == (_FIRST_IMAGE, _SECOND_IMAGE) + + +@pytest.mark.asyncio +async def test_openai_image_edit_accepts_mixed_bytes_and_bytesio( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + route: Final = respx_mock.post("https://api.openai.com/v1/images/edits").mock( + return_value=httpx.Response(200, json=_EDIT_RESPONSE) + ) + + result: Final = await litellm.aimage_edit( + prompt="Create a cohesive artistic style across all images", + model="gpt-image-1", + image=[_FIRST_IMAGE, io.BytesIO(_SECOND_IMAGE)], + api_key="fake-key", + ) + + assert isinstance(result, ImageResponse) + assert result.data is not None and len(result.data) == 1 + assert result.data[0].b64_json == _EDITED_IMAGE_B64 + assert route.call_count == 1 + assert _multipart_image_parts(route.calls[0].request) == (_FIRST_IMAGE, _SECOND_IMAGE) + + +@pytest.mark.asyncio +async def test_azure_image_edit_logs_deployment_model_and_positive_cost( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + logger: Final = _SuccessLogger() + monkeypatch.setattr(litellm, "callbacks", [logger]) + route: Final = respx_mock.post( + url__startswith="https://fake.openai.azure.com/openai/deployments/CUSTOM_AZURE_DEPLOYMENT_NAME/images/edits" + ).mock(return_value=httpx.Response(200, json=_EDIT_RESPONSE)) + + result: Final = await litellm.aimage_edit( + prompt="combine the reference images", + model="azure/CUSTOM_AZURE_DEPLOYMENT_NAME", + base_model="azure/gpt-image-1", + image=[_FIRST_IMAGE, _SECOND_IMAGE], + api_key="fake-key", + api_base="https://fake.openai.azure.com", + api_version="2025-04-01-preview", + ) + await asyncio.wait_for(logger.logged.wait(), timeout=10) + + assert isinstance(result, ImageResponse) + assert route.call_count == 1 + payload: Final = logger.payload + assert payload is not None + assert payload["model"] == "CUSTOM_AZURE_DEPLOYMENT_NAME" + assert payload["custom_llm_provider"] == "azure" + pricing: Final = litellm.model_cost["azure/gpt-image-1"] + expected_cost: Final = ( + _TEXT_TOKENS * pricing["input_cost_per_token"] + + _IMAGE_TOKENS * pricing["input_cost_per_image_token"] + + _OUTPUT_TOKENS * pricing["output_cost_per_image_token"] + ) + assert expected_cost > 0 + assert payload["response_cost"] == pytest.approx(expected_cost) + assert result._hidden_params["response_cost"] == pytest.approx(expected_cost) # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py new file mode 100644 index 00000000000..c55fa476c06 --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py @@ -0,0 +1,619 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import struct +import uuid +from collections.abc import AsyncIterable, Mapping +from typing import Final +from zlib import crc32 + +import httpx +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.router import Router +from litellm.types.llms.anthropic import ( + AnthropicMessagesTextParam, + AnthropicMessagesTool, + AnthropicMessagesUserMessageParam, + AnthropicToolSearchToolRegex, +) +from litellm.types.utils import StandardLoggingPayload + +_ALIAS: Final = "claude-special-alias" +_ANTHROPIC_MODEL: Final = "claude-haiku-4-5-20251001" +_BEDROCK_MODEL: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +_BEDROCK_SONNET: Final = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" +_OPENAI_MODEL: Final = "openai/gpt-4.1-mini" +_JOKE: Final = "why did the chicken cross the road" +_PROMPT: Final = "Hello, can you tell me a short joke?" +_ANTHROPIC_URL: Final = r".*api\.anthropic\.com/v1/messages.*" +_OPENAI_RESPONSES_URL: Final = r".*api\.openai\.com/v1/responses.*" +_BEDROCK_INVOKE_URL: Final = r".*bedrock-runtime.*/invoke$" +_BEDROCK_INVOKE_STREAM_URL: Final = r".*bedrock-runtime.*/invoke-with-response-stream$" +_BEDROCK_CONVERSE_URL: Final = r".*bedrock-runtime.*/converse$" +_BEDROCK_CONVERSE_STREAM_URL: Final = r".*bedrock-runtime.*/converse-stream$" + + +def _anthropic_body(model: str = _ANTHROPIC_MODEL, msg_id: str = "msg_1") -> Mapping[str, object]: + return { + "id": msg_id, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": _JOKE}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20}, + } + + +_CONVERSE_BODY: Final = { + "output": {"message": {"role": "assistant", "content": [{"text": _JOKE}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30}, +} + +_OPENAI_BODY: Final = { + "id": "resp_1", + "object": "response", + "status": "completed", + "created_at": 1700000000, + "model": "gpt-4.1-mini", + "output": [ + { + "type": "message", + "id": "msg_out_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _JOKE, "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30}, +} + +_STREAM_EVENTS: Final = ( + { + "type": "message_start", + "message": { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": _ANTHROPIC_MODEL, + "content": [], + "usage": { + "input_tokens": 10, + "output_tokens": 1, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + }, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JOKE}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 20}}, + {"type": "message_stop"}, +) + + +class _RecordingLogger(CustomLogger): + def __init__(self, messages: list[AnthropicMessagesUserMessageParam]) -> None: + super().__init__() + self.messages: Final = messages + self.payloads: tuple[StandardLoggingPayload, ...] = () + self.received: Final = asyncio.Event() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + payload: Final = kwargs.get("standard_logging_object") + if payload is not None and payload["messages"] == self.messages: + self.payloads = (*self.payloads, payload) + self.received.set() + + +def _unique_messages() -> list[AnthropicMessagesUserMessageParam]: + return [{"role": "user", "content": f"{_PROMPT} {uuid.uuid4().hex}"}] + + +def _router(model_name: str, model: str, **router_kwargs: object) -> Router: + return Router( + model_list=[{"model_name": model_name, "litellm_params": {"model": model, "api_key": "fake-key"}}], + **router_kwargs, + ) + + +def _event_frame(event_type: str, payload: Mapping[str, object]) -> bytes: + def header(name: str, value: str) -> bytes: + name_b: Final = name.encode() + value_b: Final = value.encode() + return ( + struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b + ) + + payload_b: Final = json.dumps(payload).encode() + headers_b: Final = ( + header(":event-type", event_type) + + header(":content-type", "application/json") + + header(":message-type", "event") + ) + prelude: Final = struct.pack("!II", 16 + len(headers_b) + len(payload_b), len(headers_b)) + prelude_crc: Final = crc32(prelude) & 0xFFFFFFFF + message: Final = struct.pack("!I", prelude_crc) + headers_b + payload_b + return prelude + message + struct.pack("!I", crc32(message, prelude_crc) & 0xFFFFFFFF) + + +def _invoke_stream_body() -> bytes: + return b"".join( + _event_frame("chunk", {"bytes": base64.b64encode(json.dumps(event).encode()).decode()}) + for event in _STREAM_EVENTS + ) + + +def _anthropic_sse_body() -> bytes: + return "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in _STREAM_EVENTS).encode() + + +def _converse_stream_body() -> bytes: + return ( + _event_frame("messageStart", {"role": "assistant"}) + + _event_frame("contentBlockDelta", {"contentBlockIndex": 0, "delta": {"text": _JOKE}}) + + _event_frame("contentBlockStop", {"contentBlockIndex": 0}) + + _event_frame("messageStop", {"stopReason": "end_turn"}) + + _event_frame( + "metadata", + { + "usage": { + "inputTokens": 10, + "outputTokens": 20, + "totalTokens": 530, + "cacheReadInputTokens": 500, + "cacheWriteInputTokens": 0, + }, + "metrics": {"latencyMs": 10}, + }, + ) + ) + + +def _set_fake_aws_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "fake") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "fake") + monkeypatch.setenv("AWS_REGION_NAME", "us-east-1") + + +def _sse_events(raw: str) -> tuple[Mapping[str, object], ...]: + return tuple(json.loads(line[len("data: ") :]) for line in raw.splitlines() if line.startswith("data: ")) + + +async def _stream_events(stream: object) -> tuple[Mapping[str, object], ...]: + assert isinstance(stream, AsyncIterable), type(stream) + chunks: Final = [chunk async for chunk in stream] + raw: Final = "".join(chunk.decode() for chunk in chunks if isinstance(chunk, bytes)) + dict_events: Final = tuple(chunk for chunk in chunks if isinstance(chunk, Mapping)) + return _sse_events(raw) + dict_events + + +async def _wait_for_payload(recorder: _RecordingLogger) -> None: + await asyncio.wait_for(recorder.received.wait(), timeout=30.0) + + +def _assert_anthropic_message(response: object, model: str) -> None: + assert isinstance(response, dict), type(response) + assert response["type"] == "message" + assert response["role"] == "assistant" + assert response["model"] == model + assert isinstance(response["id"], str) and response["id"] + block: Final = response["content"][0] + assert isinstance(block, dict), type(block) + assert block["type"] == "text" + assert block["text"] == _JOKE + + +@pytest.fixture(autouse=True) +def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_non_streaming_posts_anthropic_body(respx_mock): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body()) + ) + response: Final = await _router(_ALIAS, _ANTHROPIC_MODEL).aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model=_ALIAS, + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == _ANTHROPIC_MODEL + assert sent["max_tokens"] == 100 + assert sent["messages"] == [{"role": "user", "content": _PROMPT}] + _assert_anthropic_message(response, _ANTHROPIC_MODEL) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_latency_routing_forwards_user_id(respx_mock): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body()) + ) + response: Final = await _router( + _ALIAS, _ANTHROPIC_MODEL, routing_strategy="latency-based-routing" + ).aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model=_ALIAS, + max_tokens=100, + metadata={"user_id": "hello"}, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == _ANTHROPIC_MODEL + assert sent["metadata"] == {"user_id": "hello"} + _assert_anthropic_message(response, _ANTHROPIC_MODEL) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_falls_back_to_bedrock_after_anthropic_401(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + anthropic_route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response( + 401, json={"type": "error", "error": {"type": "authentication_error", "message": "invalid x-api-key"}} + ) + ) + bedrock_route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body(model=_BEDROCK_SONNET, msg_id="msg_bedrock")) + ) + router: Final = Router( + model_list=[ + { + "model_name": "anthropic/claude-opus-4-7", + "litellm_params": {"model": "anthropic/claude-opus-4-7", "api_key": "bad-key"}, + }, + {"model_name": f"bedrock/{_BEDROCK_SONNET}", "litellm_params": {"model": f"bedrock/{_BEDROCK_SONNET}"}}, + ], + fallbacks=[{"anthropic/claude-opus-4-7": [f"bedrock/{_BEDROCK_SONNET}"]}], + ) + response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model="anthropic/claude-opus-4-7", + max_tokens=100, + metadata={"user_id": "hello"}, + ) + assert anthropic_route.call_count == 1 + assert anthropic_route.calls.last.request.headers["x-api-key"] == "bad-key" + assert bedrock_route.call_count == 1 + assert "authorization" in bedrock_route.calls.last.request.headers + _assert_anthropic_message(response, _BEDROCK_SONNET) + assert response["id"] == "msg_bedrock" + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_bedrock_converse_and_invoke(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + converse_route: Final = respx_mock.post(url__regex=_BEDROCK_CONVERSE_URL).mock( + return_value=httpx.Response(200, json=_CONVERSE_BODY) + ) + invoke_route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body(model=_BEDROCK_SONNET)) + ) + converse_model: Final = f"bedrock/converse/{_BEDROCK_SONNET}" + invoke_model: Final = f"bedrock/{_BEDROCK_SONNET}" + router: Final = Router( + model_list=[ + {"model_name": converse_model, "litellm_params": {"model": converse_model}}, + {"model_name": invoke_model, "litellm_params": {"model": invoke_model}}, + ] + ) + converse_response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], model=converse_model, max_tokens=100 + ) + invoke_response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], model=invoke_model, max_tokens=100 + ) + assert converse_route.call_count == 1 + assert invoke_route.call_count == 1 + converse_request: Final = converse_route.calls.last.request + invoke_request: Final = invoke_route.calls.last.request + assert "authorization" in converse_request.headers + assert "authorization" in invoke_request.headers + assert json.loads(converse_request.read())["messages"] == [{"role": "user", "content": [{"text": _PROMPT}]}] + assert json.loads(invoke_request.read())["messages"] == [{"role": "user", "content": _PROMPT}] + _assert_anthropic_message(converse_response, _BEDROCK_SONNET) + _assert_anthropic_message(invoke_response, _BEDROCK_SONNET) + + +def test_sync_openai_bridge_anthropic_messages_returns_content_blocks(respx_mock): + route: Final = respx_mock.post(url__regex=_OPENAI_RESPONSES_URL).mock( + return_value=httpx.Response(200, json=_OPENAI_BODY) + ) + response: Final = litellm.anthropic.messages.create( + messages=[{"role": "user", "content": _PROMPT}], + model=_OPENAI_MODEL, + max_tokens=100, + api_key="fake-key", + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == "gpt-4.1-mini" + assert isinstance(response, dict) + assert response["content"][0]["text"] == _JOKE + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "body", "expected_model"), + [ + pytest.param(_ANTHROPIC_MODEL, _ANTHROPIC_URL, _anthropic_body(), _ANTHROPIC_MODEL, id="anthropic"), + pytest.param( + _BEDROCK_MODEL, + _BEDROCK_INVOKE_URL, + _anthropic_body(model="us.anthropic.claude-haiku-4-5-20251001-v1:0"), + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + id="bedrock-invoke", + ), + pytest.param(_OPENAI_MODEL, _OPENAI_RESPONSES_URL, _OPENAI_BODY, "gpt-4.1-mini", id="openai-bridge"), + ], +) +async def test_acreate_non_streaming_returns_dict_content_blocks( + respx_mock, monkeypatch, model: str, url: str, body: Mapping[str, object], expected_model: str +): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=url).mock(return_value=httpx.Response(200, json=body)) + response: Final = await litellm.anthropic.messages.acreate( + messages=[{"role": "user", "content": _PROMPT}], + model=model, + max_tokens=100, + api_key="fake-key", + ) + assert route.call_count == 1 + _assert_anthropic_message(response, expected_model) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_non_streaming_logs_usage_model_and_cost(respx_mock, monkeypatch): + messages: Final = _unique_messages() + recorder: Final = _RecordingLogger(messages) + monkeypatch.setattr(litellm, "callbacks", [recorder]) + respx_mock.post(url__regex=_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_anthropic_body())) + response: Final = await _router(_ALIAS, _ANTHROPIC_MODEL).aanthropic_messages( + messages=messages, model=_ALIAS, max_tokens=100 + ) + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + assert payload["status"] == "success" + assert payload["messages"] == messages + assert payload["response"] is not None + assert payload["model"] == _ANTHROPIC_MODEL + assert payload["model_group"] == _ALIAS + assert payload["response_cost"] > 0 + assert payload["prompt_tokens"] == response["usage"]["input_tokens"] == 10 + assert payload["completion_tokens"] == response["usage"]["output_tokens"] == 20 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "content_type", "body", "expected_model"), + [ + pytest.param( + _ANTHROPIC_MODEL, + _ANTHROPIC_URL, + "text/event-stream", + _anthropic_sse_body(), + _ANTHROPIC_MODEL, + id="anthropic", + ), + pytest.param( + _BEDROCK_MODEL, + _BEDROCK_INVOKE_STREAM_URL, + "application/vnd.amazon.eventstream", + _invoke_stream_body(), + _BEDROCK_MODEL, + id="bedrock-invoke", + ), + ], +) +async def test_router_aanthropic_messages_streaming_logs_usage_model_and_cost( + respx_mock, monkeypatch, model: str, url: str, content_type: str, body: bytes, expected_model: str +): + _set_fake_aws_env(monkeypatch) + messages: Final = _unique_messages() + recorder: Final = _RecordingLogger(messages) + monkeypatch.setattr(litellm, "callbacks", [recorder]) + respx_mock.post(url__regex=url).mock( + return_value=httpx.Response(200, content=body, headers={"content-type": content_type}) + ) + stream: Final = await _router(_ALIAS, model).aanthropic_messages( + messages=messages, model=_ALIAS, max_tokens=100, stream=True + ) + events: Final = await _stream_events(stream) + usages: Final = tuple( + event["usage"] if "usage" in event else event["message"]["usage"] + for event in events + if "usage" in event or (event.get("type") == "message_start" and "usage" in event["message"]) + ) + assert usages, events + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + assert payload["status"] == "success" + assert payload["messages"] == messages + assert payload["response"] is not None + assert payload["model"] == expected_model + assert payload["response_cost"] > 0 + assert payload["prompt_tokens"] == max(usage.get("input_tokens", 0) for usage in usages) == 10 + assert payload["completion_tokens"] == max(usage.get("output_tokens", 0) for usage in usages) == 20 + + +_LARGE_SYSTEM_PROMPT: Final = "This is a comprehensive legal agreement between Party A and Party B. " * 100 + + +def _cached_system() -> list[AnthropicMessagesTextParam]: + return [{"type": "text", "text": _LARGE_SYSTEM_PROMPT, "cache_control": {"type": "ephemeral"}}] + + +@pytest.mark.asyncio +async def test_bedrock_converse_system_prompt_caching_returns_cache_tokens(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=_BEDROCK_CONVERSE_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_CONVERSE_BODY, + "usage": { + "inputTokens": 10, + "outputTokens": 20, + "totalTokens": 580, + "cacheReadInputTokens": 500, + "cacheWriteInputTokens": 50, + }, + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model=f"bedrock/converse/{_BEDROCK_SONNET}", + messages=[{"role": "user", "content": "What are the key terms?"}], + system=_cached_system(), + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["system"] == [{"text": _LARGE_SYSTEM_PROMPT}, {"cachePoint": {"type": "default"}}] + assert isinstance(response, dict) + assert response["usage"]["cache_creation_input_tokens"] == 50 + assert response["usage"]["cache_read_input_tokens"] == 500 + + +@pytest.mark.asyncio +async def test_bedrock_invoke_system_prompt_caching_returns_cache_tokens(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_anthropic_body(model=_BEDROCK_SONNET), + "usage": { + "input_tokens": 10, + "output_tokens": 20, + "cache_creation_input_tokens": 50, + "cache_read_input_tokens": 500, + }, + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model=f"bedrock/invoke/{_BEDROCK_SONNET}", + messages=[{"role": "user", "content": "What are the key terms?"}], + system=_cached_system(), + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["system"] == _cached_system() + assert isinstance(response, dict) + assert response["usage"]["cache_creation_input_tokens"] == 50 + assert response["usage"]["cache_read_input_tokens"] == 500 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "body"), + [ + pytest.param( + f"bedrock/converse/{_BEDROCK_SONNET}", _BEDROCK_CONVERSE_STREAM_URL, _converse_stream_body(), id="converse" + ), + pytest.param( + f"bedrock/invoke/{_BEDROCK_SONNET}", _BEDROCK_INVOKE_STREAM_URL, _invoke_stream_body(), id="invoke" + ), + ], +) +async def test_bedrock_streaming_message_start_carries_cache_usage_fields( + respx_mock, monkeypatch, model: str, url: str, body: bytes +): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=url).mock( + return_value=httpx.Response(200, content=body, headers={"content-type": "application/vnd.amazon.eventstream"}) + ) + stream: Final = await litellm.anthropic.messages.acreate( + model=model, + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": _LARGE_SYSTEM_PROMPT, "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "What are the payment terms in this agreement?"}, + ], + } + ], + max_tokens=100, + stream=True, + ) + events: Final = await _stream_events(stream) + assert route.call_count == 1 + message_starts: Final = [event for event in events if event.get("type") == "message_start"] + assert len(message_starts) == 1, events + usage: Final = message_starts[0]["message"]["usage"] + assert "cache_creation_input_tokens" in usage, usage + assert "cache_read_input_tokens" in usage, usage + + +def _tool_search_tools() -> list[AnthropicToolSearchToolRegex | AnthropicMessagesTool]: + def deferred(name: str, description: str, field: str) -> AnthropicMessagesTool: + return { + "name": name, + "description": description, + "input_schema": {"type": "object", "properties": {field: {"type": "string"}}, "required": [field]}, + "defer_loading": True, + } + + return [ + {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}, + deferred("get_weather", "Get the current weather for a location", "location"), + deferred("get_stock_price", "Get the current stock price for a ticker symbol", "ticker"), + deferred("search_web", "Search the web for information", "query"), + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("prompt", "tool_name", "tool_input"), + [ + pytest.param( + "I need to know the current weather in New York City. Please use the appropriate tool.", + "get_weather", + {"location": "New York, NY"}, + id="discovers-weather-tool", + ), + pytest.param( + "What's the stock price of Apple (AAPL)?", "get_stock_price", {"ticker": "AAPL"}, id="multiple-deferred" + ), + ], +) +async def test_tool_search_forwards_deferred_tools_and_beta_header( + respx_mock, prompt: str, tool_name: str, tool_input: Mapping[str, str] +): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_anthropic_body(model="claude-sonnet-4-5-20250929", msg_id="msg_tool"), + "content": [ + {"type": "tool_use", "id": "toolu_1", "name": tool_name, "input": tool_input}, + ], + "stop_reason": "tool_use", + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model="anthropic/claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": prompt}], + tools=[dict(tool) for tool in _tool_search_tools()], + max_tokens=1024, + api_key="fake-key", + extra_headers={"anthropic-beta": "advanced-tool-use-2025-11-20"}, + ) + request: Final = route.calls.last.request + assert "advanced-tool-use-2025-11-20" in request.headers["anthropic-beta"].split(",") + assert json.loads(request.read())["tools"] == _tool_search_tools() + assert isinstance(response, dict) + assert response["stop_reason"] == "tool_use" + assert [block for block in response["content"] if block["type"] == "tool_use"] == [ + {"type": "tool_use", "id": "toolu_1", "name": tool_name, "input": tool_input} + ] diff --git a/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py b/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py new file mode 100644 index 00000000000..de2cb03bf48 --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +import asyncio +import json +import uuid +from collections.abc import Mapping +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from fastapi import Request, Response + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.proxy._types import UserAPIKeyAuth, hash_token +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import anthropic_proxy_route +from litellm.types.utils import StandardLoggingPayload + +_UPSTREAM: Final = "https://api.anthropic.com/v1/messages" +_MODEL: Final = "claude-sonnet-4-5-20250929" +_VIRTUAL_KEY: Final = "sk-native-passthrough" + + +class _RecordingLogger(CustomLogger): + def __init__(self, message_id: str) -> None: + super().__init__() + self.message_id: Final = message_id + self.payloads: tuple[StandardLoggingPayload, ...] = () + self.received: Final = asyncio.Event() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + payload: Final = kwargs.get("standard_logging_object") + if payload is not None and payload["id"] == self.message_id: + self.payloads = (*self.payloads, payload) + self.received.set() + + +def _proxy_request(body: Mapping[str, object]) -> Request: + request: Final = MagicMock(spec=Request) + request.method = "POST" + request.url = httpx.URL("http://proxy/anthropic/v1/messages") + request.headers = {"content-type": "application/json", "anthropic-version": "2023-06-01"} + request.scope = {"path": "/anthropic/v1/messages", "type": "http", "method": "POST", "headers": []} + request.query_params = {} + request.body = AsyncMock(return_value=json.dumps(body).encode()) + request.json = AsyncMock(return_value=body) + return request + + +def _sse(events: tuple[Mapping[str, object], ...]) -> bytes: + return "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in events).encode() + + +async def _wait_for_payload(recorder: _RecordingLogger) -> None: + GLOBAL_LOGGING_WORKER.start() + await asyncio.wait_for(recorder.received.wait(), timeout=30.0) + + +def _assert_spend_payload( + payload: StandardLoggingPayload, message_id: str, tags: list[str], prompt_tokens: int, completion_tokens: int +) -> None: + assert payload["id"] == message_id + assert payload["call_type"] == "pass_through_endpoint" + assert payload["status"] == "success" + assert payload["custom_llm_provider"] == "anthropic" + assert payload["model"] == _MODEL + assert payload["prompt_tokens"] == prompt_tokens + assert payload["completion_tokens"] == completion_tokens + assert payload["total_tokens"] == prompt_tokens + completion_tokens + assert payload["response_cost"] > 0 + assert payload["request_tags"] == tags + assert payload["cache_hit"] is not True + assert payload["startTime"] <= payload["endTime"] + assert payload["metadata"]["user_api_key_hash"] == hash_token(_VIRTUAL_KEY) + + +@pytest.fixture +def recorder(monkeypatch: pytest.MonkeyPatch) -> _RecordingLogger: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("ANTHROPIC_API_KEY", "synthetic-anthropic-key") + monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + logger: Final = _RecordingLogger(f"msg_{uuid.uuid4().hex}") + monkeypatch.setattr(litellm, "callbacks", [logger]) + monkeypatch.setattr(litellm, "_async_success_callback", [logger]) + return logger + + +@pytest.mark.asyncio +async def test_native_anthropic_passthrough_logs_usage_tags_and_spend(respx_mock, recorder: _RecordingLogger): + tags: Final = ["test-tag-1", "test-tag-2"] + route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response( + 200, + json={ + "id": recorder.message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [{"type": "text", "text": "hello test"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 11, "output_tokens": 7}, + }, + ) + ) + response: Final = await anthropic_proxy_route( + endpoint="v1/messages", + request=_proxy_request( + { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + "litellm_metadata": {"tags": tags}, + } + ), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key=_VIRTUAL_KEY, token=_VIRTUAL_KEY), + ) + assert response.status_code == 200 + assert json.loads(response.body)["id"] == recorder.message_id + outbound: Final = route.calls.last.request + assert outbound.headers["x-api-key"] == "synthetic-anthropic-key" + assert json.loads(outbound.content) == { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + } + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + _assert_spend_payload(payload, recorder.message_id, tags, prompt_tokens=11, completion_tokens=7) + assert payload["api_base"] == _UPSTREAM + + +@pytest.mark.asyncio +async def test_native_anthropic_passthrough_streaming_logs_usage_tags_and_spend(respx_mock, recorder: _RecordingLogger): + tags: Final = ["test-tag-stream-1", "test-tag-stream-2"] + events: Final = ( + { + "type": "message_start", + "message": { + "id": recorder.message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [], + "usage": {"input_tokens": 11, "output_tokens": 1}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello stream test"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 7}}, + {"type": "message_stop"}, + ) + route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, content=_sse(events), headers={"content-type": "text/event-stream"}) + ) + response: Final = await anthropic_proxy_route( + endpoint="v1/messages", + request=_proxy_request( + { + "model": _MODEL, + "max_tokens": 10, + "stream": True, + "messages": [{"role": "user", "content": "Say 'hello stream test' and nothing else"}], + "litellm_metadata": {"tags": tags, "user": "test-user-1"}, + } + ), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key=_VIRTUAL_KEY, token=_VIRTUAL_KEY), + ) + assert response.status_code == 200 + streamed: Final = b"".join([chunk async for chunk in response.body_iterator]) + assert b"hello stream test" in streamed + assert json.loads(route.calls.last.request.content)["stream"] is True + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + _assert_spend_payload(recorder.payloads[0], recorder.message_id, tags, prompt_tokens=11, completion_tokens=7) diff --git a/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py b/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py index 2802016c04a..66001ad86be 100644 --- a/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py +++ b/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py @@ -1,4 +1,11 @@ +import json +from typing import Final + +import httpx import pytest +import respx + +import litellm from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( AmazonNovaCanvasConfig, ) @@ -79,3 +86,30 @@ def test_transform_response_dict_to_openai_response(): assert hasattr(result, "data") assert len(result.data) == 2 assert result.data[0].b64_json == "b64img1" + + +_NOVA_CANVAS_PROMPT: Final = "A serene mountain landscape at sunset with a lake reflection" +_NOVA_CANVAS_IMAGES: Final = ("b64-first-image", "b64-second-image") + + +def test_nova_canvas_image_gen_reports_positive_response_cost(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post( + url__regex=r"^https://bedrock-runtime\.us-east-1\.amazonaws\.com/model/amazon\.nova-canvas-v1(:|%3A)0/invoke$" + ).mock(return_value=httpx.Response(200, json={"images": list(_NOVA_CANVAS_IMAGES)})) + + response: Final = litellm.image_generation( + model="bedrock/amazon.nova-canvas-v1:0", + prompt=_NOVA_CANVAS_PROMPT, + aws_region_name="us-east-1", + aws_access_key_id="fake-access-key", + aws_secret_access_key="fake-secret-key", + ) + + assert route.call_count == 1 + sent: Final = json.loads(route.calls[0].request.content) + assert sent["taskType"] == "TEXT_IMAGE" + assert sent["textToImageParams"]["text"] == _NOVA_CANVAS_PROMPT + assert [image.b64_json for image in response.data] == list(_NOVA_CANVAS_IMAGES) + per_image: Final = litellm.model_cost["amazon.nova-canvas-v1:0"]["output_cost_per_image"] + assert per_image > 0 + assert response._hidden_params["response_cost"] == pytest.approx(len(_NOVA_CANVAS_IMAGES) * per_image) # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params diff --git a/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py b/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py index 474ffe0e519..68419891166 100644 --- a/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py +++ b/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py @@ -1,8 +1,12 @@ +from collections.abc import Iterator from typing import Final from unittest.mock import AsyncMock, MagicMock, patch -import litellm +import httpx import pytest +import respx + +import litellm class TestDuckDuckGoSearchMocked: @@ -223,3 +227,73 @@ class TestDuckDuckGoSearchMocked: urls = [result.url for result in response.results] assert any("India" in url for url in urls) assert any("Indus" in url for url in urls) + + +_DDG_INSTANT_ANSWER: Final = { + "AbstractText": "India is a country in South Asia.", + "AbstractURL": "https://en.wikipedia.org/wiki/India", + "Heading": "India", + "RelatedTopics": [ + {"FirstURL": f"https://example.com/{index}", "Text": f"Topic {index} - snippet text for topic {index}."} + for index in range(10) + ], + "Results": [], + "Type": "D", +} + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_duckduckgo_search_response_structure_and_max_results( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + route: Final = respx_mock.get(url__startswith="https://api.duckduckgo.com/").mock( + return_value=httpx.Response(200, json=_DDG_INSTANT_ANSWER) + ) + + response: Final = await litellm.asearch(query="india", search_provider="duckduckgo", max_results=5) + + assert route.call_count == 1 + sent_params: Final = route.calls[0].request.url.params + assert sent_params["q"] == "india" + assert sent_params["format"] == "json" + assert sent_params["_max_results"] == "5" + assert response.object == "search" + assert [result.url for result in response.results] == [ + "https://en.wikipedia.org/wiki/India", + "https://example.com/0", + "https://example.com/1", + "https://example.com/2", + "https://example.com/3", + ] + first_result: Final = response.results[0] + assert first_result.title == "India" + assert first_result.snippet == "India is a country in South Asia." + assert response.results[1].title == "Topic 0" + assert response.results[1].snippet == "snippet text for topic 0." + assert response._hidden_params["response_cost"] == litellm.model_cost["duckduckgo/search"]["input_cost_per_query"] # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params + + +def test_duckduckgo_sync_search_returns_typed_results_without_a_limit(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.get(url__startswith="https://api.duckduckgo.com/").mock( + return_value=httpx.Response(200, json=_DDG_INSTANT_ANSWER) + ) + + response: Final = litellm.search(query="india", search_provider="duckduckgo") + + assert route.call_count == 1 + assert "_max_results" not in route.calls[0].request.url.params + assert response.object == "search" + assert len(response.results) == 11 + assert all( + isinstance(result.title, str) and isinstance(result.url, str) and isinstance(result.snippet, str) + for result in response.results + ) + assert response.results[-1].url == "https://example.com/9" diff --git a/tests/unit/llms/exa_ai/search/test_transformation.py b/tests/unit/llms/exa_ai/search/test_transformation.py index 5e5eb24f23b..6b556ff241f 100644 --- a/tests/unit/llms/exa_ai/search/test_transformation.py +++ b/tests/unit/llms/exa_ai/search/test_transformation.py @@ -1,11 +1,22 @@ +import json from typing import Final from unittest.mock import Mock import httpx import pytest +import respx +import litellm from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig +EXA_SEARCH_URL: Final = "https://api.exa.ai/search" + + +@pytest.fixture +def exa_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("EXA_API_KEY", "test-exa-key") + monkeypatch.delenv("EXA_API_BASE", raising=False) + @pytest.mark.parametrize( ("content_fields", "expected_snippet"), @@ -31,3 +42,49 @@ def test_transform_search_response_snippet_falls_back_through_content_modes( response: Final = ExaAISearchConfig().transform_search_response(raw_response, logging_obj=Mock()) assert response.results[0].snippet == expected_snippet + + +@pytest.mark.usefixtures("exa_api_key") +def test_search_maps_exa_results_to_search_response(respx_mock: respx.MockRouter) -> None: + respx_mock.post(EXA_SEARCH_URL).respond( + json={ + "results": [ + { + "title": "AI news roundup", + "url": "https://example.com/ai-news", + "text": "The latest in artificial intelligence.", + "publishedDate": "2026-01-15T00:00:00.000Z", + }, + {"title": "Second", "url": "https://example.com/second", "text": "Second text."}, + ] + } + ) + + response: Final = litellm.search(query="artificial intelligence recent news", search_provider="exa_ai") + + assert response.object == "search" + assert isinstance(response.results, list) + assert len(response.results) == 2 + first: Final = response.results[0] + assert first.title == "AI news roundup" + assert first.url == "https://example.com/ai-news" + assert first.snippet == "The latest in artificial intelligence." + assert first.date == "2026-01-15T00:00:00.000Z" + assert response.results[1].url == "https://example.com/second" + + +@pytest.mark.usefixtures("exa_api_key") +def test_search_sends_max_results_as_num_results(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(EXA_SEARCH_URL).respond( + json={"results": [{"title": "ML", "url": "https://example.com/ml", "text": "Machine learning."}]} + ) + + response: Final = litellm.search(query="machine learning", search_provider="exa_ai", max_results=5) + + assert json.loads(route.calls.last.request.content) == { + "query": "machine learning", + "numResults": 5, + "contents": {"text": True}, + } + assert route.calls.last.request.headers["x-api-key"] == "test-exa-key" + assert [result.url for result in response.results] == ["https://example.com/ml"] diff --git a/tests/unit/llms/firecrawl/__init__.py b/tests/unit/llms/firecrawl/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/firecrawl/search/__init__.py b/tests/unit/llms/firecrawl/search/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/firecrawl/search/test_transformation.py b/tests/unit/llms/firecrawl/search/test_transformation.py new file mode 100644 index 00000000000..7d18b40064b --- /dev/null +++ b/tests/unit/llms/firecrawl/search/test_transformation.py @@ -0,0 +1,32 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm + + +def test_firecrawl_search_request_body(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("FIRECRAWL_API_KEY", "test-api-key") + route: Final = respx_mock.post("https://api.firecrawl.dev/v2/search").mock( + return_value=httpx.Response( + 200, + json={ + "success": True, + "data": {"web": [{"title": "Test Title", "url": "https://example.com", "markdown": "Test content"}]}, + }, + ) + ) + + response: Final = litellm.search(query="test query", search_provider="firecrawl", max_results=10, country="US") + + assert route.call_count == 1 + sent: Final = route.calls[0].request + assert sent.headers["authorization"] == "Bearer test-api-key" + body: Final = json.loads(sent.content) + assert body["query"] == "test query" + assert body["limit"] == 10 + assert body["country"] == "US" + assert [(result.title, result.url) for result in response.results] == [("Test Title", "https://example.com")] diff --git a/tests/unit/llms/perplexity/search/__init__.py b/tests/unit/llms/perplexity/search/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py b/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py new file mode 100644 index 00000000000..432ed8248e6 --- /dev/null +++ b/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py @@ -0,0 +1,57 @@ +import json +from typing import Final + +import pytest +import respx + +import litellm + +PERPLEXITY_SEARCH_URL: Final = "https://api.perplexity.ai/search" + + +@pytest.fixture(autouse=True) +def perplexity_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PERPLEXITYAI_API_KEY", "test-perplexity-key") + monkeypatch.delenv("PERPLEXITY_API_BASE", raising=False) + + +def test_search_maps_perplexity_results_to_search_response(respx_mock: respx.MockRouter) -> None: + respx_mock.post(PERPLEXITY_SEARCH_URL).respond( + json={ + "results": [ + { + "title": "AI news roundup", + "url": "https://example.com/ai-news", + "snippet": "The latest in artificial intelligence.", + "date": "2026-01-15", + "last_updated": "2026-01-16", + }, + {"title": "Second", "url": "https://example.com/second", "snippet": "Second snippet."}, + ] + } + ) + + response: Final = litellm.search(query="artificial intelligence recent news", search_provider="perplexity") + + assert response.object == "search" + assert isinstance(response.results, list) + assert len(response.results) == 2 + first: Final = response.results[0] + assert first.title == "AI news roundup" + assert first.url == "https://example.com/ai-news" + assert first.snippet == "The latest in artificial intelligence." + assert first.date == "2026-01-15" + assert first.last_updated == "2026-01-16" + assert response.results[1].snippet == "Second snippet." + + +def test_search_sends_max_results_in_request_body(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(PERPLEXITY_SEARCH_URL).respond( + json={"results": [{"title": "ML", "url": "https://example.com/ml", "snippet": "Machine learning."}]} + ) + + response: Final = litellm.search(query="machine learning", search_provider="perplexity", max_results=5) + + assert json.loads(route.calls.last.request.content) == {"query": "machine learning", "max_results": 5} + assert route.calls.last.request.headers["Authorization"] == "Bearer test-perplexity-key" + assert [result.url for result in response.results] == ["https://example.com/ml"] diff --git a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py index 8da95f839b9..d1cd640e590 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py +++ b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py @@ -1,6 +1,11 @@ +import json +from types import MappingProxyType +from typing import Final from unittest.mock import MagicMock, patch +import httpx import pytest +import respx import litellm from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction @@ -197,3 +202,134 @@ async def test_litellm_cancel_batch_vertex_ai(): assert mock_instance.cancel_batch.called assert response.id == "batch_123" assert response.status == "cancelling" + + +_MOCK_GCS_FILE_RESPONSE: Final = MappingProxyType( + { + "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", + "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: Final = MappingProxyType( + { + "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}, + } +) + + +@pytest.mark.asyncio +async def test_vertex_file_upload_create_and_retrieve_batch( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +): + monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local") + monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project") + monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + mock_creds: Final = MagicMock(token="mock-token", valid=True, expiry=None) + monkeypatch.setattr("google.auth.default", lambda *args, **kwargs: (mock_creds, "mock-project")) + jobs_url: Final = ( + "https://us-central1-aiplatform.googleapis.com/v1/projects/mock-project/locations/us-central1" + "/batchPredictionJobs" + ) + upload_route: Final = respx_mock.post( + url__startswith="https://storage.googleapis.com/upload/storage/v1/b/litellm-local/o" + ).mock(return_value=httpx.Response(200, json=dict(_MOCK_GCS_FILE_RESPONSE))) + create_route: Final = respx_mock.post(jobs_url).mock( + return_value=httpx.Response(200, json=dict(_MOCK_VERTEX_BATCH_RESPONSE)) + ) + retrieve_route: Final = respx_mock.get(f"{jobs_url}/test-batch-id-456").mock( + return_value=httpx.Response(200, json=dict(_MOCK_VERTEX_BATCH_RESPONSE)) + ) + gcs_object_uri: Final = ( + "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/" + "5f7b99ad-9203-4430-98bf-3b45451af4cb" + ) + + file_obj: Final = await litellm.acreate_file( + file=( + "vertex_batch.jsonl", + b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", ' + b'"body": {"model": "gemini-1.5-flash-001", "messages": [{"role": "user", "content": "hi"}]}}\n', + "application/jsonl", + ), + purpose="batch", + custom_llm_provider="vertex_ai", + ) + + assert file_obj.id == gcs_object_uri + assert upload_route.call_count == 1 + upload_request: Final = upload_route.calls.last.request + assert upload_request.url.params["uploadType"] == "media" + assert upload_request.url.params["name"].startswith( + "litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/" + ) + assert upload_request.headers["Content-Type"] == "application/json" + uploaded_row: Final = json.loads(upload_request.content) + assert uploaded_row["request"]["contents"] == [{"role": "user", "parts": [{"text": "hi"}]}] + assert uploaded_row["request"]["labels"]["litellm_custom_id"] == "request-1" + + create_batch_response: Final = 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"}, + ) + + create_body: Final = json.loads(create_route.calls.last.request.content) + assert create_body["inputConfig"] == {"gcsSource": {"uris": [gcs_object_uri]}, "instancesFormat": "jsonl"} + assert create_body["model"] == "publishers/google/models/gemini-1.5-flash-001" + assert create_body["outputConfig"]["predictionsFormat"] == "jsonl" + assert create_body["outputConfig"]["gcsDestination"]["outputUriPrefix"].startswith("gs://litellm-local/") + assert create_batch_response.id == "test-batch-id-456" + assert create_batch_response.input_file_id == gcs_object_uri + + retrieved_batch: Final = await litellm.aretrieve_batch( + batch_id=create_batch_response.id, custom_llm_provider="vertex_ai" + ) + + assert retrieve_route.call_count == 1 + assert retrieved_batch.id == "test-batch-id-456" + assert retrieved_batch.input_file_id == gcs_object_uri diff --git a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py index 8c73b72a65a..8ca0a4a4335 100644 --- a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py +++ b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py @@ -1,9 +1,13 @@ import base64 from typing import Final +import json +from collections.abc import Iterator, Mapping from unittest.mock import MagicMock, Mock, patch import httpx import pytest +import respx +import responses from pydantic import ValidationError import litellm @@ -663,3 +667,97 @@ def test_transform_text_to_speech_response_rejects_malformed_payloads_without_ec ) assert "input_value" not in str(exc_info.value) + + +_SYNTHESIZE_URL: Final = "https://texttospeech.googleapis.com/v1/text:synthesize" +_AUTHORIZED_USER: Final = json.dumps( + { + "type": "authorized_user", + "client_id": "synthetic-client-id", + "client_secret": "synthetic-client-secret", + "refresh_token": "synthetic-refresh-token", + "quota_project_id": "test-project", + } +) +_ASYNC_INPUT: Final = "async hello what llm guardrail do you have" +_UK_VOICE: Final = {"languageCode": "en-UK", "name": "en-UK-Studio-O"} +_UK_AUDIO_CONFIG: Final = {"audioEncoding": "LINEAR22", "speakingRate": "10"} + + +@pytest.fixture +def google_token_endpoint(monkeypatch: pytest.MonkeyPatch) -> Iterator[responses.RequestsMock]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + with responses.RequestsMock(assert_all_requests_are_fired=False) as token_endpoint: + token_endpoint.post( + "https://oauth2.googleapis.com/token", + json={"access_token": "minted-google-token", "expires_in": 3600, "token_type": "Bearer"}, + ) + yield token_endpoint + litellm.in_memory_llm_clients_cache.flush_cache() + + +async def _aspeech_vertex( + respx_mock: respx.MockRouter, speech_input: str, voice_params: Mapping[str, object] +) -> httpx.Request: + route: Final = respx_mock.post(_SYNTHESIZE_URL).mock( + return_value=httpx.Response(200, json={"audioContent": base64.b64encode(b"vertex-audio").decode()}) + ) + response: Final = await litellm.aspeech( + model="vertex_ai/test", + input=speech_input, + vertex_credentials=_AUTHORIZED_USER, + **voice_params, + ) + assert response.content == b"vertex-audio" + assert route.call_count == 1 + sent: Final = route.calls[0].request + assert sent.headers["x-goog-user-project"] == "test-project" + assert sent.headers["authorization"] == "Bearer minted-google-token" + return sent + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_default_voice_posts_synthesize_request( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + sent: Final = await _aspeech_vertex(respx_mock, _ASYNC_INPUT, {}) + + assert json.loads(sent.content) == { + "input": {"text": _ASYNC_INPUT}, + "voice": {"languageCode": "en-US", "name": "en-US-Studio-O"}, + "audioConfig": {"audioEncoding": "LINEAR16", "speakingRate": "1"}, + } + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_forwards_caller_voice_and_audio_config( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + sent: Final = await _aspeech_vertex(respx_mock, _ASYNC_INPUT, {"voice": _UK_VOICE, "audioConfig": _UK_AUDIO_CONFIG}) + + assert json.loads(sent.content) == { + "input": {"text": _ASYNC_INPUT}, + "voice": _UK_VOICE, + "audioConfig": _UK_AUDIO_CONFIG, + } + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_sends_ssml_input( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + ssml: Final = """ + +

Hello, world!

+

This is a test of the text-to-speech API.

+
+ """ + + sent: Final = await _aspeech_vertex(respx_mock, ssml, {"voice": _UK_VOICE, "audioConfig": _UK_AUDIO_CONFIG}) + + assert json.loads(sent.content) == { + "input": {"ssml": ssml}, + "voice": _UK_VOICE, + "audioConfig": _UK_AUDIO_CONFIG, + } diff --git a/tests/unit/passthrough/test_main.py b/tests/unit/passthrough/test_main.py new file mode 100644 index 00000000000..78fa83d126b --- /dev/null +++ b/tests/unit/passthrough/test_main.py @@ -0,0 +1,53 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.vllm.passthrough.transformation import VLLMPassthroughConfig +from litellm.passthrough.main import allm_passthrough_route +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + +_API_BASE: Final = "http://vllm-upstream.test:8090" + + +@pytest.fixture(autouse=True) +def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +def test_hosted_vllm_resolves_vllm_passthrough_config() -> None: + cfg: Final = ProviderConfigManager.get_provider_passthrough_config( + model="hosted_vllm/my-deployment", + provider=LlmProviders.HOSTED_VLLM, + ) + assert isinstance(cfg, VLLMPassthroughConfig) + + +@pytest.mark.asyncio +async def test_allm_passthrough_route_hosted_vllm_sends_normalized_model(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(f"{_API_BASE}/v1/chat/completions").mock( + return_value=httpx.Response(200, json={"ok": True}) + ) + client: Final = AsyncHTTPHandler() + response: Final = await allm_passthrough_route( + method="POST", + endpoint="v1/chat/completions", + model="hosted_vllm/my-deployment", + api_base=_API_BASE, + json={ + "model": "anything", + "messages": [{"role": "user", "content": "Hello"}], + }, + client=client, + ) + assert response.status_code == 200 + assert route.call_count == 1 + outbound: Final = json.loads(route.calls[0].request.content) + assert outbound["model"] == "my-deployment" + assert outbound["messages"] == [{"role": "user", "content": "Hello"}] diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 01be28470d1..f9795197f90 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -4,12 +4,14 @@ Unit tests for Bedrock Guardrails import json import asyncio +from typing import Final from datetime import datetime, timezone import sys from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +import respx from fastapi import HTTPException @@ -7483,3 +7485,142 @@ async def test_should_raise_guardrail_blocked_exception_null_fields(): guardrail._should_raise_guardrail_blocked_exception(response_null_grounding) is False ) + + +_MASKING_GUARDRAIL_URL: Final = "https://bedrock-runtime.us-east-1.amazonaws.com/guardrail/wf0hkdb5x07f/version/DRAFT/apply" + + +def _masking_guardrail() -> BedrockGuardrail: + return BedrockGuardrail( + guardrailIdentifier="wf0hkdb5x07f", + guardrailVersion="DRAFT", + aws_access_key_id="fake-access-key", + aws_secret_access_key="fake-secret-key", + aws_region_name="us-east-1", + ) + + +def _anonymized_reply(masked_texts: tuple[str, ...], entity_types: tuple[str, ...]) -> httpx.Response: + return httpx.Response( + 200, + json={ + "action": "GUARDRAIL_INTERVENED", + "outputs": [{"text": text} for text in masked_texts], + "assessments": [ + { + "sensitiveInformationPolicy": { + "piiEntities": [ + {"type": entity_type, "match": "redacted", "action": "ANONYMIZED"} + for entity_type in entity_types + ] + } + } + ], + }, + ) + + +def _sent_texts(route: respx.Route) -> tuple[str, ...]: + body: Final = json.loads(route.calls[0].request.content) + assert body["source"] == "INPUT" + return tuple(item["text"]["text"] for item in body["content"]) + + +@pytest.mark.asyncio +async def test_during_call_masking_rewrites_pii_in_messages( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + guardrail: Final = _masking_guardrail() + route: Final = respx_mock.post(_MASKING_GUARDRAIL_URL).mock( + return_value=_anonymized_reply( + ( + "Hello, my phone number is {PHONE}", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}", + ), + ("PHONE", "CREDIT_DEBIT_CARD_NUMBER"), + ) + ) + + response: Final = await guardrail.async_moderation_hook( + 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"}, + ], + }, + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + assert route.call_count == 1 + assert _sent_texts(route) == ( + "Hello, my phone number is +1 412 555 1212", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is 1234-5678-9012-3456", + ) + assert response is not None + assert [message["content"] for message in response["messages"]] == [ + "Hello, my phone number is {PHONE}", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}", + ] + + +@pytest.mark.asyncio +async def test_during_call_masking_rewrites_only_pii_block_in_content_list( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + guardrail: Final = _masking_guardrail() + route: Final = respx_mock.post(_MASKING_GUARDRAIL_URL).mock( + return_value=_anonymized_reply( + ( + "Hello, my phone number is {PHONE}", + "what time is it?", + "Hello, how can I help you today?", + "who is the president of the united states?", + ), + ("PHONE",), + ) + ) + + response: Final = await guardrail.async_moderation_hook( + 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?"}, + ], + }, + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + assert route.call_count == 1 + assert _sent_texts(route) == ( + "Hello, my phone number is +1 412 555 1212", + "what time is it?", + "Hello, how can I help you today?", + "who is the president of the united states?", + ) + assert response is not None + messages: Final = response["messages"] + assert messages[0]["content"] == [ + {"type": "text", "text": "Hello, my phone number is {PHONE}"}, + {"type": "text", "text": "what time is it?"}, + ] + assert messages[1]["content"] == "Hello, how can I help you today?" + assert messages[2]["content"] == "who is the president of the united states?" diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py index 0cac6228085..bd760ed27b1 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -8,13 +8,18 @@ import copy import json import os import re +from collections.abc import Iterable, Sequence from contextlib import asynccontextmanager from typing import Final, Literal from unittest.mock import MagicMock, patch +import aiohttp from aiohttp import web +from aiohttp.client_proto import ResponseHandler from aiohttp.test_utils import TestServer import pytest +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict import litellm @@ -27,6 +32,7 @@ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( ) from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType +from litellm.types.proxy.guardrails.guardrail_hooks.presidio import PresidioAnalyzeRequest, PresidioAnalyzeResponseItem from litellm.types.utils import Choices, Delta, Message, ModelResponse, StreamingChoices from litellm.exceptions import BlockedPiiEntityError @@ -4575,3 +4581,221 @@ async def test_presidio_language_configuration_with_per_request_override(): assert analyze_request_default["language"] == "de" assert analyze_request_default["text"] == test_text + + +_CARD_NUMBER: Final = "4111-1111-1111-1111" +_EMAIL: Final = "test@example.com" +_BLOCK_CARD_MASK_EMAIL: Final = { + PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, + PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, +} + + +class _PresidioAnonymizeRequest(TypedDict): + text: ReadOnly[str] + + +class _PresidioAnonymizeReply(TypedDict): + text: ReadOnly[str] + items: ReadOnly[tuple[()]] + + +_ANALYZE_REQUEST: Final = TypeAdapter(PresidioAnalyzeRequest) +_ANONYMIZE_REQUEST: Final = TypeAdapter(_PresidioAnonymizeRequest) + + +def _card_and_email_spans(text: str) -> tuple[PresidioAnalyzeResponseItem, ...]: + return tuple( + PresidioAnalyzeResponseItem( + entity_type=entity_type, + start=text.index(value), + end=text.index(value) + len(value), + score=1.0, + analysis_explanation=None, + ) + for entity_type, value in (("CREDIT_CARD", _CARD_NUMBER), ("EMAIL_ADDRESS", _EMAIL)) + if value in text + ) + + +class _InMemoryPresidioTransport(asyncio.Transport): + def __init__(self, protocol: ResponseHandler, analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> None: + super().__init__() + self._protocol: Final = protocol + self._analyzed: Final = analyzed + self._received = b"" + self._closing = False + + def write(self, data: bytes | bytearray | memoryview) -> None: + self._received += bytes(data) + head, separator, body = self._received.partition(b"\r\n\r\n") + if not separator: + return + request_line, *header_lines = head.decode().split("\r\n") + headers: Final = {name.lower(): value.strip() for name, _, value in (line.partition(":") for line in header_lines)} + if len(body) < int(headers.get("content-length", "0")): + return + self._received = b"" + reply: Final = json.dumps(self._reply(request_line.split(" ")[1], body)).encode() + response: Final = ( + b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: close\r\n" + + f"Content-Length: {len(reply)}\r\n\r\n".encode() + + reply + ) + asyncio.get_running_loop().call_soon(self._protocol.data_received, response) + + def _reply(self, path: str, body: bytes) -> tuple[PresidioAnalyzeResponseItem, ...] | _PresidioAnonymizeReply: + if path == "/analyze": + analyze_request: Final = _ANALYZE_REQUEST.validate_json(body) + self._analyzed.put_nowait(analyze_request) + return _card_and_email_spans(analyze_request.get("text") or "") + assert path == "/anonymize", path + return _PresidioAnonymizeReply(text=_ANONYMIZE_REQUEST.validate_json(body)["text"], items=()) + + def writelines(self, list_of_data: Iterable[bytes | bytearray | memoryview]) -> None: + self.write(b"".join(bytes(chunk) for chunk in list_of_data)) + + def is_closing(self) -> bool: + return self._closing + + def close(self) -> None: + if not self._closing: + self._closing = True + asyncio.get_running_loop().call_soon(self._protocol.connection_lost, None) + + def abort(self) -> None: + self.close() + + def get_extra_info(self, name: str, default: object = None) -> object: + return default + + def can_write_eof(self) -> bool: + return False + + def get_write_buffer_size(self) -> int: + return 0 + + def pause_reading(self) -> None: + return None + + def resume_reading(self) -> None: + return None + + +class _InMemoryPresidioConnector(aiohttp.BaseConnector): + def __init__(self, analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> None: + super().__init__() + self._analyzed: Final = analyzed + + async def _create_connection( # pyright: ignore[reportImplicitOverride] # aiohttp's connector extension point + self, req: aiohttp.ClientRequest, traces: Sequence[object], timeout: aiohttp.ClientTimeout + ) -> ResponseHandler: + protocol: Final = ResponseHandler(asyncio.get_running_loop()) + protocol.connection_made(_InMemoryPresidioTransport(protocol, self._analyzed)) + return protocol + + +def _drain(analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> tuple[PresidioAnalyzeRequest, ...]: + return tuple(analyzed.get_nowait() for _ in range(analyzed.qsize())) + + +def _guardrail_with_in_memory_presidio( + analyzed: asyncio.Queue[PresidioAnalyzeRequest], +) -> OPTIONAL_PresidioPIIMasking: + guardrail: Final = OPTIONAL_PresidioPIIMasking( + pii_entities_config=_BLOCK_CARD_MASK_EMAIL, + presidio_analyzer_api_base="http://presidio-analyzer.test/", + presidio_anonymizer_api_base="http://presidio-anonymizer.test/", + ) + guardrail._http_session = aiohttp.ClientSession(connector=_InMemoryPresidioConnector(analyzed)) + return guardrail + + +@pytest.mark.asyncio +async def test_check_pii_raises_blocked_entity_for_card() -> None: + text: Final = f"My credit card number is {_CARD_NUMBER} and my email is {_EMAIL}" + analyzed: Final[asyncio.Queue[PresidioAnalyzeRequest]] = asyncio.Queue() + guardrail: Final = _guardrail_with_in_memory_presidio(analyzed) + try: + with pytest.raises(BlockedPiiEntityError) as excinfo: + await guardrail.check_pii(text=text, output_parse_pii=True, presidio_config=None, request_data={}) + finally: + await guardrail._close_http_session() + + analyze_requests: Final = _drain(analyzed) + assert len(analyze_requests) == 1 + assert analyze_requests[0].get("text") == text + assert set(analyze_requests[0].get("entities") or ()) == set(_BLOCK_CARD_MASK_EMAIL) + assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD + assert excinfo.value.guardrail_name == guardrail.guardrail_name + + +@pytest.mark.asyncio +async def test_pre_call_hook_raises_blocked_entity_for_card_message( + mock_user_api_key: UserAPIKeyAuth, mock_cache: DualCache +) -> None: + user_text: Final = f"My credit card is {_CARD_NUMBER} and my email is {_EMAIL}." + analyzed: Final[asyncio.Queue[PresidioAnalyzeRequest]] = asyncio.Queue() + guardrail: Final = _guardrail_with_in_memory_presidio(analyzed) + try: + with pytest.raises(BlockedPiiEntityError) as excinfo: + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key, + cache=mock_cache, + data={ + "messages": [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": user_text}, + ], + "model": "gpt-5-mini", + }, + call_type="completion", + ) + finally: + await guardrail._close_http_session() + + analyze_requests: Final = _drain(analyzed) + assert user_text in [payload.get("text") for payload in analyze_requests] + assert all(set(payload.get("entities") or ()) == set(_BLOCK_CARD_MASK_EMAIL) for payload in analyze_requests) + assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD + assert excinfo.value.guardrail_name == guardrail.guardrail_name + + +@pytest.mark.asyncio +async def test_legacy_pii_masking_config_registers_logging_only_guardrail(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PRESIDIO_ANALYZER_API_BASE", "http://localhost:5002") + monkeypatch.setenv("PRESIDIO_ANONYMIZER_API_BASE", "http://localhost:5001") + monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) + monkeypatch.setattr(litellm, "callbacks", []) + + from litellm.proxy.guardrails.init_guardrails import initialize_guardrails + from litellm.types.guardrails import GuardrailEventHooks + + guardrails_config: Final = [ + { + "pii_masking": { + "callbacks": ["presidio"], + "default_on": True, + "logging_only": True, + } + } + ] + assert len(litellm.guardrail_name_config_map) == 0 + initialize_guardrails( + guardrails_config=guardrails_config, + premium_user=True, + config_file_path="", + litellm_settings={"guardrails": guardrails_config}, + ) + assert len(litellm.guardrail_name_config_map) == 1 + + pii_masking_obj: Final = next( + (c for c in litellm.callbacks if isinstance(c, OPTIONAL_PresidioPIIMasking)), + None, + ) + assert pii_masking_obj is not None + assert hasattr(pii_masking_obj, "logging_only") + assert pii_masking_obj.event_hook == GuardrailEventHooks.logging_only + assert pii_masking_obj.should_run_guardrail( + data={}, event_type=GuardrailEventHooks.logging_only + ) diff --git a/tests/unit/proxy/hooks/test_batch_rate_limiter.py b/tests/unit/proxy/hooks/test_batch_rate_limiter.py index a3e60a89c9f..2f71b177705 100644 --- a/tests/unit/proxy/hooks/test_batch_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_batch_rate_limiter.py @@ -6,14 +6,19 @@ batch under a per-minute RPM/TPM budget. Scopes that configure `tpd_limit` are charged against a 24h token window instead of their minute counters. """ +import json import time -from collections.abc import Iterator +from collections.abc import Iterator, Sequence from datetime import datetime, timezone -from typing import Final +from typing import Final, Literal +import httpx import pytest +import respx from fastapi import HTTPException +from typing_extensions import ReadOnly, TypedDict +import litellm from litellm import DualCache from litellm.constants import BATCH_TPD_WINDOW_SECONDS from litellm.proxy._types import UserAPIKeyAuth @@ -22,6 +27,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token +from litellm.types.llms.openai import ChatCompletionUserMessage, LiteLLMBatchCreateRequest class _Clock: @@ -296,3 +302,242 @@ async def test_batch_rate_limit_error_reports_reset_time_in_utc_on_a_non_utc_pro assert exc.value.headers["retry-after"] == str(BATCH_TPD_WINDOW_SECONDS - 3 * 3600) assert exc.value.headers["reset_at"] == "2026-09-14 08:00:00 UTC" assert str(exc.value.detail).endswith("Limit resets at: 2026-09-14 08:00:00 UTC") + + +_BATCH_MODEL: Final = "gpt-3.5-turbo" +_MANAGED_FILE_ID: Final = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxyZWdyZXNzaW9uLXRlc3QtZmlsZQ==" +) + + +class _BatchBody(TypedDict): + model: ReadOnly[str] + messages: ReadOnly[Sequence[ChatCompletionUserMessage]] + + +class _BatchLine(TypedDict): + custom_id: ReadOnly[str] + method: ReadOnly[Literal["POST"]] + url: ReadOnly[Literal["/v1/chat/completions"]] + body: ReadOnly[_BatchBody] + + +@pytest.fixture +def openai_files(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.MockRouter: + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + return respx_mock + + +def _batch_rows(messages: Sequence[str]) -> tuple[_BatchLine, ...]: + return tuple( + _BatchLine( + custom_id=f"request-{i}", + method="POST", + url="/v1/chat/completions", + body=_BatchBody(model=_BATCH_MODEL, messages=(ChatCompletionUserMessage(role="user", content=message),)), + ) + for i, message in enumerate(messages, start=1) + ) + + +def _serve_file(router: respx.MockRouter, file_id: str, rows: Sequence[_BatchLine]) -> respx.Route: + jsonl: Final = "\n".join(json.dumps(row) for row in rows) + return router.get(f"https://api.openai.com/v1/files/{file_id}/content").mock( + return_value=httpx.Response(200, content=jsonl.encode()) + ) + + +def _token_counter_total(rows: Sequence[_BatchLine]) -> int: + return sum(litellm.token_counter(model=row["body"]["model"], messages=row["body"]["messages"]) for row in rows) + + +def _create_batch_data(input_file_id: str) -> LiteLLMBatchCreateRequest: + return LiteLLMBatchCreateRequest(model=_BATCH_MODEL, input_file_id=input_file_id) + + +@pytest.mark.asyncio +async def test_count_input_file_usage_matches_token_counter(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + rows: Final = _batch_rows(("Hello", "Hi there", "Hey")) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + usage: Final = await batch_limiter.count_input_file_usage(file_id="file-abc123", custom_llm_provider="openai") + + assert content_route.call_count == 1 + assert usage.request_count == 3 + assert usage.total_tokens == _token_counter_total(rows) + + +@pytest.mark.asyncio +async def test_batch_rate_limit_single_file_under_and_over_tpm(openai_files: respx.MockRouter): + small_rows: Final = _batch_rows(("Hello", "Hi", "Hey")) + big_rows: Final = _batch_rows( + ("This is a longer message that will consume more tokens from the rate limit. " * 100,) * 3 + ) + _serve_file(openai_files, "file-small", small_rows) + _serve_file(openai_files, "file-big", big_rows) + user_api_key_dict: Final = UserAPIKeyAuth(api_key="test-key-123", tpm_limit=200, rpm_limit=10) + _, _, small_limiter = _make_limiters() + + data_small: Final = dict(_create_batch_data("file-small")) + result: Final = await small_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data_small, + call_type="acreate_batch", + ) + + assert result is data_small + assert data_small["_batch_token_count"] == _token_counter_total(small_rows) + assert data_small["_batch_request_count"] == 3 + + _, _, big_limiter = _make_limiters() + with pytest.raises(HTTPException) as exc_info: + await big_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=dict(_create_batch_data("file-big")), + call_type="acreate_batch", + ) + assert exc_info.value.status_code == 429 + assert "tokens" in str(exc_info.value.detail).lower() + + +@pytest.mark.asyncio +async def test_batch_rate_limit_cumulative_tpm_rejects_second_request(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth(api_key="test-key-456", tpm_limit=200, rpm_limit=10) + first_rows: Final = _batch_rows(("This message has some content to reach about 100 tokens total. " * 4,) * 2) + second_rows: Final = _batch_rows( + ("This is another message with more content to exceed the remaining limit. " * 11,) * 2 + ) + _serve_file(openai_files, "file-1", first_rows) + _serve_file(openai_files, "file-2", second_rows) + first_tokens: Final = _token_counter_total(first_rows) + assert first_tokens <= 200 < first_tokens + _token_counter_total(second_rows) + + first_data: Final = dict(_create_batch_data("file-1")) + first_result: Final = await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=first_data, + call_type="acreate_batch", + ) + assert first_result is first_data + assert first_data["_batch_token_count"] == first_tokens + + with pytest.raises(HTTPException) as exc_info: + await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=dict(_create_batch_data("file-2")), + call_type="acreate_batch", + ) + assert exc_info.value.status_code == 429 + assert "tokens" in str(exc_info.value.detail).lower() + + +@pytest.mark.asyncio +async def test_batch_rate_limiter_reads_a_provider_file_with_user_context(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key="test-key-managed-files", user_id="test-user-abc123", tpm_limit=500, rpm_limit=10 + ) + rows: Final = _batch_rows(("This is a test message for batch rate limiting with managed files. " * 5,) * 3) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + data: Final = dict(_create_batch_data("file-abc123")) + result: Final = await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="acreate_batch", + ) + + assert content_route.call_count == 1 + assert result is data + assert data["_batch_token_count"] == _token_counter_total(rows) + assert data["_batch_request_count"] == 3 + + +@pytest.mark.asyncio +async def test_batch_rate_limiter_without_user_context(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + rows: Final = _batch_rows(("Hello",)) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + usage_without_context: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", custom_llm_provider="openai", user_api_key_dict=None + ) + usage_with_context: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", + custom_llm_provider="openai", + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", user_id="test-user-123"), + ) + + assert content_route.call_count == 2 + assert usage_without_context.request_count == usage_with_context.request_count == 1 + assert usage_without_context.total_tokens == usage_with_context.total_tokens == _token_counter_total(rows) + + +@pytest.mark.asyncio +async def test_managed_file_is_read_through_the_managed_files_hook_with_user_context( + openai_files: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + + from litellm import Router + from litellm.models.managed_files import LiteLLM_ManagedFileTable + from litellm.proxy import proxy_server + from litellm.proxy.openai_files_endpoints.common_utils import is_base64_encoded_unified_file_id + from litellm.proxy.utils import ProxyLogging + + assert is_base64_encoded_unified_file_id(_MANAGED_FILE_ID) + rows: Final = _batch_rows(("Test message for regression",)) + provider_route: Final = _serve_file(openai_files, "file-provider-1", rows) + standard_route: Final = _serve_file(openai_files, "file-abc123", rows) + unrouted_managed_read: Final = _serve_file(openai_files, _MANAGED_FILE_ID, rows) + file_cache: Final = InternalUsageCache(dual_cache=DualCache()) + await file_cache.async_set_cache( + key=_MANAGED_FILE_ID, + value=LiteLLM_ManagedFileTable( + unified_file_id=_MANAGED_FILE_ID, + model_mappings={"deployment-1": "file-provider-1"}, + flat_model_file_ids=["file-provider-1"], + created_by="test-user-regression", + ).model_dump(), + litellm_parent_otel_span=None, + ) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging.proxy_hook_mapping["managed_files"] = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=file_cache, prisma_client=None + ) + router: Final = Router( + model_list=[ + { + "model_name": _BATCH_MODEL, + "litellm_params": {"model": f"openai/{_BATCH_MODEL}", "api_key": "sk-test"}, + "model_info": {"id": "deployment-1"}, + } + ] + ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", proxy_logging) + monkeypatch.setattr(proxy_server, "llm_router", router) + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key="test-key-regression", user_id="test-user-regression", tpm_limit=1000, rpm_limit=10 + ) + + managed_usage: Final = await batch_limiter.count_input_file_usage( + file_id=_MANAGED_FILE_ID, custom_llm_provider="openai", user_api_key_dict=user_api_key_dict + ) + standard_usage: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", custom_llm_provider="openai", user_api_key_dict=user_api_key_dict + ) + + assert provider_route.call_count == 1 + assert standard_route.call_count == 1 + assert not unrouted_managed_read.called + assert managed_usage.request_count == standard_usage.request_count == 1 + assert managed_usage.total_tokens == standard_usage.total_tokens == _token_counter_total(rows) diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py index 50768e48d43..180ece77a6e 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py @@ -3,7 +3,9 @@ from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm.caching.caching import DualCache from litellm.proxy.openai_files_endpoints.common_utils import ( apply_unified_file_ids, get_credentials_for_model, @@ -11,7 +13,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( map_raw_file_ids_to_unified, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError -from litellm.proxy.utils import handle_exception_on_proxy +from litellm.proxy.utils import InternalUsageCache, handle_exception_on_proxy from litellm.types.utils import LiteLLMBatch _RAW_MODEL_WITH_PROMPT: Final = "opus-4.6 Please summarize my medical records\nPatient has diabetes" @@ -514,3 +516,139 @@ class TestCompletedBatchSafeToRetire: ) def test_is_litellm_executed_batch_reads_the_llm_batch_id_prefix(decoded_unified_batch_id: str, executed: bool): assert is_litellm_executed_batch(decoded_unified_batch_id) is executed + + +def _managed_files_hook(prisma_client: MagicMock) -> _PROXY_LiteLLMManagedFiles: + return _PROXY_LiteLLMManagedFiles(internal_usage_cache=InternalUsageCache(DualCache()), prisma_client=prisma_client) + + +def _managed_object_prisma(db_batch_object: MagicMock | None) -> MagicMock: + prisma_client: Final = MagicMock() + prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=db_batch_object) + prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + return prisma_client + + +@pytest.mark.asyncio +async def test_batch_status_sync_from_provider_to_database(caplog: pytest.LogCaptureFixture): + import json + import logging + + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_batch_from_database, + update_batch_in_database, + ) + + batch_id: Final = "batch_test123" + unified_batch_id: Final = "litellm_proxy:test_unified_batch" + stored_row: Final = MagicMock( + unified_object_id=batch_id, + status="validating", + file_object=json.dumps( + { + "id": batch_id, + "object": "batch", + "status": "validating", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-test123", + "completion_window": "24h", + "created_at": 1234567890, + } + ), + ) + prisma_client: Final = _managed_object_prisma(stored_row) + managed_files: Final = _managed_files_hook(prisma_client) + logger: Final = logging.getLogger("test_batch_status_sync") + + db_batch_object, response_batch = await get_batch_from_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + managed_files_obj=managed_files, + prisma_client=prisma_client, + verbose_proxy_logger=logger, + ) + + prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with( + where={"unified_object_id": batch_id} + ) + assert db_batch_object is stored_row + assert isinstance(response_batch, LiteLLMBatch) + assert response_batch.id == batch_id + assert response_batch.status == "validating" + assert response_batch.input_file_id == "file-test123" + + completed: Final = LiteLLMBatch( + id=batch_id, + object="batch", + status="completed", + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + output_file_id="file-output123", + ) + with caplog.at_level(logging.INFO, logger=logger.name): + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=completed, + managed_files_obj=managed_files, + prisma_client=prisma_client, + verbose_proxy_logger=logger, + db_batch_object=db_batch_object, + operation="retrieve", + poller_owns_accounting=False, + ) + + update: Final = prisma_client.db.litellm_managedobjecttable.update + update.assert_awaited_once() + assert update.await_args.kwargs["where"] == {"unified_object_id": batch_id} + written: Final = update.await_args.kwargs["data"] + assert written["status"] == "complete" + assert written["batch_processed"] is True + assert written["updated_at"] is not None + assert json.loads(written["file_object"])["status"] == "completed" + assert json.loads(written["file_object"])["output_file_id"] == "file-output123" + assert f"Updating batch {batch_id} status from validating to completed" in caplog.messages + + +@pytest.mark.asyncio +async def test_batch_cancel_updates_database(caplog: pytest.LogCaptureFixture): + import json + import logging + + from litellm.proxy.openai_files_endpoints.common_utils import update_batch_in_database + + batch_id: Final = "batch_cancel_test" + cancelled: Final = LiteLLMBatch( + id=batch_id, + object="batch", + status="cancelled", + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + cancelled_at=1234567999, + ) + prisma_client: Final = _managed_object_prisma(None) + logger: Final = logging.getLogger("test_batch_cancel_updates_database") + + with caplog.at_level(logging.INFO, logger=logger.name): + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id="litellm_proxy:cancel_test", + response=cancelled, + managed_files_obj=_managed_files_hook(prisma_client), + prisma_client=prisma_client, + verbose_proxy_logger=logger, + operation="cancel", + ) + + update: Final = prisma_client.db.litellm_managedobjecttable.update + update.assert_awaited_once() + assert update.await_args.kwargs["where"] == {"unified_object_id": batch_id} + written: Final = update.await_args.kwargs["data"] + assert written["status"] == "cancelled" + assert "batch_processed" not in written + assert json.loads(written["file_object"])["cancelled_at"] == 1234567999 + assert f"Updating batch {batch_id} status to cancelled after cancel" in caplog.messages diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py index 4ce31cbb449..793921255ec 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2145,6 +2145,54 @@ def test_get_file_content_streams_openai_direct_path( proxy_logging_obj.post_call_failure_hook.assert_not_called() +@respx.mock +def test_get_file_content_forwards_upstream_download_headers(monkeypatch, llm_router: Router): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + setup_proxy_logging_object(monkeypatch, llm_router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + + body: Final = b'{"prompt": "Hello", "completion": "Hi"}' + upstream_filename: Final = "mydata.jsonl" + upstream_request_id: Final = "req_upstream_123" + upstream_route: Final = respx.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response( + status_code=200, + content=body, + headers={ + "content-type": "application/octet-stream", + "content-length": str(len(body)), + "content-disposition": f'attachment; filename="{upstream_filename}"', + "x-request-id": upstream_request_id, + }, + ) + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + try: + response: Final = client.get( + "/v1/files/file-abc123/content", + headers={"Authorization": "Bearer test-key"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert upstream_route.call_count == 1 + assert response.content == body + assert response.headers["content-type"].startswith("application/octet-stream") + assert int(response.headers["content-length"]) == len(response.content) + assert upstream_filename in response.headers["content-disposition"] + assert response.headers["x-request-id"] == upstream_request_id + + def test_get_file_content_routed_provider_skips_streaming_when_resolved_provider_is_not_supported( mocker: MockerFixture, monkeypatch, llm_router: Router ): diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py new file mode 100644 index 00000000000..a862d0fa108 --- /dev/null +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py @@ -0,0 +1,172 @@ +import asyncio +import json +from collections.abc import Mapping +from datetime import datetime +from typing import Final, cast + +import httpx +import pytest +import respx +from fastapi import Request, Response +from starlette.types import Message + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import assemblyai_proxy_route +from litellm.types.utils import StandardLoggingPayload + +_KEY: Final = "synthetic-assemblyai-key" +_UPSTREAM: Final = "https://api.assemblyai.com/v2/transcript" +_AUDIO_URL: Final = "https://assembly.ai/wildfires.mp3" + + +def _is_transcript_log(payload: StandardLoggingPayload, transcript_id: str) -> bool: + response: Final = payload["response"] + return isinstance(response, dict) and response.get("id") == transcript_id + + +class _SuccessRecorder(CustomLogger): + def __init__(self, transcript_id: str, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.transcript_id: Final = transcript_id + self.loop: Final = loop + self.logged: Final = asyncio.Event() + self.payloads: tuple[StandardLoggingPayload, ...] = () + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + payload: Final = cast(StandardLoggingPayload, kwargs["standard_logging_object"]) + self.payloads = (*self.payloads, payload) + if _is_transcript_log(payload, self.transcript_id): + self.loop.call_soon_threadsafe(self.logged.set) + + +def _request(method: str, path: str, body: bytes = b"") -> Request: + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": method, + "scheme": "http", + "path": path, + "raw_path": path.encode(), + "root_path": "", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> Message: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + +def _install_recorder(monkeypatch: pytest.MonkeyPatch, transcript_id: str) -> _SuccessRecorder: + recorder: Final = _SuccessRecorder(transcript_id, asyncio.get_running_loop()) + monkeypatch.setenv("ASSEMBLYAI_API_KEY", _KEY) + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + monkeypatch.setattr(litellm, "success_callback", []) + return recorder + + +async def _wait_for_success_log(recorder: _SuccessRecorder) -> StandardLoggingPayload: + await asyncio.wait_for(recorder.logged.wait(), 30) + logged: Final = tuple( + payload for payload in recorder.payloads if _is_transcript_log(payload, recorder.transcript_id) + ) + assert len(logged) == 1, f"AssemblyAI success log for {recorder.transcript_id} emitted {len(logged)} times" + return logged[0] + + +@pytest.mark.asyncio +async def test_assemblyai_transcribe_create_poll_delete( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + recorder: Final = _install_recorder(monkeypatch, "tr_1") + create_route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "queued", "audio_url": _AUDIO_URL}) + ) + poll_route: Final = respx_mock.get(f"{_UPSTREAM}/tr_1").mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "completed", "text": "fires near town"}) + ) + delete_route: Final = respx_mock.delete(f"{_UPSTREAM}/tr_1").mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "deleted"}) + ) + admin: Final = UserAPIKeyAuth(api_key="sk-master", user_role=LitellmUserRoles.PROXY_ADMIN) + create_body: Final = {"audio_url": _AUDIO_URL, "speech_models": ["universal-2"]} + + create: Final = await assemblyai_proxy_route( + endpoint="v2/transcript", + request=_request("POST", "/assemblyai/v2/transcript", json.dumps(create_body).encode()), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + await _wait_for_success_log(recorder) + poll: Final = await assemblyai_proxy_route( + endpoint="v2/transcript/tr_1", + request=_request("GET", "/assemblyai/v2/transcript/tr_1"), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + delete: Final = await assemblyai_proxy_route( + endpoint="v2/transcript/tr_1", + request=_request("DELETE", "/assemblyai/v2/transcript/tr_1"), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + + assert isinstance(create, Response) and isinstance(poll, Response) and isinstance(delete, Response) + assert create.status_code == 200 + assert json.loads(create.body)["id"] == "tr_1" + assert poll.status_code == 200 + assert json.loads(poll.body)["status"] == "completed" + assert delete.status_code == 200 + assert json.loads(delete.body)["status"] == "deleted" + assert create_route.call_count == 1 + assert json.loads(create_route.calls[0].request.content) == create_body + assert delete_route.call_count == 1 + client_poll: Final = poll_route.calls[-1].request + assert client_poll.headers["authorization"] == _KEY + assert create_route.calls[0].request.headers["authorization"] == _KEY + assert delete_route.calls[0].request.headers["authorization"] == _KEY + + +@pytest.mark.asyncio +async def test_assemblyai_transcribe_with_non_admin_key_logs_key_identity( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + recorder: Final = _install_recorder(monkeypatch, "tr_9") + create_route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, json={"id": "tr_9", "status": "queued", "audio_url": _AUDIO_URL}) + ) + respx_mock.get(f"{_UPSTREAM}/tr_9").mock( + return_value=httpx.Response(200, json={"id": "tr_9", "status": "completed", "text": "fires near town"}) + ) + non_admin: Final = UserAPIKeyAuth( + api_key="hashed-non-admin-key", + token="hashed-non-admin-key", + user_id="non-admin-user", + team_id="non-admin-team", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + response: Final = await assemblyai_proxy_route( + endpoint="v2/transcript", + request=_request("POST", "/assemblyai/v2/transcript", json.dumps({"audio_url": _AUDIO_URL}).encode()), + fastapi_response=Response(), + user_api_key_dict=non_admin, + ) + logged: Final = await _wait_for_success_log(recorder) + + assert isinstance(response, Response) + assert response.status_code == 200, response.body + assert create_route.call_count == 1 + assert create_route.calls[0].request.headers["authorization"] == _KEY + assert logged["metadata"]["user_api_key_hash"] == "hashed-non-admin-key" + assert logged["metadata"]["user_api_key_user_id"] == "non-admin-user" + assert logged["metadata"]["user_api_key_team_id"] == "non-admin-team" + assert logged["custom_llm_provider"] == "assemblyai" diff --git a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py index f6f63aca71d..9c372380749 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py +++ b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py @@ -1,17 +1,28 @@ -from collections.abc import Sequence +import asyncio +import json +import uuid +from collections.abc import Mapping, Sequence from datetime import datetime +from typing import Final, cast from unittest.mock import MagicMock, patch +import httpx import litellm import pytest +import respx +from fastapi import Request, Response +from starlette.types import Message from typing_extensions import NotRequired, ReadOnly, TypedDict +from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import ( VertexAILivePassthroughLoggingHandler, ) +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import vertex_proxy_route from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging -from litellm.types.utils import CostBreakdown, LlmProviders, Usage +from litellm.types.utils import CostBreakdown, LlmProviders, StandardLoggingPayload, Usage class _LiveTurn(TypedDict): @@ -887,3 +898,95 @@ class TestVertexAILivePassthroughErrorHandling: assert "result" in result assert "kwargs" in result + + +class _SuccessRecorder(CustomLogger): + def __init__(self, api_base: str, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.api_base: Final = api_base + self.loop: Final = loop + self.logged: Final = asyncio.Event() + self.payloads: tuple[StandardLoggingPayload, ...] = () + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + payload: Final = cast(StandardLoggingPayload, kwargs["standard_logging_object"]) + self.payloads = (*self.payloads, payload) + if payload["api_base"] == self.api_base: + self.loop.call_soon_threadsafe(self.logged.set) + + +def _generate_content_endpoint(project: str) -> str: + return f"v1/projects/{project}/locations/us-central1/publishers/google/models/gemini-2.0-flash:generateContent" + + +def _vertex_request(endpoint: str, body: bytes) -> Request: + path: Final = f"/vertex_ai/{endpoint}" + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": path, + "raw_path": path.encode(), + "root_path": "", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer client-google-token")], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> Message: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + +@pytest.mark.asyncio +async def test_vertex_ai_generate_content_spendlog( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + endpoint: Final = _generate_content_endpoint(f"p-{uuid.uuid4().hex}") + upstream_url: Final = f"https://us-central1-aiplatform.googleapis.com/{endpoint}" + recorder: Final = _SuccessRecorder(upstream_url, asyncio.get_running_loop()) + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + monkeypatch.setattr(litellm, "success_callback", []) + contents: Final = [{"role": "user", "parts": [{"text": "hi"}]}] + route: Final = respx_mock.post(upstream_url).mock( + return_value=httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"role": "model", "parts": [{"text": "hello vertex"}]}, "finishReason": "STOP"} + ], + "usageMetadata": {"promptTokenCount": 9, "candidatesTokenCount": 6, "totalTokenCount": 15}, + }, + ) + ) + + response: Final = await vertex_proxy_route( + endpoint=endpoint, + request=_vertex_request(endpoint, json.dumps({"contents": contents}).encode()), + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + assert isinstance(response, Response) + assert response.status_code == 200 + call_id: Final = response.headers.get("x-litellm-call-id") + assert call_id + await asyncio.wait_for(recorder.logged.wait(), 30) + assert route.call_count == 1 + outbound: Final = route.calls[0].request + assert outbound.headers["authorization"] == "Bearer client-google-token" + assert json.loads(outbound.content)["contents"] == contents + matching: Final = tuple(payload for payload in recorder.payloads if payload["api_base"] == upstream_url) + assert len(matching) == 1, recorder.payloads + logged: Final = matching[0] + assert logged["id"] == call_id + assert logged["response_cost"] > 0 + assert "gemini" in logged["model"] + assert logged["custom_llm_provider"] == "vertex_ai" + assert logged["prompt_tokens"] == 9 + assert logged["completion_tokens"] == 6 diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index 0df6eecd4b9..c5bb1ed5b61 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -8,18 +8,21 @@ from typing import Any, Final, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest -from typing_extensions import ReadOnly, TypedDict +from pydantic import TypeAdapter +from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm import litellm.constants as litellm_constants import litellm.proxy.spend_tracking.spend_tracking_utils as spend_tracking_utils from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE, LITTELM_CLI_SERVICE_ACCOUNT_NAME, LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, MAX_SPEND_LOG_MODEL_NAME_LENGTH, REDACTED_BY_LITELM_STRING, SESSION_ID_OMITTED_METADATA_KEY, UNKNOWN_MODEL_SPEND_LOG_MODEL from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.llms.base_llm.ocr.transformation import OCRResponse, OCRUsageInfo from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_tracking_utils import ( + _extract_usage_for_ocr_call, _get_messages_for_spend_logs_payload, _get_proxy_server_request_for_spend_logs_payload, _get_request_duration_ms, @@ -5805,3 +5808,236 @@ def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_ ) assert payload["agent_id"] == "header-selected-agent" assert payload["billing_agent_id"] == billing_agent + + +class _OcrUsageInfoDict(TypedDict, total=False): + pages_processed: ReadOnly[int] + doc_size_bytes: ReadOnly[int] + + +class _OcrResponseDict(TypedDict): + id: ReadOnly[str] + object: ReadOnly[str] + model: ReadOnly[str] + usage_info: ReadOnly[NotRequired[_OcrUsageInfoDict]] + + +class _TokenUsageDict(TypedDict): + prompt_tokens: ReadOnly[int] + completion_tokens: ReadOnly[int] + total_tokens: ReadOnly[int] + + +class _CompletionResponseDict(TypedDict): + id: ReadOnly[str] + object: ReadOnly[str] + model: ReadOnly[str] + usage: ReadOnly[_TokenUsageDict] + + +class _LoggingMetadata(TypedDict, total=False): + user_api_key_user_id: ReadOnly[str] + user_api_key_team_id: ReadOnly[str] + + +class _LoggingLitellmParams(TypedDict, total=False): + metadata: ReadOnly[_LoggingMetadata] + + +class _LoggingKwargs(TypedDict): + model: ReadOnly[str] + call_type: ReadOnly[str] + litellm_params: ReadOnly[_LoggingLitellmParams] + response_cost: ReadOnly[float] + + +class _AdditionalUsageValues(TypedDict, total=False): + pages_processed: ReadOnly[int | None] + doc_size_bytes: ReadOnly[int | None] + + +class _SpendLogMetadata(TypedDict): + additional_usage_values: ReadOnly[_AdditionalUsageValues] + + +_SPEND_LOG_METADATA: Final = TypeAdapter(_SpendLogMetadata) +_OCR_LOGGED_AT: Final = datetime.datetime(2026, 1, 1, tzinfo=timezone.utc) +_OCR_RESPONSE_COST: Final = 0.05 + + +def _ocr_logging_kwargs(call_type: str = "ocr", metadata: _LoggingMetadata | None = None) -> _LoggingKwargs: + return _LoggingKwargs( + model="test-ocr-model", + call_type=call_type, + litellm_params=_LoggingLitellmParams() if metadata is None else _LoggingLitellmParams(metadata=metadata), + response_cost=_OCR_RESPONSE_COST, + ) + + +def _ocr_payload( + kwargs: _LoggingKwargs, response_obj: _OcrResponseDict | _CompletionResponseDict | OCRResponse +) -> SpendLogsPayload: + return get_logging_payload( + kwargs=dict(kwargs), + response_obj=response_obj, + start_time=_OCR_LOGGED_AT, + end_time=_OCR_LOGGED_AT, + ) + + +def _additional_usage_values(payload: SpendLogsPayload) -> _AdditionalUsageValues: + return _SPEND_LOG_METADATA.validate_json(payload["metadata"])["additional_usage_values"] + + +class TestExtractUsageForOCRCall: + def test_extract_usage_from_dict(self) -> None: + response_obj_dict: Final = {"usage_info": _OcrUsageInfoDict(pages_processed=5)} + + usage: Final = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 5} + + def test_extract_usage_from_pydantic_model(self) -> None: + response_obj: Final = OCRResponse( + pages=[], + model="test-ocr-model", + usage_info=OCRUsageInfo(pages_processed=10, doc_size_bytes=1024), + ) + + usage: Final = _extract_usage_for_ocr_call(response_obj, response_obj.model_dump()) + + assert usage["prompt_tokens"] == 0 + assert usage["completion_tokens"] == 0 + assert usage["total_tokens"] == 0 + assert usage["pages_processed"] == 10 + assert usage["doc_size_bytes"] == 1024 + + def test_extract_usage_with_object_attributes(self) -> None: + class _SimpleUsageInfo: + def __init__(self, pages_processed: int) -> None: + self.pages_processed = pages_processed + + class _SimpleOCRResponse: + def __init__(self) -> None: + self.usage_info = _SimpleUsageInfo(pages_processed=3) + + usage: Final = _extract_usage_for_ocr_call(_SimpleOCRResponse(), {}) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 3} + + def test_extract_usage_missing_usage_info(self) -> None: + assert _extract_usage_for_ocr_call({}, {}) == {} + + def test_extract_usage_empty_usage_info(self) -> None: + usage: Final = _extract_usage_for_ocr_call({"usage_info": {}}, {"usage_info": {}}) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 0} + + +class TestGetLoggingPayloadOCR: + def test_ocr_call_with_dict_response(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + { + "id": "ocr-test-123", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 7, "doc_size_bytes": 2048}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["request_id"] == "ocr-test-123" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 7 + assert _additional_usage_values(payload)["doc_size_bytes"] == 2048 + + def test_aocr_call_with_pydantic_response(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(call_type="aocr"), + OCRResponse(pages=[], model="test-ocr-model", usage_info=OCRUsageInfo(pages_processed=12)), + ) + + assert payload["call_type"] == "aocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 12 + + def test_ocr_call_missing_usage_info(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + {"id": "ocr-test-789", "object": "ocr", "model": "test-ocr-model"}, + ) + + assert payload["call_type"] == "ocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert "pages_processed" not in _additional_usage_values(payload) + + def test_ocr_call_with_zero_pages(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + { + "id": "ocr-test-000", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 0}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 0 + + def test_non_ocr_call_uses_token_based_usage(self) -> None: + payload: Final = _ocr_payload( + _LoggingKwargs( + model="gpt-5.5", call_type="completion", litellm_params=_LoggingLitellmParams(), response_cost=0.02 + ), + { + "id": "completion-test-123", + "object": "chat.completion", + "model": "gpt-5.5", + "usage": {"prompt_tokens": 50, "completion_tokens": 100, "total_tokens": 150}, + }, + ) + + assert payload["call_type"] == "completion" + assert payload["prompt_tokens"] == 50 + assert payload["completion_tokens"] == 100 + assert payload["total_tokens"] == 150 + assert payload["spend"] == 0.02 + assert "pages_processed" not in _additional_usage_values(payload) + + def test_ocr_with_metadata(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs( + metadata=_LoggingMetadata(user_api_key_user_id="test-user", user_api_key_team_id="test-team") + ), + { + "id": "ocr-metadata-test", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 5, "doc_size_bytes": 1024}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["user"] == "test-user" + assert payload["team_id"] == "test-team" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 5 + assert _additional_usage_values(payload)["doc_size_bytes"] == 1024 diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index 2d0d8740e12..dd7073bb381 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -25,6 +25,8 @@ import litellm from litellm import acompletion, completion from litellm import main as litellm_main from litellm.constants import CONTROL_OPTIONS_KEY +from litellm.caching.base_cache import BaseCache +from litellm.caching.caching import Cache from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs @@ -33,7 +35,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.types.litellm_params import ControlOptions -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, HttpxBinaryResponseContent from litellm.types.prompts.init_prompts import PromptSpec from litellm.types.utils import Delta, ModelResponseStream, StandardCallbackDynamicParams, StreamingChoices, Usage @@ -5831,3 +5833,148 @@ def test_completion_openai_metadata(monkeypatch, enable_preview_features): } else: assert "metadata" not in mock_completion.call_args.kwargs + + +AZURE_TTS_BASE: Final = "https://tts.example.azure.com" +SPEECH_INPUT: Final = "the quick brown fox jumped over the lazy dogs" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_speech_azure_returns_binary_audio( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, sync_mode: bool +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post( + url__regex=rf"{AZURE_TTS_BASE}/openai/deployments/tts/audio/speech\?api-version=.+" + ).mock(return_value=httpx.Response(200, content=b"ID3-fake-mp3")) + + speech_kwargs: Final = { + "model": "azure/tts", + "input": SPEECH_INPUT, + "voice": "alloy", + "api_base": AZURE_TTS_BASE, + "api_key": "fake-key", + "max_retries": 1, + "timeout": 60, + } + response: Final = litellm.speech(**speech_kwargs) if sync_mode else await litellm.aspeech(**speech_kwargs) + + assert route.call_count == 1 + assert route.calls[0].request.headers["api-key"] == "fake-key" + assert json.loads(route.calls[0].request.content) == {"model": "tts", "input": SPEECH_INPUT, "voice": "alloy"} + assert isinstance(response, HttpxBinaryResponseContent) + assert response.content == b"ID3-fake-mp3" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_speech_openai_returns_binary_audio( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, sync_mode: bool +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").mock( + return_value=httpx.Response(200, content=b"ID3-fake-mp3") + ) + + speech_kwargs: Final = { + "model": "openai/tts-1", + "input": SPEECH_INPUT, + "voice": "alloy", + "api_key": "fake-key", + "max_retries": 1, + "timeout": 60, + } + response: Final = litellm.speech(**speech_kwargs) if sync_mode else await litellm.aspeech(**speech_kwargs) + + assert route.call_count == 1 + assert route.calls[0].request.headers["authorization"] == "Bearer fake-key" + assert json.loads(route.calls[0].request.content) == {"model": "tts-1", "input": SPEECH_INPUT, "voice": "alloy"} + assert isinstance(response, HttpxBinaryResponseContent) + assert response.content == b"ID3-fake-mp3" + + +class _SignallingCache(Cache): + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.loop: Final = loop + self.written: Final = asyncio.Event() + + async def async_add_cache( + self, result: object, dynamic_cache_object: BaseCache | None = None, **kwargs: object + ) -> None: + await super().async_add_cache(result, dynamic_cache_object=dynamic_cache_object, **kwargs) + self.loop.call_soon_threadsafe(self.written.set) + + +GETTYSBURG_WAV: Final = ("gettysburg.wav", b"RIFF\x00\x00\x00\x00WAVE-gettysburg", "audio/wav") +EAGLE_WAV: Final = ("eagle.wav", b"RIFF\x00\x00\x00\x00WAVE-eagle", "audio/wav") + + +async def test_transcription_caching_hit_same_file_miss_different_file( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + cache: Final = _SignallingCache(asyncio.get_running_loop()) + monkeypatch.setattr(litellm, "cache", cache) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + side_effect=[ + httpx.Response(200, json={"text": "gettysburg transcript"}), + httpx.Response(200, json={"text": "eagle transcript"}), + ] + ) + + response_1: Final = await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + await asyncio.wait_for(cache.written.wait(), 30) + + response_2: Final = await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + assert response_2._hidden_params["cache_hit"] is True + assert response_2.text == response_1.text == "gettysburg transcript" + + response_3: Final = await litellm.atranscription(model="openai/whisper-1", file=EAGLE_WAV, api_key="fake-key") + assert response_3._hidden_params.get("cache_hit") is not True + assert response_3.text == "eagle transcript" + assert route.call_count == 2 + + +async def test_whisper_log_pre_call_fires_once(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + return_value=httpx.Response(200, json={"text": "hello"}) + ) + + class _PreCallRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.models: tuple[str, ...] = () + + def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None: + self.models = (*self.models, model) + + recorder: Final = _PreCallRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + + await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + + assert recorder.models == ("whisper-1",) + + +@pytest.mark.parametrize("model", ["gpt-4o-mini-transcribe", "gpt-4o-transcribe", "whisper-1"]) +async def test_transcription_model_names_pass_through( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, model: str +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + return_value=httpx.Response(200, json={"text": "hello"}) + ) + + response: Final = await litellm.atranscription( + model=f"openai/{model}", + file=GETTYSBURG_WAV, + api_key="fake-key", + response_format="json", + ) + + assert response._hidden_params["model"] == model + assert response._hidden_params["custom_llm_provider"] == "openai" + assert response.text == "hello" + assert route.call_count == 1 + assert f'name="model"\r\n\r\n{model}\r\n'.encode() in route.calls[0].request.content