From 22f18179f49adcc8a92e73de56ae7d59ee4b8aeb Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:47:41 +0000 Subject: [PATCH] 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 --- ..._bedrock_dynamic_auth_params_unit_tests.py | 297 --------- tests/llm_translation/test_bedrock_mantle.py | 149 ----- .../test_litellm_proxy_provider.py | 592 ------------------ 3 files changed, 1038 deletions(-) delete mode 100644 tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py delete mode 100644 tests/llm_translation/test_bedrock_mantle.py delete mode 100644 tests/llm_translation/test_litellm_proxy_provider.py diff --git a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py deleted file mode 100644 index 19662ae8ba6..00000000000 --- a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py +++ /dev/null @@ -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 - ) diff --git a/tests/llm_translation/test_bedrock_mantle.py b/tests/llm_translation/test_bedrock_mantle.py deleted file mode 100644 index 46a0c653005..00000000000 --- a/tests/llm_translation/test_bedrock_mantle.py +++ /dev/null @@ -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']}" diff --git a/tests/llm_translation/test_litellm_proxy_provider.py b/tests/llm_translation/test_litellm_proxy_provider.py deleted file mode 100644 index 8b6f37bfbc9..00000000000 --- a/tests/llm_translation/test_litellm_proxy_provider.py +++ /dev/null @@ -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)