test: make 38 legacy live base-class translation tests offline (#45351)

* test: make 38 legacy live base-class translation tests offline

* test: tighten offline base-class translation test assertions

* test: cover remaining nova invoke and xai base copies offline

* test: assert outbound bodies, restore dropped providers and fix shared tool dict mutation in base chat translation tests

* test: assert outbound thinking bodies and streamed thinking blocks across anthropic and bedrock converse

* test: restore volcengine/voyage max_retries and openai timestamp_granularities coverage, assert rerank id/meta/cost and scope the interactions transport fixture

* test: call the raw provider helper for thinking stream cases

* test: flatten thinking blocks before collecting signatures

---------

Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 23:18:25 +00:00 • committed by GitHub
parent 4f62bbfd8b
commit 805bb6888b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
30 changed files with 2609 additions and 1751 deletions

View file

@ -1,20 +1,11 @@
import httpx
import json
import pytest
from typing import Any, Dict, List
from unittest.mock import MagicMock, Mock, patch
import os
from litellm._uuid import uuid
import litellm
from litellm import transcription
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
)
from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
)
from litellm.utils import ProviderConfigManager
from abc import ABC, abstractmethod
pwd = os.path.dirname(os.path.realpath(__file__))
@ -24,7 +15,6 @@ file_path = os.path.join(pwd, "gettysburg.wav")
audio_file = open(file_path, "rb")
class BaseLLMAudioTranscriptionTest(ABC):
@abstractmethod
def get_base_audio_transcription_call_args(self) -> dict:
@ -36,65 +26,3 @@ class BaseLLMAudioTranscriptionTest(ABC):
"""Must return the custom llm provider"""
pass
def test_audio_transcription(self):
"""
Test that the audio transcription is translated correctly.
"""
litellm.set_verbose = True
transcription_call_args = self.get_base_audio_transcription_call_args()
transcript = transcription(**transcription_call_args, file=audio_file)
print(f"transcript: {transcript.model_dump()}")
print(f"transcript hidden params: {transcript._hidden_params}")
assert transcript.text is not None
@pytest.mark.asyncio
async def test_audio_transcription_async(self):
"""
Test that the audio transcription is translated correctly.
"""
litellm.set_verbose = True
litellm.turn_on_debug()
AUDIO_FILE = open(file_path, "rb")
transcription_call_args = self.get_base_audio_transcription_call_args()
transcript = await litellm.atranscription(
**transcription_call_args, file=AUDIO_FILE
)
print(f"transcript: {transcript.model_dump()}")
print(f"transcript hidden params: {transcript._hidden_params}")
assert transcript.text is not None
def test_audio_transcription_optional_params(self):
"""
Test that the audio transcription is translated correctly.
"""
transcription_args = self.get_base_audio_transcription_call_args()
model = transcription_args["model"]
custom_llm_provider = self.get_custom_llm_provider()
optional_params = get_supported_openai_params(
model=model,
custom_llm_provider=custom_llm_provider.value,
request_type="transcription",
)
print(f"optional_params: {optional_params}")
assert optional_params is not None
assert (
"max_completion_tokens" not in optional_params
) # assert default chat completion response not returned
def test_audio_transcription_config(self):
"""
Test that the audio transcription config is implemented and correctly instrumented.
"""
transcription_args = self.get_base_audio_transcription_call_args()
model = transcription_args["model"]
custom_llm_provider = self.get_custom_llm_provider()
config = ProviderConfigManager.get_provider_audio_transcription_config(
model=model,
provider=custom_llm_provider,
)
print(f"config: {config}")
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)

View file

@ -4,17 +4,14 @@ import json
import pytest
from typing import Any, Dict, List
from unittest.mock import MagicMock, Mock, patch
import os
import litellm
from litellm import embedding
from litellm.exceptions import BadRequestError
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.utils import (
CustomStreamWrapper,
get_supported_openai_params,
get_optional_params,
get_optional_params_embeddings,
)
import base64
from pathlib import Path
@ -26,7 +23,6 @@ file_data = (Path(__file__).parent.parent / "white_100x100.png").read_bytes()
encoded_file = base64.b64encode(file_data).decode("utf-8")
base64_image = f"data:image/png;base64,{encoded_file}"
class BaseLLMEmbeddingTest(ABC):
"""
Abstract base test class that enforces a common test across all test classes.
@ -66,23 +62,3 @@ class BaseLLMEmbeddingTest(ABC):
CreateEmbeddingResponse.model_validate(response.model_dump())
def test_embedding_optional_params_max_retries(self):
embedding_call_args = self.get_base_embedding_call_args()
optional_params = get_optional_params_embeddings(
**embedding_call_args, max_retries=20
)
assert optional_params["max_retries"] == 20
def test_image_embedding(self):
litellm.set_verbose = True
from litellm.utils import supports_embedding_image_input
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
base_embedding_call_args = self.get_base_embedding_call_args()
if not supports_embedding_image_input(base_embedding_call_args["model"], None):
print("Model does not support embedding image input")
pytest.skip("Model does not support embedding image input")
embedding(**base_embedding_call_args, input=[base64_image])

File diff suppressed because it is too large Load diff

View file

@ -1,10 +1,8 @@
import asyncio
import httpx
import json
import pytest
from typing import Any, Dict, List
from unittest.mock import MagicMock, Mock, patch
import os
import litellm
from litellm.exceptions import BadRequestError
@ -18,7 +16,6 @@ from litellm.utils import (
# test_example.py
from abc import ABC, abstractmethod
def assert_response_shape(response, custom_llm_provider):
expected_response_shape = {"id": str, "results": list, "meta": dict}
@ -63,7 +60,6 @@ def assert_response_shape(response, custom_llm_provider):
expected_billed_units_shape["search_units"],
)
class BaseLLMRerankTest(ABC):
"""
Abstract base test class that enforces a common test across all test classes.
@ -87,57 +83,3 @@ class BaseLLMRerankTest(ABC):
"""
return None
@pytest.mark.asyncio()
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_basic_rerank(self, sync_mode):
litellm.turn_on_debug()
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
rerank_call_args = self.get_base_rerank_call_args()
custom_llm_provider = self.get_custom_llm_provider()
if sync_mode is True:
response = litellm.rerank(
**rerank_call_args,
query="hello",
documents=["hello", "world"],
top_n=2,
)
print("re rank response: ", response)
assert response.id is not None
assert response.results is not None
assert response._hidden_params["response_cost"] is not None
# Check expected cost
expected_cost = self.get_expected_cost()
if expected_cost is not None:
# If expected cost is specified, check exact match or >= for 0
if expected_cost == 0.0:
assert response._hidden_params["response_cost"] >= 0
else:
assert response._hidden_params["response_cost"] == expected_cost
else:
# Default behavior: cost should be greater than 0
assert response._hidden_params["response_cost"] > 0
assert_response_shape(
response=response, custom_llm_provider=custom_llm_provider.value
)
else:
response = await litellm.arerank(
**rerank_call_args,
query="hello",
documents=["hello", "world"],
top_n=2,
)
print("async re rank response: ", response)
assert response.id is not None
assert response.results is not None
assert_response_shape(
response=response, custom_llm_provider=custom_llm_provider.value
)

View file

@ -12,7 +12,6 @@ import pytest
import litellm.interactions as interactions
class BaseInteractionsTest(ABC):
"""Abstract base class for interactions API tests.
@ -102,17 +101,3 @@ class BaseInteractionsTest(ABC):
assert len(chunks) > 0
@pytest.mark.asyncio
async def test_acreate_simple(self):
"""Test async interaction creation."""
api_key = self.get_api_key()
if not api_key:
pytest.skip(f"API key not set for {self.__class__.__name__}")
response = await interactions.acreate(
model=self.get_model(),
input="What is the speed of light?",
api_key=api_key,
)
assert response is not None
assert response.id is not None or response.status is not None

View file

@ -32,7 +32,6 @@ from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
from httpx import Headers
from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest
def streaming_format_tests(chunk: dict, idx: int):
"""
1st chunk - chunk.get("type") == "message_start"
@ -46,7 +45,6 @@ def streaming_format_tests(chunk: dict, idx: int):
elif idx == 2:
assert chunk.get("type") == "content_block_delta"
anthropic_chunk_list = [
{
"type": "content_block_start",
@ -234,17 +232,6 @@ anthropic_chunk_list = [
{"type": "message_stop"},
]
@pytest.mark.parametrize(
"tool_type, tool_config, message_content",
[
@ -295,16 +282,8 @@ def test_anthropic_tool_use(tool_type, tool_config, message_content):
except litellm.InternalServerError:
pass
from litellm import completion
class TestAnthropicCompletion(BaseLLMChatTest, BaseAnthropicChatTest):
def get_base_completion_call_args(self) -> dict:
return {"model": "anthropic/claude-sonnet-4-5-20250929"}
@ -380,33 +359,10 @@ class TestAnthropicCompletion(BaseLLMChatTest, BaseAnthropicChatTest):
@pytest.mark.asyncio
async def test_pdf_handling(self, pdf_messages, sync_mode):
await super().test_pdf_handling(pdf_messages, sync_mode)
test_content_list_handling = None
test_image_url = None
test_image_url_string = None
test_web_search = None
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
def test_anthropic_citations_api():
"""
Test the citations API
@ -454,7 +410,6 @@ def test_anthropic_citations_api():
assert "start_char_index" in citation
assert "end_char_index" in citation
def test_anthropic_citations_api_streaming():
resp = completion(
@ -493,7 +448,6 @@ def test_anthropic_citations_api_streaming():
assert has_citations
@pytest.mark.parametrize(
"model",
[
@ -538,7 +492,6 @@ def test_anthropic_thinking_output_stream(model):
except litellm.Timeout:
pytest.skip("Model is timing out")
def test_anthropic_custom_headers():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -576,9 +529,6 @@ def test_anthropic_custom_headers():
headers = mock_post.call_args[1]["headers"]
assert "computer-use-2025-01-24" in headers["anthropic-beta"]
@pytest.mark.parametrize(
"optional_params",
[
@ -617,7 +567,6 @@ def test_anthropic_websearch(optional_params: dict):
assert response.usage.server_tool_use is not None
assert response.usage.server_tool_use.web_search_requests >= 1
def test_anthropic_text_editor():
litellm.turn_on_debug()
params = {
@ -640,7 +589,6 @@ def test_anthropic_text_editor():
assert response is not None
@pytest.mark.parametrize("spec", ["anthropic", "openai"])
@pytest.mark.skipif(
os.getenv("ZAPIER_CI_CD_MCP_TOKEN") is None, reason="ZAPIER_CI_CD_MCP_TOKEN not set"
@ -682,7 +630,6 @@ def test_anthropic_mcp_server_tool_use(spec: str):
except litellm.InternalServerError as e:
pytest.skip(f"Skipping test due to internal server error: {e}")
@pytest.mark.parametrize(
"model", ["openai/gpt-4.1", "anthropic/claude-sonnet-4-5-20250929"]
)
@ -714,7 +661,6 @@ def test_anthropic_mcp_server_responses_api(model: str):
assert response is not None
def test_anthropic_prefix_prompt():
params = {
"model": "anthropic/claude-sonnet-4-5-20250929",
@ -729,7 +675,6 @@ def test_anthropic_prefix_prompt():
assert response is not None
assert response.choices[0].message.content.startswith("Argentina")
@pytest.mark.asyncio
async def test_claude_tool_use_with_anthropic_acreate():
response = await litellm.anthropic.messages.acreate(
@ -754,9 +699,6 @@ async def test_claude_tool_use_with_anthropic_acreate():
async for chunk in response:
print(chunk)
def test_anthropic_streaming():
request_data = {
@ -811,7 +753,6 @@ def test_anthropic_streaming():
assert role_set_count == 1
def test_anthropic_via_responses_api():
from litellm.types.llms.openai import ResponsesAPIStreamEvents
@ -934,11 +875,6 @@ def test_anthropic_via_responses_api():
print(f"✓ All {len(events_seen)} events matched expected structure")
print(f"✓ Received {text_delta_count} text delta chunks")
def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict:
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
@ -950,19 +886,6 @@ def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict
headers={},
)
def test_anthropic_basic_completion_replay():
response = litellm.completion(
model="anthropic/claude-sonnet-4-5-20250929",
@ -976,7 +899,6 @@ def test_anthropic_basic_completion_replay():
assert response.usage.completion_tokens > 0
assert response.choices[0].finish_reason in {"stop", "length"}
def test_anthropic_streaming_completion_replay():
stream = litellm.completion(
model="anthropic/claude-sonnet-4-5-20250929",

View file

@ -1,15 +1,11 @@
import os
import pytest
import litellm
from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest
class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest):
test_content_list_handling = None
test_empty_tools = None
test_function_calling_with_tool_response = None
def get_base_completion_call_args(self):
@ -34,8 +30,6 @@ class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest):
def test_basic_tool_calling(self):
pass
class TestAzureOpenAIO3(BaseOSeriesModelsTest):
def get_base_completion_call_args(self):
return {

View file

@ -35,7 +35,6 @@ litellm.success_callback = []
user_message = "Write a short poem about the sky"
messages = [{"content": user_message, "role": "user"}]
@pytest.fixture(autouse=True)
def reset_callbacks():
print("\npytest fixture - resetting callbacks")
@ -44,7 +43,6 @@ def reset_callbacks():
litellm.failure_callback = []
litellm.callbacks = []
def test_completion_bedrock_claude_completion_auth(monkeypatch):
print("calling bedrock claude completion params auth")
@ -72,16 +70,13 @@ def test_completion_bedrock_claude_completion_auth(monkeypatch):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# test_completion_bedrock_claude_completion_auth()
@pytest.mark.parametrize("streaming", [True, False])
def test_completion_bedrock_guardrails(streaming):
litellm.set_verbose = True
# verbose_logger.setLevel(logging.DEBUG)
try:
if streaming is False:
@ -146,10 +141,8 @@ def test_completion_bedrock_guardrails(streaming):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# test_completion_bedrock_claude_2_1_completion_auth()
def test_completion_bedrock_claude_external_client_auth(monkeypatch):
print("\ncalling bedrock claude external client auth")
@ -187,21 +180,10 @@ def test_completion_bedrock_claude_external_client_auth(monkeypatch):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# test_completion_bedrock_claude_external_client_auth()
# test_completion_bedrock_claude_sts_client_auth()
@pytest.mark.parametrize(
"stop",
[""],
@ -240,7 +222,6 @@ def test_bedrock_stop_value(stop, model):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@pytest.mark.parametrize(
"system",
["You are an AI", [{"type": "text", "text": "You are an AI"}], ""],
@ -282,7 +263,6 @@ def test_bedrock_system_prompt(system, model):
def test_completion_bedrock_mistral_completion_auth():
print("calling bedrock mistral completion params auth")
litellm.turn_on_debug()
# aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"]
@ -311,10 +291,8 @@ def test_completion_bedrock_mistral_completion_auth():
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# test_completion_bedrock_mistral_completion_auth()
def test_bedrock_ptu():
"""
Check if a url with 'modelId' passed in, is created correctly
@ -346,7 +324,6 @@ def test_bedrock_ptu():
)
mock_client_post.assert_called_once()
@pytest.mark.asyncio
async def test_bedrock_custom_api_base():
"""
@ -384,7 +361,6 @@ async def test_bedrock_custom_api_base():
)
mock_client_post.assert_called_once()
@pytest.mark.parametrize(
"model",
[
@ -421,7 +397,6 @@ async def test_bedrock_extra_headers(model):
)
mock_client_post.assert_called_once()
@pytest.mark.asyncio
async def test_bedrock_custom_prompt_template():
"""
@ -466,7 +441,6 @@ async def test_bedrock_custom_prompt_template():
assert prompt == "<|im_start|>user\nWhat's AWS?<|im_end|>"
mock_client_post.assert_called_once()
def test_completion_bedrock_external_client_region(monkeypatch):
print("\ncalling bedrock claude external client auth")
@ -515,32 +489,15 @@ def test_completion_bedrock_external_client_region(monkeypatch):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
)
def test_base_aws_llm_get_credentials():
import time
import boto3
start_time = time.time()
session = boto3.Session(
aws_access_key_id="test",
@ -571,15 +528,6 @@ def test_base_aws_llm_get_credentials():
)
)
def test_bedrock_converse_route():
litellm.set_verbose = True
try:
@ -593,7 +541,6 @@ def test_bedrock_converse_route():
else:
raise
def test_bedrock_mapped_converse_models():
litellm.set_verbose = True
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
@ -604,23 +551,8 @@ def test_bedrock_mapped_converse_models():
messages=[{"role": "user", "content": "Hello, world!"}],
)
class TestBedrockConverseChatCrossRegion(BaseLLMChatTest):
test_content_list_handling = None
test_developer_role_translation = None
test_function_calling_with_tool_response = None
test_image_url = None
test_json_response_format_stream = None
test_tool_call_with_empty_enum_property = None
test_tool_call_with_property_type_array = None
def get_base_completion_call_args(self) -> dict:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
@ -649,10 +581,7 @@ class TestBedrockConverseChatCrossRegion(BaseLLMChatTest):
assert cost > 0
class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest):
test_completion_thinking_with_max_tokens = None
test_completion_thinking_without_max_tokens = None
def get_base_completion_call_args(self) -> dict:
return {
@ -665,12 +594,8 @@ class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest):
"thinking": {"type": "enabled", "budget_tokens": 16000},
}
class TestBedrockConverseChatNormal(BaseLLMChatTest):
test_content_list_handling = None
test_empty_tools = None
test_function_calling_with_tool_response = None
test_image_url = None
def get_base_completion_call_args(self) -> dict:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
@ -681,12 +606,8 @@ class TestBedrockConverseChatNormal(BaseLLMChatTest):
"aws_region_name": "us-east-1",
}
class TestBedrockConverseNovaTestSuite(BaseLLMChatTest):
test_content_list_handling = None
test_function_calling_with_tool_response = None
test_image_url = None
def get_base_completion_call_args(self) -> dict:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
@ -697,9 +618,6 @@ class TestBedrockConverseNovaTestSuite(BaseLLMChatTest):
"aws_region_name": "us-east-1",
}
class TestBedrockRerank(BaseLLMRerankTest):
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.BEDROCK
@ -709,7 +627,6 @@ class TestBedrockRerank(BaseLLMRerankTest):
"model": "bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0",
}
class TestBedrockCohereRerank(BaseLLMRerankTest):
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.BEDROCK
@ -719,13 +636,6 @@ class TestBedrockCohereRerank(BaseLLMRerankTest):
"model": "bedrock/arn:aws:bedrock:us-west-2::foundation-model/cohere.rerank-v3-5:0",
}
@pytest.mark.parametrize("top_k_param", ["top_k", "topK"])
def test_bedrock_nova_topk(top_k_param):
litellm.set_verbose = True
@ -751,7 +661,6 @@ def test_bedrock_nova_topk(top_k_param):
assert "inferenceConfig" in captured_data["additionalModelRequestFields"]
assert captured_data["additionalModelRequestFields"]["inferenceConfig"]["topK"] == 10
def test_bedrock_cross_region_inference(monkeypatch):
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -777,7 +686,6 @@ def test_bedrock_cross_region_inference(monkeypatch):
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/us.meta.llama3-3-70b-instruct-v1%3A0/converse"
)
def test_bedrock_empty_content_real_call():
completion(
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
@ -795,11 +703,6 @@ def test_bedrock_empty_content_real_call():
],
)
class TestBedrockEmbedding(BaseLLMEmbeddingTest):
def get_base_embedding_call_args(self) -> dict:
return {
@ -809,8 +712,6 @@ class TestBedrockEmbedding(BaseLLMEmbeddingTest):
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.BEDROCK
@pytest.mark.asyncio
async def test_bedrock_image_url_sync_client():
import logging
@ -849,9 +750,6 @@ async def test_bedrock_image_url_sync_client():
print(e)
mock_post.assert_called_once()
def test_bedrock_custom_proxy():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -874,7 +772,6 @@ def test_bedrock_custom_proxy():
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer Token"
def test_bedrock_custom_deepseek():
import json
@ -924,13 +821,6 @@ def test_bedrock_custom_deepseek():
print(f"Error: {str(e)}")
raise e
def test_bedrock_description_param():
from litellm import completion
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -969,7 +859,6 @@ def test_bedrock_description_param():
"Find the meaning inside a poem" in request_body_str
) # assert description is passed
@pytest.mark.parametrize(
"sync_mode",
[
@ -1029,7 +918,6 @@ async def test_bedrock_thinking_in_assistant_message(sync_mode):
in json_data
)
@pytest.mark.asyncio
async def test_bedrock_stream_thinking_content_openwebui():
"""
@ -1098,7 +986,6 @@ async def test_bedrock_stream_thinking_content_openwebui():
len(response_content) > 0
), "There should be non-empty content after thinking tags"
def test_bedrock_application_inference_profile():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -1163,7 +1050,6 @@ def test_bedrock_application_inference_profile():
)
assert mock_post2.call_args.kwargs["url"] == mock_post.call_args.kwargs["url"]
def return_mocked_response(model: str):
if model == "bedrock/mistral.mistral-large-2407-v1:0":
return {
@ -1178,7 +1064,6 @@ def return_mocked_response(model: str):
"usage": {"inputTokens": 5, "outputTokens": 10, "totalTokens": 15},
}
@pytest.mark.parametrize(
"model",
[
@ -1222,9 +1107,6 @@ async def test_bedrock_max_completion_tokens(model: str):
"inferenceConfig": {"maxTokens": 10},
}
@pytest.mark.asyncio
async def test_bedrock_passthrough_router():
"""
@ -1278,7 +1160,6 @@ async def test_bedrock_passthrough_router():
assert response.status_code == 200
@pytest.mark.asyncio
async def test_bedrock_converse__streaming_passthrough(monkeypatch):
import asyncio
@ -1332,7 +1213,6 @@ async def test_bedrock_converse__streaming_passthrough(monkeypatch):
assert response_cost is not None and response_cost > 0
assert "standard_logging_object" in mock_callback.call_args.kwargs["kwargs"]
@pytest.mark.asyncio
async def test_bedrock_streaming_passthrough_test2(monkeypatch):
import asyncio
@ -1383,7 +1263,6 @@ async def test_bedrock_streaming_passthrough_test2(monkeypatch):
assert "standard_logging_object" in mock_callback.call_args.kwargs["kwargs"]
assert "response_cost" in mock_callback.call_args.kwargs["kwargs"]
def test_bedrock_openai_imported_model():
"""
Test that Bedrock imported models using OpenAI format work correctly.
@ -1485,25 +1364,6 @@ def test_bedrock_openai_imported_model():
assert request_body["max_tokens"] == 300
assert request_body["temperature"] == 0.5
def test_bedrock_openai_multiple_message_types():
"""
Test that various message content types are handled correctly.
@ -1552,14 +1412,10 @@ def test_bedrock_openai_multiple_message_types():
print("✓ Multiple message types handled correctly")
# ============================================================================
# Nova Grounding (web_search_options) Unit Tests (Mocked)
# ============================================================================
def test_bedrock_nova_grounding_web_search_options_non_streaming():
"""
Unit test for Nova grounding using web_search_options parameter (non-streaming).
@ -1618,7 +1474,6 @@ def test_bedrock_nova_grounding_web_search_options_non_streaming():
f"✓ web_search_options correctly transformed to systemTool (non-streaming)"
)
def test_bedrock_nova_grounding_with_function_tools():
"""
Unit test for Nova grounding combined with regular function tools.
@ -1703,7 +1558,6 @@ def test_bedrock_nova_grounding_with_function_tools():
assert system_tool_found, "systemTool (nova_grounding) should be present"
print(f"✓ Both function tools and web_search_options correctly combined")
@pytest.mark.asyncio
async def test_bedrock_nova_grounding_async():
"""
@ -1757,9 +1611,6 @@ async def test_bedrock_nova_grounding_async():
assert system_tool_found, "systemTool with nova_grounding should be present"
print(f"✓ Async web_search_options correctly transformed to systemTool")
def test_bedrock_nova_grounding_request_transformation():
"""
Unit test to verify that web_search_options transforms to systemTool in the request.

View file

@ -1,8 +1,6 @@
from base_llm_unit_tests import BaseLLMChatTest
class TestBedrockGPTOSS(BaseLLMChatTest):
test_json_response_format = None
def get_base_completion_call_args(self) -> dict:
return {
@ -19,8 +17,3 @@ class TestBedrockGPTOSS(BaseLLMChatTest):
"""
pass
async def test_completion_cost(self):
"""
Bedrock GPT-OSS models are flaky and occasionally report 0 token counts in api response
"""
pass

View file

@ -5,7 +5,6 @@ import os
import litellm
from litellm.types.llms.bedrock import BedrockInvokeNovaRequest
_LITELLM_LOGO_IMAGE_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/"
"ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg"
@ -15,7 +14,6 @@ _AWSMP_LOGO_IMAGE_URL = (
"c233c9ade2ccb5491072ae232c814942.png"
)
@pytest.mark.flaky(retries=3, delay=5)
class TestBedrockInvokeClaudeJson(BaseLLMChatTest):
def get_base_completion_call_args(self) -> dict:
@ -24,26 +22,9 @@ class TestBedrockInvokeClaudeJson(BaseLLMChatTest):
"model": "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0",
}
@pytest.mark.parametrize(
"image_url, detail",
[
(_LITELLM_LOGO_IMAGE_URL, None),
(_LITELLM_LOGO_IMAGE_URL, "low"),
(_LITELLM_LOGO_IMAGE_URL, "high"),
(_AWSMP_LOGO_IMAGE_URL, "low"),
(_AWSMP_LOGO_IMAGE_URL, "high"),
],
)
@pytest.mark.flaky(retries=4, delay=2)
def test_image_url(self, image_url, detail):
super().test_image_url(detail=detail, image_url=image_url)
test_content_list_handling = None
test_image_url_string = None
test_pdf_handling = None
class TestBedrockInvokeNovaJson(BaseLLMChatTest):
test_json_response_format = None
def get_base_completion_call_args(self) -> dict:
return {
@ -57,11 +38,3 @@ class TestBedrockInvokeNovaJson(BaseLLMChatTest):
f"Skipping non-JSON test: {request.function.__name__} does not contain 'json'"
)
def test_json_response_pydantic_obj(self):
if os.environ.get("LITELLM_RUN_LIVE_BEDROCK_NOVA_JSON_TESTS") != "1":
pytest.skip("Live Bedrock Nova response-schema E2E tests are opt-in")
if os.environ.get("CASSETTE_REDIS_URL"):
pytest.skip(
"Live Bedrock Nova response-schema E2E tests cannot run under VCR replay"
)
super().test_json_response_pydantic_obj()

View file

@ -3,10 +3,7 @@ import pytest
import litellm
class TestBedrockTestSuite(BaseLLMChatTest):
test_content_list_handling = None
test_empty_tools = None
test_function_calling_with_tool_response = None
def get_base_completion_call_args(self) -> dict:

View file

@ -15,20 +15,12 @@ from base_llm_unit_tests import BaseLLMChatTest
import litellm
class TestBedrockMoonshotInvoke(BaseLLMChatTest):
"""
Test suite for Bedrock Moonshot via invoke route.
Inherits all standard LLM tests from BaseLLMChatTest.
"""
test_json_response_format_stream = None
test_completion_cost = None
test_content_list_handling = None
test_developer_role_translation = None
test_message_with_name = None
test_pydantic_model_input = None
test_response_format_type_text_with_tool_calls_no_tool_choice = None
test_streaming = None
def get_base_completion_call_args(self) -> dict:

View file

@ -3,15 +3,8 @@ import pytest
import litellm
class TestBedrockNovaJson(BaseLLMChatTest):
test_content_list_handling = None
test_developer_role_translation = None
test_empty_tools = None
test_function_calling_with_tool_response = None
test_json_response_format_stream = None
test_tool_call_with_empty_enum_property = None
test_tool_call_with_property_type_array = None
def get_base_completion_call_args(self) -> dict:
litellm.turn_on_debug()
@ -19,14 +12,6 @@ class TestBedrockNovaJson(BaseLLMChatTest):
"model": "bedrock/converse/us.amazon.nova-micro-v1:0",
}
def test_json_response_nested_pydantic_obj(self):
pass
def test_json_response_nested_json_schema(self):
pass
# @pytest.fixture(autouse=True)
# def skip_non_json_tests(self, request):
# if not "json" in request.function.__name__.lower():

View file

@ -2,7 +2,6 @@ import os
import pytest
from base_llm_unit_tests import BaseLLMChatTest
from litellm.llms.vertex_ai.context_caching.transformation import (
separate_cached_messages,
@ -12,18 +11,9 @@ import litellm
from litellm import completion
import json
class TestGoogleAIStudioGemini(BaseLLMChatTest):
test_async_pdf_handling_with_file_id = None
test_content_list_handling = None
test_developer_role_translation = None
test_function_calling_with_tool_response = None
test_image_url = None
test_json_response_nested_json_schema = None
test_json_response_nested_pydantic_obj = None
test_json_response_pydantic_obj = None
test_web_search = None
def get_base_completion_call_args(self) -> dict:
@ -32,7 +22,6 @@ class TestGoogleAIStudioGemini(BaseLLMChatTest):
def get_base_completion_call_args_with_reasoning_model(self) -> dict:
return {"model": "gemini/gemini-2.5-flash"}
@pytest.mark.flaky(retries=3, delay=2)
def test_url_context(self):
from litellm.utils import supports_url_context
@ -64,11 +53,6 @@ class TestGoogleAIStudioGemini(BaseLLMChatTest):
), "URL context metadata should be present"
print(f"response={response}")
def test_gemini_image_generation():
# litellm.turn_on_debug()
response = completion(
@ -90,23 +74,6 @@ def test_gemini_image_generation():
.startswith("data:image/png;base64,")
)
def test_gemini_thinking():
litellm.turn_on_debug()
from litellm.types.utils import Message, CallTypes
@ -146,9 +113,6 @@ def test_gemini_thinking():
print(response.choices[0].message)
assert response.choices[0].message.content is not None
def test_gemini_finish_reason():
import os
from litellm import completion
@ -163,7 +127,6 @@ def test_gemini_finish_reason():
assert response.choices[0].finish_reason is not None
assert response.choices[0].finish_reason == "length"
@pytest.mark.flaky(retries=3, delay=2)
def test_gemini_url_context():
from litellm import completion
@ -189,7 +152,6 @@ def test_gemini_url_context():
assert urlMetadata["retrievedUrl"] == URL1
assert urlMetadata["urlRetrievalStatus"] == "URL_RETRIEVAL_STATUS_SUCCESS"
@pytest.mark.flaky(retries=3, delay=2)
def test_gemini_with_grounding():
from litellm import completion, Usage, stream_chunk_builder
@ -226,7 +188,6 @@ def test_gemini_with_grounding():
assert usage.prompt_tokens_details.web_search_requests is not None
assert usage.prompt_tokens_details.web_search_requests > 0
def test_gemini_with_empty_function_call_arguments():
from litellm import completion
@ -248,9 +209,6 @@ def test_gemini_with_empty_function_call_arguments():
print(response)
assert response.choices[0].message.content is not None
def test_gemini_tool_use():
data = {
"max_tokens": 8192,
@ -299,7 +257,6 @@ def test_gemini_tool_use():
assert stop_reason is not None
assert stop_reason == "tool_calls"
@pytest.mark.asyncio
async def test_gemini_image_generation_async():
litellm.turn_on_debug()
@ -332,7 +289,6 @@ async def test_gemini_image_generation_async():
assert IMAGE_URL["url"] is not None, "IMAGE_URL['url'] is not None"
assert IMAGE_URL["url"].startswith("data:image/png;base64,")
@pytest.mark.asyncio
async def test_gemini_image_generation_async_stream():
# litellm.turn_on_debug()
@ -367,7 +323,6 @@ async def test_gemini_image_generation_async_stream():
assert model_response_image is not None
assert model_response_image["url"].startswith("data:image/png;base64,")
def test_system_message_with_no_user_message():
"""
Test that the system message is translated correctly for non-OpenAI providers.
@ -387,7 +342,6 @@ def test_system_message_with_no_user_message():
assert response.choices[0].message.content is not None
def get_current_weather(location, unit="fahrenheit"):
"""Get the current weather in a given location"""
if "tokyo" in location.lower():
@ -401,7 +355,6 @@ def get_current_weather(location, unit="fahrenheit"):
else:
return json.dumps({"location": location, "temperature": "unknown"})
def test_gemini_with_thinking():
from litellm import completion
@ -493,11 +446,6 @@ def test_gemini_with_thinking():
) # get a new response from the model where it can see the function response
print("second response\n", second_response)
@pytest.mark.parametrize(
"status_code,expected_exception",
[
@ -582,7 +530,6 @@ def l(status_code, expected_exception):
"VertexAIException" not in error_message
), f"Should not contain 'VertexAIException' for status {status_code}, got: {error_message}"
def test_gemini_embedding():
litellm.turn_on_debug()
response = litellm.embedding(
@ -592,23 +539,6 @@ def test_gemini_embedding():
print("response: ", response)
assert response is not None
@pytest.mark.asyncio
async def test_gemini_openai_web_search_tool_to_google_search():
"""

View file

@ -1,6 +1,5 @@
# sys.path.insert(
# 0, os.path.abspath("../..")
# ) # noqa
@ -8,10 +7,7 @@
from base_llm_unit_tests import BaseLLMChatTest
class TestGroq(BaseLLMChatTest):
test_content_list_handling = None
test_empty_tools = None
test_web_search = None
def get_base_completion_call_args(self) -> dict:
@ -19,5 +15,3 @@ class TestGroq(BaseLLMChatTest):
"model": "groq/openai/gpt-oss-120b",
}
def test_tool_call_with_empty_enum_property(self):
pass

View file

@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
from base_llm_unit_tests import BaseLLMChatTest
import pytest
import litellm
@ -126,7 +125,6 @@ MOCK_STREAMING_CHUNKS = [
},
]
PROVIDER_MAPPING_RESPONSE = {
"fireworks-ai": {
"status": "live",
@ -145,14 +143,12 @@ PROVIDER_MAPPING_RESPONSE = {
},
}
@pytest.fixture
def mock_provider_mapping():
with patch("litellm.llms.huggingface.chat.transformation.fetch_inference_provider_mapping") as mock:
mock.return_value = PROVIDER_MAPPING_RESPONSE
yield mock
@pytest.fixture(autouse=True)
def clear_lru_cache():
from litellm.llms.huggingface.common_utils import fetch_inference_provider_mapping
@ -161,7 +157,6 @@ def clear_lru_cache():
yield
fetch_inference_provider_mapping.cache_clear()
@pytest.fixture
def mock_http_handler():
"""Fixture to mock the HTTP handler"""
@ -188,7 +183,6 @@ def mock_http_handler():
mock.side_effect = mock_side_effect
yield mock
@pytest.fixture
def mock_http_async_handler():
"""Fixture to mock the async HTTP handler"""
@ -220,7 +214,6 @@ def mock_http_async_handler():
mock.side_effect = mock_side_effect
yield mock
class TestHuggingFace(BaseLLMChatTest):
@pytest.fixture(autouse=True)
def setup(self, mock_provider_mapping, mock_http_handler, mock_http_async_handler):
@ -355,8 +348,6 @@ class TestHuggingFace(BaseLLMChatTest):
== tool_call_no_arguments["tool_calls"][0]["function"]["arguments"]
)
def test_completion_with_api_base(self):
messages = [{"role": "user", "content": "This is a test message"}]
api_base = "https://abcd123.us-east-1.aws.endpoints.huggingface.cloud"
@ -421,9 +412,3 @@ class TestHuggingFace(BaseLLMChatTest):
called_url = call_args[1]["url"]
assert called_url == f"{api_base}/v1/chat/completions"
@pytest.mark.asyncio
async def test_completion_cost(self):
pass

View file

@ -8,7 +8,6 @@ from tests.llm_translation.base_audio_transcription_unit_tests import (
BaseLLMAudioTranscriptionTest,
)
@pytest.mark.skipif(
not os.getenv("MISTRAL_API_KEY"),
reason="MISTRAL_API_KEY not set, skipping Mistral audio transcription tests",
@ -22,8 +21,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest):
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.MISTRAL
def test_audio_transcription_async(self): # type: ignore[override]
pytest.skip(
"Async audio transcription test for Mistral is skipped in this suite; "
"async test plugins (e.g. pytest-asyncio/anyio) are not configured here."
)

View file

@ -3,7 +3,6 @@ from datetime import datetime
from typing import Final
from unittest.mock import AsyncMock
import httpx
import pytest
from openai.types import CreateEmbeddingResponse, Embedding
@ -16,7 +15,6 @@ from litellm import completion
from base_rerank_unit_tests import BaseLLMRerankTest
from tests.capturing_transport import CapturingTransport
def test_completion_nvidia_nim():
from openai import OpenAI
@ -59,7 +57,6 @@ def test_completion_nvidia_nim():
assert request_body["frequency_penalty"] == 0.1
assert request_body["presence_penalty"] == 0.5
class TestNvidiaNim(BaseLLMRerankTest):
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.NVIDIA_NIM
@ -73,43 +70,3 @@ class TestNvidiaNim(BaseLLMRerankTest):
"""Nvidia NIM rerank models are free (cost = 0.0)"""
return 0.0
@pytest.mark.asyncio()
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_basic_rerank(self, sync_mode, monkeypatch):
"""
Override the base live rerank test with a mocked HTTP layer.
NVIDIA reached end-of-life for the hosted
nvidia/llama-3.2-nv-rerankqa-1b-v2 rerank API on 2026-05-18 and
published no replacement model, so a live call now returns HTTP 410
("Gone"). NVIDIA's hosted catalog rotates on a schedule, so pointing
at another live model would only defer the same failure. Mock the
transport instead (same pattern as
test_nvidia_nim_rerank_ranking_endpoint above) so the request/response
transformation and cost calculation stay covered offline.
"""
monkeypatch.setenv("NVIDIA_NIM_API_KEY", "fake-api-key")
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.text = ""
mock_response.json.return_value = {
"rankings": [
{"index": 0, "logit": 0.95},
{"index": 1, "logit": 0.75},
],
"usage": {"total_tokens": 7},
}
with (
patch(
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
return_value=mock_response,
),
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=mock_response,
),
):
await super().test_basic_rerank(sync_mode=sync_mode)

View file

@ -1,18 +1,13 @@
import os
from unittest.mock import patch
import pytest
import litellm
from litellm import ModelResponse
from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest
class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest):
test_empty_tools = None
test_tool_call_with_empty_enum_property = None
test_tool_call_with_property_type_array = None
def get_base_completion_call_args(self):
return {
@ -24,9 +19,6 @@ class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest):
return OpenAI(api_key="fake-api-key")
class TestOpenAIO3(BaseOSeriesModelsTest, BaseLLMChatTest):
test_basic_tool_calling = None
test_function_calling_with_tool_response = None
@ -41,9 +33,6 @@ class TestOpenAIO3(BaseOSeriesModelsTest, BaseLLMChatTest):
return OpenAI(api_key="fake-api-key")
def test_o3_reasoning_effort():
resp = litellm.completion(
model="o3-mini",

View file

@ -12,7 +12,6 @@ from tests.llm_translation.base_audio_transcription_unit_tests import (
BaseLLMAudioTranscriptionTest,
)
@pytest.mark.skipif(
not os.getenv("OVHCLOUD_API_KEY"),
reason="OVHCLOUD_API_KEY not set, skipping OVHCloud audio transcription tests",
@ -29,12 +28,6 @@ class TestOVHCloudAudioTranscription(BaseLLMAudioTranscriptionTest):
# Override the async base test with a sync no-op to avoid
# 'async def functions are not natively supported' failures when
# running this file in isolation without pytest-asyncio.
def test_audio_transcription_async(self): # type: ignore[override]
pytest.skip(
"Async audio transcription test for OVHCloud is skipped in this suite; "
"async test plugins (e.g. pytest-asyncio/anyio) are not configured here."
)
@pytest.mark.skipif(
not os.getenv("OVHCLOUD_API_KEY"),

View file

@ -8,21 +8,12 @@ import json
from datetime import datetime
from unittest.mock import AsyncMock
import litellm
import pytest
class TestTogetherAI(BaseLLMChatTest):
test_basic_tool_calling = None
test_empty_tools = None
test_function_calling_with_tool_response = None
test_json_response_format = None
test_json_response_nested_json_schema = None
test_json_response_nested_pydantic_obj = None
test_json_response_pydantic_obj = None
test_tool_call_with_empty_enum_property = None
test_tool_call_with_property_type_array = None
def get_base_completion_call_args(self) -> dict:
litellm.set_verbose = True

View file

@ -1,4 +1,10 @@
import base64
import json
from typing import Final
import httpx
import pytest
from respx import MockRouter
import litellm
import litellm.interactions as interactions
@ -18,3 +24,71 @@ class TestGoogleInteractionsCreate:
input="Hello",
api_key=api_key,
)
class TestInteractionsAcreateOffline:
@pytest.fixture(autouse=True)
def _httpx_only_transport(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
@pytest.mark.usefixtures("fake_provider_credentials")
@pytest.mark.asyncio
async def test_acreate_simple_gemini(self, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post("https://generativelanguage.googleapis.com/v1beta/interactions").mock(
return_value=httpx.Response(
200,
json={
"id": "interaction-offline",
"object": "interaction",
"model": "gemini-2.5-flash",
"status": "completed",
"steps": [{"type": "model_output", "content": [{"type": "text", "text": "299792458"}]}],
"usage": {"input_tokens": 6, "output_tokens": 3},
},
)
)
response: Final = await interactions.acreate(
model="gemini/gemini-2.5-flash",
input="What is the speed of light?",
api_key="gemini-offline",
)
body: Final = json.loads(route.calls.last.request.content)
assert body["model"] == "gemini-2.5-flash"
assert body["input"] == "What is the speed of light?"
assert response.id == "interaction-offline"
assert response.status == "completed"
@pytest.mark.usefixtures("fake_provider_credentials")
@pytest.mark.asyncio
async def test_acreate_simple_litellm_responses_bridge(self, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post("https://api.openai.com/v1/responses").mock(
return_value=httpx.Response(
200,
json={
"id": "resp-offline",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o",
"output": [
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "299792458 m/s"}],
}
],
"usage": {"input_tokens": 6, "output_tokens": 3, "total_tokens": 9},
},
)
)
response: Final = await interactions.acreate(
model="gpt-4o",
input="What is the speed of light?",
api_key="sk-offline",
)
body: Final = json.loads(route.calls.last.request.content)
assert body["model"] == "gpt-4o"
serialized: Final = json.dumps(body)
assert "What is the speed of light?" in serialized
assert "response_id:resp-offline" in base64.b64decode(response.id.removeprefix("resp_")).decode()
assert response.status == "completed"

View file

@ -0,0 +1,210 @@
import io
import json
from typing import Final, Mapping, Sequence, cast
import httpx
import pytest
import respx
from respx import MockRouter
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm import transcription
from litellm.litellm_core_utils.get_supported_openai_params import get_supported_openai_params
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
)
from litellm.llms.elevenlabs.audio_transcription.transformation import ElevenLabsAudioTranscriptionConfig
from litellm.llms.mistral.audio_transcription.transformation import MistralAudioTranscriptionConfig
from litellm.llms.ovhcloud.audio_transcription.transformation import OVHCloudAudioTranscriptionConfig
from litellm.utils import ProviderConfigManager
class _Kwargs(TypedDict, total=False):
model: ReadOnly[str]
api_key: ReadOnly[str]
api_base: ReadOnly[str]
timestamp_granularities: ReadOnly[Sequence[str]]
class _Case(TypedDict):
id: ReadOnly[str]
provider: ReadOnly[str]
kwargs: ReadOnly[_Kwargs]
url: ReadOnly[str]
prefix_match: ReadOnly[bool]
base_model: ReadOnly[str]
request_markers: ReadOnly[tuple[bytes, ...]]
marker_in_url: ReadOnly[bool]
config_class: ReadOnly[type[BaseAudioTranscriptionConfig]]
_AUDIO_BYTES: Final = b"RIFFFAKEWAVDATA-gettysburg"
_CASES: Final[tuple[_Case, ...]] = (
{
"id": "openai_gpt4o",
"provider": "openai",
"kwargs": {
"model": "openai/gpt-4o-transcribe",
"api_key": "sk-offline",
"timestamp_granularities": ["word"],
},
"url": "https://api.openai.com/v1/audio/transcriptions",
"prefix_match": False,
"base_model": "gpt-4o-transcribe",
"request_markers": (b"gpt-4o-transcribe", b'name="timestamp_granularities[]"\r\n\r\nword\r\n'),
"marker_in_url": False,
"config_class": litellm.OpenAIGPTAudioTranscriptionConfig,
},
{
"id": "elevenlabs_scribe",
"provider": "elevenlabs",
"kwargs": {"model": "elevenlabs/scribe_v1", "api_key": "xi-offline"},
"url": "https://api.elevenlabs.io/v1/speech-to-text",
"prefix_match": False,
"base_model": "scribe_v1",
"request_markers": (b"scribe_v1",),
"marker_in_url": False,
"config_class": ElevenLabsAudioTranscriptionConfig,
},
{
"id": "deepgram_nova",
"provider": "deepgram",
"kwargs": {"model": "deepgram/nova-2", "api_key": "dg-offline"},
"url": "https://api.deepgram.com/v1/listen",
"prefix_match": True,
"base_model": "nova-2",
"request_markers": (b"model=nova-2",),
"marker_in_url": True,
"config_class": litellm.DeepgramAudioTranscriptionConfig,
},
{
"id": "mistral_voxtral",
"provider": "mistral",
"kwargs": {"model": "mistral/voxtral-mini-latest", "api_key": "mistral-offline"},
"url": "https://api.mistral.ai/v1/audio/transcriptions",
"prefix_match": False,
"base_model": "voxtral-mini-latest",
"request_markers": (b"voxtral-mini-latest",),
"marker_in_url": False,
"config_class": MistralAudioTranscriptionConfig,
},
{
"id": "ovhcloud_whisper",
"provider": "ovhcloud",
"kwargs": {"model": "ovhcloud/whisper-large-v3-turbo", "api_key": "ovh-offline"},
"url": "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1/audio/transcriptions",
"prefix_match": False,
"base_model": "whisper-large-v3-turbo",
"request_markers": (b"whisper-large-v3-turbo",),
"marker_in_url": False,
"config_class": OVHCloudAudioTranscriptionConfig,
},
)
def _canned_response(case: _Case) -> httpx.Response:
if case["provider"] == "deepgram":
return httpx.Response(
200,
json={
"metadata": {"transaction_key": "offline", "duration": 1.5},
"results": {
"channels": [
{"alternatives": [{"transcript": "four score and seven years ago", "confidence": 0.99}]}
]
},
},
)
return httpx.Response(200, json={"text": "four score and seven years ago"})
@pytest.fixture(autouse=True)
def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
def _register(case: _Case, respx_mock: MockRouter) -> respx.Route:
if case["prefix_match"]:
return respx_mock.post(url__startswith=case["url"]).mock(return_value=_canned_response(case))
return respx_mock.post(case["url"]).mock(return_value=_canned_response(case))
def _assert_translated_request(case: _Case, request: httpx.Request) -> None:
searched: Final = request.url.query if case["marker_in_url"] else request.content
for marker in case["request_markers"]:
assert marker in searched
assert _AUDIO_BYTES in request.content
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
def test_audio_transcription(case: _Case, respx_mock: MockRouter) -> None:
route: Final = _register(case, respx_mock)
transcript: Final = transcription(**dict(case["kwargs"]), file=io.BytesIO(_AUDIO_BYTES))
_assert_translated_request(case, route.calls.last.request)
assert transcript.text == "four score and seven years ago"
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
@pytest.mark.asyncio
async def test_audio_transcription_async(case: _Case, respx_mock: MockRouter) -> None:
route: Final = _register(case, respx_mock)
transcript: Final = await litellm.atranscription(**dict(case["kwargs"]), file=io.BytesIO(_AUDIO_BYTES))
_assert_translated_request(case, route.calls.last.request)
assert transcript.text == "four score and seven years ago"
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
def test_audio_transcription_optional_params(case: _Case) -> None:
optional_params: Final = get_supported_openai_params(
model=case["kwargs"]["model"],
custom_llm_provider=case["provider"],
request_type="transcription",
)
assert isinstance(optional_params, list)
assert optional_params == case["config_class"]().get_supported_openai_params(case["base_model"])
assert "max_completion_tokens" not in optional_params
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
def test_audio_transcription_config(case: _Case) -> None:
config: Final = ProviderConfigManager.get_provider_audio_transcription_config(
model=case["kwargs"]["model"],
provider=litellm.LlmProviders(case["provider"]),
)
assert type(config) is case["config_class"]
assert isinstance(config, BaseAudioTranscriptionConfig)
if case["provider"] == "deepgram":
complete_url: Final = config.get_complete_url(
api_base=None,
api_key=None,
model=case["base_model"],
optional_params={},
litellm_params={},
)
assert "api.deepgram.com" in complete_url
assert "model=nova-2" in complete_url
else:
transformed: Final[AudioTranscriptionRequestData] = config.transform_audio_transcription_request(
model=case["base_model"],
audio_file=io.BytesIO(_AUDIO_BYTES),
optional_params={},
litellm_params={},
)
assert _AUDIO_BYTES in _transformed_payload(transformed)
def _transformed_payload(transformed: AudioTranscriptionRequestData) -> bytes:
data: Final = transformed.data
if isinstance(data, bytes):
return data
file_entry: Final = data.get("file") if isinstance(data, dict) else None
if isinstance(file_entry, io.BytesIO):
return file_entry.getvalue()
if transformed.files is not None:
first: Final = next(iter(transformed.files.values()))
blob: Final = first[1] if isinstance(first, tuple) else first
return blob.getvalue() if isinstance(blob, io.BytesIO) else cast(bytes, blob)
return b""

View file

@ -0,0 +1,385 @@
import json
from itertools import chain
from typing import Final, Mapping, cast
import httpx
import pytest
from pydantic import BaseModel, ConfigDict, JsonValue
from respx import MockRouter
import litellm
from litellm import get_llm_provider
from litellm.constants import (
DEFAULT_MAX_TOKENS,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
)
from litellm.main import stream_chunk_builder
from litellm.utils import get_optional_params
from tests.unit.llms.base_llm.chat.test_provider_chat_translation import (
_BY_ID,
_Case,
_aws_frame,
_call as _provider_call,
_request_body,
)
_THINKING_BUDGET: Final = 16000
_THINKING: Final[JsonValue] = {"type": "enabled", "budget_tokens": _THINKING_BUDGET}
_ANTHROPIC: Final = _BY_ID["anthropic_sonnet45"]
_BEDROCK_HAIKU: Final = _BY_ID["bedrock_converse_haiku"]
_BEDROCK_SONNET: Final = _BY_ID["bedrock_converse_anthropic_thinking"]
_THINKING_CASES: Final = (_ANTHROPIC, _BEDROCK_SONNET)
_RESPONSE_FORMAT_CASES: Final = (_ANTHROPIC, _BEDROCK_HAIKU)
_JSON_PREFIX: Final = '{"agent_doing": "researching '
_JSON_SUFFIX: Final = 'home automation"}'
_JSON_CONTENT: Final = _JSON_PREFIX + _JSON_SUFFIX
_REASONING: Final = "reasoning here"
_SIGNATURE: Final = "sig-1"
class _RFormat(BaseModel):
model_config = ConfigDict(frozen=True)
question: str
answer: str
_JSON_SCHEMA_ARGS: Final[Mapping[str, JsonValue]] = {
"messages": [
{"role": "system", "content": "Summarize the agent's thinking into short descriptions."},
{"role": "user", "content": "Here is the input data."},
],
"response_format": {
"type": "json_schema",
"json_schema": {
"name": "final_output",
"strict": True,
"schema": {
"properties": {"agent_doing": {"title": "Agent Doing", "type": "string"}},
"required": ["agent_doing"],
"title": "ThinkingStep",
"type": "object",
"additionalProperties": False,
},
},
},
}
_THINKING_MESSAGES: Final[Mapping[str, JsonValue]] = {
"messages": [{"role": "user", "content": "Generate 5 question + answer pairs"}],
}
def _case_id(case: _Case) -> str:
return case["id"]
def _call(case: _Case, **extra: JsonValue) -> object:
return _provider_call(case, extra)
def _anthropic_sse(events: tuple[Mapping[str, JsonValue], ...]) -> str:
return "".join(f"event: {e['type']}\ndata: {json.dumps(e)}\n\n" for e in events)
_ANTHROPIC_START: Final[Mapping[str, JsonValue]] = {
"type": "message_start",
"message": {
"id": "msg_offline",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 10, "output_tokens": 1},
},
}
_ANTHROPIC_END: Final[tuple[Mapping[str, JsonValue], ...]] = (
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}},
{"type": "message_stop"},
)
def _anthropic_json_stream() -> str:
return _anthropic_sse(
(
_ANTHROPIC_START,
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JSON_PREFIX}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JSON_SUFFIX}},
{"type": "content_block_stop", "index": 0},
*_ANTHROPIC_END,
)
)
def _anthropic_thinking_stream() -> str:
return _anthropic_sse(
(
_ANTHROPIC_START,
{"type": "content_block_start", "index": 0, "content_block": {"type": "thinking", "thinking": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "thinking_delta", "thinking": _REASONING}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "signature_delta", "signature": _SIGNATURE}},
{"type": "content_block_stop", "index": 0},
{"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 1, "delta": {"type": "text_delta", "text": "done"}},
{"type": "content_block_stop", "index": 1},
*_ANTHROPIC_END,
)
)
_CONVERSE_USAGE: Final[Mapping[str, JsonValue]] = {"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}}
def _converse_stream(frames: tuple[tuple[str, Mapping[str, JsonValue]], ...]) -> bytes:
return b"".join(_aws_frame(event_type, payload) for event_type, payload in frames)
def _converse_json_stream() -> bytes:
return _converse_stream(
(
("messageStart", {"role": "assistant"}),
("contentBlockDelta", {"delta": {"text": _JSON_PREFIX}, "contentBlockIndex": 0}),
("contentBlockDelta", {"delta": {"text": _JSON_SUFFIX}, "contentBlockIndex": 0}),
("contentBlockStop", {"contentBlockIndex": 0}),
("messageStop", {"stopReason": "end_turn"}),
("metadata", _CONVERSE_USAGE),
)
)
def _converse_thinking_stream() -> bytes:
return _converse_stream(
(
("messageStart", {"role": "assistant"}),
("contentBlockDelta", {"delta": {"reasoningContent": {"text": _REASONING}}, "contentBlockIndex": 0}),
("contentBlockDelta", {"delta": {"reasoningContent": {"signature": _SIGNATURE}}, "contentBlockIndex": 0}),
("contentBlockStop", {"contentBlockIndex": 0}),
("contentBlockDelta", {"delta": {"text": "done"}, "contentBlockIndex": 1}),
("contentBlockStop", {"contentBlockIndex": 1}),
("messageStop", {"stopReason": "end_turn"}),
("metadata", _CONVERSE_USAGE),
)
)
def _non_stream_response(case: _Case) -> httpx.Response:
if case["shape"] == "anthropic":
return httpx.Response(
200,
json={
"id": "msg_offline",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [
{"type": "thinking", "thinking": _REASONING, "signature": _SIGNATURE},
{"type": "text", "text": _JSON_CONTENT},
],
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
},
)
return httpx.Response(
200,
json={
"output": {
"message": {
"role": "assistant",
"content": [
{"reasoningContent": {"reasoningText": {"text": _REASONING, "signature": _SIGNATURE}}},
{"text": _JSON_CONTENT},
],
}
},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15},
},
)
def _stream_response(case: _Case, *, thinking: bool) -> httpx.Response:
if case["shape"] == "anthropic":
return httpx.Response(
200,
content=_anthropic_thinking_stream() if thinking else _anthropic_json_stream(),
headers={"content-type": "text/event-stream"},
)
return httpx.Response(
200,
content=_converse_thinking_stream() if thinking else _converse_json_stream(),
headers={"content-type": "application/vnd.amazon.eventstream"},
)
def _thinking_param(case: _Case, body: Mapping[str, JsonValue]) -> JsonValue:
if case["shape"] == "anthropic":
return body["thinking"]
return cast(Mapping[str, JsonValue], body["additionalModelRequestFields"])["thinking"]
def _max_tokens(case: _Case, body: Mapping[str, JsonValue]) -> JsonValue:
if case["shape"] == "anthropic":
return body["max_tokens"]
return cast(Mapping[str, JsonValue], body["inferenceConfig"])["maxTokens"]
def _json_schema_title(case: _Case, body: Mapping[str, JsonValue]) -> str:
if case["shape"] == "anthropic":
output_format: Final = cast(Mapping[str, JsonValue], body["output_format"])
assert output_format["type"] == "json_schema"
return cast(str, cast(Mapping[str, JsonValue], output_format["schema"])["title"])
text_format: Final = cast(
Mapping[str, JsonValue],
cast(Mapping[str, JsonValue], body["outputConfig"])["textFormat"],
)
assert text_format["type"] == "json_schema"
json_schema: Final = cast(
Mapping[str, JsonValue], cast(Mapping[str, JsonValue], text_format["structure"])["jsonSchema"]
)
return cast(str, json.loads(cast(str, json_schema["schema"]))["title"])
def _has_forced_tool_choice(case: _Case, body: Mapping[str, JsonValue]) -> bool:
if case["shape"] == "anthropic":
return body.get("tool_choice") is not None
return "toolConfig" in body
@pytest.mark.parametrize("case", _RESPONSE_FORMAT_CASES, ids=_case_id)
def test_anthropic_response_format_streaming_vs_non_streaming(case: _Case, respx_mock: MockRouter) -> None:
stream_route: Final = respx_mock.post(case["stream_url"]).mock(return_value=_stream_response(case, thinking=False))
chunks: Final = tuple(cast(litellm.CustomStreamWrapper, _call(case, **_JSON_SCHEMA_ARGS, stream=True)))
built: Final = stream_chunk_builder(chunks=list(chunks))
stream_body: Final = _request_body(stream_route)
non_stream_route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
non_stream: Final = cast(litellm.ModelResponse, _call(case, **_JSON_SCHEMA_ARGS))
non_stream_body: Final = _request_body(non_stream_route)
assert len(chunks) > 1
assert _json_schema_title(case, stream_body) == "ThinkingStep"
assert _json_schema_title(case, non_stream_body) == "ThinkingStep"
assert built is not None
streamed_json: Final = cast(
Mapping[str, JsonValue],
json.loads(cast(str, cast(litellm.ModelResponse, built).choices[0].message.content)),
)
non_stream_json: Final = cast(Mapping[str, JsonValue], json.loads(cast(str, non_stream.choices[0].message.content)))
assert streamed_json == non_stream_json == {"agent_doing": "researching home automation"}
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
def test_completion_thinking_with_response_format(case: _Case, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
response: Final = cast(
litellm.ModelResponse,
_call(case, thinking=_THINKING, **_THINKING_MESSAGES, response_format=cast(JsonValue, _RFormat)),
)
body: Final = _request_body(route)
assert _thinking_param(case, body) == _THINKING
assert _json_schema_title(case, body) == "_RFormat"
assert not _has_forced_tool_choice(case, body)
assert response.choices[0].message.content == _JSON_CONTENT
assert response.choices[0].message.reasoning_content == _REASONING
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
def test_completion_thinking_with_max_tokens(case: _Case, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
response: Final = cast(
litellm.ModelResponse,
_call(case, thinking=_THINKING, **_THINKING_MESSAGES, max_completion_tokens=20000),
)
body: Final = _request_body(route)
assert _max_tokens(case, body) == 20000
assert _thinking_param(case, body) == _THINKING
assert response.choices[0].message.content == _JSON_CONTENT
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
def test_completion_thinking_without_max_tokens(case: _Case, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
response: Final = cast(litellm.ModelResponse, _call(case, thinking=_THINKING, **_THINKING_MESSAGES))
body: Final = _request_body(route)
max_tokens: Final = cast(int, _max_tokens(case, body))
assert max_tokens == _THINKING_BUDGET + DEFAULT_MAX_TOKENS
assert max_tokens > _THINKING_BUDGET
assert _thinking_param(case, body) == _THINKING
assert response.choices[0].message.content == _JSON_CONTENT
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
def test_anthropic_thinking_output_stream(case: _Case, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post(case["stream_url"]).mock(return_value=_stream_response(case, thinking=True))
chunks: Final = tuple(
cast(
litellm.CustomStreamWrapper,
_call(
case,
thinking=_THINKING,
messages=[{"role": "user", "content": "Tell me a joke."}],
stream=True,
),
)
)
deltas: Final = tuple(chunk.choices[0].delta for chunk in chunks)
thinking_deltas: Final = tuple(
delta
for delta in deltas
if isinstance(getattr(delta, "thinking_blocks", None), list)
and delta.thinking_blocks
and isinstance(getattr(delta, "reasoning_content", None), str)
)
blocks: Final = chain.from_iterable(cast(list[object], delta.thinking_blocks) for delta in thinking_deltas)
signatures: Final = tuple(cast(Mapping[str, JsonValue], block).get("signature") for block in blocks)
assert _thinking_param(case, _request_body(route)) == _THINKING
assert not any(delta.tool_calls for delta in deltas)
assert "".join(cast(str, delta.reasoning_content) for delta in thinking_deltas) == _REASONING
assert _SIGNATURE in signatures
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
def test_anthropic_reasoning_effort_thinking_translation(case: _Case, respx_mock: MockRouter) -> None:
model: Final = case["kwargs"].get("model", "")
_, provider, _, _ = get_llm_provider(model=model)
optional_params: Final = get_optional_params(model=model, custom_llm_provider=provider, reasoning_effort="high")
assert optional_params["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
}
assert "reasoning_effort" not in optional_params
route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
_call(case, reasoning_effort="high", messages=[{"role": "user", "content": "hi"}])
body: Final = _request_body(route)
assert _thinking_param(case, body) == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
}
assert "reasoning_effort" not in json.dumps(body)
assert _max_tokens(case, body) == DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET + DEFAULT_MAX_TOKENS
@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id)
@pytest.mark.parametrize(
("effort", "budget"),
(
("low", DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET),
("medium", DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET),
),
)
def test_reasoning_effort_maps_to_distinct_thinking_budgets(
case: _Case, effort: str, budget: int, respx_mock: MockRouter
) -> None:
route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case))
_call(case, reasoning_effort=effort, messages=[{"role": "user", "content": "hi"}])
body: Final = _request_body(route)
assert _thinking_param(case, body) == {"type": "enabled", "budget_tokens": budget}
assert _max_tokens(case, body) == budget + DEFAULT_MAX_TOKENS

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,108 @@
import base64
import json
from typing import Final, cast
import httpx
import pytest
from respx import MockRouter
from typing_extensions import ReadOnly, TypedDict
from litellm import embedding
from litellm.utils import get_optional_params_embeddings
class _Kwargs(TypedDict, total=False):
model: ReadOnly[str]
api_key: ReadOnly[str]
api_base: ReadOnly[str]
api_version: ReadOnly[str]
aws_access_key_id: ReadOnly[str]
aws_secret_access_key: ReadOnly[str]
aws_region_name: ReadOnly[str]
class _Case(TypedDict):
id: ReadOnly[str]
provider: ReadOnly[str]
kwargs: ReadOnly[_Kwargs]
url: ReadOnly[str]
_AZURE_BASE: Final = "https://offline-embed.openai.azure.com"
_AZURE_URL: Final = f"{_AZURE_BASE}/openai/deployments/text-embedding-ada-002/embeddings?api-version=2024-02-15-preview"
_TITAN_URL: Final = "https://bedrock-runtime.us-west-2.amazonaws.com/model/amazon.titan-embed-image-v1/invoke"
_CASES: Final[tuple[_Case, ...]] = (
{
"id": "azure_text_embedding",
"provider": "azure",
"kwargs": {
"model": "azure/text-embedding-ada-002",
"api_key": "azure-offline-key",
"api_base": _AZURE_BASE,
"api_version": "2024-02-15-preview",
},
"url": _AZURE_URL,
},
{
"id": "bedrock_titan_image",
"provider": "bedrock",
"kwargs": cast(
_Kwargs,
{
"model": "bedrock/amazon.titan-embed-image-v1",
"aws_access_key_id": "AKIAFAKE",
"aws_secret_access_key": "fakesecret",
"aws_region_name": "us-west-2",
},
),
"url": _TITAN_URL,
},
)
_MAX_RETRIES_KWARGS: Final[tuple[_Kwargs, ...]] = (
*(case["kwargs"] for case in _CASES),
{"model": "volcengine/doubao-embedding-text-240715"},
{"model": "voyage/voyage-3-lite"},
)
def _canned_response(case: _Case) -> httpx.Response:
vector: Final = [0.11, 0.22, 0.33]
if case["provider"] == "azure":
return httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": vector}],
"model": "text-embedding-ada-002",
"usage": {"prompt_tokens": 2, "total_tokens": 2},
},
)
return httpx.Response(200, json={"embedding": vector, "inputTextTokenCount": 4})
@pytest.fixture(autouse=True)
def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
@pytest.mark.parametrize("kwargs", _MAX_RETRIES_KWARGS, ids=lambda k: k["model"])
def test_embedding_optional_params_max_retries(kwargs: _Kwargs) -> None:
optional_params: Final = get_optional_params_embeddings(**dict(kwargs), max_retries=20)
assert optional_params["max_retries"] == 20
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
def test_image_embedding(case: _Case, respx_mock: MockRouter) -> None:
png_b64: Final = base64.b64encode(b"\x89PNG\r\n\x1a\nFAKEPIXELS").decode()
data_url: Final = f"data:image/png;base64,{png_b64}"
route: Final = respx_mock.post(case["url"]).mock(return_value=_canned_response(case))
response: Final = embedding(**dict(case["kwargs"]), input=[data_url])
body: Final = json.loads(route.calls.last.request.content)
if case["provider"] == "azure":
assert body["input"] == [data_url]
else:
assert body["inputImage"] == png_b64
assert response.data[0]["embedding"] == [0.11, 0.22, 0.33]

View file

@ -0,0 +1,187 @@
import json
from typing import Final, Mapping, cast
import httpx
import pytest
import respx
from pydantic import JsonValue
from respx import MockRouter
from typing_extensions import ReadOnly, TypedDict
import litellm
class _Kwargs(TypedDict, total=False):
model: ReadOnly[str]
api_key: ReadOnly[str]
aws_access_key_id: ReadOnly[str]
aws_secret_access_key: ReadOnly[str]
aws_region_name: ReadOnly[str]
class _Case(TypedDict):
id: ReadOnly[str]
provider: ReadOnly[str]
kwargs: ReadOnly[_Kwargs]
url: ReadOnly[str]
expected_cost_zero: ReadOnly[bool]
billed_units: ReadOnly[Mapping[str, int]]
response_id: ReadOnly[str | None]
_AWS: Final[Mapping[str, str]] = {
"aws_access_key_id": "AKIAFAKE",
"aws_secret_access_key": "fakesecret",
"aws_region_name": "us-west-2",
}
def _bedrock_arn(model_id: str) -> str:
return f"bedrock/arn:aws:bedrock:us-west-2::foundation-model/{model_id}"
_CASES: Final[tuple[_Case, ...]] = (
{
"id": "jina_reranker",
"provider": "cohere",
"kwargs": {"model": "jina_ai/jina-reranker-v2-base-multilingual", "api_key": "jina-offline"},
"url": "https://api.jina.ai/v1/rerank",
"expected_cost_zero": False,
"billed_units": {"total_tokens": 4},
"response_id": "rerank-offline",
},
{
"id": "bedrock_amazon_rerank",
"provider": "bedrock",
"kwargs": cast(_Kwargs, {"model": _bedrock_arn("amazon.rerank-v1:0"), **dict(_AWS)}),
"url": "https://bedrock-agent-runtime.us-west-2.amazonaws.com/rerank",
"expected_cost_zero": False,
"billed_units": {"search_units": 1},
"response_id": "rerank-offline",
},
{
"id": "bedrock_cohere_rerank",
"provider": "bedrock",
"kwargs": cast(_Kwargs, {"model": _bedrock_arn("cohere.rerank-v3-5:0"), **dict(_AWS)}),
"url": "https://bedrock-agent-runtime.us-west-2.amazonaws.com/rerank",
"expected_cost_zero": False,
"billed_units": {"search_units": 1},
"response_id": "rerank-offline",
},
{
"id": "nvidia_nim_rerank",
"provider": "nvidia_nim",
"kwargs": {"model": "nvidia_nim/nvidia/llama-3_2-nv-rerankqa-1b-v2", "api_key": "nvapi-offline"},
"url": "https://ai.api.nvidia.com/v1/retrieval/nvidia/llama-3_2-nv-rerankqa-1b-v2/reranking",
"expected_cost_zero": True,
"billed_units": {"total_tokens": 4},
"response_id": None,
},
)
def _canned_response(case: _Case) -> httpx.Response:
if case["provider"] == "cohere":
return httpx.Response(
200,
json={
"id": "rerank-offline",
"results": [
{"index": 0, "relevance_score": 0.95},
{"index": 1, "relevance_score": 0.4},
],
"usage": {"total_tokens": 4},
},
)
if case["provider"] == "bedrock":
return httpx.Response(
200,
json={
"id": "rerank-offline",
"results": [
{"index": 0, "relevanceScore": 0.95},
{"index": 1, "relevanceScore": 0.4},
],
"usage": {"search_units": 1},
},
)
return httpx.Response(
200,
json={
"rankings": [
{"index": 0, "logit": 0.95},
{"index": 1, "logit": 0.4},
],
"usage": {"prompt_tokens": 4, "total_tokens": 4},
},
)
@pytest.fixture(autouse=True)
def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
def _request_body(route: respx.Route) -> Mapping[str, JsonValue]:
return json.loads(route.calls.last.request.content)
def _assert_translated_request(case: _Case, body: Mapping[str, JsonValue]) -> None:
if case["provider"] == "bedrock":
queries: Final = body["queries"]
assert queries == [{"textQuery": {"text": "hello"}, "type": "TEXT"}]
config: Final = body["rerankingConfiguration"]["bedrockRerankingConfiguration"]["modelConfiguration"]
assert config["modelArn"].endswith((".rerank-v1:0", ".rerank-v3-5:0"))
sources: Final = body["sources"]
assert len(sources) == 2
elif case["provider"] == "nvidia_nim":
assert body["model"] == "nvidia/llama-3.2-nv-rerankqa-1b-v2"
assert body["query"] == {"text": "hello"}
assert body["passages"] == [{"text": "hello"}, {"text": "world"}]
assert body["top_k"] == 2
else:
assert body["model"] == "jina-reranker-v2-base-multilingual"
assert body["query"] == "hello"
assert body["documents"] == ["hello", "world"]
assert body["top_n"] == 2
@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"])
@pytest.mark.parametrize("sync_mode", (True, False))
@pytest.mark.asyncio
async def test_basic_rerank(case: _Case, sync_mode: bool, respx_mock: MockRouter) -> None:
route: Final = respx_mock.post(case["url"]).mock(return_value=_canned_response(case))
if sync_mode:
response: Final = litellm.rerank(
**dict(case["kwargs"]),
query="hello",
documents=["hello", "world"],
top_n=2,
)
else:
response: Final = await litellm.arerank(
**dict(case["kwargs"]),
query="hello",
documents=["hello", "world"],
top_n=2,
)
body: Final = _request_body(route)
_assert_translated_request(case, body)
assert route.call_count == 1
assert isinstance(response.id, str)
if case["response_id"] is not None:
assert response.id == case["response_id"]
assert response.meta["billed_units"] == case["billed_units"]
assert len(response.results) == 2
assert response.results[0]["index"] == 0
assert response.results[0]["relevance_score"] == 0.95
assert response.results[1]["index"] == 1
assert response.results[1]["relevance_score"] == 0.4
if case["provider"] == "nvidia_nim":
assert response.results[0]["document"] == {"text": "hello"}
assert response.results[1]["document"] == {"text": "world"}
cost: Final = response._hidden_params["response_cost"]
if case["expected_cost_zero"]:
assert cost == 0.0
else:
assert cost > 0