test: make 77 legacy live tests offline in litellm_utils, router_unit and responses dirs (#45298)

* test: make 77 legacy live tests offline in litellm_utils, router_unit and responses dirs

* test: restore lost coverage in litellm_utils offline replacements

Bound live health-check tasks to the concurrency limit so eager task
creation behind a semaphore fails, route the Langfuse trace-id check
through completion -> Logging.get_trace_id with all four metadata
combinations, run the callback dedup checks through acompletion, pass
the dynamic key via litellm_params, cover the env-inferred default
model list, and add the per-provider audio transcription config lookup

* test: restore router coverage lost in the offline move

Add the missing non-stream LIT-3058 header/count-once test, pin the UTC minute on every
usage-counter test, assert router-level client reuse for transcription, and replace the
private selector-attribute checks with routing outcomes per strategy. Assert outbound bodies
for speech, rerank, image, assistants and moderation, add timeouts to event waits, and drop
the router_unit_tests husk files that no longer hold tests

* test: restore sync stream, router sync stream, error event and field-type coverage for migrated responses tests

Add the sync streaming logging and sync Router.responses streaming cases the
offline replacements dropped, port the legacy per-event and response field-type
validation, assert raw headers, the float created_at conversion, MCP tool headers,
the search_context_size input, and add an offline replacement for the in-stream
context-window error event. Remove helpers left dead in the legacy file.

* test: run numpydoc-backed unit tests in GHA and drop a test that pinned a crash

* test: restore sequence number, item id and content part checks in the responses stream validator

* test: match router logging events to the test's own deployment so late events from other tests are ignored

* test: restore the s3 cold storage history test to its original assertions

---------

Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 02:37:55 -07:00 • committed by GitHub
parent 7c7b0ea85b
commit 04d97abffb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 2752 additions and 3051 deletions

View file

@ -143,9 +143,9 @@ jobs:
run: |
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
if [ -z "$RUST_BRIDGE_ARTIFACT" ]; then
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils
else
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --no-install-project
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils --no-install-project
uv pip install --no-deps --python .venv/bin/python rust-bridge-dist/*.whl
cp rust-bridge-dist/litellm/rust_bridge/_native.abi3.so litellm/rust_bridge/_native.abi3.so
uv run --no-sync python -c "import importlib.metadata; import litellm.rust_bridge._native; print(importlib.metadata.version('litellm'))"

View file

@ -6,17 +6,20 @@ import uuid
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Literal, TypeAlias
import anthropic
import httpx
import openai
import pytest
import yaml
from _pytest.mark.structures import ParameterSet
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway, Scenario, 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
from openai.types.responses import ResponseCompletedEvent
from pydantic import JsonValue, TypeAdapter
@ -463,6 +466,26 @@ def _spend_row(*request_ids: str) -> Mapping[str, JsonValue]:
return rows[0]
def _spend_proxy_server_request(request_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
'SELECT proxy_server_request FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(request_id,),
),
lambda found: len(found) == 1,
seconds=70,
)
return _JSON_OBJECT.validate_python(rows[0]["proxy_server_request"])
def _config_storing_prompts(directory: Path) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"]["store_prompts_in_spend_logs"] = True
path: Final = directory / "store-prompts.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _assert_caller_id(endpoint: Endpoint, stream: bool, caller_id: str, response_id: str) -> None:
if endpoint == "messages" and stream:
assert caller_id.startswith("msg_"), caller_id
@ -654,6 +677,74 @@ def test_gemini_own_tool_call_signature_is_still_replayed_on_the_function_call_p
assert _spend_row(rounds[1])["status"] == "success"
def test_responses_previous_response_id_replays_gemini_session_history(
gateway: Gateway, tmp_path: Path
) -> None:
first_prompt: Final = f"gemini-session-first-{uuid.uuid4().hex}"
second_prompt: Final = f"gemini-session-second-{uuid.uuid4().hex}"
first_answer: Final = f"gemini-answer-first-{uuid.uuid4().hex}"
second_answer: Final = f"gemini-answer-second-{uuid.uuid4().hex}"
first_provider_id: Final = f"gemini-response-first-{uuid.uuid4().hex}"
second_provider_id: Final = f"gemini-response-second-{uuid.uuid4().hex}"
calls: Final = itertools.count()
def respond(request: Request) -> Reply:
if next(calls) == 0:
return _reply(first_provider_id, stream=False, parts=({"text": first_answer},))
return _reply(second_provider_id, stream=False, parts=({"text": second_answer},))
with (
wire_server(respond) as wire,
owned_proxy(gateway, tmp_path, {}, config=_config_storing_prompts(tmp_path), workers=1) as prompt_gateway,
prompt_gateway.scenario() as scenario,
):
model: Final = _register(prompt_gateway, scenario, "gemini", wire)
first_response: Final = prompt_gateway.request(
"POST",
"/v1/responses",
{"model": model, "input": first_prompt, **_CACHE_BUST},
)
assert first_response.status_code == 200, first_response.text
first_body: Final = _JSON_OBJECT.validate_json(first_response.content)
first_response_id: Final = str(first_body["id"])
assert _output_text(first_body) == first_answer
first_spend_id: Final = _upstream_response_id("responses", first_response_id)
assert _spend_row(first_spend_id)["status"] == "success"
first_proxy_request: Final = _spend_proxy_server_request(first_spend_id)
assert first_proxy_request["input"] == first_prompt, first_proxy_request
second_response: Final = prompt_gateway.request(
"POST",
"/v1/responses",
{
"model": model,
"input": second_prompt,
"previous_response_id": first_response_id,
**_CACHE_BUST,
},
)
assert second_response.status_code == 200, second_response.text
second_body: Final = _JSON_OBJECT.validate_json(second_response.content)
second_response_id: Final = str(second_body["id"])
assert _output_text(second_body) == second_answer
assert _spend_row(_upstream_response_id("responses", second_response_id))["status"] == "success"
requests: Final = wire.drain()
assert [(request.method, request.target) for request in requests] == [
("POST", _target("gemini", False)),
("POST", _target("gemini", False)),
]
first_provider_body: Final = _JSON_OBJECT.validate_json(requests[0].body)
second_provider_body: Final = _JSON_OBJECT.validate_json(requests[1].body)
assert first_provider_body == _expected_body("responses", _user(first_prompt))
assert second_provider_body == _expected_body(
"responses",
_user(first_prompt),
_model_turn({"text": first_answer}),
_user(second_prompt),
)
_FIVE_KB: Final = "s" * 5000
_ASSISTANT_SHAPES: Final = (
pytest.param({"reasoning_content": _REASONING, "thinking_blocks": 5}, _REPLAYED_TURN, id="blocks-int"),

View file

@ -17,7 +17,6 @@ import pytest
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.types.utils import TokenCountResponse
class BaseTokenCounterTest(ABC):
@ -71,69 +70,3 @@ class BaseTokenCounterTest(ABC):
):
pytest.skip(f"Missing or invalid credentials: {e}")
raise
@pytest.mark.asyncio
async def test_count_tokens_basic(self):
"""
Test basic token counting functionality.
Verifies that:
- Token counter returns a TokenCountResponse
- total_tokens is greater than 0
- tokenizer_type is set
- No error occurred
"""
token_counter = self.get_token_counter()
model = self.get_test_model()
messages = self.get_test_messages()
deployment = self.get_deployment_config()
result = await token_counter.count_tokens(
model_to_use=model,
messages=messages,
contents=None,
deployment=deployment,
request_model=model,
)
print(f"Token count result: {result}")
assert result is not None, "Token counter should return a result"
assert isinstance(
result, TokenCountResponse
), "Result should be TokenCountResponse"
assert (
result.total_tokens > 0
), f"Token count should be > 0, got {result.total_tokens}"
assert result.tokenizer_type is not None, "tokenizer_type should be set"
assert (
result.error is not True
), f"Token counting should not error: {result.error_message}"
def test_should_use_token_counting_api(self):
"""
Test that should_use_token_counting_api returns True for the correct provider.
Verifies that the token counter correctly identifies when it should be used
based on the custom_llm_provider.
"""
token_counter = self.get_token_counter()
provider = self.get_custom_llm_provider()
result = token_counter.should_use_token_counting_api(
custom_llm_provider=provider
)
assert (
result is True
), f"should_use_token_counting_api should return True for {provider}"
# Also verify it returns False for other providers
other_provider = "some_other_provider_that_doesnt_exist"
result_other = token_counter.should_use_token_counting_api(
custom_llm_provider=other_provider
)
assert (
result_other is False
), f"should_use_token_counting_api should return False for {other_provider}"

View file

@ -1,326 +0,0 @@
#### What this tests ####
# This tests if ahealth_check() actually works
import asyncio
import os
from unittest.mock import AsyncMock, patch
import pytest
import litellm
@pytest.mark.asyncio
async def test_azure_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "azure/gpt-4.1-mini",
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
"api_version": os.getenv("AZURE_AI_API_VERSION"),
}
)
print(f"response: {response}")
assert "x-ratelimit-remaining-tokens" in response
return response
# asyncio.run(test_azure_health_check())
@pytest.mark.asyncio
async def test_azure_embedding_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "azure/text-embedding-ada-002",
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
"api_version": os.getenv("AZURE_AI_API_VERSION"),
},
input=["test for litellm"],
mode="embedding",
)
print(f"response: {response}")
assert "x-ratelimit-remaining-tokens" in response
return response
@pytest.mark.asyncio
async def test_openai_img_gen_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "gpt-image-1",
"api_key": os.getenv("OPENAI_API_KEY"),
},
mode="image_generation",
prompt="cute baby sea otter",
)
print(f"response: {response}")
assert isinstance(response, dict) and "error" not in response
return response
# asyncio.run(test_openai_img_gen_health_check())
# asyncio.run(test_sagemaker_embedding_health_check())
@pytest.mark.asyncio
async def test_groq_health_check():
"""
This should not fail
ensure that provider wildcard model passes health check
"""
litellm.set_verbose = True
response = await litellm.ahealth_check(
model_params={
"api_key": os.environ.get("GROQ_API_KEY"),
"model": "groq/*",
"messages": [{"role": "user", "content": "What's 1 + 1?"}],
},
mode=None,
prompt="What's 1 + 1?",
input=["test from litellm"],
)
print(f"response: {response}")
assert response == {}
return response
@pytest.mark.asyncio
async def test_cohere_rerank_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "cohere/rerank-english-v3.0",
"api_key": os.getenv("COHERE_API_KEY"),
},
mode="rerank",
prompt="Hey, how's it going",
)
assert "error" not in response
print(response)
@pytest.mark.asyncio
async def test_audio_speech_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "openai/tts-1",
"api_key": os.getenv("OPENAI_API_KEY"),
},
mode="audio_speech",
prompt="Hey",
)
assert "error" not in response
print(response)
@pytest.mark.asyncio
async def test_audio_speech_health_check_with_another_voice():
response = await litellm.ahealth_check(
model_params={
"model": "openai/tts-1",
"api_key": os.getenv("OPENAI_API_KEY"),
"health_check_voice": "en-US-JennyNeural",
},
mode="audio_speech",
prompt="Hey",
)
assert "error" not in response
print(response)
@pytest.mark.asyncio
async def test_audio_transcription_health_check():
litellm.set_verbose = True
response = await litellm.ahealth_check(
model_params={
"model": "openai/whisper-1",
"api_key": os.getenv("OPENAI_API_KEY"),
},
mode="audio_transcription",
)
print(f"response: {response}")
assert "error" not in response
print(response)
@pytest.mark.asyncio
async def test_health_check_bad_model():
import time
from litellm.proxy.health_check import _perform_health_check
model_list = [
{
"model_name": "openai-gpt-4o",
"litellm_params": {
"api_key": "sk-9876",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app",
"model": "openai/my-fake-openai-endpoint",
"mock_timeout": True,
"timeout": 60,
},
"model_info": {
"id": "ca27ca2eeea2f9e38bb274ead831948a26621a3738d06f1797253f0e6c4278c0",
"db_model": False,
"health_check_timeout": 1,
},
},
]
details = None
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
model_list, details
)
print(f"healthy_endpoints: {healthy_endpoints}")
print(f"unhealthy_endpoints: {unhealthy_endpoints}")
# Track which model is actually used in the health check
health_check_calls = []
async def mock_health_check(litellm_params, **kwargs):
health_check_calls.append(litellm_params["model"])
await asyncio.sleep(10)
return {"status": "healthy"}
with patch(
"litellm.ahealth_check", side_effect=mock_health_check
) as mock_health_check:
start_time = time.time()
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
model_list
)
end_time = time.time()
print("health check calls: ", health_check_calls)
assert len(healthy_endpoints) == 0
assert len(unhealthy_endpoints) == 1
assert (
end_time - start_time < 2
), "Health check took longer than health_check_timeout"
@pytest.mark.asyncio
async def test_health_check_respects_concurrency_limit():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(6)
]
active = 0
max_active = 0
async def mock_health_check(litellm_params, **kwargs):
nonlocal active, max_active
active += 1
max_active = max(max_active, active)
await asyncio.sleep(0.05)
active -= 1
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
await _perform_health_check(model_list, max_concurrency=2)
assert max_active <= 2
@pytest.mark.asyncio
async def test_health_check_creates_only_bounded_initial_tasks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(10)
]
release_event = asyncio.Event()
create_task_call_count = 0
real_create_task = asyncio.create_task
async def mock_health_check(litellm_params, **kwargs):
await release_event.wait()
return {"status": "healthy"}
def tracked_create_task(coro):
nonlocal create_task_call_count
create_task_call_count += 1
return real_create_task(coro)
with (
patch("litellm.ahealth_check", side_effect=mock_health_check),
patch(
"litellm.proxy.health_check.asyncio.create_task",
side_effect=tracked_create_task,
),
):
perform_task = real_create_task(
_perform_health_check(model_list, max_concurrency=2)
)
await asyncio.sleep(0.05)
assert create_task_call_count == 2
release_event.set()
await perform_task
@pytest.mark.asyncio
async def test_timeout_does_not_cancel_other_health_checks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{
"litellm_params": {"model": "openai/slow-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 0.05},
},
{
"litellm_params": {"model": "openai/fast-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 1},
},
]
async def mock_health_check(litellm_params, **kwargs):
if litellm_params["model"] == "openai/slow-model":
await asyncio.sleep(0.2)
return {"status": "healthy"}
await asyncio.sleep(0.01)
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
model_list, max_concurrency=1
)
healthy_models = {endpoint["model"] for endpoint in healthy_endpoints}
unhealthy_models = {endpoint["model"] for endpoint in unhealthy_endpoints}
assert "openai/fast-model" in healthy_models
assert "openai/slow-model" in unhealthy_models

View file

@ -1,4 +1,3 @@
import base64
import hashlib
import json
import os
@ -7,7 +6,6 @@ from dotenv import load_dotenv
load_dotenv()
import tempfile
from unittest.mock import MagicMock, patch
from typing import Final
import pytest
@ -93,8 +91,6 @@ def test_oidc_google():
)
print(f"secret_val: {redact_oidc_signature(secret_val)}")
@pytest.mark.skipif(
os.environ.get("ACTIONS_ID_TOKEN_REQUEST_TOKEN") is None,
reason="Cannot run without being in GitHub Actions",
@ -105,112 +101,3 @@ def test_oidc_github():
)
print(f"secret_val: {redact_oidc_signature(secret_val)}")
@pytest.mark.skipif(
os.environ.get("CIRCLE_OIDC_TOKEN") is None,
reason="Cannot run without being in CircleCI Runner",
)
def test_oidc_circleci():
secret_val = get_secret("oidc/circleci/")
print(f"secret_val: {redact_oidc_signature(secret_val)}")
@pytest.mark.skipif(
os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None,
reason="Cannot run without being in CircleCI Runner",
)
def test_oidc_circleci_v2():
secret_val = get_secret(
"oidc/circleci_v2/https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.titan-text-express-v1/invoke"
)
print(f"secret_val: {redact_oidc_signature(secret_val)}")
def test_google_secret_manager():
"""
Test that we can get a secret from Google Secret Manager
"""
os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd"
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"payload": {
"data": base64.b64encode(b"anything").decode("utf-8"),
}
}
with (
patch("litellm.proxy.proxy_server.premium_user", True),
patch.object(
GoogleSecretManager,
"sync_construct_request_headers",
return_value={"Authorization": "Bearer mock_token"},
),
):
secret_manager = GoogleSecretManager()
secret_manager.sync_httpx_client = MagicMock()
secret_manager.sync_httpx_client.get.return_value = mock_response
secret_val = secret_manager.get_secret_from_google_secret_manager(
secret_name="OPENAI_API_KEY"
)
print("secret_val: {}".format(secret_val))
assert (
secret_val == "anything"
), "did not get expected secret value. expect 'anything', got '{}'".format(
secret_val
)
secret_manager.sync_httpx_client.get.assert_called_once()
call_url = secret_manager.sync_httpx_client.get.call_args[1]["url"]
assert "projects/litellm-ci-cd/secrets/OPENAI_API_KEY" in call_url
def test_google_secret_manager_read_in_memory():
"""
Test that Google Secret manager returns in memory value when it exists
"""
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd"
with (
patch("litellm.proxy.proxy_server.premium_user", True),
patch.object(
GoogleSecretManager,
"sync_construct_request_headers",
return_value={"Authorization": "Bearer mock_token"},
),
):
secret_manager = GoogleSecretManager()
secret_manager.cache.cache_dict["UNIQUE_KEY"] = None
secret_manager.cache.cache_dict["UNIQUE_KEY_2"] = "lite-llm"
secret_val = secret_manager.get_secret_from_google_secret_manager(
secret_name="UNIQUE_KEY"
)
print("secret_val: {}".format(secret_val))
assert secret_val is None
secret_val = secret_manager.get_secret_from_google_secret_manager(
secret_name="UNIQUE_KEY_2"
)
print("secret_val: {}".format(secret_val))
assert secret_val == "lite-llm"

View file

@ -1,7 +1,5 @@
import copy
import logging
import time
from datetime import datetime
from unittest import mock
from dotenv import load_dotenv
@ -15,22 +13,14 @@ import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, headers
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.duration_parser import (
get_last_day_of_month,
_extract_from_regex,
)
from litellm.utils import (
check_valid_key,
get_llm_provider,
get_supported_openai_params,
get_token_count,
get_valid_models,
trim_messages,
validate_environment,
)
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from unittest.mock import AsyncMock, MagicMock, patch
# Assuming your trim_messages, shorten_message_to_fit_limit, and get_token_count functions are all in a module named 'message_utils'
@ -47,676 +37,29 @@ def reset_mock_cache():
# test_basic_trimming()
# test_basic_trimming_no_max_tokens_specified()
# test_multiple_messages_trimming()
# test_multiple_messages_no_trimming()
# test_large_trimming()
@pytest.mark.parametrize("custom_llm_provider", ["anthropic", "xai"])
def test_get_valid_models_with_custom_llm_provider(custom_llm_provider):
from litellm.utils import ProviderConfigManager
from litellm.types.utils import LlmProviders
provider_config = ProviderConfigManager.get_provider_model_info(
model=None,
provider=LlmProviders(custom_llm_provider),
)
assert provider_config is not None
valid_models = get_valid_models(
check_provider_endpoint=True, custom_llm_provider=custom_llm_provider
)
print(valid_models)
assert len(valid_models) > 0
assert set(provider_config.get_models()) == set(valid_models)
# test_get_valid_models()
def test_bad_key():
key = "bad-key"
response = check_valid_key(model="gpt-5-mini", api_key=key)
print(response, key)
assert response == False
def test_good_key():
key = os.environ["OPENAI_API_KEY"]
response = check_valid_key(model="gpt-5-mini", api_key=key)
assert response == True
# test validate environment
def test_function_to_dict():
print("testing function to dict for get current weather")
def get_current_weather(location: str, unit: str):
"""Get the current weather in a given location
Parameters
----------
location : str
The city and state, e.g. San Francisco, CA
unit : {'celsius', 'fahrenheit'}
Temperature unit
Returns
-------
str
a sentence indicating the weather
"""
if location == "Boston, MA":
return "The weather is 12F"
function_json = litellm.utils.function_to_dict(get_current_weather)
print(function_json)
expected_output = {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {
"type": "string",
"description": "Temperature unit",
"enum": "['fahrenheit', 'celsius']",
},
},
"required": ["location", "unit"],
},
}
print(expected_output)
assert function_json["name"] == expected_output["name"]
assert function_json["description"] == expected_output["description"]
assert function_json["parameters"]["type"] == expected_output["parameters"]["type"]
assert (
function_json["parameters"]["properties"]["location"]
== expected_output["parameters"]["properties"]["location"]
)
# the enum can change it can be - which is why we don't assert on unit
# {'type': 'string', 'description': 'Temperature unit', 'enum': "['fahrenheit', 'celsius']"}
# {'type': 'string', 'description': 'Temperature unit', 'enum': "['celsius', 'fahrenheit']"}
assert (
function_json["parameters"]["required"]
== expected_output["parameters"]["required"]
)
print("passed")
# test_function_to_dict()
def test_duration_in_seconds():
"""
Test if duration int is correctly calculated for different str
"""
import time
now = time.time()
current_time = datetime.fromtimestamp(now)
if current_time.month == 12:
target_year = current_time.year + 1
target_month = 1
else:
target_year = current_time.year
target_month = current_time.month + 1
# Determine the day to set for next month
target_day = current_time.day
last_day_of_target_month = get_last_day_of_month(target_year, target_month)
if target_day > last_day_of_target_month:
target_day = last_day_of_target_month
next_month = datetime(
year=target_year,
month=target_month,
day=target_day,
hour=current_time.hour,
minute=current_time.minute,
second=current_time.second,
microsecond=current_time.microsecond,
)
# Calculate the duration until the first day of the next month
duration_until_next_month = next_month - current_time
expected_duration = int(duration_until_next_month.total_seconds())
value = duration_in_seconds(duration="1mo")
assert value - expected_duration < 2
@pytest.mark.parametrize("langfuse_trace_id", [None, "my-unique-trace-id"])
@pytest.mark.parametrize(
"langfuse_existing_trace_id", [None, "my-unique-existing-trace-id"]
)
def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id):
"""
- Unit test for `_get_trace_id` function in Logging obj
"""
from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id
from litellm.litellm_core_utils.litellm_logging import Logging
litellm.success_callback = ["langfuse"]
litellm_call_id = "my-unique-call-id"
litellm_logging_obj = Logging(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="acompletion",
litellm_call_id=litellm_call_id,
start_time=datetime.now(),
function_id="1234",
)
metadata = {}
if langfuse_trace_id is not None:
metadata["trace_id"] = langfuse_trace_id
if langfuse_existing_trace_id is not None:
metadata["existing_trace_id"] = langfuse_existing_trace_id
litellm.completion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hey how's it going?"}],
mock_response="Hey!",
litellm_logging_obj=litellm_logging_obj,
metadata=metadata,
)
time.sleep(3)
assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None
# langfuse addresses a trace by a 32-hex id, so the id litellm reports back is the
# resolved form of whichever source won; that is what the alerting deep link needs
if langfuse_existing_trace_id is not None:
expected_source = langfuse_existing_trace_id
elif langfuse_trace_id is not None:
expected_source = langfuse_trace_id
else:
expected_source = litellm_logging_obj.litellm_trace_id
assert litellm_logging_obj.get_trace_id(service_name="langfuse") == resolve_trace_id(
expected_source
)
def test_is_base64_encoded():
import base64
import requests
litellm.set_verbose = True
url = "https://dummyimage.com/100/100/fff&text=Test+image"
response = requests.get(url)
file_data = response.content
encoded_file = base64.b64encode(file_data).decode("utf-8")
base64_image = f"data:image/png;base64,{encoded_file}"
from litellm.utils import is_base64_encoded
assert is_base64_encoded(s=base64_image) is True
def test_is_prompt_caching_enabled_return_default_image_dimensions():
"""
Assert that `is_prompt_caching_valid_prompt` counts tokens with use_default_image_token_count=True
when processing messages containing images
IMPORTANT: Ensures Get token counter does not make a GET request to the image url
"""
mock_token_counter = MagicMock(return_value=False)
with patch(
"litellm.utils.get_messages_reach_token_count",
return_value=mock_token_counter,
):
litellm.utils.is_prompt_caching_valid_prompt(
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this image?"},
{
"type": "image_url",
"image_url": {
"url": "https://www.gstatic.com/webp/gallery/1.webp",
"detail": "high",
},
},
],
}
],
tools=None,
custom_llm_provider="openai",
model="gpt-4o-mini",
)
# Assert token_counter was called with use_default_image_token_count=True
args_to_mock_token_counter = mock_token_counter.call_args[1]
print("args_to_mock", args_to_mock_token_counter)
assert args_to_mock_token_counter["use_default_image_token_count"] is True
def test_get_valid_models_fireworks_ai(monkeypatch):
from litellm.utils import get_valid_models
import litellm
litellm.turn_on_debug()
monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876")
monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "1234")
monkeypatch.setattr(litellm, "provider_list", ["fireworks_ai"])
mock_response_data = {
"models": [
{
"name": "accounts/fireworks/models/llama-3.1-8b-instruct",
"displayName": "<string>",
"description": "<string>",
"createTime": "2023-11-07T05:31:56Z",
"createdBy": "<string>",
"state": "STATE_UNSPECIFIED",
"status": {"code": "OK", "message": "<string>"},
"kind": "KIND_UNSPECIFIED",
"githubUrl": "<string>",
"huggingFaceUrl": "<string>",
"baseModelDetails": {
"worldSize": 123,
"checkpointFormat": "CHECKPOINT_FORMAT_UNSPECIFIED",
"parameterCount": "<string>",
"moe": True,
"tunable": True,
},
"peftDetails": {
"baseModel": "<string>",
"r": 123,
"targetModules": ["<string>"],
},
"teftDetails": {},
"public": True,
"conversationConfig": {
"style": "<string>",
"system": "<string>",
"template": "<string>",
},
"contextLength": 123,
"supportsImageInput": True,
"supportsTools": True,
"importedFrom": "<string>",
"fineTuningJob": "<string>",
"defaultDraftModel": "<string>",
"defaultDraftTokenCount": 123,
"precisions": ["PRECISION_UNSPECIFIED"],
"deployedModelRefs": [
{
"name": "<string>",
"deployment": "<string>",
"state": "STATE_UNSPECIFIED",
"default": True,
"public": True,
}
],
"cluster": "<string>",
"deprecationDate": {"year": 123, "month": 123, "day": 123},
}
],
"nextPageToken": "<string>",
"totalSize": 123,
}
# Create a mock response object
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_response_data
with patch.object(
litellm.module_level_client, "get", return_value=mock_response
) as mock_post:
valid_models = get_valid_models(check_provider_endpoint=True)
print("valid_models", valid_models)
mock_post.assert_called_once()
assert (
"fireworks_ai/accounts/fireworks/models/llama-3.1-8b-instruct"
in valid_models
)
def test_get_valid_models_default(monkeypatch):
"""
Ensure that the default models is used when error retrieving from model api.
Prevent regression for existing usage.
"""
from litellm.utils import get_valid_models
monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876")
valid_models = get_valid_models()
assert len(valid_models) > 0
def test_add_custom_logger_callback_to_specific_event(monkeypatch):
from litellm.utils import add_custom_logger_callback_to_specific_event
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
add_custom_logger_callback_to_specific_event("langfuse", "success")
assert len(litellm.success_callback) == 1
assert len(litellm.failure_callback) == 0
@pytest.mark.asyncio
async def test_add_custom_logger_callback_to_specific_event_with_duplicates(
monkeypatch,
):
"""
Test that when a callback exists in both success_callback and _async_success_callback,
it's not added again
"""
from litellm.integrations.langfuse.langfuse_prompt_management import (
LangfusePromptManagement,
)
# Reset all callback lists
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
# Add logger to both success_callback and _async_success_callback
langfuse_logger = LangfusePromptManagement()
litellm.success_callback.append(langfuse_logger)
litellm._async_success_callback.append(langfuse_logger)
# Get initial lengths
initial_success_callback_len = len(litellm.success_callback)
initial_async_success_callback_len = len(litellm._async_success_callback)
# Make a completion call
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hello, world!"}],
mock_response="Testing duplicate callbacks",
)
# Assert no new callbacks were added
assert len(litellm.success_callback) == initial_success_callback_len
assert len(litellm._async_success_callback) == initial_async_success_callback_len
@pytest.mark.asyncio
async def test_add_custom_logger_callback_to_specific_event_with_duplicates_success_callback(
monkeypatch,
):
"""
Test that when a callback exists in both success_callback and _async_success_callback,
it's not added again
"""
from litellm.integrations.langfuse.langfuse_prompt_management import (
LangfusePromptManagement,
)
# Reset all callback lists
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
# Add logger to both success_callback and _async_success_callback
langfuse_logger = LangfusePromptManagement()
litellm.success_callback.append(langfuse_logger)
# Get initial lengths
initial_success_callback_len = len(litellm.success_callback)
initial_async_success_callback_len = len(litellm._async_success_callback)
# Make a completion call
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hello, world!"}],
mock_response="Testing duplicate callbacks",
)
# Assert no new callbacks were added
assert len(litellm.success_callback) == initial_success_callback_len
assert len(litellm._async_success_callback) == initial_async_success_callback_len
@pytest.mark.asyncio
async def test_add_custom_logger_callback_to_specific_event_with_duplicates_callbacks(
monkeypatch,
):
"""
Test that when a callback exists in both success_callback and _async_success_callback,
it's not added again
"""
from litellm.integrations.langfuse.langfuse_prompt_management import (
LangfusePromptManagement,
)
# Reset all callback lists
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
# Add logger to both success_callback and _async_success_callback
langfuse_logger = LangfusePromptManagement()
litellm.callbacks.append(langfuse_logger)
# Make a completion call
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hello, world!"}],
mock_response="Testing duplicate callbacks",
)
# Assert no new callbacks were added
initial_callbacks_len = len(litellm.callbacks)
initial_async_success_callback_len = len(litellm._async_success_callback)
initial_success_callback_len = len(litellm.success_callback)
print(
f"Num callbacks before: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}"
)
for _ in range(10):
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hello, world!"}],
mock_response="Testing duplicate callbacks",
)
assert len(litellm.callbacks) == initial_callbacks_len
assert len(litellm._async_success_callback) == initial_async_success_callback_len
assert len(litellm.success_callback) == initial_success_callback_len
print(
f"Num callbacks after 10 mock calls: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}"
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.utils import get_applied_guardrails
from unittest.mock import Mock
def test_get_provider_audio_transcription_config():
from litellm.utils import ProviderConfigManager
from litellm.types.utils import LlmProviders
for provider in LlmProviders:
config = ProviderConfigManager.get_provider_audio_transcription_config(
model="whisper-1", provider=provider
)
def test_get_valid_models_from_dynamic_api_key():
"""
Test that get_valid_models returns the correct models for a given provider
"""
from litellm.utils import get_valid_models
from litellm.types.router import CredentialLiteLLMParams
creds = CredentialLiteLLMParams(api_key="123")
valid_models = get_valid_models(
custom_llm_provider="anthropic",
litellm_params=creds,
check_provider_endpoint=True,
)
assert len(valid_models) == 0
creds = CredentialLiteLLMParams(api_key=os.getenv("ANTHROPIC_API_KEY"))
valid_models = get_valid_models(
custom_llm_provider="anthropic",
litellm_params=creds,
check_provider_endpoint=True,
)
assert len(valid_models) > 0
assert "anthropic/claude-sonnet-4-6" in valid_models
def test_get_whitelisted_models():
"""
Snapshot of all bedrock models as of 12/24/2024.

View file

@ -347,31 +347,6 @@ class BaseResponsesAPITest(ABC):
else:
raise ValueError("response is not a ResponsesAPIResponse")
@pytest.mark.asyncio
async def test_multiturn_responses_api(self):
litellm.turn_on_debug()
litellm.set_verbose = True
try:
base_completion_call_args = self.get_base_completion_call_args()
response_1 = await litellm.aresponses(
input="Basic ping", max_output_tokens=20, **base_completion_call_args
)
# follow up with a second request
response_1_id = response_1.id
response_2 = await litellm.aresponses(
input="Basic ping",
max_output_tokens=20,
previous_response_id=response_1_id,
**base_completion_call_args,
)
# assert the response is not None
assert response_1 is not None
assert response_2 is not None
except litellm.InternalServerError:
pytest.skip("Skipping test due to litellm.InternalServerError")
@pytest.mark.asyncio
async def test_responses_api_with_tool_calls(self):
"""Test that calls the Responses API with tool calls including function call and output"""
@ -447,72 +422,6 @@ class BaseResponsesAPITest(ABC):
else:
assert len(response["output"]) > 0
def test_openai_responses_api_dict_input_filtering(self):
"""
Test that regular dict inputs with status fields are properly filtered
to replicate exclude_unset=True behavior for non-Pydantic objects.
"""
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
# Test input with regular dict objects (like from JSON)
test_input = [
{"role": "user", "content": "test"},
{
"id": "rs_123",
"summary": [{"text": "test", "type": "summary_text"}],
"type": "reasoning",
"content": None, # Should be filtered out
"encrypted_content": None, # Should be filtered out
"status": None, # Should be filtered out
},
{
"arguments": "{}",
"call_id": "call_123",
"name": "get_today",
"type": "function_call",
"id": "fc_123",
"status": "completed", # Should be preserved (not a default field)
},
]
config = OpenAIResponsesAPIConfig()
validated_input = config._validate_input_param(test_input)
# Verify the results
assert len(validated_input) == 3
# Check reasoning item (index 1)
reasoning_item = validated_input[1]
assert reasoning_item["type"] == "reasoning"
assert (
"status" not in reasoning_item
), "status field should be filtered out from reasoning item"
assert (
"content" not in reasoning_item
), "content field should be filtered out from reasoning item"
assert (
"encrypted_content" not in reasoning_item
), "encrypted_content field should be filtered out from reasoning item"
# Note: ID auto-generation was disabled, so reasoning items may not have IDs
# Only check for ID if it was present in the original input
if "id" in reasoning_item:
assert reasoning_item["id"] == "rs_123", "ID should be preserved if present"
assert "summary" in reasoning_item, "summary field should be preserved"
# Check function call item (index 2)
function_call_item = validated_input[2]
assert function_call_item["type"] == "function_call"
assert (
"status" in function_call_item
), "status field should be preserved in function call item"
assert (
function_call_item["status"] == "completed"
), "status value should be preserved"
print("✅ OpenAI Responses API dict input filtering test passed")
@pytest.mark.parametrize("sync_mode", [False, True])
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
@ -565,23 +474,6 @@ class BaseResponsesAPITest(ABC):
else:
raise e
@pytest.mark.parametrize("sync_mode", [False, True])
@pytest.mark.asyncio
async def test_cancel_responses_invalid_response_id(self, sync_mode):
"""Test cancel_responses with invalid response ID should raise appropriate error"""
base_completion_call_args = self.get_base_completion_call_args()
if sync_mode:
with pytest.raises(openai.APIError):
litellm.cancel_responses(
response_id="invalid_response_id_12345", **base_completion_call_args
)
else:
with pytest.raises(openai.APIError):
await litellm.acancel_responses(
response_id="invalid_response_id_12345", **base_completion_call_args
)
@pytest.mark.asyncio
async def test_responses_api_context_management_server_side_compaction(self):
"""

View file

@ -1,35 +1,20 @@
import pytest
import litellm
from base_responses_api import BaseResponsesAPITest
from openai.types.responses.function_tool import FunctionTool
class TestAnthropicResponsesAPITest(BaseResponsesAPITest):
test_basic_openai_responses_delete_endpoint = None
test_basic_openai_responses_streaming_delete_endpoint = None
test_basic_openai_responses_get_endpoint = None
test_basic_openai_responses_cancel_endpoint = None
def get_base_completion_call_args(self):
# litellm.turn_on_debug()
return {
"model": "anthropic/claude-sonnet-4-5",
}
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
pytest.skip("DELETE responses is not supported for anthropic")
async def test_basic_openai_responses_streaming_delete_endpoint(
self, sync_mode=False
):
pytest.skip("DELETE responses is not supported for anthropic")
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
pytest.skip("GET responses is not supported for anthropic")
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for anthropic")
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for anthropic")
def test_multiturn_tool_calls():
# Test streaming response with tools for Anthropic
litellm.turn_on_debug()

View file

@ -6,7 +6,6 @@ from base_responses_api import BaseResponsesAPITest
class TestAzureResponsesAPITest(BaseResponsesAPITest):
test_multiturn_responses_api = None
test_responses_api_with_tool_calls = None
def get_base_completion_call_args(self):

View file

@ -1,233 +1,12 @@
import os
import pytest
import litellm
import json
from base_responses_api import BaseResponsesAPITest
@pytest.mark.asyncio
async def test_basic_google_ai_studio_responses_api_with_tools():
litellm.turn_on_debug()
litellm.set_verbose = True
request_model = "gemini/gemini-2.5-flash"
response = await litellm.aresponses(
model=request_model,
input="what is the latest version of supabase python package and when was it released?",
tools=[{"type": "web_search_preview", "search_context_size": "low"}],
)
print("litellm response=", json.dumps(response, indent=4, default=str))
@pytest.mark.asyncio
async def test_gemini_3_responses_api_with_thought_signatures():
"""
Test that Gemini 3 Responses API preserves thought signatures in function calls.
This test verifies that provider_specific_fields with thought_signature are correctly
preserved when using the Responses API with Gemini 3.
"""
if not os.getenv("GEMINI_API_KEY"):
pytest.skip("GEMINI_API_KEY not set")
litellm.set_verbose = False
request_model = "gemini/gemini-3.1-pro-preview"
tools = [
{
"type": "function",
"name": "get_weather",
"description": "Get the current weather for a location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "City and country e.g. Mumbai, India",
},
"units": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
"description": "Units the temperature will be returned in.",
},
},
"required": ["location", "units"],
"additionalProperties": False,
},
"strict": True,
}
]
# Step 1: Initial request with tools
response = await litellm.aresponses(
model=request_model,
input="What is the weather in Mumbai?",
tools=tools,
reasoning_effort="low",
)
# Validate response structure
from litellm.types.llms.openai import ResponsesAPIResponse
assert isinstance(
response, ResponsesAPIResponse
), "Response should be a ResponsesAPIResponse"
assert (
hasattr(response, "output") or "output" in response
), "Response should have 'output' field"
assert isinstance(response.output, list), "Output should be a list"
# Find function call in output
function_call_item = None
for item in response.output:
# Convert to dict if it's a Pydantic model for easier access
if hasattr(item, "model_dump"):
item_dict = item.model_dump()
elif hasattr(item, "__dict__"):
item_dict = dict(item) if not isinstance(item, dict) else item
else:
item_dict = item if isinstance(item, dict) else {}
if isinstance(item_dict, dict) and item_dict.get("type") == "function_call":
function_call_item = item_dict
break
# Verify function call exists
assert (
function_call_item is not None
), "Response should contain a function_call item"
assert (
function_call_item.get("name") == "get_weather"
), "Function call should be for get_weather"
# Verify thought signature is present in provider_specific_fields
provider_specific_fields = function_call_item.get("provider_specific_fields")
assert (
provider_specific_fields is not None
), "Function call should have provider_specific_fields"
assert (
"thought_signature" in provider_specific_fields
), "provider_specific_fields should contain thought_signature"
assert isinstance(
provider_specific_fields["thought_signature"], str
), "thought_signature should be a string"
assert (
len(provider_specific_fields["thought_signature"]) > 0
), "thought_signature should not be empty"
print(
f"✅ Thought signature preserved: {provider_specific_fields['thought_signature'][:50]}..."
)
@pytest.mark.asyncio
async def test_gemini_3_responses_api_streaming_with_thought_signatures():
"""
Test that Gemini 3 Responses API preserves thought signatures in streaming mode.
This test verifies that provider_specific_fields with thought_signature are correctly
preserved when using streaming Responses API with Gemini 3.
"""
if not os.getenv("GEMINI_API_KEY"):
pytest.skip("GEMINI_API_KEY not set")
litellm.set_verbose = False
request_model = "gemini/gemini-3.1-pro-preview"
tools = [
{
"type": "function",
"name": "get_weather",
"description": "Get the current weather for a location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "City and country e.g. Mumbai, India",
},
"units": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
"description": "Units the temperature will be returned in.",
},
},
"required": ["location", "units"],
"additionalProperties": False,
},
"strict": True,
}
]
# Step 1: Streaming request with tools
response_stream = await litellm.aresponses(
model=request_model,
input="What is the weather in Mumbai?",
tools=tools,
stream=True,
reasoning_effort="low",
)
# Collect all chunks
chunks = []
completed_response = None
async for chunk in response_stream:
chunks.append(chunk)
# Check if this is the completed response event
if hasattr(chunk, "type") and chunk.type == "response.completed":
completed_response = chunk.response
elif isinstance(chunk, dict) and chunk.get("type") == "response.completed":
completed_response = chunk.get("response")
# Verify we got chunks
assert len(chunks) > 0, "Should receive at least one chunk"
# If we have a completed response, check for thought signatures
if completed_response:
output = completed_response.get("output", [])
function_call_item = None
for item in output:
if isinstance(item, dict) and item.get("type") == "function_call":
function_call_item = item
break
if function_call_item:
provider_specific_fields = function_call_item.get(
"provider_specific_fields"
)
if provider_specific_fields:
thought_signature = provider_specific_fields.get("thought_signature")
if thought_signature:
assert isinstance(
thought_signature, str
), "thought_signature should be a string"
assert (
len(thought_signature) > 0
), "thought_signature should not be empty"
print(
f"✅ Streaming thought signature preserved: {thought_signature[:50]}..."
)
print(f"✅ Collected {len(chunks)} streaming chunks")
class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest):
test_basic_openai_responses_delete_endpoint = None
test_basic_openai_responses_streaming_delete_endpoint = None
test_basic_openai_responses_get_endpoint = None
test_basic_openai_responses_cancel_endpoint = None
def get_base_completion_call_args(self):
# litellm.turn_on_debug()
return {"model": "gemini/gemini-2.5-flash-lite"}
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
pytest.skip("DELETE responses is not supported for Google AI Studio")
async def test_basic_openai_responses_streaming_delete_endpoint(
self, sync_mode=False
):
pytest.skip("DELETE responses is not supported for Google AI Studio")
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
pytest.skip("GET responses is not supported for Google AI Studio")
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for Google AI Studio")
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for Google AI Studio")

View file

@ -1,20 +1,9 @@
import asyncio
import json
import os
import time
from typing import Optional, cast
import pytest
from base_responses_api import BaseResponsesAPITest, validate_responses_api_response
from base_responses_api import BaseResponsesAPITest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
from litellm.types.utils import StandardLoggingPayload
class TestOpenAIResponsesAPITest(BaseResponsesAPITest):
@ -29,736 +18,6 @@ class TestOpenAIResponsesAPITest(BaseResponsesAPITest):
return "openai/gpt-5.2"
class TestCustomLogger(CustomLogger):
def __init__(
self,
):
self.standard_logging_object: Optional[StandardLoggingPayload] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
print("in async_log_success_event")
print("kwargs=", json.dumps(kwargs, indent=4, default=str))
self.standard_logging_object = kwargs["standard_logging_object"]
pass
def validate_standard_logging_payload(
slp: StandardLoggingPayload, response: ResponsesAPIResponse, request_model: str
):
"""
Validate that a StandardLoggingPayload object matches the expected response
Args:
slp (StandardLoggingPayload): The standard logging payload object to validate
response (dict): The litellm response to compare against
request_model (str): The model name that was requested
"""
# Validate payload exists
assert slp is not None, "Standard logging payload should not be None"
# Validate token counts
print(
"VALIDATING STANDARD LOGGING PAYLOAD. response=",
json.dumps(response, indent=4, default=str),
)
print("FIELDS IN SLP=", json.dumps(slp, indent=4, default=str))
print("SLP PROMPT TOKENS=", slp["prompt_tokens"])
print("RESPONSE PROMPT TOKENS=", response["usage"]["input_tokens"])
assert (
slp["prompt_tokens"] == response["usage"]["input_tokens"]
), "Prompt tokens mismatch"
assert (
slp["completion_tokens"] == response["usage"]["output_tokens"]
), "Completion tokens mismatch"
assert (
slp["total_tokens"]
== response["usage"]["input_tokens"] + response["usage"]["output_tokens"]
), "Total tokens mismatch"
# Validate spend and response metadata
assert slp["response_cost"] > 0, "Response cost should be greater than 0"
assert slp["id"] == response["id"], "Response ID mismatch"
assert slp["model"] == request_model, "Model name mismatch"
# Validate messages
assert slp["messages"] == [{"content": "hi", "role": "user"}], "Messages mismatch"
# Validate complete response structure
validate_responses_match(slp["response"], response)
@pytest.mark.asyncio
def test_basic_openai_responses_api_streaming_with_logging():
litellm.turn_on_debug()
litellm.set_verbose = True
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
request_model = "gpt-5.5"
response = litellm.responses(
model=request_model,
input="hi",
stream=True,
)
final_response: Optional[ResponseCompletedEvent] = None
for event in response:
if event.type == "response.completed":
final_response = event
print("litellm response=", json.dumps(event, indent=4, default=str))
print("sleeping for 2 seconds...")
time.sleep(2)
print(
"standard logging payload=",
json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str),
)
assert final_response is not None
assert test_custom_logger.standard_logging_object is not None
validate_standard_logging_payload(
slp=test_custom_logger.standard_logging_object,
response=final_response.response,
request_model=request_model,
)
def validate_responses_match(slp_response, litellm_response):
"""Validate that the standard logging payload OpenAI response matches the litellm response"""
# Validate core fields
assert slp_response["id"] == litellm_response["id"], "ID mismatch"
assert slp_response["model"] == litellm_response["model"], "Model mismatch"
assert (
slp_response["created_at"] == litellm_response["created_at"]
), "Created at mismatch"
# Validate usage
assert (
slp_response["usage"]["prompt_tokens"]
== litellm_response["usage"]["input_tokens"]
), "Input tokens mismatch"
assert (
slp_response["usage"]["completion_tokens"]
== litellm_response["usage"]["output_tokens"]
), "Output tokens mismatch"
assert (
slp_response["usage"]["total_tokens"]
== litellm_response["usage"]["total_tokens"]
), "Total tokens mismatch"
# Validate output/messages
assert len(slp_response["output"]) == len(
litellm_response["output"]
), "Output length mismatch"
for slp_msg, litellm_msg in zip(slp_response["output"], litellm_response["output"]):
assert slp_msg["role"] == litellm_msg.role, "Message role mismatch"
# Access the content's text field for the litellm response
litellm_content = litellm_msg.content[0].text if litellm_msg.content else ""
assert (
slp_msg["content"][0]["text"] == litellm_content
), f"Message content mismatch. Expected {litellm_content}, Got {slp_msg['content']}"
assert slp_msg["status"] == litellm_msg.status, "Message status mismatch"
@pytest.mark.asyncio
async def test_basic_openai_responses_api_non_streaming_with_logging():
litellm.turn_on_debug()
litellm.set_verbose = True
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
request_model = "gpt-5.5"
response = await litellm.aresponses(
model=request_model,
input="hi",
)
print("litellm response=", json.dumps(response, indent=4, default=str))
print("response hidden params=", response._hidden_params)
print("sleeping for 2 seconds...")
await asyncio.sleep(5)
print(
"standard logging payload=",
json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str),
)
print("response usage=", response.usage)
assert response is not None
assert test_custom_logger.standard_logging_object is not None
validate_standard_logging_payload(
test_custom_logger.standard_logging_object, response, request_model
)
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_openai_responses_api_returns_headers(sync_mode):
"""
Test that OpenAI responses API returns OpenAI headers in _hidden_params.
This ensures the proxy can forward these headers to clients.
Related issue: LiteLLM responses API should return OpenAI headers like chat completions does
"""
litellm.turn_on_debug()
litellm.set_verbose = True
if sync_mode:
response = litellm.responses(
model="gpt-5.5",
input="Say hello",
max_output_tokens=20,
)
else:
response = await litellm.aresponses(
model="gpt-5.5",
input="Say hello",
max_output_tokens=20,
)
# Verify response is valid
assert response is not None
assert isinstance(response, ResponsesAPIResponse)
# Verify _hidden_params exists
assert hasattr(
response, "_hidden_params"
), "Response should have _hidden_params attribute"
assert response._hidden_params is not None, "_hidden_params should not be None"
# Verify additional_headers exists in _hidden_params
assert (
"additional_headers" in response._hidden_params
), "_hidden_params should contain 'additional_headers' key"
additional_headers = response._hidden_params["additional_headers"]
assert isinstance(
additional_headers, dict
), "additional_headers should be a dictionary"
assert len(additional_headers) > 0, "additional_headers should not be empty"
# Check for expected OpenAI rate limit headers
# These can be either direct (x-ratelimit-*) or prefixed (llm_provider-x-ratelimit-*)
rate_limit_headers = [
"x-ratelimit-remaining-tokens",
"x-ratelimit-limit-tokens",
"x-ratelimit-remaining-requests",
"x-ratelimit-limit-requests",
]
found_headers = []
for header_name in rate_limit_headers:
if header_name in additional_headers:
found_headers.append(header_name)
elif f"llm_provider-{header_name}" in additional_headers:
found_headers.append(f"llm_provider-{header_name}")
assert (
len(found_headers) > 0
), f"Should find at least one OpenAI rate limit header. Headers found: {list(additional_headers.keys())}"
# Verify headers key also exists (raw headers)
assert (
"headers" in response._hidden_params
), "_hidden_params should contain 'headers' key with raw response headers"
print(
f"✓ Successfully validated OpenAI headers in {'sync' if sync_mode else 'async'} mode"
)
print(f" Found {len(additional_headers)} headers total")
print(f" Rate limit headers found: {found_headers}")
def validate_stream_event(event):
"""
Validate that a streaming event from litellm.responses() or litellm.aresponses()
with stream=True conforms to the expected structure based on its event type.
Args:
event: The streaming event object to validate
Raises:
AssertionError: If the event doesn't match the expected structure for its type
"""
# Common validation for all event types
assert hasattr(event, "type"), "Event should have a 'type' attribute"
# Type-specific validation
if event.type == "response.created" or event.type == "response.in_progress":
assert hasattr(
event, "response"
), f"{event.type} event should have a 'response' attribute"
validate_responses_api_response(event.response, final_chunk=False)
elif event.type == "response.completed":
assert hasattr(
event, "response"
), "response.completed event should have a 'response' attribute"
validate_responses_api_response(event.response, final_chunk=True)
# Usage is guaranteed only on the completed event
assert (
"usage" in event.response
), "response.completed event should have usage information"
print("Usage in event.response=", event.response["usage"])
assert isinstance(event.response["usage"], ResponseAPIUsage)
elif event.type == "response.failed" or event.type == "response.incomplete":
assert hasattr(
event, "response"
), f"{event.type} event should have a 'response' attribute"
elif (
event.type == "response.output_item.added"
or event.type == "response.output_item.done"
):
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "item"
), f"{event.type} event should have an 'item' attribute"
elif (
event.type == "response.content_part.added"
or event.type == "response.content_part.done"
):
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "part"
), f"{event.type} event should have a 'part' attribute"
elif event.type == "response.output_text.delta":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "delta"
), f"{event.type} event should have a 'delta' attribute"
elif event.type == "response.output_text.annotation.added":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "annotation_index"
), f"{event.type} event should have an 'annotation_index' attribute"
assert hasattr(
event, "annotation"
), f"{event.type} event should have an 'annotation' attribute"
elif event.type == "response.output_text.done":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "text"
), f"{event.type} event should have a 'text' attribute"
elif event.type == "response.refusal.delta":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "delta"
), f"{event.type} event should have a 'delta' attribute"
elif event.type == "response.refusal.done":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "content_index"
), f"{event.type} event should have a 'content_index' attribute"
assert hasattr(
event, "refusal"
), f"{event.type} event should have a 'refusal' attribute"
elif event.type == "response.function_call_arguments.delta":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "delta"
), f"{event.type} event should have a 'delta' attribute"
elif event.type == "response.function_call_arguments.done":
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "arguments"
), f"{event.type} event should have an 'arguments' attribute"
elif event.type in [
"response.file_search_call.in_progress",
"response.file_search_call.searching",
"response.file_search_call.completed",
"response.web_search_call.in_progress",
"response.web_search_call.searching",
"response.web_search_call.completed",
]:
assert hasattr(
event, "output_index"
), f"{event.type} event should have an 'output_index' attribute"
assert hasattr(
event, "item_id"
), f"{event.type} event should have an 'item_id' attribute"
elif event.type == "error":
assert hasattr(
event, "message"
), "Error event should have a 'message' attribute"
return True # Return True if validation passes
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_openai_responses_api_streaming_validation(sync_mode):
"""Test that validates each streaming event from the responses API"""
litellm.turn_on_debug()
event_types_seen = set()
if sync_mode:
response = litellm.responses(
model="gpt-5.5",
input="Tell me about artificial intelligence in 3 sentences.",
stream=True,
)
for event in response:
print(f"Validating event type: {event.type}")
validate_stream_event(event)
event_types_seen.add(event.type)
else:
response = await litellm.aresponses(
model="gpt-5.5",
input="Tell me about artificial intelligence in 3 sentences.",
stream=True,
)
async for event in response:
print(f"Validating event type: {event.type}")
validate_stream_event(event)
event_types_seen.add(event.type)
# At minimum, we should see these core event types
required_events = {"response.created", "response.completed"}
missing_events = required_events - event_types_seen
assert not missing_events, f"Missing required event types: {missing_events}"
print(f"Successfully validated all event types: {event_types_seen}")
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_openai_responses_litellm_router(sync_mode):
"""
Test the OpenAI responses API with LiteLLM Router in both sync and async modes
"""
litellm.turn_on_debug()
router = litellm.Router(
model_list=[
{
"model_name": "gpt4o-special-alias",
"litellm_params": {
"model": "gpt-5.5",
"api_key": os.getenv("OPENAI_API_KEY"),
},
}
]
)
# Call the handler
if sync_mode:
response = router.responses(
model="gpt4o-special-alias",
input="Hello, can you tell me a short joke?",
max_output_tokens=100,
)
print("SYNC MODE RESPONSE=", response)
else:
response = await router.aresponses(
model="gpt4o-special-alias",
input="Hello, can you tell me a short joke?",
max_output_tokens=100,
)
print(
f"Router {'sync' if sync_mode else 'async'} response=",
json.dumps(response, indent=4, default=str),
)
# Use the helper function to validate the response
validate_responses_api_response(response, final_chunk=True)
return response
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_openai_responses_litellm_router_streaming(sync_mode):
"""
Test the OpenAI responses API with streaming through LiteLLM Router
"""
litellm.turn_on_debug()
router = litellm.Router(
model_list=[
{
"model_name": "gpt4o-special-alias",
"litellm_params": {
"model": "gpt-5.5",
"api_key": os.getenv("OPENAI_API_KEY"),
},
}
]
)
event_types_seen = set()
if sync_mode:
response = router.responses(
model="gpt4o-special-alias",
input="Tell me about artificial intelligence in 2 sentences.",
stream=True,
)
for event in response:
print(f"Validating event type: {event.type}")
validate_stream_event(event)
event_types_seen.add(event.type)
else:
response = await router.aresponses(
model="gpt4o-special-alias",
input="Tell me about artificial intelligence in 2 sentences.",
stream=True,
)
async for event in response:
print(f"Validating event type: {event.type}")
validate_stream_event(event)
event_types_seen.add(event.type)
# At minimum, we should see these core event types
required_events = {"response.created", "response.completed"}
missing_events = required_events - event_types_seen
assert not missing_events, f"Missing required event types: {missing_events}"
print(f"Successfully validated all event types: {event_types_seen}")
def test_mcp_tools_with_responses_api():
litellm.turn_on_debug()
MCP_TOOLS = [
{
"type": "mcp",
"server_label": "zapier",
"server_url": "https://mcp.zapier.com/api/mcp/mcp",
"headers": {
"Authorization": f"Bearer {os.getenv('ZAPIER_CI_CD_MCP_TOKEN')}"
},
}
]
MODEL = "openai/gpt-4.1"
USER_QUERY = "how does tiktoken work?"
#########################################################
# Step 1: OpenAI will use MCP LIST, and return a list of MCP calls for our approval
try:
response = litellm.responses(model=MODEL, tools=MCP_TOOLS, input=USER_QUERY)
print(response)
response = cast(ResponsesAPIResponse, response)
mcp_approval_id: Optional[str] = None
for output in response.output:
if output.type == "mcp_approval_request":
mcp_approval_id = output.id
break
# Step 2: Send followup with approval for the MCP call
if mcp_approval_id:
response_with_mcp_call = litellm.responses(
model=MODEL,
tools=MCP_TOOLS,
input=[
{
"type": "mcp_approval_response",
"approve": True,
"approval_request_id": mcp_approval_id,
}
],
previous_response_id=response.id,
)
print(response_with_mcp_call)
except litellm.APIError as e:
if (
"424" in str(e)
or "Failed Dependency" in str(e)
or "external_connector_error" in str(e)
):
pytest.skip(f"Skipping test due to external MCP server error: {e}")
else:
raise e
except litellm.InternalServerError as e:
if "500" in str(e) or "server_error" in str(e):
pytest.skip(
f"Skipping test due to OpenAI server error (likely MCP server unavailable): {e}"
)
else:
raise e
@pytest.mark.asyncio
async def test_openai_responses_api_field_types():
"""Test that specific fields in the response have the correct types"""
litellm.turn_on_debug()
litellm.set_verbose = True
# Test with store=True
response = await litellm.aresponses(
model="gpt-5.5",
input="hi",
)
# Verify created_at is an integer
assert isinstance(response.created_at, int), "created_at should be an integer"
# Verify store field is present and matches input
assert hasattr(response, "store"), "store field should be present"
assert response.store is True, "store field should match input value"
# Test without store parameter
response_without_store = await litellm.aresponses(model="gpt-5.5", input="hi")
# Verify created_at is still an integer
assert isinstance(
response_without_store.created_at, int
), "created_at should be an integer"
# Verify store field is present but None when not specified
assert hasattr(response_without_store, "store"), "store field should be present"
@pytest.mark.asyncio
async def test_openai_responses_api_token_limit_error():
"""
Relevant issue: https://github.com/BerriAI/litellm/issues/15785
Parsing the in-stream ErrorEvent must not raise
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent".
The iterator routes the event through litellm.exception_type, so it surfaces as
the typed 400 client error the non-streaming path raises (litellm.BadRequestError)
carrying the provider's message. invalid_request_error is a non-retriable client
error, so there is no MidStreamFallbackError wrapping.
"""
litellm.turn_on_debug()
# Generate text with >400k tokens to trigger token limit error
oversized_text = "This is a test sentence. " * 50000 # ~400k tokens
response = await litellm.aresponses(
model="gpt-5-mini", input=oversized_text, stream=True
)
async def _drain():
async for event in response:
print(event)
with pytest.raises(litellm.BadRequestError) as exc_info:
await _drain()
assert exc_info.value.status_code == 400
assert "exceeds the context window" in str(exc_info.value)
async def test_openai_streaming_logging():
"""Test that OpenAI Responses API streaming logging is working correctly."""
litellm.turn_on_debug()
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import Usage
class TestCustomLogger(CustomLogger):
validate_usage = False
def __init__(self):
self.standard_logging_object: Optional[StandardLoggingPayload] = None
async def async_log_success_event(
self, kwargs, response_obj, start_time, end_time
):
print(f"response_obj: {response_obj.usage}")
assert isinstance(
response_obj.usage, (Usage, dict)
), f"Expected response_obj.usage to be of type Usage or dict, but got {type(response_obj.usage)}"
# Verify it has the chat completion format fields
if isinstance(response_obj.usage, dict):
assert (
"prompt_tokens" in response_obj.usage
), "Usage dict should have prompt_tokens"
assert (
"completion_tokens" in response_obj.usage
), "Usage dict should have completion_tokens"
print("\n\nVALIDATED USAGE\n\n")
self.validate_usage = True
tcl = TestCustomLogger()
litellm.callbacks = [tcl]
request_model = "gpt-5-mini"
response = await litellm.aresponses(
model=request_model,
input="What is the capital of France?",
stream=True,
)
print("response=", json.dumps(response, indent=4, default=str))
async for event in response:
if event.type == "response.completed":
final_response = event
print("litellm response=", json.dumps(event, indent=4, default=str))
await asyncio.sleep(2)
assert tcl.validate_usage, "Usage should be validated"
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_openai_compact_responses_api(sync_mode):

View file

@ -1,261 +0,0 @@
import os
import json
import traceback
from typing import Optional
from dotenv import load_dotenv
from fastapi import Request
from datetime import datetime
from unittest.mock import AsyncMock, patch, MagicMock
from litellm import Router, CustomLogger
from litellm.types.utils import StandardLoggingPayload
## 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")
audio_file = open(file_path, "rb")
from pathlib import Path
import litellm
import pytest
import asyncio
@pytest.fixture
def model_list():
return [
{
"model_name": "gpt-5-mini",
"litellm_params": {
"model": "gpt-5-mini",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "gpt-5.5",
"litellm_params": {
"model": "gpt-5.5",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "gpt-image-1",
"litellm_params": {
"model": "gpt-image-1",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "cohere-rerank",
"litellm_params": {
"model": "cohere/rerank-english-v3.0",
"api_key": os.getenv("COHERE_API_KEY"),
},
},
{
"model_name": "claude-sonnet-4-5-20250929",
"litellm_params": {
"model": "gpt-5-mini",
"mock_response": "hi this is macintosh.",
},
},
]
# This file includes the custom callbacks for LiteLLM Proxy
# Once defined, these can be passed in proxy_config.yaml
class MyCustomHandler(CustomLogger):
def __init__(self):
self.openai_client = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
# init logging config
print("logging a transcript kwargs: ", kwargs)
print("openai client=", kwargs.get("client"))
self.openai_client = kwargs.get("client")
self.standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
except Exception:
pass
# Set litellm.callbacks = [proxy_handler_instance] on the proxy
@pytest.mark.asyncio
@pytest.mark.flaky(retries=6, delay=10)
async def test_transcription_on_router():
proxy_handler_instance = MyCustomHandler()
litellm.set_verbose = True
litellm.callbacks = [proxy_handler_instance]
print("\n Testing async transcription on router\n")
try:
model_list = [
{
"model_name": "whisper",
"litellm_params": {
"model": "whisper-1",
},
},
{
"model_name": "whisper",
"litellm_params": {
"model": "azure/azure-whisper",
"api_base": "https://my-endpoint-europe-berri-992.openai.azure.com/",
"api_key": os.getenv("AZURE_EUROPE_API_KEY"),
"api_version": "2024-02-15-preview",
},
},
]
router = Router(model_list=model_list)
router_level_clients = []
for deployment in router.model_list:
_deployment_openai_client = router._get_client(
deployment=deployment,
kwargs={"model": "whisper-1"},
client_type="async",
)
router_level_clients.append(str(_deployment_openai_client))
## test 1: user facing function
response = await router.atranscription(
model="whisper",
file=audio_file,
)
## test 2: underlying function
response = await router._atranscription(
model="whisper",
file=audio_file,
)
print(response)
# PROD Test
# Ensure we ONLY use OpenAI/Azure client initialized on the router level
await asyncio.sleep(5)
print("OpenAI Client used= ", proxy_handler_instance.openai_client)
print("all router level clients= ", router_level_clients)
assert proxy_handler_instance.openai_client in router_level_clients
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
@pytest.mark.parametrize("mode", ["iterator"]) # "file",
@pytest.mark.asyncio
async def test_audio_speech_router(mode):
litellm.set_verbose = True
test_logger = MyCustomHandler()
litellm.callbacks = [test_logger]
from litellm import Router
client = Router(
model_list=[
{
"model_name": "tts",
"litellm_params": {
"model": "openai/tts-1",
},
},
]
)
response = await client.aspeech(
model="tts",
voice="alloy",
input="the quick brown fox jumped over the lazy dogs",
api_base=None,
api_key=None,
organization=None,
project=None,
max_retries=1,
timeout=600,
client=None,
optional_params={},
)
await asyncio.sleep(3)
from litellm.llms.openai.openai import HttpxBinaryResponseContent
assert isinstance(response, HttpxBinaryResponseContent)
assert test_logger.standard_logging_object is not None
print(
"standard_logging_object=",
json.dumps(test_logger.standard_logging_object, indent=4),
)
assert test_logger.standard_logging_object["model_group"] == "tts"
@pytest.mark.asyncio()
async def test_rerank_endpoint(model_list):
from litellm.types.utils import RerankResponse
router = Router(model_list=model_list)
## Test 1: user facing function
response = await router.arerank(
model="cohere-rerank",
query="hello",
documents=["hello", "world"],
top_n=3,
)
## Test 2: underlying function
response = await router._arerank(
model="cohere-rerank",
query="hello",
documents=["hello", "world"],
top_n=3,
)
print("async re rank response: ", response)
assert response.id is not None
assert response.results is not None
RerankResponse.model_validate(response)
@pytest.mark.asyncio()
@pytest.mark.parametrize(
"model", ["omni-moderation-latest", "openai/omni-moderation-latest", None]
)
async def test_moderation_endpoint(model):
litellm.set_verbose = True
router = Router(
model_list=[
{
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
},
},
{
"model_name": "*",
"litellm_params": {
"model": "openai/*",
},
},
]
)
if model is None:
response = await router.amoderation(input="hello this is a test")
else:
response = await router.amoderation(model=model, input="hello this is a test")
print("moderation response: ", response)

View file

@ -1,526 +0,0 @@
import asyncio
import json
import os
import traceback
from dotenv import load_dotenv
from fastapi import Request
from datetime import datetime, timezone
from litellm import Router
import pytest
import litellm
from unittest.mock import patch, MagicMock, AsyncMock
from litellm.types.utils import ModelResponse, StandardLoggingPayload
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS, ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
@pytest.fixture
def model_list():
return [
{
"model_name": "gpt-5-mini",
"litellm_params": {
"model": "gpt-5-mini",
"api_key": os.getenv("OPENAI_API_KEY"),
"tpm": 1000, # Add TPM limit so async method doesn't return early
"rpm": 100, # Add RPM limit so async method doesn't return early
},
"model_info": {
"access_groups": ["group1", "group2"],
},
},
{
"model_name": "gpt-5.5",
"litellm_params": {
"model": "gpt-5.5",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "gpt-image-1",
"litellm_params": {
"model": "gpt-image-1",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "*",
"litellm_params": {
"model": "openai/*",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "claude-*",
"litellm_params": {
"model": "anthropic/*",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
},
]
def test_validate_fallbacks(model_list):
router = Router(model_list=model_list, fallbacks=[{"gpt-5.5": "gpt-5-mini"}])
router.validate_fallbacks(fallback_param=[{"gpt-5.5": "gpt-5-mini"}])
def test_routing_strategy_init(model_list):
"""Test if all routing strategies are initialized correctly"""
from litellm.types.router import RoutingStrategy
router = Router(model_list=model_list)
for strategy in RoutingStrategy:
router.routing_strategy_init(
routing_strategy=strategy, routing_strategy_args={}
)
def test_routing_strategy_init_valid_string_strategies(model_list):
"""Test that all valid string routing strategies work without error.
Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'.
"""
from litellm.types.router import RoutingStrategy
router = Router(model_list=model_list)
# All strategies from enum + simple-shuffle (default, not in enum)
valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
for strategy in valid_strategies:
# Should not raise
router.routing_strategy_init(
routing_strategy=strategy, routing_strategy_args={}
)
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.flaky(retries=6, delay=1)
@pytest.mark.asyncio
async def test_image_generation(model_list, sync_mode):
"""Test if the underlying '_image_generation' function is working correctly"""
from litellm.types.utils import ImageResponse
router = Router(model_list=model_list)
if sync_mode:
response = router._image_generation(
model="gpt-image-1",
prompt="A cute baby sea otter",
)
else:
response = await router._aimage_generation(
model="gpt-image-1",
prompt="A cute baby sea otter",
)
ImageResponse.model_validate(response)
def _rpm_tpm_router(model_id: str) -> Router:
return Router(
model_list=[
{
"model_name": "gpt-5-mini",
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
"model_info": {"id": model_id},
}
]
)
@pytest.fixture
def router_minute_pinned(monkeypatch):
pinned = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc)
monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned)
def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]:
return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")}
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_acompletion_headers_read_post_increment_counter_and_count_once():
router = _rpm_tpm_router("lit-3058-async")
response = await router.acompletion(
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
)
total_tokens = response.usage.total_tokens
assert total_tokens > 0
headers = _ratelimit_headers(response)
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
assert headers["x-ratelimit-remaining-requests"] == 99
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
await asyncio.sleep(0.5)
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_acompletion_stream_counts_request_before_headers_and_tokens_once_on_completion():
router = _rpm_tpm_router("lit-3058-stream")
stream = await router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
mock_response="pong pong pong",
stream=True,
stream_options={"include_usage": True},
)
headers = _ratelimit_headers(stream)
assert headers["x-ratelimit-remaining-tokens"] == 1000
assert headers["x-ratelimit-remaining-requests"] == 99
assert await router.get_model_group_usage("gpt-5-mini") == (0, 1)
chunks = [chunk async for chunk in stream]
total_tokens = chunks[-1].usage.total_tokens
assert total_tokens > 0
await asyncio.sleep(0.5)
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
class _GatedIncrementCache(DualCache):
def __init__(self) -> None:
super().__init__(in_memory_cache=InMemoryCache())
self.first_increment_started = asyncio.Event()
self.release_first_increment = asyncio.Event()
self.increment_calls = 0
async def async_increment_cache_pipeline(
self,
increment_list: list[RedisPipelineIncrementOperation],
local_only: bool = False,
parent_otel_span: object = None,
**kwargs: object,
) -> list[float] | None:
self.increment_calls += 1
if self.increment_calls == 1:
self.first_increment_started.set()
await self.release_first_increment.wait()
return await super().async_increment_cache_pipeline(
increment_list, local_only=local_only, parent_otel_span=parent_otel_span, **kwargs
)
@pytest.mark.asyncio
async def test_success_callback_running_during_pre_header_increment_does_not_double_count():
router = _rpm_tpm_router("lit-3058-race")
cache = _GatedIncrementCache()
router.cache = cache
request = asyncio.ensure_future(
router.acompletion(model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong")
)
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
for _ in range(50):
if get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1:
break
await asyncio.sleep(0.1)
assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1
assert cache.increment_calls == 1
cache.release_first_increment.set()
response = await request
assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1)
class _UnavailableIncrementCache(DualCache):
def __init__(self) -> None:
super().__init__(in_memory_cache=InMemoryCache())
self.first_increment_started = asyncio.Event()
self.release_first_increment = asyncio.Event()
self.increment_calls = 0
async def async_increment_cache_pipeline(
self,
increment_list: list[RedisPipelineIncrementOperation],
local_only: bool = False,
parent_otel_span: object = None,
**kwargs: object,
) -> list[float] | None:
self.increment_calls += 1
if self.increment_calls == 1:
self.first_increment_started.set()
await self.release_first_increment.wait()
raise RuntimeError("cache unavailable")
@pytest.mark.asyncio
async def test_callback_observing_stamp_before_pre_header_increment_fails_leaves_no_stamp_behind():
router = _rpm_tpm_router("lit-3058-fail")
cache = _UnavailableIncrementCache()
router.cache = cache
metadata: dict[str, object] = {}
request = asyncio.ensure_future(
router.acompletion(
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong", metadata=metadata
)
)
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30
for _ in range(50):
if get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1:
break
await asyncio.sleep(0.1)
assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1
assert cache.increment_calls == 1
cache.release_first_increment.set()
response = await request
assert response.usage.total_tokens == 30
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata
assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
def test_track_deployment_metrics(model_list):
"""Test if the 'track_deployment_metrics' function is working correctly"""
from litellm.types.utils import ModelResponse
router = Router(model_list=model_list)
router._track_deployment_metrics(
deployment=router.get_deployment_by_model_group_name(
model_group_name="gpt-5-mini"
),
response=ModelResponse(
model="gpt-5-mini",
usage={"total_tokens": 100},
),
parent_otel_span=None,
)
def test_pass_through_assistants_endpoint_factory(model_list):
"""Test if the 'pass_through_assistants_endpoint_factory' function is working correctly"""
router = Router(model_list=model_list)
router._pass_through_assistants_endpoint_factory(
original_function=litellm.acreate_assistants,
custom_llm_provider="openai",
client=None,
**{},
)
def test_factory_function(model_list):
"""Test if the 'factory_function' function is working correctly"""
router = Router(model_list=model_list)
router.factory_function(litellm.acreate_assistants)
# def test_pattern_match_deployments(model_list):
# from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
# import re
# patter_router = PatternMatchRouter()
# request = "fo::hi::static::hello"
# model_name = "fo::*:static::*"
# model_name_regex = patter_router._pattern_to_regex(model_name)
# # Match against the request
# match = re.match(model_name_regex, request)
# print(f"match: {match}")
# print(f"match.end: {match.end()}")
# if match is None:
# raise ValueError("Match not found")
# updated_model = patter_router.set_deployment_model_name(
# matched_pattern=match, litellm_deployment_litellm_model="openai/*"
# )
# assert updated_model == "openai/fo::hi:static::hello"
@pytest.mark.asyncio
async def test_pass_through_moderation_endpoint_factory(model_list):
router = Router(model_list=model_list)
response = await router._pass_through_moderation_endpoint_factory(
original_function=litellm.amoderation,
input="this is valid good text",
model=None,
)
assert response is not None
def test_handle_clientside_credential_no_metadata(model_list):
"""Test that _handle_clientside_credential handles cases where no metadata is provided"""
router = Router(model_list=model_list)
# Mock deployment
deployment = {
"model_name": "gpt-4.1",
"litellm_params": {"model": "gpt-4.1", "api_key": "test_key"},
"model_info": {"id": "original-id-789"},
}
# Mock kwargs with clientside credentials but NO metadata
kwargs = {
"api_key": "client_side_key",
"api_base": "https://api.openai.com/v1",
# No metadata key at all
}
# This should fail because there's no model_group in metadata
# The function expects to find model_group in the metadata
try:
result_deployment = router._handle_clientside_credential(
deployment=deployment, kwargs=kwargs, function_name="acompletion"
)
# If we get here, the function should have used deployment.model_name as fallback
assert result_deployment.model_name == "gpt-4.1"
print("✓ Success with no metadata - used deployment.model_name as fallback")
except Exception as e:
# This is expected behavior - the function needs model_group to generate model_id
print(f"✓ Correctly handled no metadata case: {e}")
# Test with empty metadata
kwargs_with_empty_metadata = {
"api_key": "client_side_key",
"api_base": "https://api.openai.com/v1",
"metadata": {}, # Empty metadata
}
try:
result_deployment = router._handle_clientside_credential(
deployment=deployment,
kwargs=kwargs_with_empty_metadata,
function_name="acompletion",
)
# Should fail because empty metadata has no model_group
pytest.fail("Expected failure with empty metadata")
except Exception as e:
print(f"✓ Correctly handled empty metadata case: {e}")

View file

@ -2197,6 +2197,58 @@ def test_log_event_returns_the_v2_dict_shape_for_the_alerting_trace_id_cache():
assert returned["generation_id"]
@pytest.mark.parametrize(
("metadata", "expected_source"),
[
({}, None),
({"trace_id": "my-unique-trace-id"}, "my-unique-trace-id"),
({"existing_trace_id": "my-unique-existing-trace-id"}, "my-unique-existing-trace-id"),
(
{"trace_id": "my-unique-trace-id", "existing_trace_id": "my-unique-existing-trace-id"},
"my-unique-existing-trace-id",
),
],
)
def test_logging_get_trace_id_reports_the_langfuse_trace_that_won_precedence(
monkeypatch: pytest.MonkeyPatch,
metadata: dict[str, str],
expected_source: str | None,
) -> None:
from litellm.litellm_core_utils import litellm_logging
from litellm.litellm_core_utils.litellm_logging import Logging
logger, exporter = _steering_logger()
monkeypatch.setattr(litellm_logging, "langFuseLogger", logger)
monkeypatch.setattr(litellm, "success_callback", ["langfuse"])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "callbacks", [])
call_id: Final = f"trace-precedence-{len(metadata)}-{expected_source}"
logging_obj: Final = Logging(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="completion",
litellm_call_id=call_id,
start_time=datetime.datetime.now(),
function_id="trace-precedence",
)
litellm.completion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "trace precedence"}],
mock_response="ok",
litellm_logging_obj=logging_obj,
metadata=dict(metadata),
)
deadline: Final = time.monotonic() + 5
while logging_obj.get_trace_id(service_name="langfuse") is None and time.monotonic() < deadline:
time.sleep(0.01)
expected_trace_id: Final = resolve_trace_id(expected_source or logging_obj.litellm_trace_id)
assert logging_obj.get_trace_id(service_name="langfuse") == expected_trace_id
assert _span_trace_id(_exported_span(logger, exporter)) == expected_trace_id
def test_parse_langfuse_debug_only_enables_on_true_strings():
"""v4 treats any truthy value as debug=on, so the raw env string "false" would enable debug."""
assert langfuse_module.parse_langfuse_debug("true") is True

View file

@ -1,7 +1,6 @@
"""Test health check helper functions"""
import json
import os
import socket
import struct
import zlib
@ -1235,3 +1234,193 @@ async def test_health_check_with_custom_llm_provider(
assert "error" not in response, response
assert upstream.called
assert json.loads(upstream.calls[0].request.content)["model"] == "deepseek-r1-distill-qwen-1.5B-q4"
@pytest.mark.asyncio
async def test_azure_chat_health_check_surfaces_provider_rate_limit_headers(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post(
url__regex=r"https://resource\.example/openai/deployments/gpt-4\.1-mini/chat/completions.*"
).respond(
json={
"id": "chatcmpl-health",
"object": "chat.completion",
"created": 1,
"model": "gpt-4.1-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
headers={"x-ratelimit-remaining-tokens": "42"},
)
response: Final = await ahealth_check(
{
"model": "azure/gpt-4.1-mini",
"api_key": "fake-key",
"api_base": "https://resource.example",
"api_version": "2024-06-01",
},
mode="chat",
)
assert response["x-ratelimit-remaining-tokens"] == "42"
assert upstream.called
@pytest.mark.asyncio
async def test_azure_embedding_health_check_surfaces_provider_rate_limit_headers(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post(
url__regex=r"https://resource\.example/openai/deployments/text-embedding-ada-002/embeddings.*"
).respond(
json={
"object": "list",
"data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
"model": "text-embedding-ada-002",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
},
headers={"x-ratelimit-remaining-tokens": "84"},
)
response: Final = await ahealth_check(
{
"model": "azure/text-embedding-ada-002",
"api_key": "fake-key",
"api_base": "https://resource.example",
"api_version": "2024-06-01",
},
input=["health check"],
mode="embedding",
)
assert response["x-ratelimit-remaining-tokens"] == "84"
assert upstream.called
@pytest.mark.asyncio
async def test_image_generation_health_check_returns_a_successful_response(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond(
json={"created": 1, "data": [{"b64_json": "AA=="}]}
)
response: Final = await ahealth_check(
{"model": "gpt-image-1", "api_key": "fake-key"},
mode="image_generation",
prompt="health check",
)
assert "error" not in response
assert upstream.called
@pytest.mark.asyncio
async def test_groq_wildcard_health_check_uses_a_concrete_model(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(
litellm,
"models_by_provider",
{"groq": ["groq/openai/gpt-oss-20b"]},
)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.groq.com/openai/v1/chat/completions").respond(
json={
"id": "chatcmpl-health",
"object": "chat.completion",
"created": 1,
"model": "groq/openai/gpt-oss-20b",
"service_tier": "on_demand",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "2"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
)
response: Final = await ahealth_check(
{
"model": "groq/*",
"api_key": "fake-key",
"messages": [{"role": "user", "content": "What is 1 + 1?"}],
}
)
assert upstream.called
assert json.loads(upstream.calls.last.request.content)["model"] == "openai/gpt-oss-20b"
assert response == {}
@pytest.mark.asyncio
async def test_cohere_rerank_health_check_returns_a_successful_response(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond(
json={
"id": "rerank-health",
"results": [{"index": 0, "relevance_score": 0.7}],
"meta": {"billed_units": {"search_units": 1}},
}
)
response: Final = await ahealth_check(
{"model": "cohere/rerank-english-v3.0", "api_key": "fake-key"},
mode="rerank",
prompt="health check",
)
assert "error" not in response
assert upstream.called
assert json.loads(upstream.calls.last.request.content)["query"] == "health check"
@pytest.mark.asyncio
async def test_audio_speech_health_check_returns_audio(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond(
content=b"audio",
headers={"content-type": "audio/mpeg"},
)
response: Final = await ahealth_check(
{"model": "openai/tts-1", "api_key": "fake-key"},
mode="audio_speech",
prompt="health check",
)
assert "error" not in response
assert upstream.called
assert json.loads(upstream.calls.last.request.content)["input"] == "health check"
@pytest.mark.asyncio
async def test_audio_transcription_health_check_returns_transcribed_text(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").respond(
json={"text": "health check audio"}
)
response: Final = await ahealth_check(
{"model": "openai/whisper-1", "api_key": "fake-key"},
mode="audio_transcription",
)
assert "error" not in response
assert upstream.called
assert b'name="file"' in upstream.calls.last.request.content

View file

@ -11,6 +11,7 @@ Regression tests for https://github.com/BerriAI/litellm/issues/22040 and for the
import os
import sys
from typing import Final
import httpx
import pytest
@ -21,10 +22,15 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.
import litellm
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.count_tokens.token_counter import AnthropicTokenCounter
from litellm.llms.anthropic.count_tokens.transformation import (
AnthropicCountTokensConfig,
)
from litellm.llms.azure_ai.anthropic.count_tokens.token_counter import (
AzureAIAnthropicTokenCounter,
)
from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_BETA_HEADER
from litellm.types.utils import TokenCountResponse
# Fake tokens for testing (not real secrets)
FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef"
@ -269,3 +275,80 @@ class TestCountTokensUsesWorkloadIdentity:
assert result is not None
assert result.total_tokens == 7
assert seen["auth_header"] == {"x-api-key": vault_key}
@pytest.mark.parametrize(
("counter_type", "api_base", "endpoint", "tokenizer_type"),
(
(
AnthropicTokenCounter,
"https://gateway.example",
"https://gateway.example/v1/messages/count_tokens",
"anthropic_api",
),
(
AzureAIAnthropicTokenCounter,
"https://resource.example",
"https://resource.example/anthropic/v1/messages/count_tokens",
"azure_ai_anthropic_api",
),
),
)
@pytest.mark.parametrize(("status_code", "expected_error"), ((200, False), (401, True)))
@pytest.mark.asyncio
async def test_count_token_counters_return_typed_success_and_error_responses(
counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter],
api_base: str,
endpoint: str,
tokenizer_type: str,
status_code: int,
expected_error: bool,
httpx_transport_clients: None,
) -> None:
response_body: Final = {"input_tokens": 17} if status_code == 200 else {"error": {"message": "invalid key"}}
router: Final = respx.mock
with router:
count_route: Final = router.post(endpoint).mock(return_value=httpx.Response(status_code, json=response_body))
result: Final = await counter_type().count_tokens(
model_to_use="claude-test",
messages=[{"role": "user", "content": "hi"}],
contents=None,
deployment={"litellm_params": {"api_key": "sk-ant-api03-test-key", "api_base": api_base}},
request_model="claude-test",
)
requests: Final = tuple(router.calls)
assert count_route.called
assert len(requests) == 1
assert requests[0].request.url == httpx.URL(endpoint)
assert isinstance(result, TokenCountResponse)
assert result.request_model == "claude-test"
assert result.model_used == "claude-test"
assert result.tokenizer_type == tokenizer_type
if expected_error:
assert result.error is True
assert result.status_code == status_code
assert result.total_tokens == 0
return
assert result.error is not True
assert result.total_tokens == 17
@pytest.mark.parametrize(
("counter_type", "provider"),
(
(AnthropicTokenCounter, "anthropic"),
(AzureAIAnthropicTokenCounter, "azure_ai"),
),
)
def test_count_token_counters_select_their_own_provider(
counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter],
provider: str,
) -> None:
counter: Final = counter_type()
assert counter.should_use_token_counting_api(custom_llm_provider=provider) is True
assert counter.should_use_token_counting_api(custom_llm_provider="unknown") is False

View file

@ -0,0 +1,531 @@
import asyncio
import json
import threading
from collections.abc import Iterable
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
import httpx
import pytest
import respx
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import (
IncompleteDetails,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
from litellm.types.utils import StandardLoggingPayload, Usage
from tests.unit.proxy.conftest import httpx_transport
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
_OPENAI_URL: Final = "https://api.openai.com/v1/responses"
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
_STANDARD_LOGGING_PAYLOAD: Final = TypeAdapter(StandardLoggingPayload)
_OUTPUT_TEXT: Final = "Hello from the mocked response"
_CREATED_AT: Final = 1750000000
_AsyncLoggingMode: TypeAlias = Literal["non_stream", "stream"]
_RESPONSE_FIELD_TYPES: Final = MappingProxyType(
{
"error": (dict, type(None)),
"incomplete_details": (IncompleteDetails, type(None)),
"instructions": (str, type(None)),
"metadata": (dict,),
"model": (str,),
"object": (str,),
"parallel_tool_calls": (bool, type(None)),
"temperature": (int, float, type(None)),
"tool_choice": (dict, str, type(None)),
"tools": (list, type(None)),
"top_p": (int, float, type(None)),
"max_output_tokens": (int, type(None)),
"previous_response_id": (str, type(None)),
"reasoning": (dict, type(None)),
"status": (str,),
"text": (dict,),
"truncation": (str, type(None)),
"user": (str, type(None)),
"store": (bool, type(None)),
}
)
_STREAM_EVENT_FIELDS: Final = MappingProxyType(
{
"response.created": ("response",),
"response.in_progress": ("response",),
"response.output_item.added": ("output_index", "item"),
"response.content_part.added": ("item_id", "output_index", "content_index", "part"),
"response.output_text.delta": ("item_id", "output_index", "content_index", "delta"),
"response.output_text.done": ("item_id", "output_index", "content_index", "text"),
"response.content_part.done": ("item_id", "output_index", "content_index", "part"),
"response.output_item.done": ("output_index", "item"),
"response.completed": ("response",),
}
)
class _ResponsesLoggingCapture(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.completed: Final = threading.Event()
self.payload: StandardLoggingPayload | None = None
self.usage: object = None
async def async_log_success_event(
self,
kwargs: dict[str, object],
response_obj: object,
start_time: object,
end_time: object,
) -> None:
assert isinstance(response_obj, ResponsesAPIResponse)
self.payload = _STANDARD_LOGGING_PAYLOAD.validate_python(kwargs["standard_logging_object"])
self.usage = response_obj.usage
self.completed.set()
def _response_body(response_id: str, store: bool = False, created_at: float = _CREATED_AT) -> dict[str, object]:
return {
"id": response_id,
"object": "response",
"created_at": created_at,
"status": "completed",
"error": None,
"incomplete_details": None,
"instructions": None,
"max_output_tokens": None,
"model": "gpt-4o",
"output": [
{
"id": f"msg_{response_id}",
"type": "message",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []}],
}
],
"parallel_tool_calls": True,
"previous_response_id": None,
"reasoning": {"effort": None, "summary": None},
"store": store,
"temperature": 1.0,
"text": {"format": {"type": "text"}},
"tool_choice": "auto",
"tools": [],
"top_p": 1.0,
"truncation": "disabled",
"usage": {
"input_tokens": 2,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens": 3,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 5,
},
"user": None,
"metadata": {},
}
def _sse(events: Iterable[dict[str, object]]) -> str:
return "".join(f"data: {json.dumps(event)}\n\n" for event in events)
def _response_events(response_id: str) -> tuple[dict[str, object], ...]:
completed_response: Final = _response_body(response_id)
in_progress_response: Final = {**completed_response, "status": "in_progress", "output": [], "usage": None}
item_id: Final = f"msg_{response_id}"
text_part: Final = {"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []}
position: Final = {"item_id": item_id, "output_index": 0, "content_index": 0}
return (
{"type": "response.created", "sequence_number": 0, "response": in_progress_response},
{"type": "response.in_progress", "sequence_number": 1, "response": in_progress_response},
{
"type": "response.output_item.added",
"sequence_number": 2,
"output_index": 0,
"item": {"id": item_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []},
},
{
"type": "response.content_part.added",
"sequence_number": 3,
**position,
"part": {"type": "output_text", "text": "", "annotations": []},
},
{"type": "response.output_text.delta", "sequence_number": 4, **position, "delta": _OUTPUT_TEXT},
{"type": "response.output_text.done", "sequence_number": 5, **position, "text": _OUTPUT_TEXT},
{"type": "response.content_part.done", "sequence_number": 6, **position, "part": text_part},
{
"type": "response.output_item.done",
"sequence_number": 7,
"output_index": 0,
"item": completed_response["output"][0],
},
{"type": "response.completed", "sequence_number": 8, "response": completed_response},
)
def _response_sse(response_id: str) -> str:
return _sse(_response_events(response_id))
def _sse_reply(body: str) -> httpx.Response:
return httpx.Response(status_code=200, content=body, headers={"content-type": "text/event-stream"})
def _assert_valid_response(response: object, final_chunk: bool) -> None:
assert isinstance(response, ResponsesAPIResponse)
assert isinstance(response.id, str)
assert isinstance(response.created_at, int)
assert isinstance(response.usage, ResponseAPIUsage if final_chunk else type(None))
mistyped: Final = {
field: type(response[field]).__name__
for field, expected in _RESPONSE_FIELD_TYPES.items()
if not isinstance(response[field], expected)
}
assert mistyped == {}
if final_chunk and response.status == "completed":
assert len(response.output) > 0
def _json_object(value: object) -> dict[str, object]:
return _JSON_OBJECT.validate_python(value.model_dump(mode="json") if isinstance(value, BaseModel) else value)
def _assert_valid_stream(events: tuple[object, ...], item_id: str) -> None:
event_types: Final = tuple(getattr(event, "type", None) for event in events)
assert event_types == tuple(_STREAM_EVENT_FIELDS)
missing_fields: Final = {
event_type: tuple(name for name in fields if getattr(event, name, None) is None)
for (event_type, fields), event in zip(_STREAM_EVENT_FIELDS.items(), events)
}
assert all(fields == () for fields in missing_fields.values()), missing_fields
created: Final = getattr(events[0], "response", None)
_assert_valid_response(created, final_chunk=False)
_assert_valid_response(getattr(events[1], "response", None), final_chunk=False)
completed: Final = events[-1]
assert isinstance(completed, ResponseCompletedEvent)
_assert_valid_response(completed.response, final_chunk=True)
assert completed.response.id == getattr(created, "id", None)
assert completed.response.output_text == _OUTPUT_TEXT
assert tuple(getattr(event, "sequence_number", None) for event in events) == tuple(range(len(events)))
item_ids: Final = (
getattr(getattr(events[2], "item", None), "id", None),
*(getattr(event, "item_id", None) for event in events[3:7]),
getattr(getattr(events[7], "item", None), "id", None),
)
assert item_ids == (item_id,) * 6
assert tuple(getattr(event, "output_index", None) for event in events[2:8]) == (0,) * 6
assert tuple(getattr(event, "content_index", None) for event in events[3:7]) == (0,) * 4
assert _json_object(getattr(events[3], "part", None)) == {
"type": "output_text",
"text": "",
"annotations": [],
}
assert getattr(events[4], "delta", None) == _OUTPUT_TEXT
assert getattr(events[5], "text", None) == _OUTPUT_TEXT
assert _json_object(getattr(events[6], "part", None)) == {
"type": "output_text",
"text": _OUTPUT_TEXT,
"annotations": [],
}
assert getattr(getattr(events[7], "item", None), "type", None) == "message"
assert tuple(item.id for item in completed.response.output) == (item_id,)
def _completed_response(events: tuple[object, ...]) -> ResponsesAPIResponse:
completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent))
assert len(completed_events) == 1
return completed_events[0].response
def _install_capture(monkeypatch: pytest.MonkeyPatch) -> _ResponsesLoggingCapture:
capture: Final = _ResponsesLoggingCapture()
monkeypatch.setattr(litellm, "callbacks", [capture])
monkeypatch.setattr(litellm, "success_callback", [capture])
monkeypatch.setattr(litellm, "_async_success_callback", [capture])
return capture
def _assert_logged_payload_matches(capture: _ResponsesLoggingCapture, response: ResponsesAPIResponse) -> None:
payload: Final = capture.payload
usage: Final = capture.usage
assert payload is not None
assert response.usage is not None
assert payload["prompt_tokens"] == response.usage.input_tokens
assert payload["completion_tokens"] == response.usage.output_tokens
assert payload["total_tokens"] == response.usage.input_tokens + response.usage.output_tokens
assert payload["response_cost"] > 0
assert payload["id"] == response.id
assert payload["model"] == "gpt-4o-mini"
assert payload["messages"] == [{"content": "hi", "role": "user"}]
callback_usage: Final = usage.model_dump() if isinstance(usage, Usage) else _JSON_OBJECT.validate_python(usage)
assert callback_usage["prompt_tokens"] == response.usage.input_tokens
assert callback_usage["completion_tokens"] == response.usage.output_tokens
logged_response: Final = _JSON_OBJECT.validate_python(payload["response"])
final_response: Final = response.model_dump(mode="json")
assert {key: value for key, value in logged_response.items() if key != "usage"} == {
key: value for key, value in final_response.items() if key != "usage"
}
logged_usage: Final = _JSON_OBJECT.validate_python(logged_response["usage"])
assert logged_usage["prompt_tokens"] == response.usage.input_tokens
assert logged_usage["completion_tokens"] == response.usage.output_tokens
assert logged_usage["total_tokens"] == response.usage.total_tokens
async def _async_logged_openai_response(mode: _AsyncLoggingMode) -> ResponsesAPIResponse:
match mode:
case "stream":
response_stream: Final = await litellm.aresponses(
model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True
)
return _completed_response(tuple([event async for event in response_stream]))
case "non_stream":
response: Final = await litellm.aresponses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi")
assert isinstance(response, ResponsesAPIResponse)
return response
@pytest.mark.parametrize("sync_mode", (True, False))
@pytest.mark.asyncio
async def test_responses_exposes_provider_rate_limit_headers(sync_mode: bool) -> None:
response_headers: Final = {
"x-ratelimit-limit-requests": "500",
"x-ratelimit-remaining-requests": "499",
"x-ratelimit-reset-requests": "1s",
}
with respx.mock() as mock_router:
mock_router.post(_OPENAI_URL).mock(
return_value=httpx.Response(
status_code=200,
json=_response_body("resp_rate_limit"),
headers=response_headers,
)
)
response: Final = (
litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi")
if sync_mode
else await litellm.aresponses(model="openai/gpt-4o", api_key="sk-test", input="hi")
)
requests: Final = tuple(mock_router.calls)
assert isinstance(response, ResponsesAPIResponse)
assert len(requests) == 1
additional_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["additional_headers"])
raw_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["headers"])
assert {name: additional_headers[f"llm_provider-{name}"] for name in response_headers} == response_headers
assert {name: raw_headers[name] for name in response_headers} == response_headers
@pytest.mark.asyncio
async def test_responses_converts_created_at_to_int_and_forwards_store() -> None:
response_body: Final = _response_body("resp_store", store=True, created_at=_CREATED_AT + 0.75)
with respx.mock() as mock_router:
mock_router.post(_OPENAI_URL).mock(return_value=httpx.Response(status_code=200, json=response_body))
response: Final = await litellm.aresponses(
model="openai/gpt-4o",
api_key="sk-test",
input="hi",
store=True,
)
requests: Final = tuple(mock_router.calls)
assert isinstance(response, ResponsesAPIResponse)
assert type(response.created_at) is int
assert response.created_at == _CREATED_AT
assert response.store is True
assert len(requests) == 1
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
assert request_body["store"] is True
@pytest.mark.asyncio
async def test_responses_mcp_followup_forwards_approval_and_previous_id() -> None:
mcp_tools: Final = [
{
"type": "mcp",
"server_label": "weather",
"server_url": "https://mcp.example.test",
"headers": {"Authorization": "Bearer mcp-test-token"},
}
]
approval: Final = [{"type": "mcp_approval_response", "approve": True, "approval_request_id": "approval_123"}]
with respx.mock() as mock_router:
route: Final = mock_router.post(_OPENAI_URL).mock(
return_value=httpx.Response(status_code=200, json=_response_body("resp_mcp"))
)
first_response: Final = await litellm.aresponses(
model="openai/gpt-4o",
api_key="sk-test",
input="Search for recent weather",
tools=mcp_tools,
)
assert isinstance(first_response, ResponsesAPIResponse)
second_response: Final = await litellm.aresponses(
model="openai/gpt-4o",
api_key="sk-test",
input=approval,
tools=mcp_tools,
previous_response_id=first_response.id,
)
requests: Final = tuple(route.calls)
assert isinstance(second_response, ResponsesAPIResponse)
assert len(requests) == 2
first_request: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
second_request: Final = _JSON_OBJECT.validate_json(requests[1].request.content)
assert first_request["tools"] == mcp_tools
assert second_request["tools"] == mcp_tools
assert second_request["input"] == approval
assert second_request["previous_response_id"] == "resp_mcp"
@pytest.mark.parametrize("mode", ("non_stream", "stream"))
@pytest.mark.asyncio
async def test_responses_standard_logging_matches_final_response(
mode: _AsyncLoggingMode, monkeypatch: pytest.MonkeyPatch
) -> None:
capture: Final = _install_capture(monkeypatch)
with respx.mock() as mock_router:
mock_router.post(_OPENAI_URL).mock(
return_value=(
_sse_reply(_response_sse("resp_logging"))
if mode == "stream"
else httpx.Response(status_code=200, json=_response_body("resp_logging"))
)
)
response: Final = await _async_logged_openai_response(mode)
requests: Final = tuple(mock_router.calls)
logged: Final = await asyncio.to_thread(capture.completed.wait, 10)
assert logged
assert len(requests) == 1
_assert_logged_payload_matches(capture, response)
def test_sync_stream_standard_logging_matches_final_response(monkeypatch: pytest.MonkeyPatch) -> None:
capture: Final = _install_capture(monkeypatch)
with respx.mock() as mock_router:
mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_sync_logging")))
stream: Final = litellm.responses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True)
response: Final = _completed_response(tuple(stream))
requests: Final = tuple(mock_router.calls)
logged: Final = capture.completed.wait(10)
assert logged
assert len(requests) == 1
_assert_logged_payload_matches(capture, response)
@pytest.mark.parametrize("sync_mode", (True, False))
@pytest.mark.asyncio
async def test_responses_stream_emits_valid_events(sync_mode: bool) -> None:
with respx.mock() as mock_router:
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_stream")))
events: Final = (
tuple(litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True))
if sync_mode
else tuple(
[
event
async for event in await litellm.aresponses(
model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True
)
]
)
)
requests: Final = tuple(route.calls)
assert len(requests) == 1
_assert_valid_stream(events, "msg_resp_stream")
@pytest.mark.asyncio
async def test_responses_stream_error_event_raises_bad_request() -> None:
message: Final = "Your input exceeds the context window of this model."
created: Final = _response_events("resp_too_long")[0]
error_event: Final = {
"type": "error",
"sequence_number": 1,
"error": {
"type": "invalid_request_error",
"code": "context_length_exceeded",
"message": message,
"param": "input",
},
}
with respx.mock() as mock_router:
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_sse((created, error_event))))
stream: Final = await litellm.aresponses(
model="openai/gpt-5-mini", api_key="sk-test", input="oversized prompt", stream=True
)
with pytest.raises(litellm.BadRequestError) as exc_info:
_ = [event async for event in stream]
requests: Final = tuple(route.calls)
assert len(requests) == 1
assert exc_info.value.status_code == 400
assert message in str(exc_info.value)
def _alias_router() -> litellm.Router:
return litellm.Router(
model_list=[
{
"model_name": "openai-offline-alias",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
}
]
)
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", (True, False))
async def test_router_responses_alias_uses_underlying_model(sync_mode: bool) -> None:
router: Final = _alias_router()
with respx.mock() as mock_router:
route: Final = mock_router.post(_OPENAI_URL).mock(
return_value=httpx.Response(status_code=200, json=_response_body("resp_router_alias"))
)
response: Final = (
router.responses(model="openai-offline-alias", input="hi")
if sync_mode
else await router.aresponses(model="openai-offline-alias", input="hi")
)
requests: Final = tuple(route.calls)
_assert_valid_response(response, final_chunk=True)
assert len(requests) == 1
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
assert request_body["model"] == "gpt-4o"
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", (True, False))
async def test_router_responses_alias_stream_uses_underlying_model(sync_mode: bool) -> None:
router: Final = _alias_router()
with respx.mock() as mock_router:
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_router_stream")))
events: Final = (
tuple(router.responses(model="openai-offline-alias", input="hi", stream=True))
if sync_mode
else tuple(
[
event
async for event in await router.aresponses(model="openai-offline-alias", input="hi", stream=True)
]
)
)
requests: Final = tuple(route.calls)
_assert_valid_stream(events, "msg_resp_router_stream")
assert len(requests) == 1
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
assert request_body["model"] == "gpt-4o"

View file

@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import pytest
import respx
from pydantic import JsonValue, TypeAdapter
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -24,6 +26,7 @@ from litellm.types.llms.openai import (
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import Choices, Message, ModelResponse
from tests.unit.proxy.conftest import httpx_transport
import time
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
@ -3445,3 +3448,127 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data
assert "temperature" in request_body
assert "custom_field" in request_body
assert request_body["custom_field"] == "custom_value"
@pytest.mark.asyncio
@pytest.mark.usefixtures(httpx_transport.__name__)
async def test_aresponses_forwards_previous_response_id_to_openai() -> None:
first_input: Final = "remember the first turn"
second_input: Final = "continue the conversation"
first_id: Final = "resp_previous_turn"
second_id: Final = "resp_follow_up"
response_payloads: Final = (
{
"id": first_id,
"object": "response",
"created_at": 1734366691,
"status": "completed",
"model": "gpt-4o",
"output": [
{
"type": "message",
"id": "msg_first",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "first answer", "annotations": []}],
}
],
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"total_tokens": 2,
"output_tokens_details": {"reasoning_tokens": 0},
},
},
{
"id": second_id,
"object": "response",
"created_at": 1734366692,
"status": "completed",
"model": "gpt-4o",
"output": [
{
"type": "message",
"id": "msg_second",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "second answer", "annotations": []}],
}
],
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"total_tokens": 2,
"output_tokens_details": {"reasoning_tokens": 0},
},
},
)
provider_url: Final = "https://api.openai.com/v1/responses"
with respx.mock() as router:
route: Final = router.post(provider_url).mock(
side_effect=[
httpx.Response(status_code=200, json=response_payloads[0]),
httpx.Response(status_code=200, json=response_payloads[1]),
]
)
first_response: Final = await litellm.aresponses(
model="openai/gpt-4o",
api_key="sk-test",
input=first_input,
)
assert isinstance(first_response, ResponsesAPIResponse)
second_response: Final = await litellm.aresponses(
model="openai/gpt-4o",
api_key="sk-test",
input=second_input,
previous_response_id=first_response.id,
)
calls: Final = route.calls
assert isinstance(second_response, ResponsesAPIResponse)
assert first_response.output[0].content[0].text == "first answer"
assert second_response.output[0].content[0].text == "second answer"
assert len(calls) == 2
request_adapter: Final = TypeAdapter(dict[str, JsonValue])
request_bodies: Final = tuple(request_adapter.validate_json(call.request.content) for call in calls)
assert request_bodies[0]["input"] == first_input
assert request_bodies[1]["input"] == second_input
assert request_bodies[1]["previous_response_id"] == first_id
def test_dict_responses_input_filters_unset_reasoning_fields() -> None:
test_input: Final = [
{"role": "user", "content": "test"},
{
"id": "rs_123",
"summary": [{"text": "test", "type": "summary_text"}],
"type": "reasoning",
"content": None,
"encrypted_content": None,
"status": None,
},
{
"arguments": "{}",
"call_id": "call_123",
"name": "get_today",
"type": "function_call",
"id": "fc_123",
"status": "completed",
},
]
validated_input: Final = OpenAIResponsesAPIConfig()._validate_input_param(test_input)
assert len(validated_input) == 3
reasoning_item: Final = validated_input[1]
assert reasoning_item["type"] == "reasoning"
assert "status" not in reasoning_item
assert "content" not in reasoning_item
assert "encrypted_content" not in reasoning_item
assert reasoning_item["id"] == "rs_123"
assert reasoning_item["summary"] == [{"text": "test", "type": "summary_text"}]
function_call_item: Final = validated_input[2]
assert function_call_item["type"] == "function_call"
assert function_call_item["status"] == "completed"

View file

@ -1,7 +1,10 @@
import asyncio
import json
import logging
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
@ -1131,3 +1134,166 @@ def test_transitive_probe_expansion_terminates_on_a_router_cycle():
probes = hc_module._dependency_deployments_to_probe(a_only, router.model_list, router)
assert {d["model_info"]["id"] for d in probes} == {"b-1"}
@pytest.mark.asyncio
async def test_background_audio_speech_health_check_uses_model_info_voice(
httpx_transport: None, respx_mock: respx.MockRouter
) -> None:
upstream: Final = respx_mock.post("https://speech.example/v1/audio/speech").respond(
content=b"audio",
headers={"content-type": "audio/mpeg"},
)
healthy, unhealthy, _ = await hc_module.perform_health_check(
[
{
"litellm_params": {
"model": "openai/tts-1",
"api_key": "fake-key",
"api_base": "https://speech.example/v1",
},
"model_info": {"id": "speech", "mode": "audio_speech", "health_check_voice": "nova"},
}
],
max_concurrency=1,
)
assert len(healthy) == 1
assert unhealthy == []
assert upstream.called
assert json.loads(upstream.calls.last.request.content)["voice"] == "nova"
@pytest.mark.asyncio
async def test_background_health_check_observes_the_concurrency_limit_and_queue(
httpx_transport: None, respx_mock: respx.MockRouter
) -> None:
request_started: Final = asyncio.Queue[None]()
release: Final = asyncio.Event()
async def complete_request(_: httpx.Request) -> httpx.Response:
request_started.put_nowait(None)
await release.wait()
return httpx.Response(
200,
content=b"audio",
headers={"content-type": "audio/mpeg"},
)
upstream: Final = respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=complete_request)
model_list: Final = [
{
"litellm_params": {
"model": "openai/tts-1",
"api_key": "fake-key",
"api_base": "https://health.example/v1",
},
"model_info": {"id": f"audio-{index}", "mode": "audio_speech"},
}
for index in range(10)
]
tasks_before: Final = len(asyncio.all_tasks())
perform_task: Final = asyncio.create_task(hc_module.perform_health_check(model_list, max_concurrency=2))
try:
await asyncio.wait_for(request_started.get(), timeout=1)
await asyncio.wait_for(request_started.get(), timeout=1)
for _ in range(20):
await asyncio.sleep(0)
extra_requests_started: Final = request_started.qsize()
tasks_while_blocked: Final = len(asyncio.all_tasks()) - tasks_before
finally:
release.set()
healthy, unhealthy, _ = await perform_task
assert extra_requests_started == 0
assert tasks_while_blocked <= 5
assert upstream.call_count == 10
assert len(healthy) == 10
assert unhealthy == []
@pytest.mark.asyncio
async def test_background_health_check_timeout_marks_a_blocked_provider_unhealthy(
httpx_transport: None, respx_mock: respx.MockRouter
) -> None:
never_release: Final = asyncio.Event()
request_started: Final = asyncio.Event()
async def blocked_response(_: httpx.Request) -> httpx.Response:
request_started.set()
await never_release.wait()
return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"})
respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=blocked_response)
model_list: Final = [
{
"litellm_params": {
"model": "openai/tts-1",
"api_key": "fake-key",
"api_base": "https://health.example/v1",
},
"model_info": {"id": "blocked", "mode": "audio_speech", "health_check_timeout": 2},
}
]
healthy, unhealthy, _ = await asyncio.wait_for(
hc_module.perform_health_check(model_list),
timeout=4,
)
assert request_started.is_set()
assert unhealthy[0]["error"] == "Timeout exceeded"
assert healthy == []
assert len(unhealthy) == 1
assert unhealthy[0]["model"] == "openai/tts-1"
@pytest.mark.asyncio
async def test_background_health_check_timeout_does_not_cancel_a_sibling(
httpx_transport: None, respx_mock: respx.MockRouter
) -> None:
never_release: Final = asyncio.Event()
slow_request_started: Final = asyncio.Event()
async def blocked_response(_: httpx.Request) -> httpx.Response:
slow_request_started.set()
await never_release.wait()
return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"})
respx_mock.post("https://slow.example/v1/audio/speech").mock(side_effect=blocked_response)
fast_upstream: Final = respx_mock.post("https://fast.example/v1/audio/speech").respond(
content=b"audio",
headers={"content-type": "audio/mpeg"},
)
model_list: Final = [
{
"litellm_params": {
"model": "openai/tts-1",
"api_key": "fake-key",
"api_base": "https://slow.example/v1",
},
"model_info": {"id": "slow", "mode": "audio_speech", "health_check_timeout": 1},
},
{
"litellm_params": {
"model": "openai/tts-1",
"api_key": "fake-key",
"api_base": "https://fast.example/v1",
},
"model_info": {"id": "fast", "mode": "audio_speech", "health_check_timeout": 2},
},
]
healthy, unhealthy, _ = await asyncio.wait_for(
hc_module.perform_health_check(model_list, max_concurrency=1),
timeout=4,
)
healthy_model_ids: Final = {endpoint["model_id"] for endpoint in healthy}
unhealthy_model_ids: Final = {endpoint["model_id"] for endpoint in unhealthy}
assert slow_request_started.is_set()
assert fast_upstream.called
assert healthy_model_ids == {"fast"}
assert unhealthy_model_ids == {"slow"}

View file

@ -0,0 +1,158 @@
import json
from typing import Final
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
from litellm.types.llms.openai import ResponseCompletedEvent, ResponsesAPIResponse
from tests.unit.proxy.conftest import httpx_transport
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
_GEMINI_URL: Final = (
"https://generativelanguage.googleapis.com/v1beta/models/"
"gemini-2.5-flash:(?:generateContent|streamGenerateContent).*"
)
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
def _function_call_response(signature: str) -> dict[str, object]:
return {
"candidates": [
{
"content": {
"role": "model",
"parts": [
{
"functionCall": {
"name": "get_weather",
"args": {"location": "San Francisco"},
},
"thoughtSignature": signature,
}
],
},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 4,
"candidatesTokenCount": 2,
"totalTokenCount": 6,
},
}
@pytest.mark.asyncio
async def test_web_search_preview_is_sent_as_google_search() -> None:
response_body: Final = {
"candidates": [
{
"content": {"role": "model", "parts": [{"text": "Search completed"}]},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 2,
"candidatesTokenCount": 2,
"totalTokenCount": 4,
},
}
with respx.mock() as mock_router:
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
return_value=httpx.Response(status_code=200, json=response_body)
)
response: Final = await litellm.aresponses(
model="gemini/gemini-2.5-flash",
api_key="test-key",
input="Find current weather",
tools=[{"type": "web_search_preview", "search_context_size": "low"}],
)
requests: Final = tuple(route.calls)
assert isinstance(response, ResponsesAPIResponse)
assert response.output_text == "Search completed"
assert len(requests) == 1
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
assert request_body["tools"] == [{"googleSearch": {}}]
@pytest.mark.asyncio
async def test_gemini_function_call_preserves_thought_signature() -> None:
signature: Final = "gemini-thought-signature"
response_body: Final = _function_call_response(signature)
with respx.mock() as mock_router:
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
return_value=httpx.Response(status_code=200, json=response_body)
)
response: Final = await litellm.aresponses(
model="gemini/gemini-2.5-flash",
api_key="test-key",
input="What is the weather in San Francisco?",
tools=[
{
"type": "function",
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
},
}
],
)
requests: Final = tuple(route.calls)
assert isinstance(response, ResponsesAPIResponse)
function_calls: Final = tuple(item for item in response.output if item.type == "function_call")
assert len(function_calls) == 1
assert function_calls[0].name == "get_weather"
assert function_calls[0].provider_specific_fields == {"thought_signature": signature}
assert len(requests) == 1
@pytest.mark.asyncio
async def test_gemini_streaming_function_call_preserves_thought_signature() -> None:
signature: Final = "gemini-stream-thought-signature"
response_body: Final = _function_call_response(signature)
event_stream: Final = f"data: {json.dumps(response_body)}\n\n"
with respx.mock() as mock_router:
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
return_value=httpx.Response(
status_code=200,
content=event_stream,
headers={"content-type": "text/event-stream"},
)
)
stream: Final = await litellm.aresponses(
model="gemini/gemini-2.5-flash",
api_key="test-key",
input="What is the weather in San Francisco?",
stream=True,
tools=[
{
"type": "function",
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
},
}
],
)
events: Final = tuple([event async for event in stream])
requests: Final = tuple(route.calls)
completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent))
assert len(completed_events) == 1
function_calls: Final = tuple(item for item in completed_events[0].response.output if item.type == "function_call")
assert len(function_calls) == 1
assert function_calls[0].name == "get_weather"
assert function_calls[0].provider_specific_fields == {"thought_signature": signature}
assert len(requests) == 1

View file

@ -0,0 +1,109 @@
from typing import Final, Literal, TypeAlias
import httpx
import openai
import pytest
import respx
import litellm
from tests.unit.proxy.conftest import httpx_transport
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
Provider: TypeAlias = Literal["anthropic", "gemini"]
Operation: TypeAlias = Literal["delete", "get", "cancel"]
CancelProvider: TypeAlias = Literal["openai", "azure"]
def _invoke_unsupported_response_operation(provider: Provider, operation: Operation) -> object:
match operation:
case "delete":
return litellm.delete_responses(
response_id="resp_unsupported",
custom_llm_provider=provider,
api_key="sk-test",
)
case "get":
return litellm.get_responses(
response_id="resp_unsupported",
custom_llm_provider=provider,
api_key="sk-test",
)
case "cancel":
return litellm.cancel_responses(
response_id="resp_unsupported",
custom_llm_provider=provider,
api_key="sk-test",
)
async def _invoke_cancel_response(sync_mode: bool, call_kwargs: dict[str, object]) -> object:
if sync_mode:
return litellm.cancel_responses(**call_kwargs)
return await litellm.acancel_responses(**call_kwargs)
@pytest.mark.parametrize(
("provider", "operation"),
(
("anthropic", "delete"),
("anthropic", "get"),
("anthropic", "cancel"),
("gemini", "delete"),
("gemini", "get"),
("gemini", "cancel"),
),
)
def test_unsupported_response_lifecycle_operation_fails_before_http(provider: Provider, operation: Operation) -> None:
with respx.mock() as mock_router:
with pytest.raises(litellm.APIConnectionError) as exc_info:
_invoke_unsupported_response_operation(provider, operation)
calls: Final = tuple(mock_router.calls)
assert exc_info.value.status_code == 500
assert f"not supported for {provider}" in str(exc_info.value)
assert calls == ()
@pytest.mark.parametrize(
("provider", "sync_mode"),
(("openai", True), ("openai", False), ("azure", True), ("azure", False)),
)
@pytest.mark.asyncio
async def test_cancel_responses_404_surfaces_openai_api_error_with_exact_url(
provider: CancelProvider, sync_mode: bool
) -> None:
response_id: Final = "resp_missing"
error_body: Final = {
"error": {
"message": "Response was not found",
"type": "invalid_request_error",
"code": "not_found",
}
}
api_base: Final = (
"https://api.openai.com/v1" if provider == "openai" else "https://example-resource.openai.azure.com"
)
expected_url: Final = (
f"{api_base}/responses/{response_id}/cancel"
if provider == "openai"
else f"{api_base}/openai/responses/{response_id}/cancel?api-version=2025-03-01-preview"
)
with respx.mock() as mock_router:
route: Final = mock_router.post(expected_url).mock(
return_value=httpx.Response(status_code=404, json=error_body)
)
call_kwargs: Final = {
"custom_llm_provider": provider,
"api_key": "sk-test",
"api_base": api_base,
"api_version": "2025-03-01-preview",
"response_id": response_id,
}
with pytest.raises(openai.APIError) as exc_info:
await _invoke_cancel_response(sync_mode, call_kwargs)
requests: Final = tuple(route.calls)
assert exc_info.value.status_code == 404
assert len(requests) == 1
assert str(requests[0].request.url) == expected_url

View file

@ -0,0 +1,83 @@
import base64
from dataclasses import dataclass
from typing import Final
import litellm
import pytest
import respx
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
@dataclass(frozen=True, slots=True)
class _CachedVertexCredentials:
token: str
quota_project_id: str | None
expired: bool = False
def refresh(self, request: object) -> None:
return None
def _google_secret_manager(
monkeypatch: pytest.MonkeyPatch,
project_id: str,
) -> GoogleSecretManager:
monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", project_id)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
secret_manager: Final = GoogleSecretManager()
credentials: Final = _CachedVertexCredentials(
token="test-gsm-token",
quota_project_id=project_id,
)
vertex_chat_completion: Final = litellm.vertex_chat_completion
monkeypatch.setitem(
vertex_chat_completion._credentials_project_mapping,
(None, project_id),
(credentials, project_id),
)
return secret_manager
def test_google_secret_manager_decodes_secret_and_requests_latest_version(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "test-secret-project"
secret_manager: Final = _google_secret_manager(monkeypatch, project_id)
secret_url: Final = (
f"https://secretmanager.googleapis.com/v1/projects/{project_id}/secrets/OPENAI_API_KEY/versions/latest:access"
)
encoded_secret: Final = base64.b64encode(b"anything").decode("ascii")
upstream: Final[respx.MockRouter]
with respx.mock(assert_all_called=True) as upstream:
secret_route: Final = upstream.get(secret_url).respond(
200,
json={"payload": {"data": encoded_secret}},
)
result: Final = secret_manager.get_secret_from_google_secret_manager("OPENAI_API_KEY")
assert result == "anything"
assert secret_route.called
assert len(upstream.calls) == 1
assert str(upstream.calls.last.request.url) == secret_url
assert upstream.calls.last.request.headers["Authorization"] == "Bearer test-gsm-token"
def test_google_secret_manager_returns_cached_values_without_http(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "test-secret-project"
secret_manager: Final = _google_secret_manager(monkeypatch, project_id)
secret_manager.cache.set_cache("cached-none", None)
secret_manager.cache.set_cache("cached-string", "lite-llm")
upstream: Final[respx.MockRouter]
with respx.mock() as upstream:
missing_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-none")
cached_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-string")
assert missing_value is None
assert cached_value == "lite-llm"
assert upstream.calls.call_count == 0

View file

@ -3,6 +3,7 @@ import json
import logging
import os
import time
from typing import Final
from unittest.mock import Mock, patch
import pytest
@ -230,6 +231,13 @@ def test_oidc_circleci_success(monkeypatch):
assert result == "circleci_token"
def test_oidc_circleci_v2_returns_the_environment_token(monkeypatch: pytest.MonkeyPatch) -> None:
token: Final = "circleci-v2-token"
monkeypatch.setenv("CIRCLE_OIDC_TOKEN_V2", token)
assert get_secret("oidc/circleci_v2/test-audience") == token
def test_oidc_circleci_failure(monkeypatch):
monkeypatch.delenv("CIRCLE_OIDC_TOKEN", raising=False)
secret_name = "oidc/circleci/test-audience"

View file

@ -0,0 +1,732 @@
import asyncio
import contextlib
import io
import json
import uuid
from collections.abc import Callable, Iterator
from datetime import datetime, timezone
from typing import Final
import httpx
import litellm
import pytest
import respx
from litellm import CustomLogger, Router
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
get_deployment_successes_for_current_minute,
)
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.router import RoutingStrategy
from litellm.types.utils import ImageResponse, ModelResponse, RerankResponse
from openai import AsyncAzureOpenAI, AsyncOpenAI
EVENT_TIMEOUT_SECONDS: Final = 5
ROUTING_SELECTIONS: Final = 40
ROUTING_MESSAGES: Final = ({"role": "user", "content": "route this request"},)
EXPENSIVE_COSTS: Final = {"input_cost_per_token": 1.0, "output_cost_per_token": 1.0}
CHEAP_COSTS: Final = {"input_cost_per_token": 1e-9, "output_cost_per_token": 1e-9}
CHAT_RESPONSE: Final = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
}
@pytest.fixture(autouse=True)
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.fixture
def router_minute_pinned(monkeypatch: pytest.MonkeyPatch) -> None:
pinned: Final = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc)
monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned)
class _RouterLoggingCapture(CustomLogger):
def __init__(self, model_id: str) -> None:
super().__init__()
self.model_id: Final = model_id
self.success_events: asyncio.Queue[tuple[object | None, object | None]] = asyncio.Queue()
async def async_log_success_event(
self,
kwargs: dict[str, object],
response_obj: object,
start_time: object,
end_time: object,
) -> None:
standard_logging_object: Final = kwargs.get("standard_logging_object")
if not isinstance(standard_logging_object, dict) or standard_logging_object.get("model_id") != self.model_id:
return
self.success_events.put_nowait((kwargs.get("client"), standard_logging_object))
async def next_event(self) -> tuple[object | None, object | None]:
return await asyncio.wait_for(self.success_events.get(), timeout=EVENT_TIMEOUT_SECONDS)
def _rpm_tpm_router(model_id: str) -> Router:
return Router(
model_list=[
{
"model_name": "gpt-5-mini",
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
"model_info": {"id": model_id},
}
]
)
def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]:
return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")}
def _openai_router_client() -> AsyncOpenAI:
return AsyncOpenAI(api_key="sk-fake")
def _azure_router_client() -> AsyncAzureOpenAI:
return AsyncAzureOpenAI(api_key="sk-fake", azure_endpoint="https://azure.test", api_version="2025-02-01-preview")
@pytest.mark.asyncio
@pytest.mark.parametrize(
("litellm_params", "build_client", "route_url"),
[
(
{"model": "whisper-1", "api_key": "sk-fake"},
_openai_router_client,
"https://api.openai.com/v1/audio/transcriptions",
),
(
{
"model": "azure/whisper",
"api_base": "https://azure.test",
"api_key": "sk-fake",
"api_version": "2025-02-01-preview",
},
_azure_router_client,
"https://azure.test/openai/deployments/whisper/audio/transcriptions?api-version=2025-02-01-preview",
),
],
ids=["openai", "azure"],
)
async def test_router_transcription_reuses_router_level_client_for_each_deployment(
litellm_params: dict[str, str],
build_client: Callable[[], AsyncOpenAI],
route_url: str,
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
model_id: Final = f"whisper-{uuid.uuid4().hex}"
capture: Final = _RouterLoggingCapture(model_id)
monkeypatch.setattr(litellm, "callbacks", [capture])
router: Final = Router(
model_list=[{"model_name": "whisper", "litellm_params": dict(litellm_params), "model_info": {"id": model_id}}]
)
router_level_client: Final = build_client()
router.cache.set_cache(key=f"{model_id}_async_client", value=router_level_client, local_only=True)
route: Final = respx_mock.post(route_url).respond(200, json={"text": "hello"})
response: Final = await router.atranscription(
model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav")
)
public_client, public_logging = await capture.next_event()
internal_response: Final = await router._atranscription(
model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav")
)
internal_client, _ = await capture.next_event()
upstream_requests: Final = tuple(call.request for call in route.calls)
assert public_client == str(router_level_client)
assert internal_client == str(router_level_client)
assert tuple(str(request.url) for request in upstream_requests) == (route_url, route_url)
assert all(
request.headers.get("content-type", "").startswith("multipart/form-data")
and b'filename="speech.wav"' in request.content
and b"offline audio" in request.content
for request in upstream_requests
)
assert isinstance(public_logging, dict)
assert public_logging.get("model_group") == "whisper"
assert response.text == "hello"
assert internal_response.text == "hello"
@pytest.mark.asyncio
async def test_router_speech_returns_binary_content_and_logs_model_group(
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
model_id: Final = f"tts-{uuid.uuid4().hex}"
capture: Final = _RouterLoggingCapture(model_id)
monkeypatch.setattr(litellm, "callbacks", [capture])
router: Final = Router(
model_list=[
{
"model_name": "tts",
"litellm_params": {"model": "openai/tts-1", "api_key": "sk-fake"},
"model_info": {"id": model_id},
}
]
)
route: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond(200, content=b"audio")
response: Final = await router.aspeech(model="tts", input="hello", voice="alloy")
_, standard_logging_object = await capture.next_event()
assert route.call_count == 1
assert json.loads(route.calls[0].request.content) == {"model": "tts-1", "input": "hello", "voice": "alloy"}
assert isinstance(response, HttpxBinaryResponseContent)
assert response.content == b"audio"
assert isinstance(standard_logging_object, dict)
assert standard_logging_object["model_group"] == "tts"
@pytest.mark.asyncio
async def test_router_rerank_returns_valid_response_from_public_and_underlying_calls(
respx_mock: respx.MockRouter,
) -> None:
router: Final = Router(
model_list=[
{
"model_name": "cohere-rerank",
"litellm_params": {"model": "cohere/rerank-english-v3.0", "api_key": "sk-fake"},
}
]
)
route: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond(
200,
json={
"id": "rerank-1",
"results": [{"index": 0, "relevance_score": 0.9}],
"meta": {"api_version": {"version": "2"}},
},
)
public_response: Final = await router.arerank(
model="cohere-rerank",
query="hello",
documents=["hello", "world"],
top_n=1,
)
underlying_response: Final = await router._arerank(
model="cohere-rerank",
query="hello",
documents=["hello", "world"],
top_n=1,
)
assert route.call_count == 2
request_bodies: Final = tuple(json.loads(call.request.content) for call in route.calls)
assert all(
body["model"] == "rerank-english-v3.0"
and body["query"] == "hello"
and body["documents"] == ["hello", "world"]
and body["top_n"] == 1
for body in request_bodies
)
public_validated: Final = RerankResponse.model_validate(public_response)
assert public_validated.id == "rerank-1"
assert public_validated.results[0]["relevance_score"] == 0.9
underlying_validated: Final = RerankResponse.model_validate(underlying_response)
assert underlying_validated.id == "rerank-1"
assert underlying_validated.results[0]["relevance_score"] == 0.9
@pytest.mark.asyncio
@pytest.mark.parametrize(
("model", "expected_model", "expected_api_key"),
[
("omni-moderation-latest", "omni-moderation-latest", "sk-catch-all"),
("openai/omni-moderation-latest", "omni-moderation-latest", "sk-openai-wildcard"),
(None, None, "sk-env"),
],
)
async def test_router_moderation_routes_through_wildcard_deployments(
model: str | None,
expected_model: str | None,
expected_api_key: str,
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "sk-env")
router: Final = Router(
model_list=[
{
"model_name": "openai/*",
"litellm_params": {"model": "openai/*", "api_key": "sk-openai-wildcard"},
},
{
"model_name": "*",
"litellm_params": {"model": "openai/*", "api_key": "sk-catch-all"},
},
]
)
route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond(
200,
json={
"id": "modr-1",
"model": "omni-moderation-latest",
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
},
)
response: Final = await router.amoderation(model=model, input="hello")
assert route.call_count == 1
upstream_request: Final = route.calls[0].request
expected_body: Final = {"input": "hello"} if expected_model is None else {"input": "hello", "model": expected_model}
assert json.loads(upstream_request.content) == expected_body
assert upstream_request.headers["authorization"] == f"Bearer {expected_api_key}"
assert response.id == "modr-1"
assert response.model == "omni-moderation-latest"
assert response.results[0].flagged is False
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_router_image_generation_returns_valid_image_response(
sync_mode: bool,
respx_mock: respx.MockRouter,
) -> None:
router: Final = Router(
model_list=[
{
"model_name": "gpt-image-1",
"litellm_params": {"model": "openai/gpt-image-1", "api_key": "sk-fake"},
}
]
)
route: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond(
200,
json={"created": 1700000000, "data": [{"url": "https://images.test/result.png"}]},
)
response: Final = (
router._image_generation(model="gpt-image-1", prompt="a cat")
if sync_mode
else await router._aimage_generation(model="gpt-image-1", prompt="a cat")
)
assert route.call_count == 1
request_body: Final = json.loads(route.calls[0].request.content)
assert request_body["model"] == "gpt-image-1"
assert request_body["prompt"] == "a cat"
validated: Final = ImageResponse.model_validate(response)
assert validated.data[0].url == "https://images.test/result.png"
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_router_acompletion_headers_read_post_increment_counter_and_count_once(
monkeypatch: pytest.MonkeyPatch,
) -> None:
capture: Final = _RouterLoggingCapture("lit-3058-async")
monkeypatch.setattr(litellm, "callbacks", [capture])
router: Final = _rpm_tpm_router("lit-3058-async")
response: Final = await router.acompletion(
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
)
total_tokens: Final = response.usage.total_tokens
headers: Final = _ratelimit_headers(response)
assert total_tokens > 0
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
assert headers["x-ratelimit-remaining-requests"] == 99
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
await capture.next_event()
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_router_stream_counts_request_before_headers_and_tokens_once_on_completion(
monkeypatch: pytest.MonkeyPatch,
) -> None:
capture: Final = _RouterLoggingCapture("lit-3058-stream")
monkeypatch.setattr(litellm, "callbacks", [capture])
router: Final = _rpm_tpm_router("lit-3058-stream")
stream: Final = await router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
mock_response="pong pong pong",
stream=True,
stream_options={"include_usage": True},
)
headers: Final = _ratelimit_headers(stream)
assert headers["x-ratelimit-remaining-tokens"] == 1000
assert headers["x-ratelimit-remaining-requests"] == 99
assert await router.get_model_group_usage("gpt-5-mini") == (0, 1)
chunks: Final = [chunk async for chunk in stream]
total_tokens: Final = chunks[-1].usage.total_tokens
await capture.next_event()
assert total_tokens > 0
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
def test_router_validate_fallbacks_accepts_well_formed_and_rejects_malformed_entries() -> None:
router: Final = Router(model_list=[])
assert router.validate_fallbacks([{"gpt-5.5": ["gpt-5-mini"]}, {"gpt-5-mini": ["gpt-5.5"]}]) is None
with pytest.raises(ValueError, match="must have exactly one key"):
router.validate_fallbacks([{"primary": "fallback", "other": "fallback"}])
with pytest.raises(ValueError, match="is not a dictionary"):
router.validate_fallbacks(["primary"])
def _routing_deployment(deployment_id: str, extra_params: dict[str, float]) -> dict[str, object]:
return {
"model_name": "gpt",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-fake",
"api_base": f"https://{deployment_id}.test/v1",
**extra_params,
},
"model_info": {"id": deployment_id},
}
def _routing_router(
strategy: RoutingStrategy | str,
a_params: dict[str, float] | None = None,
b_params: dict[str, float] | None = None,
) -> Router:
return Router(
model_list=[_routing_deployment("a", a_params or {}), _routing_deployment("b", b_params or {})],
routing_strategy=strategy,
disable_cooldowns=True,
num_retries=0,
)
async def _selected_deployment_ids(router: Router) -> frozenset[str]:
deployments: Final = [
await router.async_get_available_deployment(model="gpt", messages=list(ROUTING_MESSAGES), request_kwargs={})
for _ in range(ROUTING_SELECTIONS)
]
return frozenset(deployment["model_info"]["id"] for deployment in deployments)
def _enum_and_string(strategy: RoutingStrategy) -> list[RoutingStrategy | str]:
return [strategy, strategy.value]
@pytest.mark.asyncio
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.COST_BASED))
async def test_router_cost_based_routing_selects_cheapest_deployment(strategy: RoutingStrategy | str) -> None:
router: Final = _routing_router(strategy, EXPENSIVE_COSTS, CHEAP_COSTS)
assert await _selected_deployment_ids(router) == frozenset({"b"})
@pytest.mark.asyncio
@pytest.mark.parametrize(
"strategy",
[*_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING), *_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING_V2)],
)
async def test_router_usage_based_routing_skips_deployment_over_tpm_limit(strategy: RoutingStrategy | str) -> None:
router: Final = _routing_router(strategy, {"tpm": 1}, {"tpm": 1_000_000})
assert await _selected_deployment_ids(router) == frozenset({"b"})
@pytest.mark.asyncio
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LATENCY_BASED))
async def test_router_latency_based_routing_avoids_deployment_that_timed_out(
strategy: RoutingStrategy | str,
respx_mock: respx.MockRouter,
) -> None:
router: Final = _routing_router(strategy)
timed_out: Final = respx_mock.post("https://a.test/v1/chat/completions").mock(
side_effect=httpx.ReadTimeout("upstream timed out")
)
respx_mock.post("https://b.test/v1/chat/completions").respond(200, json=CHAT_RESPONSE)
for _ in range(5):
if timed_out.called:
break
with contextlib.suppress(litellm.Timeout):
await router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES), max_retries=0)
assert timed_out.called
assert await _selected_deployment_ids(router) == frozenset({"b"})
@pytest.mark.asyncio
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LEAST_BUSY))
async def test_router_least_busy_routing_avoids_deployment_with_request_in_flight(
strategy: RoutingStrategy | str,
respx_mock: respx.MockRouter,
) -> None:
router: Final = _routing_router(strategy)
upstream_hosts: Final = asyncio.Queue[str]()
release_upstream: Final = asyncio.Event()
async def hold_request(request: httpx.Request) -> httpx.Response:
upstream_hosts.put_nowait(request.url.host.split(".")[0])
await release_upstream.wait()
return httpx.Response(200, json=CHAT_RESPONSE)
respx_mock.post(url__regex=r"https://[ab]\.test/v1/chat/completions").mock(side_effect=hold_request)
in_flight: Final = asyncio.create_task(router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES)))
try:
busy_id: Final = await asyncio.wait_for(upstream_hosts.get(), timeout=EVENT_TIMEOUT_SECONDS)
selected: Final = await _selected_deployment_ids(router)
finally:
release_upstream.set()
await in_flight
assert selected == frozenset({"a", "b"} - {busy_id})
@pytest.mark.asyncio
async def test_router_simple_shuffle_ignores_cost_and_spreads_across_deployments() -> None:
router: Final = _routing_router("simple-shuffle", EXPENSIVE_COSTS, CHEAP_COSTS)
assert await _selected_deployment_ids(router) == frozenset({"a", "b"})
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.PROVIDER_BUDGET_LIMITING))
def test_router_routing_strategy_init_accepts_provider_budget_strategy(strategy: RoutingStrategy | str) -> None:
router: Final = _routing_router(strategy)
router.routing_strategy_init(routing_strategy=strategy, routing_strategy_args={})
assert router.get_settings()["routing_strategy"] == "provider-budget-routing"
def test_router_track_deployment_metrics_updates_observable_usage() -> None:
router: Final = Router(
model_list=[
{
"model_name": "gpt-5-mini",
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake"},
"model_info": {"id": "metrics-deployment"},
}
]
)
deployment: Final = router.model_list[0]
router._track_deployment_metrics(deployment=deployment, parent_otel_span=None)
router._track_deployment_metrics(deployment=deployment, parent_otel_span=None)
assert router.cache.get_cache(key="metrics-deployment", local_only=True) == 2
class _GatedIncrementCache(DualCache):
def __init__(self) -> None:
super().__init__(in_memory_cache=InMemoryCache())
self.first_increment_started = asyncio.Event()
self.release_first_increment = asyncio.Event()
self.deployment_success_incremented = asyncio.Event()
self.increment_calls = 0
def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int:
result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs)
if key.endswith(":successes"):
self.deployment_success_incremented.set()
return result
async def async_increment_cache_pipeline(
self,
increment_list: list[RedisPipelineIncrementOperation],
local_only: bool = False,
parent_otel_span: object = None,
**kwargs: object,
) -> list[float] | None:
self.increment_calls += 1
if self.increment_calls == 1:
self.first_increment_started.set()
await self.release_first_increment.wait()
return await super().async_increment_cache_pipeline(
increment_list=increment_list,
local_only=local_only,
parent_otel_span=parent_otel_span,
**kwargs,
)
class _UnavailableIncrementCache(DualCache):
def __init__(self) -> None:
super().__init__(in_memory_cache=InMemoryCache())
self.first_increment_started = asyncio.Event()
self.release_first_increment = asyncio.Event()
self.deployment_success_incremented = asyncio.Event()
self.increment_calls = 0
def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int:
result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs)
if key.endswith(":successes"):
self.deployment_success_incremented.set()
return result
async def async_increment_cache_pipeline(
self,
increment_list: list[RedisPipelineIncrementOperation],
local_only: bool = False,
parent_otel_span: object = None,
**kwargs: object,
) -> list[float] | None:
self.increment_calls += 1
if self.increment_calls == 1:
self.first_increment_started.set()
await self.release_first_increment.wait()
raise RuntimeError("cache unavailable")
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_router_success_callback_during_pre_header_increment_does_not_double_count() -> None:
router: Final = _rpm_tpm_router("lit-3058-race")
cache: Final = _GatedIncrementCache()
router.cache = cache
request: Final = asyncio.create_task(
router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
mock_response="pong",
)
)
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS)
await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS)
assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1
assert cache.increment_calls == 1
cache.release_first_increment.set()
response: Final = await request
assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1)
@pytest.mark.asyncio
@pytest.mark.usefixtures("router_minute_pinned")
async def test_router_failed_pre_header_increment_clears_counted_tokens_stamp() -> None:
router: Final = _rpm_tpm_router("lit-3058-fail")
cache: Final = _UnavailableIncrementCache()
router.cache = cache
metadata: Final[dict[str, object]] = {}
request: Final = asyncio.create_task(
router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
mock_response="pong",
metadata=metadata,
)
)
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS)
await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS)
assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30
assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1
assert cache.increment_calls == 1
cache.release_first_increment.set()
response: Final = await request
assert response.usage.total_tokens == 30
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata
assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
ASSISTANT_RESPONSE: Final = {
"object": "assistant",
"created_at": 1700000000,
"name": "offline",
"description": None,
"model": "gpt-4o-mini",
"instructions": "hello",
"tools": [],
"metadata": {},
}
@pytest.mark.asyncio
async def test_router_assistants_endpoint_factory_invokes_provider(
respx_mock: respx.MockRouter,
) -> None:
router: Final = Router(model_list=[])
route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond(
200, json={**ASSISTANT_RESPONSE, "id": "asst-1"}
)
response: Final = await router._pass_through_assistants_endpoint_factory(
original_function=litellm.acreate_assistants,
custom_llm_provider="openai",
model="gpt-4o-mini",
api_key="sk-fake",
name="offline",
)
assert route.call_count == 1
assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"}
assert response.id == "asst-1"
@pytest.mark.asyncio
async def test_router_factory_function_returns_invokable_assistants_wrapper(
respx_mock: respx.MockRouter,
) -> None:
router: Final = Router(model_list=[])
route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond(
200, json={**ASSISTANT_RESPONSE, "id": "asst-2"}
)
wrapper: Final = router.factory_function(litellm.acreate_assistants, call_type="assistants")
response: Final = await wrapper(
custom_llm_provider="openai",
model="gpt-4o-mini",
api_key="sk-fake",
name="offline",
)
assert route.call_count == 1
assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"}
assert response.id == "asst-2"
@pytest.mark.asyncio
async def test_router_moderation_endpoint_factory_invokes_default_model(
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "sk-fake")
router: Final = Router(model_list=[])
route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond(
200,
json={
"id": "modr-2",
"model": "omni-moderation-latest",
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
},
)
response: Final = await router._pass_through_moderation_endpoint_factory(
original_function=litellm.amoderation,
custom_llm_provider="openai",
input="hello",
model=None,
api_key="sk-fake",
)
assert route.call_count == 1
assert json.loads(route.calls[0].request.content) == {"input": "hello"}
assert response.id == "modr-2"

View file

@ -37,6 +37,10 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.humanloop import HumanloopLogger
from litellm.llms.base_llm.audio_transcription.transformation import BaseAudioTranscriptionConfig
from litellm.integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement
from litellm.litellm_core_utils import litellm_logging
from litellm.litellm_core_utils.duration_parser import (
_extract_from_regex,
duration_in_seconds,
@ -50,7 +54,7 @@ from litellm.proxy.utils import is_valid_api_key
from litellm.types.caching import CachingSupportedCallTypes
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams, LiteLLM_Params
from litellm.types.utils import (
ADDRESSED_RESPONSE_ID_FIELD,
CallTypes,
@ -74,6 +78,7 @@ from litellm.types.utils import (
from litellm.types.videos.main import VideoObject
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
add_custom_logger_callback_to_specific_event,
check_valid_key,
CustomStreamWrapper,
filter_out_litellm_params,
@ -6802,6 +6807,7 @@ def setup_and_teardown():
MODEL: Final = "anthropic/claude-haiku-4-5"
@pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown")
def test_validate_tool_choice_none():
"""Test that None is returned as-is."""
@ -6946,6 +6952,7 @@ _SCALAR_DEFAULTS = {
"api_key": getattr(litellm, "api_key", None),
}
@pytest.fixture(scope="module")
def setup_and_teardown_local_testing():
"""
@ -8786,6 +8793,407 @@ def test_get_valid_models_from_provider():
assert "gpt-5-mini" in valid_models
@pytest.mark.parametrize(
("provider", "api_base", "api_key", "model_id"),
[
("anthropic", "https://anthropic.models.test", "anthropic-test-key", "claude-test-model"),
("xai", "https://xai.models.test", "xai-test-key", "grok-test-model"),
],
)
def test_get_valid_models_discovers_provider_models_from_http(
provider: str,
api_base: str,
api_key: str,
model_id: str,
) -> None:
response_body: Final = {
"data": [
{
"id": model_id,
"type": "model",
"display_name": "Test model",
"created_at": "2024-01-01T00:00:00Z",
}
],
"has_more": False,
"first_id": model_id,
"last_id": model_id,
}
models_url: Final = f"{api_base}/v1/models"
upstream: Final[respx.MockRouter]
with respx.mock(assert_all_called=True) as upstream:
model_list_route: Final = upstream.get(models_url).respond(200, json=response_body)
discovered_models: Final = get_valid_models(
check_provider_endpoint=True,
custom_llm_provider=provider,
api_key=api_key,
api_base=api_base,
)
assert discovered_models == [f"{provider}/{model_id}"]
assert model_list_route.called
assert len(upstream.calls) == 1
def test_check_valid_key_returns_false_for_http_unauthorized() -> None:
response_body: Final = {
"error": {
"message": "Invalid API key",
"type": "invalid_request_error",
"param": None,
"code": "invalid_api_key",
}
}
upstream: Final[respx.MockRouter]
with respx.mock(assert_all_called=True) as upstream:
invalid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond(
401,
json=response_body,
)
valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="invalid-test-key")
assert valid_key is False
assert invalid_key_route.called
assert upstream.calls.last.request.headers["Authorization"] == "Bearer invalid-test-key"
def test_check_valid_key_returns_true_for_successful_completion() -> None:
response_body: Final = {
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 0,
"model": "gpt-5-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "ok"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
upstream: Final[respx.MockRouter]
with respx.mock(assert_all_called=True) as upstream:
valid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond(
200,
json=response_body,
)
valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="valid-test-key")
assert valid_key is True
assert valid_key_route.called
assert upstream.calls.last.request.headers["Authorization"] == "Bearer valid-test-key"
def test_function_to_dict_parses_numpy_docstring_schema() -> None:
pytest.importorskip("numpydoc")
def get_current_weather(location: str, unit: str) -> str:
"""Get the current weather in a given location
Parameters
----------
location : str
The city and state, e.g. San Francisco, CA
unit : {'celsius', 'fahrenheit'}
Temperature unit
Returns
-------
str
A sentence indicating the weather
"""
return f"Weather for {location} in {unit}"
schema: Final = litellm.utils.function_to_dict(get_current_weather)
assert schema["name"] == "get_current_weather"
assert schema["description"] == "Get the current weather in a given location"
assert schema["parameters"]["type"] == "object"
assert schema["parameters"]["properties"]["location"] == {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
}
assert schema["parameters"]["properties"]["unit"]["type"] == "string"
assert schema["parameters"]["properties"]["unit"]["description"] == "Temperature unit"
assert schema["parameters"]["required"] == ["location", "unit"]
def test_duration_in_seconds_one_month_uses_the_fixed_calendar_interval(
monkeypatch: pytest.MonkeyPatch,
) -> None:
fixed_start: Final = datetime(2025, 2, 15, 12, 0, 0, 123456)
fixed_timestamp: Final = fixed_start.timestamp()
monkeypatch.setattr(
"litellm.litellm_core_utils.duration_parser.time_module.time",
lambda: fixed_timestamp,
)
assert duration_in_seconds("1mo") == 28 * 24 * 60 * 60
def test_prompt_caching_image_check_uses_default_image_dimensions() -> None:
image_bytes: Final = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/x+AAwMCAO+ip1sAAAAASUVORK5CYII="
)
image_url: Final = "https://93.184.216.34/test.png"
messages: Final = [
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this image?"},
{
"type": "image_url",
"image_url": {"url": image_url, "detail": "high"},
},
],
}
]
with respx.mock(assert_all_called=False) as upstream:
image_route: Final = upstream.get(image_url).respond(200, content=image_bytes)
cacheable: Final = is_prompt_caching_valid_prompt(
model="gpt-4o-mini",
messages=messages,
custom_llm_provider="openai",
min_token_count=100_000,
)
assert cacheable is False
assert image_route.called is False
assert len(upstream.calls) == 0
def test_get_valid_models_discovers_fireworks_models_from_http(
monkeypatch: pytest.MonkeyPatch,
) -> None:
api_base: Final = "https://fireworks.models.test/v1"
api_key: Final = "fireworks-test-key"
account_id: Final = "fireworks-test-account"
model_name: Final = "accounts/fireworks/models/llama-test-model"
models_url: Final = f"https://fireworks.models.test/v1/accounts/{account_id}/models"
monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", account_id)
upstream: Final[respx.MockRouter]
with respx.mock(assert_all_called=True) as upstream:
model_list_route: Final = upstream.get(models_url).respond(
200,
json={"models": [{"name": model_name}]},
)
discovered_models: Final = get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="fireworks_ai",
api_key=api_key,
api_base=api_base,
)
assert discovered_models == [f"fireworks_ai/{model_name}"]
assert model_list_route.called
assert upstream.calls.last.request.headers["Authorization"] == f"Bearer {api_key}"
def test_get_valid_models_returns_static_fireworks_models_without_endpoint_check(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(litellm, "check_provider_endpoint", False)
monkeypatch.setenv("FIREWORKS_AI_API_KEY", "fireworks-test-key")
expected_models: Final = litellm.models_by_provider["fireworks_ai"]
actual_models: Final = get_valid_models(
check_provider_endpoint=False,
custom_llm_provider="fireworks_ai",
)
env_inferred_models: Final = get_valid_models()
assert set(actual_models) == expected_models
assert actual_models
assert expected_models <= set(env_inferred_models)
def test_get_valid_models_uses_the_litellm_params_anthropic_api_key() -> None:
model_id: Final = "claude-test-model"
models_url: Final = "https://api.anthropic.com/v1/models"
response_body: Final = {
"data": [
{
"id": model_id,
"type": "model",
"display_name": "Test Claude",
"created_at": "2024-01-01T00:00:00Z",
}
],
"has_more": False,
"first_id": model_id,
"last_id": model_id,
}
upstream: Final[respx.MockRouter]
def response_for_api_key(request: httpx.Request) -> httpx.Response:
if request.headers["x-api-key"] == "bad-test-key":
return httpx.Response(401, json={"error": {"message": "invalid key"}}, request=request)
return httpx.Response(200, json=response_body, request=request)
with respx.mock(assert_all_called=True) as upstream:
model_list_route: Final = upstream.get(models_url).mock(side_effect=response_for_api_key)
bad_key_models: Final = get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="anthropic",
litellm_params=LiteLLM_Params(model="anthropic/*", api_key="bad-test-key"),
)
good_key_models: Final = get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="anthropic",
litellm_params=LiteLLM_Params(model="anthropic/*", api_key="good-test-key"),
)
assert bad_key_models == []
assert good_key_models == [f"anthropic/{model_id}"]
assert model_list_route.called
assert [call.request.headers["x-api-key"] for call in upstream.calls] == [
"bad-test-key",
"good-test-key",
]
def test_add_custom_logger_to_success_callback_registers_once(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm_logging, "_in_memory_loggers", [])
add_custom_logger_callback_to_specific_event("langfuse", "success")
assert len(litellm.success_callback) == 1
assert isinstance(litellm.success_callback[0], LangfusePromptManagement)
assert len(litellm._async_success_callback) == 1
assert isinstance(litellm._async_success_callback[0], LangfusePromptManagement)
assert litellm.failure_callback == []
assert litellm._async_failure_callback == []
@pytest.mark.parametrize(
"registered_lists",
[
("success_callback", "_async_success_callback"),
("success_callback",),
],
)
def test_add_custom_logger_callback_does_not_duplicate_existing_success_logger(
monkeypatch: pytest.MonkeyPatch,
registered_lists: tuple[str, ...],
) -> None:
logger: Final = HumanloopLogger()
async_success_callbacks: Final = [logger] if "_async_success_callback" in registered_lists else []
success_callbacks: Final = [logger] if "success_callback" in registered_lists else []
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "_async_success_callback", async_success_callbacks)
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", success_callbacks)
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm_logging, "_in_memory_loggers", [])
add_custom_logger_callback_to_specific_event("humanloop", "success")
assert sum(type(callback) is HumanloopLogger for callback in litellm.success_callback) == int(
"success_callback" in registered_lists
)
assert sum(type(callback) is HumanloopLogger for callback in litellm._async_success_callback) == int(
"_async_success_callback" in registered_lists
)
assert litellm.failure_callback == []
assert litellm._async_failure_callback == []
@pytest.mark.parametrize(
("registered_lists", "expected_async_success_callback_count"),
[
(("success_callback", "_async_success_callback"), 1),
(("success_callback",), 0),
],
)
@pytest.mark.asyncio
async def test_acompletion_does_not_duplicate_a_logger_already_in_success_callbacks(
monkeypatch: pytest.MonkeyPatch,
registered_lists: tuple[str, ...],
expected_async_success_callback_count: int,
) -> None:
logger: Final = HumanloopLogger()
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "input_callback", [])
monkeypatch.setattr(litellm, "success_callback", [logger])
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_input_callback", [])
monkeypatch.setattr(
litellm, "_async_success_callback", [logger] if "_async_success_callback" in registered_lists else []
)
monkeypatch.setattr(litellm, "_async_failure_callback", [])
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "callback registration"}],
mock_response="ok",
)
assert litellm.success_callback == [logger]
assert litellm._async_success_callback == [logger] * expected_async_success_callback_count
@pytest.mark.asyncio
async def test_custom_logger_in_global_callbacks_registers_once_across_completion_calls(
monkeypatch: pytest.MonkeyPatch,
) -> None:
logger: Final = HumanloopLogger()
monkeypatch.setattr(litellm, "callbacks", [logger])
monkeypatch.setattr(litellm, "input_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_input_callback", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
for _ in range(11):
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "callback registration"}],
mock_response="ok",
)
assert litellm.callbacks == [logger]
assert litellm.input_callback == [logger]
assert litellm.success_callback == [logger]
assert litellm.failure_callback == [logger]
assert litellm._async_input_callback == []
assert litellm._async_success_callback == [logger]
assert litellm._async_failure_callback == [logger]
def test_get_provider_audio_transcription_config_resolves_for_every_provider() -> None:
configs: Final = {
provider: ProviderConfigManager.get_provider_audio_transcription_config(model="whisper-1", provider=provider)
for provider in LlmProviders
}
unexpected: Final = {
provider: config
for provider, config in configs.items()
if config is not None and not isinstance(config, BaseAudioTranscriptionConfig)
}
assert unexpected == {}
assert isinstance(configs[LlmProviders.OPENAI], litellm.OpenAIWhisperAudioTranscriptionConfig)
def test_get_valid_models_from_provider_cache_invalidation(monkeypatch):
"""
Test that get_valid_models returns the correct models for a given provider