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..ed2ec3d86f8 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -496,6 +496,10 @@ 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="ExternalId sent on sts:AssumeRole, for target roles whose trust policy requires one", + ) 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/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index 82363302d2e..8cf881fa037 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -70,6 +70,40 @@ def test_initialize_bedrock_forwards_chunk_budget_chars(): assert initialized[-1].chunk_budget_chars == 60_000 +def test_initialize_bedrock_forwards_aws_external_id(): + """Regression: `aws_external_id` set in config.yaml must reach the guardrail. + + Cross-account roles whose trust policy requires an ExternalId could not be assumed by the + bedrock guardrail, because the config field was dropped before the sts:AssumeRole call. + """ + import litellm + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + test_guardrail = { + "guardrail_name": "test_bedrock_external_id", + "litellm_params": { + "guardrail": SupportedGuardrailIntegrations.BEDROCK.value, + "mode": "pre_call", + "guardrailIdentifier": "test-guardrail", + "guardrailVersion": "DRAFT", + "aws_region_name": "us-east-1", + "aws_role_name": "arn:aws:iam::999999999999:role/litellm-guardrail", + "aws_external_id": "external-id-from-config", + }, + } + + guardrail_handler = InMemoryGuardrailHandler() + guardrail_handler.initialize_guardrail(guardrail=test_guardrail) + + initialized = [ + callback + for callback in litellm.callbacks + if isinstance(callback, BedrockGuardrail) and callback.guardrail_name == "test_bedrock_external_id" + ] + assert initialized, "bedrock guardrail was not registered as a callback" + assert initialized[-1].optional_params["aws_external_id"] == "external-id-from-config" + + def test_initialize_guardrail_preserves_guardrail_info(): """ Regression (LIT-2529): initialize_guardrail must carry guardrail_info into the