diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index c70a2ee8a74..dd76a27c80f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -686,6 +686,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): aws_profile_name: Final = self.optional_params.get("aws_profile_name", None) aws_web_identity_token: Final = self.optional_params.get("aws_web_identity_token", None) aws_sts_endpoint: Final = self.optional_params.get("aws_sts_endpoint", None) + aws_external_id: Final = self.optional_params.get("aws_external_id", None) ### SET REGION NAME ### aws_region_name = self.get_aws_region_name_for_non_llm_api_calls( @@ -702,6 +703,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): aws_role_name=aws_role_name, aws_web_identity_token=aws_web_identity_token, aws_sts_endpoint=aws_sts_endpoint, + aws_external_id=aws_external_id, ) return credentials, aws_region_name 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 dd339d4e51f..8b5e9b79f1b 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 @@ -5274,3 +5274,55 @@ async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monke assert logged["guardrail_cost"] == pytest.approx(0.0003) assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1} assert "error" in logged["guardrail_response"] + + +def test_load_credentials_forwards_aws_external_id_to_assume_role(monkeypatch): + """Regression: cross-account guardrail role assumption must send sts:ExternalId.""" + from datetime import datetime, timedelta, timezone + + import boto3 + + captured_params: list[dict[str, str]] = [] + + class FakeSTSClient: + def assume_role(self, **kwargs: str) -> dict: + captured_params.append(kwargs) + return { + "Credentials": { + "AccessKeyId": "ASIATESTKEY", + "SecretAccessKey": "test-secret", + "SessionToken": "test-token", + "Expiration": datetime.now(timezone.utc) + timedelta(hours=1), + } + } + + real_boto3_client = boto3.client + + def fake_boto3_client(service_name: str, *args, **kwargs): + if service_name == "sts": + return FakeSTSClient() + return real_boto3_client(service_name, *args, **kwargs) + + monkeypatch.setattr(boto3, "client", fake_boto3_client) + + guardrail = BedrockGuardrail( + guardrail_name="bedrock-cross-account-guard", + guardrailIdentifier="amgllac6xf3r", + guardrailVersion="1", + aws_region_name="us-east-1", + aws_access_key_id="AKIATESTKEY", + aws_secret_access_key="test-secret", + aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail", + aws_session_name="litellm-guardrail-session", + aws_external_id="external-id-for-guardrail-hook", + ) + + guardrail._load_credentials() + + assert captured_params == [ + { + "RoleArn": "arn:aws:iam::999999999999:role/litellm-guardrail", + "RoleSessionName": "litellm-guardrail-session", + "ExternalId": "external-id-for-guardrail-hook", + } + ]