mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(guardrails): forward aws_external_id when the bedrock guardrail assumes a role
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
c11c654b8e
commit
f3c1e2e2a7
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_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_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_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 ###
|
### SET REGION NAME ###
|
||||||
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
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_role_name=aws_role_name,
|
||||||
aws_web_identity_token=aws_web_identity_token,
|
aws_web_identity_token=aws_web_identity_token,
|
||||||
aws_sts_endpoint=aws_sts_endpoint,
|
aws_sts_endpoint=aws_sts_endpoint,
|
||||||
|
aws_external_id=aws_external_id,
|
||||||
)
|
)
|
||||||
return credentials, aws_region_name
|
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_role_name=litellm_params.aws_role_name,
|
||||||
aws_web_identity_token=litellm_params.aws_web_identity_token,
|
aws_web_identity_token=litellm_params.aws_web_identity_token,
|
||||||
aws_sts_endpoint=litellm_params.aws_sts_endpoint,
|
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,
|
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
|
||||||
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
|
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,
|
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_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_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_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")
|
aws_bedrock_runtime_endpoint: str | None = Field(default=None, description="AWS Bedrock runtime endpoint URL")
|
||||||
checks: BedrockChecksConfigModel | None = Field(
|
checks: BedrockChecksConfigModel | None = Field(
|
||||||
default=None,
|
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_cost"] == pytest.approx(0.0003)
|
||||||
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
|
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
|
||||||
assert "error" in logged["guardrail_response"]
|
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