diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 0ddf8896fd..d9b7eb6410 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 9e44a3eb41..a3f17a1424 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -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}")