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/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 0d23e19f88d..b377b272e5b 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -34,6 +34,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail): aws_role_name=litellm_params.aws_role_name, aws_web_identity_token=litellm_params.aws_web_identity_token, aws_sts_endpoint=litellm_params.aws_sts_endpoint, + aws_external_id=litellm_params.aws_external_id, aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint, experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only, only_scan_new_messages=litellm_params.only_scan_new_messages or False, diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index c7cdfaad780..5e21dd2f60c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -496,6 +496,9 @@ class BedrockGuardrailConfigModel(BaseModel): aws_role_name: str | None = Field(default=None, description="AWS role name for assuming roles") aws_web_identity_token: str | None = Field(default=None, description="Web identity token for AWS role assumption") aws_sts_endpoint: str | None = Field(default=None, description="AWS STS endpoint URL") + aws_external_id: str | None = Field( + default=None, description="External ID required by the target role's trust policy on sts:AssumeRole" + ) aws_bedrock_runtime_endpoint: str | None = Field(default=None, description="AWS Bedrock runtime endpoint URL") checks: BedrockChecksConfigModel | None = Field( default=None, 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..36b356e34d0 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,74 @@ 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_assumes_role_with_external_id(): + """A trust policy requiring sts:ExternalId must be satisfied by the guardrail's aws_external_id.""" + import datetime + + import boto3 + from botocore.exceptions import ClientError + + class FakeSTSClient: + """STS that mirrors a cross-account role whose trust policy requires an ExternalId.""" + + def get_caller_identity(self): + return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"} + + def assume_role(self, **params): + if params.get("ExternalId") != "external-id-123": + raise ClientError( + {"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}}, + "AssumeRole", + ) + return { + "Credentials": { + "AccessKeyId": "ASIAASSUMEDROLEKEY", + "SecretAccessKey": "assumed-secret", + "SessionToken": "assumed-session-token", + "Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30), + } + } + + guardrail = BedrockGuardrail( + guardrail_name="bedrock-external-id", + event_hook=GuardrailEventHooks.pre_call, + guardrailIdentifier="gr-1", + guardrailVersion="DRAFT", + aws_region_name="us-east-1", + aws_access_key_id="AKIAPODCALLERKEY", + aws_secret_access_key="pod-caller-secret", + aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role", + aws_session_name="litellm-session", + aws_external_id="external-id-123", + ) + + with patch.object(boto3, "client", return_value=FakeSTSClient()): + credentials, aws_region_name = guardrail._load_credentials() + + assert credentials.access_key == "ASIAASSUMEDROLEKEY" + assert credentials.token == "assumed-session-token" + assert aws_region_name == "us-east-1" + + +def test_initialize_bedrock_forwards_aws_external_id(): + """aws_external_id configured on the guardrail must survive LitellmParams and the initializer.""" + from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="bedrock", + mode="pre_call", + guardrailIdentifier="gr-1", + guardrailVersion="DRAFT", + aws_region_name="us-east-1", + aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role", + aws_external_id="external-id-123", + ) + + guardrail = initialize_bedrock(litellm_params, {"guardrail_name": "bedrock-external-id"}) + try: + assert guardrail.optional_params["aws_external_id"] == "external-id-123" + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, guardrail)