test: add unit testing

This commit is contained in:
Krrish Dholakia
2025-09-17 18:00:59 -07:00
parent 1598d3e955
commit fc18f4decf
2 changed files with 298 additions and 57 deletions
+108 -57
View File
@@ -189,23 +189,32 @@ class BaseAWSLLM:
# Check if we're in IRSA and trying to assume the same role we already have
current_role_arn = os.getenv("AWS_ROLE_ARN")
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
# In IRSA environments, we should skip role assumption if we're already running as the target role
# This is true when:
# 1. We have AWS_ROLE_ARN set (current role)
# 2. We have AWS_WEB_IDENTITY_TOKEN_FILE set (IRSA environment)
# 3. The current role matches the requested role
if (current_role_arn and web_identity_token_file and
current_role_arn == aws_role_name):
verbose_logger.debug("Using IRSA same-role optimization: calling _auth_with_env_vars")
if (
current_role_arn
and web_identity_token_file
and current_role_arn == aws_role_name
):
verbose_logger.debug(
"Using IRSA same-role optimization: calling _auth_with_env_vars"
)
# We're already running as this role via IRSA, no need to assume it again
# Use the default boto3 credentials (which will use the IRSA credentials)
credentials, _cache_ttl = self._auth_with_env_vars()
else:
verbose_logger.debug("Using role assumption: calling _auth_with_aws_role")
verbose_logger.debug(
"Using role assumption: calling _auth_with_aws_role"
)
# If aws_session_name is not provided, generate a default one
if aws_session_name is None:
aws_session_name = f"litellm-session-{int(datetime.now().timestamp())}"
aws_session_name = (
f"litellm-session-{int(datetime.now().timestamp())}"
)
credentials, _cache_ttl = self._auth_with_aws_role(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
@@ -479,55 +488,67 @@ class BaseAWSLLM:
iam_creds = session.get_credentials()
return iam_creds, self._get_default_ttl_for_boto3_credentials()
def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str,
aws_session_name: str, region: str, web_identity_token_file: str,
aws_external_id: Optional[str] = None) -> dict:
def _handle_irsa_cross_account(
self,
irsa_role_arn: str,
aws_role_name: str,
aws_session_name: str,
region: str,
web_identity_token_file: str,
aws_external_id: Optional[str] = None,
) -> dict:
"""Handle cross-account role assumption for IRSA."""
import boto3
verbose_logger.debug("Cross-account role assumption detected")
# Read the web identity token
with open(web_identity_token_file, 'r') as f:
with open(web_identity_token_file, "r") as f:
web_identity_token = f.read().strip()
# Create an STS client without credentials
with tracer.trace("boto3.client(sts) for manual IRSA"):
sts_client = boto3.client('sts', region_name=region)
sts_client = boto3.client("sts", region_name=region)
# Manually assume the IRSA role with the session name
verbose_logger.debug(f"Manually assuming IRSA role {irsa_role_arn} with session {aws_session_name}")
verbose_logger.debug(
f"Manually assuming IRSA role {irsa_role_arn} with session {aws_session_name}"
)
irsa_response = sts_client.assume_role_with_web_identity(
RoleArn=irsa_role_arn,
RoleSessionName=aws_session_name,
WebIdentityToken=web_identity_token
WebIdentityToken=web_identity_token,
)
# Extract the credentials from the IRSA assumption
irsa_creds = irsa_response["Credentials"]
# Create a new STS client with the IRSA credentials
with tracer.trace("boto3.client(sts) with manual IRSA credentials"):
sts_client_with_creds = boto3.client(
'sts',
"sts",
region_name=region,
aws_access_key_id=irsa_creds["AccessKeyId"],
aws_secret_access_key=irsa_creds["SecretAccessKey"],
aws_session_token=irsa_creds["SessionToken"]
aws_session_token=irsa_creds["SessionToken"],
)
# Get current caller identity for debugging
try:
caller_identity = sts_client_with_creds.get_caller_identity()
verbose_logger.debug(f"Current identity after manual IRSA assumption: {caller_identity.get('Arn', 'unknown')}")
verbose_logger.debug(
f"Current identity after manual IRSA assumption: {caller_identity.get('Arn', 'unknown')}"
)
except Exception as e:
verbose_logger.debug(f"Failed to get caller identity: {e}")
# Now assume the target role
verbose_logger.debug(f"Attempting to assume target role: {aws_role_name} with session: {aws_session_name}")
verbose_logger.debug(
f"Attempting to assume target role: {aws_role_name} with session: {aws_session_name}"
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
@@ -536,27 +557,36 @@ class BaseAWSLLM:
return sts_client_with_creds.assume_role(**assume_role_params)
def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str,
aws_external_id: Optional[str] = None) -> dict:
def _handle_irsa_same_account(
self,
aws_role_name: str,
aws_session_name: str,
region: str,
aws_external_id: Optional[str] = None,
) -> dict:
"""Handle same-account role assumption for IRSA."""
import boto3
verbose_logger.debug("Same account role assumption, using automatic IRSA")
with tracer.trace("boto3.client(sts) with automatic IRSA"):
sts_client = boto3.client("sts", region_name=region)
# Get current caller identity for debugging
try:
caller_identity = sts_client.get_caller_identity()
verbose_logger.debug(f"Current IRSA identity: {caller_identity.get('Arn', 'unknown')}")
verbose_logger.debug(
f"Current IRSA identity: {caller_identity.get('Arn', 'unknown')}"
)
except Exception as e:
verbose_logger.debug(f"Failed to get caller identity: {e}")
# Assume the role
verbose_logger.debug(f"Attempting to assume role: {aws_role_name} with session: {aws_session_name}")
verbose_logger.debug(
f"Attempting to assume role: {aws_role_name} with session: {aws_session_name}"
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
@@ -565,20 +595,24 @@ class BaseAWSLLM:
return sts_client.assume_role(**assume_role_params)
def _extract_credentials_and_ttl(self, sts_response: dict) -> Tuple[Credentials, Optional[int]]:
def _extract_credentials_and_ttl(
self, sts_response: dict
) -> Tuple[Credentials, Optional[int]]:
"""Extract credentials and TTL from STS response."""
from botocore.credentials import Credentials
sts_credentials = sts_response["Credentials"]
credentials = Credentials(
access_key=sts_credentials["AccessKeyId"],
secret_key=sts_credentials["SecretAccessKey"],
token=sts_credentials["SessionToken"],
)
expiration_time = sts_credentials["Expiration"]
ttl = int((expiration_time - datetime.now(expiration_time.tzinfo)).total_seconds())
ttl = int(
(expiration_time - datetime.now(expiration_time.tzinfo)).total_seconds()
)
return credentials, ttl
@tracer.wrap()
@@ -600,34 +634,51 @@ class BaseAWSLLM:
# Check if we're in an EKS/IRSA environment
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
# If we have IRSA environment variables and no explicit credentials,
# we need to use the web identity token flow
if (web_identity_token_file and irsa_role_arn and
aws_access_key_id is None and aws_secret_access_key is None):
if (
web_identity_token_file
and irsa_role_arn
and aws_access_key_id is None
and aws_secret_access_key is None
):
# For cross-account role assumption with specific session names,
# we need to manually assume the IRSA role first with the correct session name
verbose_logger.debug(f"IRSA detected: using web identity token from {web_identity_token_file}")
verbose_logger.debug(
f"IRSA detected: using web identity token from {web_identity_token_file}"
)
try:
# Get region from environment
region = os.getenv("AWS_REGION") or os.getenv("AWS_DEFAULT_REGION") or "us-east-1"
region = (
os.getenv("AWS_REGION")
or os.getenv("AWS_DEFAULT_REGION")
or "us-east-1"
)
# Check if we need to do cross-account role assumption
if aws_role_name != irsa_role_arn:
sts_response = self._handle_irsa_cross_account(
irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file, aws_external_id
irsa_role_arn,
aws_role_name,
aws_session_name,
region,
web_identity_token_file,
aws_external_id,
)
else:
sts_response = self._handle_irsa_same_account(
aws_role_name, aws_session_name, region, aws_external_id
)
return self._extract_credentials_and_ttl(sts_response)
except Exception as e:
verbose_logger.debug(f"Failed to assume role via IRSA: {e}")
if "AccessDenied" in str(e) and "is not authorized to perform: sts:AssumeRole" in str(e):
if "AccessDenied" in str(
e
) and "is not authorized to perform: sts:AssumeRole" in str(e):
# Provide a more helpful error message for trust policy issues
verbose_logger.error(
f"Access denied when trying to assume role {aws_role_name}. "
@@ -636,7 +687,7 @@ class BaseAWSLLM:
)
# Re-raise the exception instead of falling through
raise
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
# This allows the web identity token to work automatically
if aws_access_key_id is None and aws_secret_access_key is None:
@@ -653,7 +704,7 @@ class BaseAWSLLM:
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
@@ -782,14 +833,14 @@ class BaseAWSLLM:
)
# Determine proxy_endpoint_url
if env_aws_bedrock_runtime_endpoint and isinstance(
env_aws_bedrock_runtime_endpoint, str
):
proxy_endpoint_url = env_aws_bedrock_runtime_endpoint
elif aws_bedrock_runtime_endpoint is not None and isinstance(
if aws_bedrock_runtime_endpoint is not None and isinstance(
aws_bedrock_runtime_endpoint, str
):
proxy_endpoint_url = aws_bedrock_runtime_endpoint
elif env_aws_bedrock_runtime_endpoint and isinstance(
env_aws_bedrock_runtime_endpoint, str
):
proxy_endpoint_url = env_aws_bedrock_runtime_endpoint
else:
proxy_endpoint_url = endpoint_url
@@ -859,3 +859,193 @@ async def test__redact_pii_matches_comprehensive_coverage():
)
print("Comprehensive coverage redaction test passed")
@pytest.mark.asyncio
async def test_bedrock_guardrail_respects_custom_runtime_endpoint(monkeypatch):
"""Test that BedrockGuardrail respects aws_bedrock_runtime_endpoint when set"""
# Clear any existing environment variable to ensure clean test
monkeypatch.delenv("AWS_BEDROCK_RUNTIME_ENDPOINT", raising=False)
# Create guardrail with custom runtime endpoint
custom_endpoint = "https://custom-bedrock.example.com"
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
aws_bedrock_runtime_endpoint=custom_endpoint,
)
# Mock credentials
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
# Test data
data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
optional_params = {}
aws_region_name = "us-east-1"
# Mock the _load_credentials method to avoid actual AWS credential loading
with patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, aws_region_name)
):
# Call _prepare_request which internally calls get_runtime_endpoint
prepped_request = guardrail._prepare_request(
credentials=mock_credentials,
data=data,
optional_params=optional_params,
aws_region_name=aws_region_name,
)
# Verify that the custom endpoint is used in the URL
expected_url = f"{custom_endpoint}/guardrail/{guardrail.guardrailIdentifier}/version/{guardrail.guardrailVersion}/apply"
assert (
prepped_request.url == expected_url
), f"Expected URL to contain custom endpoint. Got: {prepped_request.url}"
print(f"Custom runtime endpoint test passed. URL: {prepped_request.url}")
@pytest.mark.asyncio
async def test_bedrock_guardrail_respects_env_runtime_endpoint(monkeypatch):
"""Test that BedrockGuardrail respects AWS_BEDROCK_RUNTIME_ENDPOINT environment variable"""
custom_endpoint = "https://env-bedrock.example.com"
# Set the environment variable
monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", custom_endpoint)
# Create guardrail without explicit aws_bedrock_runtime_endpoint
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT"
)
# Mock credentials
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
# Test data
data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
optional_params = {}
aws_region_name = "us-east-1"
# Mock the _load_credentials method
with patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, aws_region_name)
):
# Call _prepare_request which internally calls get_runtime_endpoint
prepped_request = guardrail._prepare_request(
credentials=mock_credentials,
data=data,
optional_params=optional_params,
aws_region_name=aws_region_name,
)
# Verify that the custom endpoint from environment is used in the URL
expected_url = f"{custom_endpoint}/guardrail/{guardrail.guardrailIdentifier}/version/{guardrail.guardrailVersion}/apply"
assert (
prepped_request.url == expected_url
), f"Expected URL to contain env endpoint. Got: {prepped_request.url}"
print(f"Environment runtime endpoint test passed. URL: {prepped_request.url}")
@pytest.mark.asyncio
async def test_bedrock_guardrail_uses_default_endpoint_when_no_custom_set(monkeypatch):
"""Test that BedrockGuardrail uses default endpoint when no custom endpoint is set"""
# Ensure no environment variable is set
monkeypatch.delenv("AWS_BEDROCK_RUNTIME_ENDPOINT", raising=False)
# Create guardrail without any custom endpoint
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT"
)
# Mock credentials
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
# Test data
data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
optional_params = {}
aws_region_name = "us-west-2"
# Mock the _load_credentials method
with patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, aws_region_name)
):
# Call _prepare_request which internally calls get_runtime_endpoint
prepped_request = guardrail._prepare_request(
credentials=mock_credentials,
data=data,
optional_params=optional_params,
aws_region_name=aws_region_name,
)
# Verify that the default endpoint is used
expected_url = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com/guardrail/{guardrail.guardrailIdentifier}/version/{guardrail.guardrailVersion}/apply"
assert (
prepped_request.url == expected_url
), f"Expected default URL. Got: {prepped_request.url}"
print(f"Default endpoint test passed. URL: {prepped_request.url}")
@pytest.mark.asyncio
async def test_bedrock_guardrail_parameter_takes_precedence_over_env(monkeypatch):
"""Test that aws_bedrock_runtime_endpoint parameter takes precedence over environment variable
This test verifies the corrected behavior where the parameter should take precedence
over the environment variable, consistent with the endpoint_url logic.
"""
param_endpoint = "https://param-bedrock.example.com"
env_endpoint = "https://env-bedrock.example.com"
# Set environment variable
monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", env_endpoint)
# Create guardrail with explicit aws_bedrock_runtime_endpoint
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
aws_bedrock_runtime_endpoint=param_endpoint,
)
# Mock credentials
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
# Test data
data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
optional_params = {}
aws_region_name = "us-east-1"
# Mock the _load_credentials method
with patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, aws_region_name)
):
# Call _prepare_request which internally calls get_runtime_endpoint
prepped_request = guardrail._prepare_request(
credentials=mock_credentials,
data=data,
optional_params=optional_params,
aws_region_name=aws_region_name,
)
# Verify that the parameter takes precedence over environment variable
expected_url = f"{param_endpoint}/guardrail/{guardrail.guardrailIdentifier}/version/{guardrail.guardrailVersion}/apply"
assert (
prepped_request.url == expected_url
), f"Expected parameter endpoint to take precedence. Got: {prepped_request.url}"
print(f"Parameter precedence test passed. URL: {prepped_request.url}")