From 0e26edfdb95c8288a610cba4746af2bc1f2fe312 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 02:38:12 -0700 Subject: [PATCH] test: move 81 legacy live tests in pass-through, spend, batches, audio, search, guardrails, image and ocr dirs offline (#45288) * test: make legacy live tests in spend, batches, openai endpoints and audio dirs offline (partial) * test: migrate wave-1b live tests offline (guardrails, images, ocr, search, openai endpoints) * test: fix wave-1b review items, add responses/ocr integration tests and firecrawl unit test * test: anthropic messages router/bedrock/openai-bridge unit tests for wave-1b nodes * test: anthropic messages logging, prompt-caching and tool-search unit tests; drop migrated base nodes * test: finish pass_through_unit_tests nodes, logging drain fix and mutations * test: migrate anthropic passthrough tests to integration wire tests * test: fix passthrough migration wire spend row lookup and wildcard config * test: migrate hosted vllm and openai file passthrough tests offline * test: move assemblyai and vertex passthrough nodes to in-process unit tests * test: restore unlisted router node and fix logging worker drain in passthrough unit tests * test: drop spend-row BUG skip and sharpen non-streaming skip reason for anthropic messages * test: use public presidio alias after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: restore the batch and file unit tests the migration rewrote away * test: restore full legacy intent in anthropic messages router unit tests Drop the false BUG skip on non-streaming aanthropic_messages logging (success callbacks do fire; the skipped body filtered on the wrong model), assert the logged model_group, messages, cost and usage, cover streaming logging for both Anthropic and Bedrock invoke, assert dict content blocks for Anthropic, Bedrock invoke and the OpenAI bridge, add the Bedrock invoke leg of the router test, fall back from a real 401, test system-prompt caching and streaming message_start cache fields on converse and invoke, send the legacy tool-search tools and beta header, remove the type: ignore and bare dict helper, and add in-process native /anthropic passthrough spend logging tests * test: fix passthrough migration wire tests and drop false native spend BUG skip The native spend row test read rows by the shared master-key digest, so it matched other tests' rows; read each row by its own message id instead and assert tokens, total, spend, tags, provider, api_base and end_user on both the non-streaming and streaming native routes. Merge the streaming test that never checked spend, assert exact tags from litellm_metadata, stop rebinding a Final in a loop, require cost > 0, and add the chat-completions bridge cost case the legacy test covered * test: harden the spend, OCR, image, search and guardrail migration replacements The OCR spend tests now build fresh kwargs per case instead of mutating a shared fixture, and every payload case asserts the exact logged spend. The OCR wire test reads its spend row by request id and checks exact page pricing. Image edit, Nova Canvas, DuckDuckGo, Firecrawl, Bedrock guardrail and Presidio replacements now fake only the provider HTTP boundary (respx or an in-process aiohttp server) and assert the outbound request, so the DuckDuckGo limit, Azure base_model pricing and guardrail masking are proven rather than assumed * test: cover Exa and Perplexity search structure and max_results offline and retire the two base search methods * test: drive batch and file replacements through the provider HTTP boundary and real logging callback Replace monkeypatched litellm.afile_content and AsyncHTTPHandler doubles with respx routes, read batch logging metadata from a registered success callback instead of get_logging_payload, use the real managed-files hook for the GEN-2166 regression, assert outbound request bodies, pin poller ownership explicitly in the migrated DB-sync tests, and require the scripted upstream to be hit in the responses error-status wire tests * test: assert file content download headers pass through the proxy * test: point spend coverage references at the tests that replaced the retired spend job * test: assert passthrough identity, spend and route dispatch from the code under test The AssemblyAI non-admin test asserted metadata it wrote itself and leaked a background poll to the real AssemblyAI host. It now drives assemblyai_proxy_route with a real Request and waits for the success callback for its own transcript id. The Vertex spend test matches its log by call id instead of taking the first event. The OpenAI files wire test hit the native /{provider}/v1/files route; it now calls /openai/files so the passthrough is what forwards the upload and delete. * test: mock only the HTTP boundary in the migrated audio tests Vertex TTS tests no longer replace _ensure_access_token or AsyncHTTPHandler.post; the token comes from a mocked Google OAuth endpoint and the synthesize call from respx. Speech tests assert the outbound body, the transcription cache test polls for the cache write instead of relying on test ordering, and the model pass-through test checks the multipart model field per model. * test: wait on a logger event instead of polling the clock in anthropic messages unit tests Recorders keep payloads in a rebound tuple and set an asyncio.Event; tests await it with asyncio.wait_for instead of a sleep-and-deadline poll loop * test: freeze module-level batch and file response fixtures as Final MappingProxyType * test: fake the presidio analyzer with an in-memory aiohttp connector The blocked-entity tests started an aiohttp TestServer, which binds a local socket. They now hand the guardrail a ClientSession whose connector answers /analyze and /anonymize in process, so no socket is opened and the outbound analyze text and entities are still asserted * test: type anthropic messages router test helpers with LiteLLM's Anthropic TypedDicts Messages, cached system blocks and tool-search tools now use AnthropicMessagesUserMessageParam, AnthropicMessagesTextParam, AnthropicToolSearchToolRegex and AnthropicMessagesTool instead of bare dict shapes; tools are converted to plain dicts only at the acreate call, whose tools parameter is list[dict] * test: type batch limiter helpers with TypedDicts and wait on the logging callback event instead of polling * test: give the migrated OCR, image and presidio helpers precise types OCR spend helpers take ReadOnly TypedDicts for kwargs and responses and use LiteLLM's OCRResponse/OCRUsageInfo instead of local pydantic stand-ins; spend metadata is validated with a TypeAdapter. The presidio fake uses LiteLLM's PresidioAnalyzeRequest/ResponseItem types, and the image-edit logger validates the logged payload instead of storing an untyped dict * test: signal callback and cache events instead of polling Recorders keep tuples and set an asyncio.Event, thread-safely, when the payload for this test's transcript id or upstream URL arrives. The transcription cache test waits on a Cache subclass that signals after async_add_cache. No clock polling or sleeps remain in these tests. * test: assert the batch limiter hook updates the caller's request in place * test: tolerate model-list probes and read native passthrough rows by owned key The router's OpenAI-compatible model-info refresh (litellm/router.py:10710) sends GET /v1/models to configured openai api_bases, so the wire answers it with an empty list and excludes it from the provider-call assertions. Native /anthropic spend rows are now read by a per-request virtual key digest and call_type, then the row's request_id is checked against the message id * test: expect the OCR alias in the proxy response model The proxy restamps every OpenAI-compatible response model to the name the client requested (_override_openai_response_model), so /v1/ocr returns the scenario alias. The upstream model is now checked on the drained request body instead of inside the peer, where a failed assert never reached the test --------- Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .circleci/config.yml | 95 -- tests/audio_tests/test_audio_speech.py | 236 ----- tests/audio_tests/test_whisper.py | 109 --- tests/batches_tests/test_batch_rate_limits.py | 896 ------------------ .../test_openai_batches_and_files.py | 580 ------------ .../SPEND_TRACKING_COVERAGE_MATRIX.md | 2 +- .../test_bedrock_guardrails.py | 100 -- tests/guardrails_tests/test_presidio_pii.py | 211 ----- .../test_bedrock_image_gen_unit_tests.py | 112 --- tests/image_gen_tests/test_image_edits.py | 179 ---- .../test_passthrough_migration_wire.py | 584 ++++++++++++ .../providers/test_ocr_router_wire.py | 69 ++ .../test_openai_passthrough_files_wire.py | 49 + .../test_responses_error_status_wire.py | 71 ++ tests/ocr_tests/test_ocr_matrix.py | 15 - .../test_e2e_openai_responses_api.py | 115 --- .../test_openai_batches_endpoint.py | 463 --------- .../test_openai_files_endpoints.py | 113 --- .../base_anthropic_messages_test.py | 90 -- .../test_anthropic_passthrough.py | 472 --------- .../test_anthropic_passthrough_basic.py | 8 - tests/pass_through_tests/test_assembly_ai.py | 102 -- .../test_hosted_vllm_passthrough.py | 71 -- .../test_openai_assistants_passthrough.py | 23 - tests/pass_through_tests/test_vertex_ai.py | 87 -- ..._anthropic_messages_prompt_caching_test.py | 135 --- ...ase_anthropic_messages_tool_search_test.py | 78 -- .../base_anthropic_unified_messages_test.py | 231 ----- .../test_anthropic_messages_passthrough.py | 287 ------ .../test_bedrock_anthropic_messages_test.py | 48 - tests/search_tests/base_search_unit_tests.py | 59 -- tests/search_tests/test_duckduckgo_search.py | 138 --- tests/search_tests/test_firecrawl_search.py | 42 - .../test_ocr_spend_tracking.py | 296 ------ tests/test_spend_logs.py | 2 +- tests/unit/batches/test_main.py | 218 +++++ tests/unit/files/test_main.py | 66 ++ tests/unit/images/test_image_edit.py | 152 +++ .../test_anthropic_messages_router.py | 619 ++++++++++++ ...hropic_native_passthrough_spend_logging.py | 181 ++++ .../test_amazon_nova_canvas_transformation.py | 34 + .../test_duckduckgo_search_transformation.py | 76 +- .../llms/exa_ai/search/test_transformation.py | 57 ++ tests/unit/llms/firecrawl/__init__.py | 0 tests/unit/llms/firecrawl/search/__init__.py | 0 .../firecrawl/search/test_transformation.py | 32 + tests/unit/llms/perplexity/search/__init__.py | 0 .../test_perplexity_search_transformation.py | 57 ++ .../test_vertex_ai_batch_transformation.py | 136 +++ .../text_to_speech/test_transformation.py | 98 ++ tests/unit/passthrough/test_main.py | 53 ++ .../test_bedrock_guardrails.py | 141 +++ .../guardrail_hooks/test_presidio.py | 224 +++++ .../proxy/hooks/test_batch_rate_limiter.py | 249 ++++- .../test_files_common_utils.py | 140 ++- .../test_files_endpoint.py | 48 + .../test_assembly_passthrough_route.py | 172 ++++ .../test_vertex_ai_live_passthrough.py | 107 ++- .../test_spend_tracking_utils.py | 238 ++++- tests/unit/test_main.py | 149 ++- 60 files changed, 4014 insertions(+), 5401 deletions(-) delete mode 100644 tests/batches_tests/test_batch_rate_limits.py delete mode 100644 tests/batches_tests/test_openai_batches_and_files.py delete mode 100644 tests/guardrails_tests/test_bedrock_guardrails.py delete mode 100644 tests/guardrails_tests/test_presidio_pii.py delete mode 100644 tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py create mode 100644 tests/integration/providers/test_ocr_router_wire.py create mode 100644 tests/integration/providers/test_openai_passthrough_files_wire.py create mode 100644 tests/integration/providers/test_responses_error_status_wire.py delete mode 100644 tests/openai_endpoints_tests/test_e2e_openai_responses_api.py delete mode 100644 tests/openai_endpoints_tests/test_openai_batches_endpoint.py delete mode 100644 tests/openai_endpoints_tests/test_openai_files_endpoints.py delete mode 100644 tests/pass_through_tests/test_anthropic_passthrough.py delete mode 100644 tests/pass_through_tests/test_assembly_ai.py delete mode 100644 tests/pass_through_tests/test_hosted_vllm_passthrough.py delete mode 100644 tests/pass_through_tests/test_openai_assistants_passthrough.py delete mode 100644 tests/search_tests/test_duckduckgo_search.py delete mode 100644 tests/search_tests/test_firecrawl_search.py delete mode 100644 tests/spend_tracking_tests/test_ocr_spend_tracking.py create mode 100644 tests/unit/images/test_image_edit.py create mode 100644 tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py create mode 100644 tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py create mode 100644 tests/unit/llms/firecrawl/__init__.py create mode 100644 tests/unit/llms/firecrawl/search/__init__.py create mode 100644 tests/unit/llms/firecrawl/search/test_transformation.py create mode 100644 tests/unit/llms/perplexity/search/__init__.py create mode 100644 tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py create mode 100644 tests/unit/passthrough/test_main.py create mode 100644 tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py 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