mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
22f18179f4
commit
00eb3dbbaf
8 changed files with 6 additions and 1351 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue