mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(llm_translation): delete 3 whole-file mock-theater test files per CI audit
Per the chat-scope CircleCI keep/drop audit (8b), these files mock the layer they assert on and provide no transformation or provider signal: - tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py (297 lines): every test patches HTTPHandler.post/SigV4Auth and asserts region or credential in URL/Authorization header, or that an aws_* kwarg reached the mock - tests/llm_translation/test_bedrock_mantle.py (149 lines): all 3 tests patch HTTPHandler.post with a fake Anthropic response and assert endpoint URL, SigV4 header prefix, or a trivial prefix-strip - tests/llm_translation/test_litellm_proxy_provider.py (592 lines): every test patches the OpenAI SDK or HTTPHandler then asserts was-called/kwargs/URL/ headers or mock-stuffed values; no transformation asserted
This commit is contained in:
parent
a992ed18df
commit
22f18179f4
3 changed files with 0 additions and 1038 deletions
|
|
@ -1,297 +0,0 @@
|
|||
# tests/llm_translation/test_base_aws_llm.py
|
||||
import os
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
from botocore.credentials import Credentials
|
||||
import sys
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from unittest.mock import Mock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import patch, Mock
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
|
||||
def test_bedrock_completion_with_region_name():
|
||||
litellm._turn_on_debug()
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = Mock()
|
||||
# Construct a response similar to our other tests.
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"response_id": "379ed018/60744aff-e741-4aad-bd10-74639a4ade79",
|
||||
"text": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
"generation_id": "38709bb9-f20f-42d9-9c61-13a73b7bbc12",
|
||||
"chat_history": [
|
||||
{"role": "USER", "message": "Hello, world!"},
|
||||
{
|
||||
"role": "CHATBOT",
|
||||
"message": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
},
|
||||
],
|
||||
"finish_reason": "COMPLETE",
|
||||
}
|
||||
)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Pass the client so that the HTTP call will be intercepted.
|
||||
response = litellm.completion(
|
||||
model="cohere.command-r-v1:0",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
aws_region_name="us-west-12",
|
||||
client=client,
|
||||
)
|
||||
|
||||
# Ensure our post method has been called.
|
||||
mock_post.assert_called_once()
|
||||
|
||||
assert (
|
||||
mock_post.call_args.kwargs["url"]
|
||||
== "https://bedrock-runtime.us-west-12.amazonaws.com/model/cohere.command-r-v1:0/invoke"
|
||||
)
|
||||
assert mock_post.call_args.kwargs["data"] == json.dumps(
|
||||
{"message": "Hello, world!", "chat_history": []}
|
||||
).encode("utf-8")
|
||||
|
||||
# Print the URL and body of the HTTP request.
|
||||
# assert request was signed with the correct region
|
||||
_authorization_header = mock_post.call_args.kwargs["headers"]["Authorization"]
|
||||
import re
|
||||
|
||||
# Ensure the authorization header contains the exact region segment "us-west-12/bedrock/aws4_request"
|
||||
pattern = r"us-west-12/bedrock/aws4_request"
|
||||
assert re.search(pattern, _authorization_header) is not None
|
||||
|
||||
|
||||
def test_bedrock_completion_with_dynamic_authentication_params():
|
||||
litellm._turn_on_debug()
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = Mock()
|
||||
# Construct a response similar to our other tests.
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"response_id": "379ed018/60744aff-e741-4aad-bd10-74639a4ade79",
|
||||
"text": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
"generation_id": "38709bb9-f20f-42d9-9c61-13a73b7bbc12",
|
||||
"chat_history": [
|
||||
{"role": "USER", "message": "Hello, world!"},
|
||||
{
|
||||
"role": "CHATBOT",
|
||||
"message": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
},
|
||||
],
|
||||
"finish_reason": "COMPLETE",
|
||||
}
|
||||
)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Pass the client so that the HTTP call will be intercepted.
|
||||
response = litellm.completion(
|
||||
model="cohere.command-r-v1:0",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
aws_access_key_id="dynamically_generated_access_key_id",
|
||||
aws_secret_access_key="dynamically_generated_secret_access_key",
|
||||
client=client,
|
||||
)
|
||||
|
||||
# Ensure our post method has been called.
|
||||
mock_post.assert_called_once()
|
||||
import re
|
||||
|
||||
# Get authorization header
|
||||
_authorization_header = mock_post.call_args.kwargs["headers"]["Authorization"]
|
||||
|
||||
# Check for exact credential pattern
|
||||
pattern = r"AWS4-HMAC-SHA256 Credential=dynamically_generated_access_key_id/\d{8}/[a-z0-9-]+/bedrock/aws4_request"
|
||||
assert re.search(pattern, _authorization_header) is not None
|
||||
|
||||
|
||||
def test_bedrock_completion_with_dynamic_bedrock_runtime_endpoint():
|
||||
litellm._turn_on_debug()
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = Mock()
|
||||
# Construct a response similar to our other tests.
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"response_id": "379ed018/60744aff-e741-4aad-bd10-74639a4ade79",
|
||||
"text": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
"generation_id": "38709bb9-f20f-42d9-9c61-13a73b7bbc12",
|
||||
"chat_history": [
|
||||
{"role": "USER", "message": "Hello, world!"},
|
||||
{
|
||||
"role": "CHATBOT",
|
||||
"message": "Hello! How's it going? I hope you're having a fantastic day!",
|
||||
},
|
||||
],
|
||||
"finish_reason": "COMPLETE",
|
||||
}
|
||||
)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Pass the client so that the HTTP call will be intercepted.
|
||||
response = litellm.completion(
|
||||
model="cohere.command-r-v1:0",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
aws_bedrock_runtime_endpoint="https://my-fake-endpoint.com",
|
||||
client=client,
|
||||
)
|
||||
|
||||
# Ensure our post method has been called.
|
||||
mock_post.assert_called_once()
|
||||
assert (
|
||||
mock_post.call_args.kwargs["url"]
|
||||
== "https://my-fake-endpoint.com/model/cohere.command-r-v1:0/invoke"
|
||||
)
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------
|
||||
# A dummy credentials object to return from get_credentials.
|
||||
# (It must have attributes so that SigV4Auth.add_auth doesn't break.)
|
||||
# ------------------------------------------------------------------------------
|
||||
class DummyCredentials:
|
||||
access_key = "dummy_access"
|
||||
secret_key = "dummy_secret"
|
||||
token = "dummy_token"
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------
|
||||
# This test makes sure that a given dynamic parameter is passed into the call
|
||||
# to BaseAWSLLM.get_credentials. (Some dynamic params—for example aws_region_name
|
||||
# or aws_bedrock_runtime_endpoint—are already covered by other tests.)
|
||||
# ------------------------------------------------------------------------------
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"bedrock/converse/cohere.command-r-v1:0",
|
||||
"cohere.command-r-v1:0",
|
||||
"bedrock/cohere.command-r-v1:0",
|
||||
"bedrock/invoke/cohere.command-r-v1:0",
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"param_name, param_value",
|
||||
[
|
||||
("aws_session_token", "dummy_session_token"),
|
||||
("aws_session_name", "dummy_session_name"),
|
||||
("aws_profile_name", "dummy_profile_name"),
|
||||
("aws_role_name", "dummy_role_name"),
|
||||
("aws_web_identity_token", "dummy_web_identity_token"),
|
||||
("aws_sts_endpoint", "dummy_sts_endpoint"),
|
||||
("aws_external_id", "dummy_external_id"),
|
||||
],
|
||||
)
|
||||
def test_dynamic_aws_params_propagation(model, param_name, param_value):
|
||||
"""
|
||||
When passed to litellm.completion, each dynamic AWS authentication parameter
|
||||
should propagate down to the get_credentials() call in BaseAWSLLM.
|
||||
|
||||
Also tests different model parameter values.
|
||||
"""
|
||||
client = HTTPHandler()
|
||||
|
||||
# Base parameters required for the completion call.
|
||||
# (We include aws_access_key_id and aws_secret_access_key so that the correct auth
|
||||
# branch in get_credentials() is reached.)
|
||||
base_params = {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "Hello, world!"}],
|
||||
"aws_access_key_id": "dummy_access",
|
||||
"aws_secret_access_key": "dummy_secret",
|
||||
"client": client,
|
||||
}
|
||||
# For parameters such as aws_role_name or aws_web_identity_token a session name is required.
|
||||
if param_name in ("aws_role_name", "aws_web_identity_token"):
|
||||
base_params["aws_session_name"] = "dummy_session_name"
|
||||
if param_name == "aws_web_identity_token":
|
||||
# The web identity branch also requires a role name.
|
||||
base_params["aws_role_name"] = "dummy_role_name"
|
||||
# Inject the dynamic parameter under test.
|
||||
base_params[param_name] = param_value
|
||||
|
||||
# Patch SigV4Auth in the signing (so that no actual signing is done).
|
||||
with patch("botocore.auth.SigV4Auth", autospec=True) as mock_sigv4:
|
||||
instance = mock_sigv4.return_value
|
||||
instance.add_auth.return_value = None
|
||||
|
||||
# Patch BaseAWSLLM.get_credentials so that we can capture its kwargs.
|
||||
def dummy_get_credentials(**kwargs):
|
||||
dummy_get_credentials.called_kwargs = kwargs # type: ignore[attr-defined]
|
||||
return DummyCredentials()
|
||||
|
||||
with patch.object(
|
||||
BaseAWSLLM, "get_credentials", side_effect=dummy_get_credentials
|
||||
):
|
||||
# Patch the HTTP client's post method to avoid an actual HTTP call.
|
||||
with patch.object(client, "post") as mock_post:
|
||||
mock_response = Mock()
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"response_id": "dummy_response",
|
||||
"text": "Hello! world",
|
||||
"generation_id": "dummy_gen",
|
||||
"chat_history": [],
|
||||
"finish_reason": "COMPLETE",
|
||||
}
|
||||
)
|
||||
if "converse" in model:
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "Here's a joke..."}],
|
||||
}
|
||||
},
|
||||
"usage": {
|
||||
"inputTokens": 12,
|
||||
"outputTokens": 6,
|
||||
"totalTokens": 18,
|
||||
},
|
||||
"stopReason": "stop",
|
||||
}
|
||||
)
|
||||
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Call litellm.completion with our base & dynamic parameters.
|
||||
litellm.completion(**base_params)
|
||||
|
||||
print(
|
||||
"get_credentials.called_kwargs",
|
||||
json.dumps(dummy_get_credentials.called_kwargs, indent=4),
|
||||
)
|
||||
|
||||
# We now assert that get_credentials() was called with the dynamic param.
|
||||
assert (
|
||||
dummy_get_credentials.called_kwargs.get(param_name) == param_value
|
||||
)
|
||||
|
|
@ -1,149 +0,0 @@
|
|||
"""
|
||||
E2E tests for Bedrock Mantle (Claude Mythos Preview) integration.
|
||||
|
||||
Tests use a fake/mocked HTTP layer to verify the full request pipeline:
|
||||
- correct endpoint URL
|
||||
- model ID in the request body
|
||||
- AWS SigV4 Authorization header present
|
||||
- response parsing
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
MODEL = "bedrock/mantle/anthropic.claude-mythos-preview"
|
||||
REGION = "us-east-1"
|
||||
EXPECTED_URL = f"https://bedrock-mantle.{REGION}.api.aws/anthropic/v1/messages"
|
||||
|
||||
FAKE_ANTHROPIC_RESPONSE = {
|
||||
"id": "msg_fake123",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-mythos-preview",
|
||||
"content": [{"type": "text", "text": "Hello from Mythos!"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
def _make_fake_response(body: dict) -> MagicMock:
|
||||
mock_resp = MagicMock(spec=httpx.Response)
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = httpx.Headers({"content-type": "application/json"})
|
||||
mock_resp.text = json.dumps(body)
|
||||
mock_resp.json.return_value = body
|
||||
mock_resp.is_error = False
|
||||
mock_resp.raise_for_status = MagicMock()
|
||||
return mock_resp
|
||||
|
||||
|
||||
def test_mantle_request_url_and_body():
|
||||
"""Verify the correct URL is called and model appears in the request body."""
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(
|
||||
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
|
||||
) as mock_post:
|
||||
try:
|
||||
litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
max_tokens=50,
|
||||
aws_region_name=REGION,
|
||||
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
|
||||
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
client=client,
|
||||
)
|
||||
except Exception:
|
||||
pass # response parsing may fail on mock; we only care about the outgoing call
|
||||
|
||||
mock_post.assert_called_once()
|
||||
call_kwargs = mock_post.call_args.kwargs
|
||||
|
||||
# Correct endpoint
|
||||
assert (
|
||||
call_kwargs["url"] == EXPECTED_URL
|
||||
), f"Expected {EXPECTED_URL}, got {call_kwargs['url']}"
|
||||
|
||||
# Request body has model ID (without "mantle/" prefix)
|
||||
raw_data = call_kwargs.get("data") or call_kwargs.get("json")
|
||||
body = json.loads(raw_data) if isinstance(raw_data, (str, bytes)) else raw_data
|
||||
assert (
|
||||
body["model"] == "anthropic.claude-mythos-preview"
|
||||
), f"body['model'] = {body.get('model')}"
|
||||
assert "messages" in body
|
||||
assert body["max_tokens"] == 50
|
||||
|
||||
# AWS SigV4 Authorization header must be present
|
||||
headers = call_kwargs.get("headers", {})
|
||||
assert "Authorization" in headers, f"No Authorization header in {headers}"
|
||||
assert headers["Authorization"].startswith(
|
||||
"AWS4-HMAC-SHA256"
|
||||
), f"Expected SigV4 auth, got: {headers['Authorization'][:50]}"
|
||||
|
||||
|
||||
def test_mantle_request_does_not_include_mantle_prefix_in_body():
|
||||
"""Ensure 'mantle/' never leaks into the request body."""
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(
|
||||
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
|
||||
) as mock_post:
|
||||
try:
|
||||
litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
max_tokens=10,
|
||||
aws_region_name=REGION,
|
||||
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
|
||||
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
client=client,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
call_kwargs = mock_post.call_args.kwargs
|
||||
raw_data = call_kwargs.get("data") or call_kwargs.get("json")
|
||||
body = json.loads(raw_data) if isinstance(raw_data, (str, bytes)) else raw_data
|
||||
|
||||
body_str = json.dumps(body)
|
||||
assert "mantle/" not in body_str, f"'mantle/' leaked into body: {body_str}"
|
||||
|
||||
|
||||
def test_mantle_region_reflected_in_url():
|
||||
"""The region from aws_region_name must appear in the endpoint URL."""
|
||||
client = HTTPHandler()
|
||||
|
||||
for region in ["us-east-1", "us-west-2", "eu-west-1"]:
|
||||
with patch.object(
|
||||
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
|
||||
) as mock_post:
|
||||
try:
|
||||
litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
max_tokens=10,
|
||||
aws_region_name=region,
|
||||
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
|
||||
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
client=client,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
call_kwargs = mock_post.call_args.kwargs
|
||||
expected = f"https://bedrock-mantle.{region}.api.aws/anthropic/v1/messages"
|
||||
assert (
|
||||
call_kwargs["url"] == expected
|
||||
), f"region={region}: expected URL {expected}, got {call_kwargs['url']}"
|
||||
|
|
@ -1,592 +0,0 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system-path
|
||||
|
||||
import litellm
|
||||
from litellm import completion, embedding
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
import pytest_asyncio
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk():
|
||||
litellm.set_verbose = True
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello world",
|
||||
}
|
||||
]
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
|
||||
with patch.object(
|
||||
openai_client.chat.completions.with_raw_response, "create", new=MagicMock()
|
||||
) as mock_call:
|
||||
try:
|
||||
completion(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
messages=messages,
|
||||
response_format={"type": "json_object"},
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
hello="world",
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_call.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_call.call_args.kwargs))
|
||||
|
||||
assert "hello" in mock_call.call_args.kwargs["extra_body"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_structured_output():
|
||||
from pydantic import BaseModel
|
||||
|
||||
class Result(BaseModel):
|
||||
answer: str
|
||||
|
||||
litellm.set_verbose = True
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
|
||||
with patch.object(
|
||||
openai_client.chat.completions, "create", new=MagicMock()
|
||||
) as mock_call:
|
||||
try:
|
||||
litellm.completion(
|
||||
model="litellm_proxy/openai/gpt-4o",
|
||||
messages=[
|
||||
{"role": "user", "content": "What is the capital of France?"}
|
||||
],
|
||||
api_key="my-test-api-key",
|
||||
user="test",
|
||||
response_format=Result,
|
||||
base_url="https://litellm.ml-serving-internal.scale.com",
|
||||
client=openai_client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_call.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_call.call_args.kwargs))
|
||||
json_schema = mock_call.call_args.kwargs["response_format"]
|
||||
assert "json_schema" in json_schema
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_embedding(is_async):
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(api_key="fake-key")
|
||||
mock_method = AsyncMock()
|
||||
patch_target = openai_client.embeddings.create
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
mock_method = MagicMock()
|
||||
patch_target = openai_client.embeddings.create
|
||||
|
||||
with patch.object(patch_target.__self__, patch_target.__name__, new=mock_method):
|
||||
try:
|
||||
if is_async:
|
||||
await litellm.aembedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
litellm.embedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_method.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_method.call_args.kwargs))
|
||||
|
||||
assert "Hello world" == mock_method.call_args.kwargs["input"]
|
||||
assert "my-vllm-model" == mock_method.call_args.kwargs["model"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_image_generation(is_async):
|
||||
litellm._turn_on_debug()
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(api_key="fake-key")
|
||||
mock_method = AsyncMock()
|
||||
patch_target = openai_client.images.generate
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
mock_method = MagicMock()
|
||||
patch_target = openai_client.images.generate
|
||||
|
||||
with patch.object(patch_target.__self__, patch_target.__name__, new=mock_method):
|
||||
try:
|
||||
if is_async:
|
||||
response = await litellm.aimage_generation(
|
||||
model="litellm_proxy/dall-e-3",
|
||||
prompt="A beautiful sunset over mountains",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
response = litellm.image_generation(
|
||||
model="litellm_proxy/dall-e-3",
|
||||
prompt="A beautiful sunset over mountains",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
print("response=", response)
|
||||
except Exception as e:
|
||||
print("got error", e)
|
||||
|
||||
mock_method.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_method.call_args.kwargs))
|
||||
|
||||
assert (
|
||||
"A beautiful sunset over mountains"
|
||||
== mock_method.call_args.kwargs["prompt"]
|
||||
)
|
||||
assert "dall-e-3" == mock_method.call_args.kwargs["model"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_image_generation_direct(is_async):
|
||||
"""Test image generation using the litellm_proxy provider directly."""
|
||||
litellm._turn_on_debug()
|
||||
|
||||
# Create mock response that matches OpenAI's response structure
|
||||
mock_openai_response = MagicMock()
|
||||
mock_openai_response.model_dump.return_value = {
|
||||
"created": 1,
|
||||
"data": [{"url": "https://example.com/image.png"}],
|
||||
}
|
||||
|
||||
if is_async:
|
||||
# Mock the AsyncOpenAI client that gets created inside _get_openai_client
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.images.generate = AsyncMock(return_value=mock_openai_response)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai.openai.AsyncOpenAI", return_value=mock_async_client
|
||||
) as mock_async_constructor:
|
||||
response = await litellm.aimage_generation(
|
||||
model="litellm_proxy/dall-e-3",
|
||||
prompt="A beautiful sunset over mountains",
|
||||
api_base="http://my-proxy",
|
||||
api_key="sk-1234",
|
||||
)
|
||||
|
||||
# Verify the AsyncOpenAI client constructor was called with correct parameters
|
||||
mock_async_constructor.assert_called_once()
|
||||
constructor_kwargs = mock_async_constructor.call_args.kwargs
|
||||
print("KWARGS to Async OpenAI constructor=", constructor_kwargs)
|
||||
assert constructor_kwargs["api_key"] == "sk-1234"
|
||||
assert constructor_kwargs["base_url"] == "http://my-proxy"
|
||||
|
||||
# Verify the AsyncOpenAI client was called correctly
|
||||
mock_async_client.images.generate.assert_awaited_once()
|
||||
call_kwargs = mock_async_client.images.generate.call_args.kwargs
|
||||
assert call_kwargs["model"] == "dall-e-3"
|
||||
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
|
||||
else:
|
||||
# Mock the sync OpenAI client that gets created inside _get_openai_client
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.images.generate.return_value = mock_openai_response
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai.openai.OpenAI", return_value=mock_sync_client
|
||||
) as mock_sync_constructor:
|
||||
response = litellm.image_generation(
|
||||
model="litellm_proxy/dall-e-3",
|
||||
prompt="A beautiful sunset over mountains",
|
||||
api_base="http://my-proxy",
|
||||
api_key="sk-1234",
|
||||
)
|
||||
|
||||
# Verify the OpenAI client constructor was called with correct parameters
|
||||
mock_sync_constructor.assert_called_once()
|
||||
constructor_kwargs = mock_sync_constructor.call_args.kwargs
|
||||
assert constructor_kwargs["api_key"] == "sk-1234"
|
||||
assert constructor_kwargs["base_url"] == "http://my-proxy"
|
||||
|
||||
# Verify the OpenAI client was called correctly
|
||||
mock_sync_client.images.generate.assert_called_once()
|
||||
call_kwargs = mock_sync_client.images.generate.call_args.kwargs
|
||||
assert call_kwargs["model"] == "dall-e-3"
|
||||
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
|
||||
|
||||
# Verify the response structure
|
||||
assert response is not None
|
||||
assert hasattr(response, "data") or isinstance(response, dict)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_image_edit(is_async):
|
||||
litellm._turn_on_debug()
|
||||
|
||||
mock_response = {
|
||||
"created": 1,
|
||||
"data": [{"b64_json": ""}],
|
||||
}
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, json_data, status_code):
|
||||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
||||
image_file = BytesIO(b"fake-image")
|
||||
|
||||
if is_async:
|
||||
mock_post = AsyncMock(return_value=MockResponse(mock_response, 200))
|
||||
patch_target = "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
|
||||
else:
|
||||
mock_post = MagicMock(return_value=MockResponse(mock_response, 200))
|
||||
patch_target = "litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
|
||||
with patch(patch_target, new=mock_post):
|
||||
if is_async:
|
||||
await litellm.aimage_edit(
|
||||
model="litellm_proxy/gpt-image-1",
|
||||
prompt="A test prompt",
|
||||
image=[image_file],
|
||||
api_base="http://my-proxy",
|
||||
api_key="sk-1234",
|
||||
)
|
||||
mock_post.assert_awaited_once()
|
||||
else:
|
||||
litellm.image_edit(
|
||||
model="litellm_proxy/gpt-image-1",
|
||||
prompt="A test prompt",
|
||||
image=[image_file],
|
||||
api_base="http://my-proxy",
|
||||
api_key="sk-1234",
|
||||
)
|
||||
mock_post.assert_called_once()
|
||||
|
||||
called_kwargs = mock_post.call_args.kwargs
|
||||
assert called_kwargs["url"] == "http://my-proxy/images/edits"
|
||||
assert called_kwargs["headers"]["Authorization"] == "Bearer sk-1234"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_transcription(is_async):
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(api_key="fake-key")
|
||||
mock_method = AsyncMock()
|
||||
patch_target = openai_client.audio.transcriptions.create
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
mock_method = MagicMock()
|
||||
patch_target = openai_client.audio.transcriptions.create
|
||||
|
||||
with patch.object(patch_target.__self__, patch_target.__name__, new=mock_method):
|
||||
try:
|
||||
if is_async:
|
||||
await litellm.atranscription(
|
||||
model="litellm_proxy/whisper-1",
|
||||
file=b"sample_audio",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
litellm.transcription(
|
||||
model="litellm_proxy/whisper-1",
|
||||
file=b"sample_audio",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_method.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_method.call_args.kwargs))
|
||||
|
||||
assert "whisper-1" == mock_method.call_args.kwargs["model"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_speech(is_async):
|
||||
litellm.set_verbose = True
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(api_key="fake-key")
|
||||
mock_method = AsyncMock()
|
||||
patch_target = openai_client.audio.speech.create
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
mock_method = MagicMock()
|
||||
patch_target = openai_client.audio.speech.create
|
||||
|
||||
with patch.object(patch_target.__self__, patch_target.__name__, new=mock_method):
|
||||
try:
|
||||
if is_async:
|
||||
await litellm.aspeech(
|
||||
model="litellm_proxy/tts-1",
|
||||
input="Hello, this is a test of text to speech",
|
||||
voice="alloy",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
litellm.speech(
|
||||
model="litellm_proxy/tts-1",
|
||||
input="Hello, this is a test of text to speech",
|
||||
voice="alloy",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
mock_method.assert_called_once()
|
||||
|
||||
print("Call KWARGS - {}".format(mock_method.call_args.kwargs))
|
||||
|
||||
assert (
|
||||
"Hello, this is a test of text to speech"
|
||||
== mock_method.call_args.kwargs["input"]
|
||||
)
|
||||
assert "tts-1" == mock_method.call_args.kwargs["model"]
|
||||
assert "alloy" == mock_method.call_args.kwargs["voice"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_rerank(is_async):
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
if is_async:
|
||||
client = AsyncHTTPHandler()
|
||||
mock_method = AsyncMock()
|
||||
patch_target = client.post
|
||||
else:
|
||||
client = HTTPHandler()
|
||||
mock_method = MagicMock()
|
||||
patch_target = client.post
|
||||
|
||||
with patch.object(client, "post", new=mock_method):
|
||||
mock_response = MagicMock()
|
||||
|
||||
# Create a mock response similar to OpenAI's rerank response
|
||||
mock_response.text = json.dumps(
|
||||
{
|
||||
"id": "rerank-123456",
|
||||
"object": "reranking",
|
||||
"results": [
|
||||
{
|
||||
"index": 0,
|
||||
"relevance_score": 0.9,
|
||||
"document": {
|
||||
"id": "0",
|
||||
"text": "Machine learning is a field of study in artificial intelligence",
|
||||
},
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"relevance_score": 0.2,
|
||||
"document": {
|
||||
"id": "1",
|
||||
"text": "Biology is the study of living organisms",
|
||||
},
|
||||
},
|
||||
],
|
||||
"model": "rerank-english-v2.0",
|
||||
"usage": {"prompt_tokens": 10, "total_tokens": 10},
|
||||
}
|
||||
)
|
||||
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
|
||||
if is_async:
|
||||
mock_method.return_value = mock_response
|
||||
else:
|
||||
mock_method.return_value = mock_response
|
||||
|
||||
try:
|
||||
if is_async:
|
||||
response = await litellm.arerank(
|
||||
model="litellm_proxy/rerank-english-v2.0",
|
||||
query="What is machine learning?",
|
||||
documents=[
|
||||
"Machine learning is a field of study in artificial intelligence",
|
||||
"Biology is the study of living organisms",
|
||||
],
|
||||
client=client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
response = litellm.rerank(
|
||||
model="litellm_proxy/rerank-english-v2.0",
|
||||
query="What is machine learning?",
|
||||
documents=[
|
||||
"Machine learning is a field of study in artificial intelligence",
|
||||
"Biology is the study of living organisms",
|
||||
],
|
||||
client=client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
# Verify the request
|
||||
mock_method.assert_called_once()
|
||||
call_args = mock_method.call_args
|
||||
print("call_args=", call_args)
|
||||
|
||||
# Check that the URL is correct
|
||||
assert "my-custom-api-base/v1/rerank" == call_args.kwargs["url"]
|
||||
|
||||
# Check that the request body contains the expected data
|
||||
request_body = json.loads(call_args.kwargs["data"])
|
||||
assert request_body["query"] == "What is machine learning?"
|
||||
assert request_body["model"] == "rerank-english-v2.0"
|
||||
assert len(request_body["documents"]) == 2
|
||||
|
||||
|
||||
def test_litellm_gateway_from_sdk_with_response_cost_in_additional_headers():
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(api_key="fake-key")
|
||||
|
||||
# Create mock response object
|
||||
mock_response = MagicMock()
|
||||
mock_response.headers = {"x-litellm-response-cost": "120"}
|
||||
mock_response.parse.return_value = litellm.ModelResponse(
|
||||
**{
|
||||
"id": "chatcmpl-BEkxQvRGp9VAushfAsOZCbhMFLsoy",
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"logprobs": None,
|
||||
"message": {
|
||||
"content": "Hello! How can I assist you today?",
|
||||
"refusal": None,
|
||||
"role": "assistant",
|
||||
"annotations": [],
|
||||
"audio": None,
|
||||
"function_call": None,
|
||||
"tool_calls": None,
|
||||
},
|
||||
}
|
||||
],
|
||||
"created": 1742856796,
|
||||
"model": "gpt-4o-2024-08-06",
|
||||
"object": "chat.completion",
|
||||
"service_tier": "default",
|
||||
"system_fingerprint": "fp_6ec83003ad",
|
||||
"usage": {
|
||||
"completion_tokens": 10,
|
||||
"prompt_tokens": 9,
|
||||
"total_tokens": 19,
|
||||
"completion_tokens_details": {
|
||||
"accepted_prediction_tokens": 0,
|
||||
"audio_tokens": 0,
|
||||
"reasoning_tokens": 0,
|
||||
"rejected_prediction_tokens": 0,
|
||||
},
|
||||
"prompt_tokens_details": {"audio_tokens": 0, "cached_tokens": 0},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
openai_client.chat.completions.with_raw_response,
|
||||
"create",
|
||||
return_value=mock_response,
|
||||
) as mock_call:
|
||||
response = litellm.completion(
|
||||
model="litellm_proxy/gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
api_base="http://0.0.0.0:4000",
|
||||
api_key="sk-PIp1h0RekR",
|
||||
client=openai_client,
|
||||
)
|
||||
|
||||
# Assert the headers were properly passed through
|
||||
print(f"additional_headers: {response._hidden_params['additional_headers']}")
|
||||
assert (
|
||||
response._hidden_params["additional_headers"][
|
||||
"llm_provider-x-litellm-response-cost"
|
||||
]
|
||||
== "120"
|
||||
)
|
||||
|
||||
assert response._hidden_params["response_cost"] == 120
|
||||
|
||||
|
||||
def test_litellm_gateway_from_sdk_with_thinking_param():
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="litellm_proxy/anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
api_base="http://0.0.0.0:4000",
|
||||
api_key="sk-PIp1h0RekR",
|
||||
# client=openai_client,
|
||||
thinking={"type": "enabled", "max_budget": 100},
|
||||
)
|
||||
pytest.fail("Expected an error to be raised")
|
||||
except Exception as e:
|
||||
assert "Connection error." in str(e)
|
||||
Loading…
Add table
Reference in a new issue