test(bedrock): drop mock-theater tests from bedrock/aws split files per CI audit

Function-level deletions per the keep/drop audit (5), patterns: (a)
patch-then-assert-called, (b) URL/header asserted on mocked HTTP, (c)
kwarg-reached-a-mock, (d) asserting values the test itself stuffed in.

- test_aws_base_llm.py: test_auth_with_web_identity_token, _aws_role,
  _aws_profile, _access_key_and_secret_key, _env_vars (mock-STS/boto3
  passthrough, a/d); mock_credentials fixture now unused
- test_bedrock_completion.py: test_bedrock_ptu (calibrated drop),
  test_bedrock_custom_api_base, test_bedrock_extra_headers,
  test_completion_bedrock_external_client_region,
  test_bedrock_cross_region_inference (monkeypatch duplicate that shadowed
  the live variant; deleting it un-shadows the keeper),
  test_bedrock_image_url_sync_client, test_bedrock_custom_proxy,
  test_bedrock_application_inference_profile (b),
  test_bedrock_meta_llama_function_calling (zero assertions)
- test_bedrock_agentcore.py: 6 with_custom_params/runtime_user_id/
  session_and_user/api_key_bearer_token/all_parameters/sigv4 tests (b)
- test_bedrock_agents.py: test_bedrock_agents_with_custom_params (a)
- test_bedrock_embedding.py: 2 async-invoke marengo tests (stuffed
  invocationArn, d); 2 region-in-URL tests (b)
- test_bedrock_govcloud.py: test_govcloud_client_initialization (a)
- test_bedrock_moonshot.py: 6 TestBedrockMoonshotInvoke mock overrides of
  live base tests + TestBedrockMoonshotToolCalling.
  test_tool_response_message_format (asserts dict it constructed, d)
- test_unit_test_bedrock_invoke.py: test_sign_request_basic (mocks botocore,
  asserts mocks called, a)

Unused imports left dangling by the deletions removed via ruff F401.
This commit is contained in:
mateo-berri 2026-06-11 18:52:21 +00:00
parent 22f18179f4
commit 00eb3dbbaf
8 changed files with 6 additions and 1351 deletions

View file

@ -1,13 +1,7 @@
import pytest
import os
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
from botocore.credentials import Credentials
from typing import Dict, Any
from litellm.llms.bedrock.base_aws_llm import (
BaseAWSLLM,
AwsAuthError,
Boto3CredentialsInfo,
)
@ -17,13 +11,6 @@ def base_aws_llm():
return BaseAWSLLM()
@pytest.fixture
def mock_credentials():
return Credentials(
access_key="test_access", secret_key="test_secret", token="test_token"
)
# Test cache key generation
def test_get_cache_key(base_aws_llm):
test_args = {
@ -35,82 +22,6 @@ def test_get_cache_key(base_aws_llm):
assert len(cache_key) == 64 # SHA-256 produces 64 character hex string
# Test web identity token authentication
@patch("boto3.client")
@patch("litellm.llms.bedrock.base_aws_llm.get_secret") # Add this patch
def test_auth_with_web_identity_token(mock_get_secret, mock_boto3_client, base_aws_llm):
# Mock get_secret to return a token
mock_get_secret.return_value = "mocked_oidc_token"
# Mock the STS client and response
mock_sts = MagicMock()
mock_sts.assume_role_with_web_identity.return_value = {
"Credentials": {
"AccessKeyId": "test_access",
"SecretAccessKey": "test_secret",
"SessionToken": "test_token",
},
"PackedPolicySize": 10,
}
mock_boto3_client.return_value = mock_sts
credentials, ttl = base_aws_llm._auth_with_web_identity_token(
aws_web_identity_token="test_token",
aws_role_name="test_role",
aws_session_name="test_session",
aws_region_name="us-west-2",
aws_sts_endpoint=None,
)
# Verify get_secret was called with the correct argument
mock_get_secret.assert_called_once_with("test_token")
assert isinstance(credentials, Credentials)
assert ttl == 3540 # default TTL (3600 - 60)
# Test AWS role authentication
@patch("boto3.client")
def test_auth_with_aws_role(mock_boto3_client, base_aws_llm):
# Mock the STS client and response
mock_sts = MagicMock()
expiry_time = datetime.now(timezone.utc)
mock_sts.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "test_access",
"SecretAccessKey": "test_secret",
"SessionToken": "test_token",
"Expiration": expiry_time,
}
}
mock_boto3_client.return_value = mock_sts
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id="test_access",
aws_secret_access_key="test_secret",
aws_session_token="test_token",
aws_role_name="test_role",
aws_session_name="test_session",
)
assert isinstance(credentials, Credentials)
assert isinstance(ttl, float)
# Test AWS profile authentication
@patch("boto3.Session")
def test_auth_with_aws_profile(mock_session, base_aws_llm, mock_credentials):
# Mock the session
mock_session_instance = MagicMock()
mock_session_instance.get_credentials.return_value = mock_credentials
mock_session.return_value = mock_session_instance
credentials, ttl = base_aws_llm._auth_with_aws_profile("test_profile")
assert credentials == mock_credentials
assert ttl is None
# Test session token authentication
def test_auth_with_aws_session_token(base_aws_llm):
credentials, ttl = base_aws_llm._auth_with_aws_session_token(
@ -126,40 +37,6 @@ def test_auth_with_aws_session_token(base_aws_llm):
assert ttl is None
# Test access key and secret key authentication
@patch("boto3.Session")
def test_auth_with_access_key_and_secret_key(
mock_session, base_aws_llm, mock_credentials
):
# Mock the session
mock_session_instance = MagicMock()
mock_session_instance.get_credentials.return_value = mock_credentials
mock_session.return_value = mock_session_instance
credentials, ttl = base_aws_llm._auth_with_access_key_and_secret_key(
aws_access_key_id="test_access",
aws_secret_access_key="test_secret",
aws_region_name="us-west-2",
)
assert credentials == mock_credentials
assert ttl == 3540 # default TTL (3600 - 60)
# Test environment variables authentication
@patch("boto3.Session")
def test_auth_with_env_vars(mock_session, base_aws_llm, mock_credentials):
# Mock the session
mock_session_instance = MagicMock()
mock_session_instance.get_credentials.return_value = mock_credentials
mock_session.return_value = mock_session_instance
credentials, ttl = base_aws_llm._auth_with_env_vars()
assert credentials == mock_credentials
assert ttl is None
# Test runtime endpoint resolution
def test_get_runtime_endpoint(base_aws_llm):
endpoint_url, proxy_endpoint_url = base_aws_llm.get_runtime_endpoint(

View file

@ -68,325 +68,6 @@ async def test_bedrock_agentcore_with_streaming(model):
print("chunk=", chunk)
def test_bedrock_agentcore_with_custom_params():
"""
Test AgentCore request structure with custom parameters
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Explain machine learning in simple terms",
}
],
runtimeSessionId="litellm-test-session-id-12345678901234567890",
qualifier="DEFAULT",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify URL structure - should include ARN and qualifier
assert "url" in call_kwargs
url = call_kwargs["url"]
print(f"URL: {url}")
assert (
"/runtimes/arn%3Aaws%3Abedrock-agentcore%3Aus-west-2%3A888602223428%3Aruntime%2Fhosted_agent_r9jvp-3ySZuRHjLC/invocations"
in url
)
assert "qualifier=DEFAULT" in url
# Verify headers - session ID should be in header
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
assert "X-Amzn-Bedrock-AgentCore-Runtime-Session-Id" in headers
assert (
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"]
== "litellm-test-session-id-12345678901234567890"
)
# Verify the request body - should just be the payload
assert "data" in call_kwargs or "json" in call_kwargs
# Parse the request data
if "data" in call_kwargs:
request_data = json.loads(call_kwargs["data"])
else:
request_data = call_kwargs["json"]
print(f"Request data: {json.dumps(request_data, indent=2)}")
# Body should just contain the prompt
assert "prompt" in request_data
assert request_data["prompt"] == "Explain machine learning in simple terms"
def test_bedrock_agentcore_with_runtime_user_id():
"""
Test AgentCore with runtimeUserId parameter
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Hello",
}
],
runtimeUserId="test-user-123",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify headers - user ID should be in header
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
assert "X-Amzn-Bedrock-AgentCore-Runtime-User-Id" in headers
assert headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] == "test-user-123"
def test_bedrock_agentcore_with_session_and_user():
"""
Test AgentCore with both runtimeSessionId and runtimeUserId
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Test message",
}
],
runtimeSessionId="session-abc-123",
runtimeUserId="user-xyz-789",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify headers contain both session and user IDs
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
assert "X-Amzn-Bedrock-AgentCore-Runtime-Session-Id" in headers
assert (
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] == "session-abc-123"
)
assert "X-Amzn-Bedrock-AgentCore-Runtime-User-Id" in headers
assert headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] == "user-xyz-789"
def test_bedrock_agentcore_with_api_key_bearer_token():
"""
Test AgentCore with api_key parameter for JWT/Bearer token authentication
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
test_jwt_token = "test-jwt-token-header.payload.signature"
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Test JWT authentication",
}
],
api_key=test_jwt_token,
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify Authorization header with Bearer token
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
assert "Authorization" in headers
assert headers["Authorization"] == f"Bearer {test_jwt_token}"
assert headers["Content-Type"] == "application/json"
# Verify the request body is JSON-encoded (not SigV4 signed)
assert "data" in call_kwargs
request_data = json.loads(call_kwargs["data"])
print(f"Request data: {json.dumps(request_data, indent=2)}")
assert "prompt" in request_data
assert request_data["prompt"] == "Test JWT authentication"
def test_bedrock_agentcore_with_all_parameters():
"""
Test AgentCore with all parameters: api_key, runtimeSessionId, runtimeUserId
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
test_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.signature"
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Complete test",
}
],
api_key=test_jwt_token,
runtimeSessionId="full-test-session-id",
runtimeUserId="full-test-user-id",
qualifier="LATEST",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify URL includes qualifier
assert "url" in call_kwargs
url = call_kwargs["url"]
print(f"URL: {url}")
assert "qualifier=LATEST" in url
# Verify all headers are present
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
# Check Bearer token authorization
assert "Authorization" in headers
assert headers["Authorization"] == f"Bearer {test_jwt_token}"
# Check session and user IDs
assert "X-Amzn-Bedrock-AgentCore-Runtime-Session-Id" in headers
assert (
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"]
== "full-test-session-id"
)
assert "X-Amzn-Bedrock-AgentCore-Runtime-User-Id" in headers
assert (
headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] == "full-test-user-id"
)
# Verify JSON body
assert "data" in call_kwargs
request_data = json.loads(call_kwargs["data"])
print(f"Request data: {json.dumps(request_data, indent=2)}")
assert "prompt" in request_data
assert request_data["prompt"] == "Complete test"
def test_bedrock_agentcore_without_api_key_uses_sigv4():
"""
Test that AgentCore uses AWS SigV4 signing when api_key is not provided
"""
import json
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
messages=[
{
"role": "user",
"content": "Test SigV4",
}
],
# No api_key provided - should use SigV4
runtimeSessionId="sigv4-test-session",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
print(f"mock_post.call_args.kwargs: {call_kwargs}")
# Verify headers - should have AWS SigV4 headers, not Bearer token
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
print(f"Headers: {headers}")
# Should NOT have Bearer Authorization when using SigV4
if "Authorization" in headers:
assert not headers["Authorization"].startswith("Bearer ")
# Should have AWS4-HMAC-SHA256 signature
assert "AWS4-HMAC-SHA256" in headers["Authorization"]
# Session ID should still be present
assert "X-Amzn-Bedrock-AgentCore-Runtime-Session-Id" in headers
assert (
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"]
== "sigv4-test-session"
)
def test_agentcore_parse_json_response():
"""
Unit test for JSON response parsing (non-streaming)

View file

@ -1,20 +1,15 @@
import os
import sys
import traceback
from dotenv import load_dotenv
import litellm.types
load_dotenv()
import io
import os
import json
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
from unittest.mock import AsyncMock, Mock, patch
import pytest
@ -65,28 +60,3 @@ async def test_bedrock_agents_with_streaming():
pass
def test_bedrock_agents_with_custom_params():
litellm._turn_on_debug()
from unittest.mock import MagicMock, patch
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agent/L1RT58GYRW/MFPSBCXYTW",
messages=[
{
"role": "user",
"content": "Hi who is ishaan cto of litellm, tell me 10 things about him",
}
],
invocationId="my-test-invocation-id",
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_post.assert_called_once()
print(f"mock_post.call_args.kwargs: {mock_post.call_args.kwargs}")

View file

@ -5,15 +5,12 @@ Tests Bedrock Completion + Rerank endpoints
# @pytest.mark.skip(reason="AWS Suspended Account")
import os
import sys
import traceback
from dotenv import load_dotenv
import litellm.types
load_dotenv()
import io
import os
import json
sys.path.insert(
@ -28,10 +25,8 @@ from litellm import (
ModelResponse,
RateLimitError,
ServiceUnavailableError,
Timeout,
completion,
completion_cost,
embedding,
)
from litellm.llms.bedrock.chat import BedrockLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
@ -95,12 +90,9 @@ def test_completion_bedrock_claude_completion_auth():
@pytest.mark.parametrize("streaming", [True, False])
def test_completion_bedrock_guardrails(streaming):
import os
litellm.set_verbose = True
import logging
from litellm._logging import verbose_logger
# verbose_logger.setLevel(logging.DEBUG)
try:
@ -253,7 +245,6 @@ def bedrock_session_token_creds():
def process_stream_response(res, messages):
import types
if isinstance(res, litellm.utils.CustomStreamWrapper):
chunks = []
@ -690,7 +681,6 @@ def test_completion_claude_3_base64():
def test_completion_bedrock_mistral_completion_auth():
print("calling bedrock mistral completion params auth")
import os
litellm._turn_on_debug()
@ -724,116 +714,6 @@ def test_completion_bedrock_mistral_completion_auth():
# test_completion_bedrock_mistral_completion_auth()
def test_bedrock_ptu():
"""
Check if a url with 'modelId' passed in, is created correctly
Reference: https://github.com/BerriAI/litellm/issues/3805
"""
client = HTTPHandler()
with patch.object(client, "post", new=Mock()) as mock_client_post:
litellm.set_verbose = True
from openai.types.chat import ChatCompletion
model_id = (
"arn:aws:bedrock:us-west-2:888602223428:provisioned-model/8fxff74qyhs3"
)
try:
response = litellm.completion(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "What's AWS?"}],
model_id=model_id,
client=client,
)
except Exception as e:
pass
assert "url" in mock_client_post.call_args.kwargs
assert (
mock_client_post.call_args.kwargs["url"]
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-west-2%3A888602223428%3Aprovisioned-model%2F8fxff74qyhs3/converse"
)
mock_client_post.assert_called_once()
@pytest.mark.asyncio
async def test_bedrock_custom_api_base():
"""
Check if a url with 'modelId' passed in, is created correctly
Reference: https://github.com/BerriAI/litellm/issues/3805, https://github.com/BerriAI/litellm/issues/5389#issuecomment-2313677977
"""
client = AsyncHTTPHandler()
with patch.object(client, "post", new=AsyncMock()) as mock_client_post:
litellm.set_verbose = True
from openai.types.chat import ChatCompletion
try:
response = await litellm.acompletion(
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "What's AWS?"}],
client=client,
extra_headers={"test": "hello world", "Authorization": "my-test-key"},
api_base="https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/aws-bedrock/bedrock-runtime/us-east-1",
)
except Exception as e:
pass
print(f"mock_client_post.call_args.kwargs: {mock_client_post.call_args.kwargs}")
assert (
mock_client_post.call_args.kwargs["url"]
== "https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/aws-bedrock/bedrock-runtime/us-east-1/model/anthropic.claude-3-sonnet-20240229-v1%3A0/converse"
)
assert "test" in mock_client_post.call_args.kwargs["headers"]
assert mock_client_post.call_args.kwargs["headers"]["test"] == "hello world"
assert (
mock_client_post.call_args.kwargs["headers"]["Authorization"]
== "my-test-key"
)
mock_client_post.assert_called_once()
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0",
],
)
@pytest.mark.asyncio
async def test_bedrock_extra_headers(model):
"""
Relevant Issue: https://github.com/BerriAI/litellm/issues/9106
"""
client = AsyncHTTPHandler()
with patch.object(client, "post", new=AsyncMock()) as mock_client_post:
litellm.set_verbose = True
from openai.types.chat import ChatCompletion
try:
response = await litellm.acompletion(
model=model,
messages=[{"role": "user", "content": "What's AWS?"}],
client=client,
extra_headers={"test": "hello world", "Authorization": "my-test-key"},
)
except Exception as e:
print(f"error: {e}")
print(f"mock_client_post.call_args.kwargs: {mock_client_post.call_args.kwargs}")
assert "test" in mock_client_post.call_args.kwargs["headers"]
assert mock_client_post.call_args.kwargs["headers"]["test"] == "hello world"
assert (
mock_client_post.call_args.kwargs["headers"]["Authorization"]
== "my-test-key"
)
mock_client_post.assert_called_once()
@pytest.mark.asyncio
async def test_bedrock_custom_prompt_template():
"""
@ -879,59 +759,6 @@ async def test_bedrock_custom_prompt_template():
mock_client_post.assert_called_once()
def test_completion_bedrock_external_client_region():
print("\ncalling bedrock claude external client auth")
import os
aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"]
aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"]
aws_region_name = "us-east-1"
os.environ.pop("AWS_ACCESS_KEY_ID", None)
os.environ.pop("AWS_SECRET_ACCESS_KEY", None)
client = HTTPHandler()
try:
import boto3
litellm.set_verbose = True
bedrock = boto3.client(
service_name="bedrock-runtime",
region_name=aws_region_name,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
endpoint_url=f"https://bedrock-runtime.{aws_region_name}.amazonaws.com",
)
with patch.object(client, "post", new=Mock()) as mock_client_post:
try:
response = completion(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=messages,
max_tokens=10,
temperature=0.1,
aws_bedrock_client=bedrock,
client=client,
)
# Add any assertions here to check the response
print(response)
except Exception as e:
pass
print(f"mock_client_post.call_args: {mock_client_post.call_args}")
assert "us-east-1" in mock_client_post.call_args.kwargs["url"]
mock_client_post.assert_called_once()
os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id
os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key
except RateLimitError:
pass
except Exception as e:
pytest.fail(f"Error occurred: {e}")
def test_bedrock_tool_calling():
"""
# related issue: https://github.com/BerriAI/litellm/issues/5007
@ -1233,7 +1060,6 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
def test_bedrock_converse_translation_tool_message():
from litellm.types.utils import ChatCompletionMessageToolCall, Function
litellm.set_verbose = True
@ -2299,32 +2125,6 @@ def test_bedrock_nova_topk(top_k_param):
)
def test_bedrock_cross_region_inference(monkeypatch):
from litellm.llms.custom_httpx.http_handler import HTTPHandler
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.add_known_models()
litellm.set_verbose = True
client = HTTPHandler()
with patch.object(client, "post") as mock_post:
try:
completion(
model="bedrock/us.meta.llama3-3-70b-instruct-v1:0",
messages=[{"role": "user", "content": "Hello, world!"}],
client=client,
)
except Exception as e:
print(e)
assert (
mock_post.call_args.kwargs["url"]
== "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/anthropic.claude-3-sonnet-20240229-v1:0",
@ -2464,50 +2264,12 @@ class TestBedrockEmbedding(BaseLLMEmbeddingTest):
] == "iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII="
@pytest.mark.asyncio
async def test_bedrock_image_url_sync_client():
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
import logging
from litellm import verbose_logger
verbose_logger.setLevel(level=logging.DEBUG)
litellm._turn_on_debug()
client = AsyncHTTPHandler()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {
"url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png"
},
},
],
}
]
with patch.object(client, "post") as mock_post:
try:
await litellm.acompletion(
model="bedrock/us.amazon.nova-pro-v1:0",
messages=messages,
client=client,
)
except Exception as e:
print(e)
mock_post.assert_called_once()
def test_bedrock_error_handling_streaming():
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
BedrockError,
)
from unittest.mock import patch, Mock
from unittest.mock import Mock
event = Mock()
event.to_response_dict = Mock(
@ -2569,29 +2331,6 @@ async def test_bedrock_document_understanding(image_url):
pytest.skip("Skipping test due to ServiceUnavailableError")
def test_bedrock_custom_proxy():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post") as mock_post:
try:
response = completion(
model="bedrock/converse_like/us.amazon.nova-pro-v1:0",
messages=[{"content": "Tell me a joke", "role": "user"}],
api_key="Token",
client=client,
api_base="https://some-api-url/models",
)
except Exception as e:
print(e)
print(mock_post.call_args.kwargs)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["url"] == "https://some-api-url/models"
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer Token"
def test_bedrock_custom_deepseek():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
import json
@ -2943,71 +2682,6 @@ async def test_bedrock_stream_thinking_content_openwebui():
), "There should be non-empty content after thinking tags"
def test_bedrock_application_inference_profile():
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
client = HTTPHandler()
client2 = HTTPHandler()
tools = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
},
},
"required": ["location"],
},
},
}
]
with (
patch.object(client, "post") as mock_post,
patch.object(client2, "post") as mock_post2,
):
try:
resp = completion(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "Hello, how are you?"}],
model_id="arn:aws:bedrock:eu-central-1:000000000000:application-inference-profile/a0a0a0a0a0a0",
client=client,
tools=tools,
)
except Exception as e:
print(e)
try:
resp = completion(
model="bedrock/converse/arn:aws:bedrock:eu-central-1:000000000000:application-inference-profile/a0a0a0a0a0a0",
messages=[{"role": "user", "content": "Hello, how are you?"}],
client=client2,
tools=tools,
)
except Exception as e:
print(e)
mock_post.assert_called_once()
mock_post2.assert_called_once()
print(mock_post.call_args.kwargs)
json_data = mock_post.call_args.kwargs["data"]
assert mock_post.call_args.kwargs["url"].startswith(
"https://bedrock-runtime.eu-central-1.amazonaws.com/"
)
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 {
@ -3067,57 +2741,6 @@ async def test_bedrock_max_completion_tokens(model: str):
}
def test_bedrock_meta_llama_function_calling():
"""
Tests that:
- meta llama models support function calling
"""
from litellm.utils import return_raw_request
from litellm.types.utils import CallTypes
tools = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
},
},
"required": ["location"],
},
},
}
]
messages = [
{
"role": "user",
"content": "What's the weather like in Boston today in fahrenheit?",
}
]
request_args = {
"messages": messages,
"tools": tools,
"model": "bedrock/us.meta.llama4-scout-17b-instruct-v1:0",
}
response = return_raw_request(
endpoint=CallTypes.completion,
kwargs=request_args,
)
print(response)
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_bedrock_passthrough(sync_mode: bool):
@ -3272,9 +2895,7 @@ async def test_bedrock_converse__streaming_passthrough(monkeypatch):
@pytest.mark.asyncio
async def test_bedrock_streaming_passthrough_test2(monkeypatch):
import litellm
import time
import asyncio
from unittest.mock import MagicMock
from litellm.integrations.custom_logger import CustomLogger
class MockCustomLogger(CustomLogger):
@ -3324,9 +2945,7 @@ async def test_bedrock_streaming_passthrough_test2(monkeypatch):
@pytest.mark.asyncio
async def test_bedrock_streaming_passthrough_test1(monkeypatch):
import litellm
import time
import asyncio
from unittest.mock import MagicMock
from litellm.integrations.custom_logger import CustomLogger
class MockCustomLogger(CustomLogger):
@ -3839,7 +3458,6 @@ def test_bedrock_openai_error_handling():
from litellm import ModelResponse
from litellm.llms.bedrock.common_utils import BedrockError
from unittest.mock import Mock
import json
bedrock_llm = BedrockLLM()
@ -3886,7 +3504,7 @@ def test_bedrock_nova_grounding_web_search_options_non_streaming():
Related: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html
"""
from unittest.mock import patch, MagicMock
from unittest.mock import patch
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()

View file

@ -1,18 +1,16 @@
import json
import os
import sys
from datetime import datetime
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import Mock, patch
import pytest
import base64
import httpx
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, AsyncHTTPHandler
from litellm.llms.custom_httpx.http_handler import HTTPHandler
titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10}
@ -222,240 +220,6 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo():
os.environ["AWS_REGION_NAME"] = original_region_name
def test_e2e_bedrock_async_invoke_embedding_twelvelabs_marengo():
"""
Test async invoke embedding with TwelveLabs Marengo.
Validates that async invoke responses include job ID in hidden parameters.
"""
print("Testing async invoke embedding...")
original_region_name = os.environ.get("AWS_REGION_NAME")
os.environ["AWS_REGION_NAME"] = "us-east-1"
litellm._turn_on_debug()
# Mock the HTTP call to return async invoke response
with patch(
"litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_sync_call"
) as mock_call:
mock_call.return_value = {
"invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123"
}
response = litellm.embedding(
model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0",
input=["Hello world from LiteLLM async invoke!"],
aws_region_name="us-east-1",
inputType="text",
output_s3_uri="s3://test-bucket/async-invoke-output/",
)
# Validate response structure
assert isinstance(
response, litellm.EmbeddingResponse
), "Response should be EmbeddingResponse type"
assert hasattr(
response, "_hidden_params"
), "Response should have _hidden_params"
assert response._hidden_params is not None, "Hidden params should not be None"
# Validate hidden params contain invocation ARN
assert hasattr(
response._hidden_params, "_invocation_arn"
), "Hidden params should have _invocation_arn"
assert (
response._hidden_params._invocation_arn
== "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123"
), "Invocation ARN should be preserved"
# Validate embedding structure
assert len(response.data) == 1, "Should have one embedding"
assert (
response.data[0].object == "embedding"
), "Embedding object should be 'embedding'"
assert (
response.data[0].embedding == []
), "Embedding should be empty for async jobs"
print(
f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}"
)
# Restore original region name
if original_region_name:
os.environ["AWS_REGION_NAME"] = original_region_name
@pytest.mark.asyncio
async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo():
"""
Test async invoke embedding with async calls.
Validates that async invoke responses work with aembedding.
"""
print("Testing async invoke embedding with async calls...")
original_region_name = os.environ.get("AWS_REGION_NAME")
os.environ["AWS_REGION_NAME"] = "us-east-1"
litellm._turn_on_debug()
# Mock the async HTTP call to return async invoke response
with patch(
"litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_async_call"
) as mock_call:
mock_call.return_value = {
"invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456"
}
response = await litellm.aembedding(
model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0",
input=["Hello world from LiteLLM async invoke async!"],
aws_region_name="us-east-1",
inputType="text",
output_s3_uri="s3://test-bucket/async-invoke-output/",
)
# Validate response structure
assert isinstance(
response, litellm.EmbeddingResponse
), "Response should be EmbeddingResponse type"
assert hasattr(
response, "_hidden_params"
), "Response should have _hidden_params"
assert response._hidden_params is not None, "Hidden params should not be None"
# Validate hidden params contain invocation ARN
assert hasattr(
response._hidden_params, "_invocation_arn"
), "Hidden params should have _invocation_arn"
assert (
response._hidden_params._invocation_arn
== "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456"
), "Invocation ARN should be preserved"
print(
f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}"
)
# Restore original region name
if original_region_name:
os.environ["AWS_REGION_NAME"] = original_region_name
titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10}
def test_bedrock_embedding_uses_correct_region_when_specified():
"""
Test that when aws_region_name is explicitly passed, it's used correctly
even if AWS_REGION_NAME env var is set to a different region.
relevant issue: https://github.com/BerriAI/litellm/issues/16517
"""
# Save original env var
original_region_name = os.environ.get("AWS_REGION_NAME")
# Set env var to a different region (this should NOT be used)
os.environ["AWS_REGION_NAME"] = "ap-northeast-1"
try:
client = HTTPHandler()
with patch.object(client, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(titan_embedding_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
# Call with explicit region
response = litellm.embedding(
model="bedrock/amazon.titan-embed-image-v1",
input=["test input"],
client=client,
aws_region_name="us-east-1", # Explicitly set to us-east-1
)
# Verify the request was made to the correct region
assert mock_post.called, "HTTP post should have been called"
# Get the URL from the call
call_args = mock_post.call_args
url = call_args.kwargs.get("url", "")
# The URL should contain us-east-1, NOT ap-northeast-1
assert "us-east-1" in url, f"URL should contain us-east-1, but got: {url}"
assert (
"ap-northeast-1" not in url
), f"URL should NOT contain ap-northeast-1, but got: {url}"
print(f"✓ Test passed: URL contains correct region: {url}")
finally:
# Restore original env var
if original_region_name:
os.environ["AWS_REGION_NAME"] = original_region_name
else:
os.environ.pop("AWS_REGION_NAME", None)
def test_bedrock_embedding_region_bug_reproduction():
"""
Reproduces the bug where aws_region_name is ignored when passed explicitly.
relevant issue: https://github.com/BerriAI/litellm/issues/16517
"""
# Save original env var
original_region_name = os.environ.get("AWS_REGION_NAME")
# Set env var to ap-northeast-1 (this is what the bug report shows)
os.environ["AWS_REGION_NAME"] = "ap-northeast-1"
try:
client = HTTPHandler()
with patch.object(client, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(titan_embedding_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
# Call with explicit region (as in the bug report)
response = litellm.embedding(
model="bedrock/amazon.titan-embed-image-v1",
input=["test input"],
client=client,
aws_region_name="us-east-1", # Explicitly set to us-east-1
)
# Verify the request was made
assert mock_post.called, "HTTP post should have been called"
# Get the URL from the call
call_args = mock_post.call_args
url = call_args.kwargs.get("url", "")
print(f"Request URL: {url}")
print(f"Expected region in URL: us-east-1")
print(f"Environment AWS_REGION_NAME: {os.environ.get('AWS_REGION_NAME')}")
# This assertion will FAIL if the bug exists (it will use ap-northeast-1)
# This assertion will PASS if the bug is fixed (it will use us-east-1)
if "ap-northeast-1" in url:
print(
"❌ BUG REPRODUCED: Using wrong region from env var instead of explicit parameter"
)
assert (
False
), f"Bug reproduced: URL contains ap-northeast-1 instead of us-east-1. URL: {url}"
else:
print(
"✓ Bug NOT reproduced: Using correct region from explicit parameter"
)
assert (
"us-east-1" in url
), f"URL should contain us-east-1, but got: {url}"
finally:
# Restore original env var
if original_region_name:
os.environ["AWS_REGION_NAME"] = original_region_name
else:
os.environ.pop("AWS_REGION_NAME", None)

View file

@ -115,32 +115,6 @@ class TestBedrockGovCloudSupport:
)
assert base_model == "meta.llama3-8b-instruct-v1:0"
@patch("litellm.llms.bedrock.common_utils.init_bedrock_client")
def test_govcloud_client_initialization(self, mock_init_client):
"""Test that Bedrock client can be initialized with GovCloud regions"""
mock_client = Mock()
mock_init_client.return_value = mock_client
# Test that init_bedrock_client accepts GovCloud regions
from litellm.llms.bedrock.common_utils import init_bedrock_client
# This should not raise an error
client = init_bedrock_client(
region_name="us-gov-east-1",
aws_access_key_id=None,
aws_secret_access_key=None,
aws_region_name="us-gov-east-1",
aws_bedrock_runtime_endpoint=None,
aws_session_name=None,
aws_profile_name=None,
aws_role_name=None,
aws_web_identity_token=None,
extra_headers=None,
timeout=None,
)
assert mock_init_client.called
def test_govcloud_model_in_bedrock_models_list(self):
"""Test that GovCloud models are NOT included in bedrock_models list (they are pricing-only)"""
# Regional models including GovCloud should be excluded from bedrock_models list
@ -475,8 +449,6 @@ class TestBedrockGovCloudSupport:
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_govcloud_completion_with_cost_tracking(self, mock_post):
"""Test that completion requests with cost tracking use correct pricing for GovCloud models"""
from litellm import completion
from unittest.mock import Mock
import json
# Mock the HTTP client's post method to return responses

View file

@ -12,17 +12,16 @@ This test suite verifies:
"""
from base_llm_unit_tests import BaseLLMChatTest
import pytest
import sys
import os
import json
from typing import Optional
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import Mock, patch
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.llms.bedrock.common_utils import get_bedrock_chat_config
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.http_handler import HTTPHandler
class TestBedrockMoonshotInvoke(BaseLLMChatTest):
@ -121,165 +120,6 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest):
assert body["messages"][1]["role"] == "user"
assert response.choices[0].message.content is not None
def test_message_with_name(self):
"""Verify a user message carrying a ``name`` field is serialized into
the outgoing Bedrock invoke request without breaking the call."""
mock_post, response = self._invoke_with_mocked_post(
messages=[{"role": "user", "content": "Hello", "name": "test_name"}],
)
mock_post.assert_called_once()
body = json.loads(mock_post.call_args.kwargs["data"])
assert body["messages"][0]["role"] == "user"
assert body["messages"][0]["content"] == "Hello"
assert response is not None
def test_content_list_handling(self):
"""Verify the inherited content-list-handling test passes against a
mocked moonshot response (no network)."""
mock_post, response = self._invoke_with_mocked_post(
messages=[
{
"role": "user",
"content": [{"type": "text", "text": "Hello, how are you?"}],
}
],
)
mock_post.assert_called_once()
assert response.choices[0].message.content is not None
def test_pydantic_model_input(self):
"""Verify a completion call with a pydantic ``Message`` as input does
not raise and produces a parseable response."""
from litellm import Message
mock_post, response = self._invoke_with_mocked_post(
messages=[Message(content="Hello, how are you?", role="user")],
)
mock_post.assert_called_once()
assert response is not None
@pytest.mark.parametrize("response_format", [{"type": "text"}])
def test_response_format_type_text_with_tool_calls_no_tool_choice(
self, response_format
):
"""Verify response_format + tools + drop_params sends a valid request
and produces a response object."""
tools = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
},
},
"required": ["location"],
},
},
}
]
mock_post, response = self._invoke_with_mocked_post(
messages=[
{"role": "user", "content": "What's the weather like in Boston today?"}
],
extra_kwargs={
"response_format": response_format,
"tools": tools,
"drop_params": True,
},
)
mock_post.assert_called_once()
body = json.loads(mock_post.call_args.kwargs["data"])
assert "tools" in body
assert body["tools"][0]["function"]["name"] == "get_current_weather"
assert response is not None
def test_streaming(self):
"""Verify stream=True routes to the invoke-with-response-stream
endpoint with the messages body. Iteration of the stream itself is
not exercised here — moonshot streaming delegates to the OpenAI
parser and is covered by the OpenAI test suite.
Note: bedrock invoke streaming cannot be intercepted by patching
the caller-supplied client, because ``CustomStreamWrapper.fetch_sync_stream``
at streaming_handler.py invokes the stored ``make_call`` partial with
``client=litellm.module_level_client``, which overrides any client the
caller passed. Patch ``make_sync_call`` at its import site in
``base_invoke_transformation`` so we observe the exact kwargs the
partial was built with at stream-wrapper construction time.
"""
from litellm.utils import CustomStreamWrapper
captured: dict = {}
def fake_make_sync_call(**kwargs):
captured.update(kwargs)
# Return an empty iterator so the stream wrapper's iteration
# doesn't try to parse real bytes.
return iter([])
with patch(
"litellm.llms.bedrock.chat.invoke_transformations."
"base_invoke_transformation.make_sync_call",
new=fake_make_sync_call,
):
response = litellm.completion(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=[
{
"role": "user",
"content": [{"type": "text", "text": "Hello, how are you?"}],
}
],
stream=True,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-west-2",
)
assert isinstance(response, CustomStreamWrapper)
# Trigger fetch_sync_stream → make_call(...) → fake_make_sync_call.
try:
next(iter(response))
except StopIteration:
pass
assert captured, "make_sync_call was never invoked"
assert captured["api_base"].endswith("/invoke-with-response-stream")
body = json.loads(captured["data"])
# Bedrock invoke does not put stream=true in the body (the URL
# carries the streaming flag); verify the user message is present.
assert body["messages"][0]["role"] == "user"
async def test_completion_cost(self):
"""Verify LiteLLM computes a positive cost from a mocked Bedrock
Moonshot response, using the local model cost map."""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
mock_response = self._make_moonshot_response()
client = AsyncHTTPHandler()
with patch.object(client, "post", new=AsyncMock(return_value=mock_response)):
response = await litellm.acompletion(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=[{"role": "user", "content": "Hello, how are you?"}],
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-west-2",
client=client,
)
assert response._hidden_params["response_cost"] > 0
class TestBedrockMoonshotBasic:
"""Unit tests for Bedrock Moonshot configuration and transformations."""
@ -434,21 +274,6 @@ class TestBedrockMoonshotToolCalling:
assert len(transformed["tools"]) == 1
assert transformed["tools"][0]["function"]["name"] == "get_weather"
def test_tool_response_message_format(self):
"""Test that tool response messages are formatted correctly."""
# This tests the proper format for sending tool responses back
tool_response_message = {
"role": "tool",
"tool_call_id": "call_123",
"content": json.dumps({"temperature": 72, "condition": "sunny"}),
}
# Verify the message structure
assert tool_response_message["role"] == "tool"
assert "tool_call_id" in tool_response_message
assert "content" in tool_response_message
class TestBedrockMoonshotParameterValidation:
"""Tests for parameter validation and edge cases."""

View file

@ -1,18 +1,13 @@
import os
import sys
import traceback
from dotenv import load_dotenv
import litellm.types
import pytest
from litellm import AmazonInvokeConfig
import json
load_dotenv()
import io
import os
sys.path.insert(0, os.path.abspath("../.."))
from unittest.mock import AsyncMock, Mock, patch
# Initialize the transformer
@ -71,53 +66,6 @@ def test_transform_request_invalid_provider(bedrock_transformer):
assert "Unknown provider" in str(exc_info.value)
@patch("botocore.auth.SigV4Auth")
@patch("botocore.awsrequest.AWSRequest")
def test_sign_request_basic(mock_aws_request, mock_sigv4_auth, bedrock_transformer):
"""Test basic request signing without extra headers"""
# Mock credentials
mock_credentials = Mock()
bedrock_transformer.get_credentials = Mock(return_value=mock_credentials)
# Setup mock SigV4Auth instance
mock_auth_instance = Mock()
mock_sigv4_auth.return_value = mock_auth_instance
# Setup mock AWSRequest instance
mock_request = Mock()
mock_request.headers = {
"Authorization": "AWS4-HMAC-SHA256 Credential=...",
"X-Amz-Date": "20240101T000000Z",
"Content-Type": "application/json",
}
mock_aws_request.return_value = mock_request
# Test parameters
headers = {}
optional_params = {"aws_region_name": "us-east-1"}
request_data = {"prompt": "Hello"}
api_base = "https://bedrock-runtime.us-east-1.amazonaws.com"
# Call the method
result, _ = bedrock_transformer.sign_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
)
# Verify the results
mock_sigv4_auth.assert_called_once_with(mock_credentials, "bedrock", "us-east-1")
mock_aws_request.assert_called_once_with(
method="POST",
url=api_base,
data='{"prompt": "Hello"}',
headers={"Content-Type": "application/json"},
)
mock_auth_instance.add_auth.assert_called_once_with(mock_request)
assert result == mock_request.headers
def test_transform_request_cohere_command(bedrock_transformer):
"""Test request transformation for Cohere Command model"""
messages = [{"role": "user", "content": "Hello"}]