mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #38376 from BerriAI/devin_ai_bedrock_guardrail_external_id
fix(guardrails): forward aws_external_id when the bedrock guardrail assumes a role
This commit is contained in:
commit
55c1537497
4 changed files with 77 additions and 0 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue