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 <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 02:38:12 -07:00 • committed by GitHub
parent 04d97abffb
commit 0e26edfdb9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
60 changed files with 4014 additions and 5401 deletions

View file

@ -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

View file

@ -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 = """
<speak>
<p>Hello, world!</p>
<p>This is a test of the <break strength="medium" /> text-to-speech API.</p>
</speak>
"""
# 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():

View file

@ -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

View file

@ -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")

View file

@ -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")

View file

@ -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)

View file

@ -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?"
)

View file

@ -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 <PERSON>. My number is: <PHONE_NUMBER>",
"items": [
{
"start": 48,
"end": 62,
"entity_type": "PHONE_NUMBER",
"text": "<PHONE_NUMBER>",
"operator": "replace",
},
{
"start": 24,
"end": 32,
"entity_type": "PERSON",
"text": "<PERSON>",
"operator": "replace",
},
],
}
input_b_anonymizer_results = {
"text": "My name is <PERSON>, who are you? Say my name in your response",
"items": [
{
"start": 11,
"end": 19,
"entity_type": "PERSON",
"text": "<PERSON>",
"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
)

View file

@ -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

View file

@ -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}")

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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")

View file

@ -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")

View file

@ -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}")

View file

@ -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}")

View file

@ -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"
)

View file

@ -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)

View file

@ -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")

View file

@ -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

View file

@ -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)

View file

@ -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():

View file

@ -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."
)

View file

@ -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}")

View file

@ -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
)

View file

@ -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

View file

@ -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():
"""

View file

@ -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)}")

View file

@ -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")

View file

@ -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"

View file

@ -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

View file

@ -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():

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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}
]

View file

@ -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)

View file

@ -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

View file

@ -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"

View file

@ -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"]

View file

View file

@ -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")]

View file

@ -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"]

View file

@ -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

View file

@ -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 = """
<speak>
<p>Hello, world!</p>
<p>This is a test of the <break strength="medium" /> text-to-speech API.</p>
</speak>
"""
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,
}

View file

@ -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"}]

View file

@ -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?"

View file

@ -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
)

View file

@ -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)

View file

@ -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

View file

@ -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
):

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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