diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index ce196757f94..0ddf8896fdd 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -66,6 +66,7 @@ class BaseAWSLLM: "aws_web_identity_token", "aws_sts_endpoint", "aws_bedrock_runtime_endpoint", + "aws_external_id", ] def get_cache_key(self, credential_args: Dict[str, Optional[str]]) -> str: @@ -88,6 +89,7 @@ class BaseAWSLLM: aws_role_name: Optional[str] = None, aws_web_identity_token: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, + aws_external_id: Optional[str] = None, ): """ Return a boto3.Credentials object @@ -103,6 +105,7 @@ class BaseAWSLLM: aws_role_name, aws_web_identity_token, aws_sts_endpoint, + aws_external_id, ] # Iterate over parameters and update if needed @@ -127,6 +130,7 @@ class BaseAWSLLM: aws_role_name, aws_web_identity_token, aws_sts_endpoint, + aws_external_id, ) = params_to_check verbose_logger.debug( @@ -139,7 +143,8 @@ class BaseAWSLLM: "aws_profile_name=%s\n" "aws_role_name=%s\n" "aws_web_identity_token=%s\n" - "aws_sts_endpoint=%s", + "aws_sts_endpoint=%s\n" + "aws_external_id=%s", aws_access_key_id, aws_secret_access_key, aws_session_token, @@ -149,6 +154,7 @@ class BaseAWSLLM: aws_role_name, aws_web_identity_token, aws_sts_endpoint, + aws_external_id, ) # create cache key for non-expiring auth flows @@ -177,6 +183,7 @@ class BaseAWSLLM: aws_session_name=aws_session_name, aws_region_name=aws_region_name, aws_sts_endpoint=aws_sts_endpoint, + aws_external_id=aws_external_id, ) elif aws_role_name is not None: # Check if we're in IRSA and trying to assume the same role we already have @@ -205,6 +212,7 @@ class BaseAWSLLM: aws_session_token=aws_session_token, aws_role_name=aws_role_name, aws_session_name=aws_session_name, + aws_external_id=aws_external_id, ) elif aws_profile_name is not None: ### CHECK SESSION ### @@ -406,6 +414,7 @@ class BaseAWSLLM: aws_session_name: str, aws_region_name: Optional[str], aws_sts_endpoint: Optional[str], + aws_external_id: Optional[str] = None, ) -> Tuple[Credentials, Optional[int]]: """ Authenticate with AWS Web Identity Token @@ -438,13 +447,19 @@ class BaseAWSLLM: # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html # https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html - sts_response = sts_client.assume_role_with_web_identity( - RoleArn=aws_role_name, - RoleSessionName=aws_session_name, - WebIdentityToken=oidc_token, - DurationSeconds=3600, - Policy='{"Version":"2012-10-17","Statement":[{"Sid":"BedrockLiteLLM","Effect":"Allow","Action":["bedrock:InvokeModel","bedrock:InvokeModelWithResponseStream"],"Resource":"*","Condition":{"Bool":{"aws:SecureTransport":"true"},"StringLike":{"aws:UserAgent":"litellm/*"}}}]}', - ) + assume_role_params = { + "RoleArn": aws_role_name, + "RoleSessionName": aws_session_name, + "WebIdentityToken": oidc_token, + "DurationSeconds": 3600, + "Policy": '{"Version":"2012-10-17","Statement":[{"Sid":"BedrockLiteLLM","Effect":"Allow","Action":["bedrock:InvokeModel","bedrock:InvokeModelWithResponseStream"],"Resource":"*","Condition":{"Bool":{"aws:SecureTransport":"true"},"StringLike":{"aws:UserAgent":"litellm/*"}}}]}', + } + + # Add ExternalId parameter if provided + if aws_external_id is not None: + assume_role_params["ExternalId"] = aws_external_id + + sts_response = sts_client.assume_role_with_web_identity(**assume_role_params) iam_creds_dict = { "aws_access_key_id": sts_response["Credentials"]["AccessKeyId"], @@ -464,8 +479,9 @@ class BaseAWSLLM: iam_creds = session.get_credentials() return iam_creds, self._get_default_ttl_for_boto3_credentials() - def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str, - aws_session_name: str, region: str, web_identity_token_file: str) -> dict: + def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str, + aws_session_name: str, region: str, web_identity_token_file: str, + aws_external_id: Optional[str] = None) -> dict: """Handle cross-account role assumption for IRSA.""" import boto3 @@ -509,11 +525,19 @@ class BaseAWSLLM: # Now assume the target role verbose_logger.debug(f"Attempting to assume target role: {aws_role_name} with session: {aws_session_name}") - return sts_client_with_creds.assume_role( - RoleArn=aws_role_name, RoleSessionName=aws_session_name - ) + assume_role_params = { + "RoleArn": aws_role_name, + "RoleSessionName": aws_session_name + } - def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str) -> dict: + # Add ExternalId parameter if provided + if aws_external_id is not None: + assume_role_params["ExternalId"] = aws_external_id + + return sts_client_with_creds.assume_role(**assume_role_params) + + def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str, + aws_external_id: Optional[str] = None) -> dict: """Handle same-account role assumption for IRSA.""" import boto3 @@ -530,9 +554,16 @@ class BaseAWSLLM: # Assume the role verbose_logger.debug(f"Attempting to assume role: {aws_role_name} with session: {aws_session_name}") - return sts_client.assume_role( - RoleArn=aws_role_name, RoleSessionName=aws_session_name - ) + assume_role_params = { + "RoleArn": aws_role_name, + "RoleSessionName": aws_session_name + } + + # Add ExternalId parameter if provided + if aws_external_id is not None: + assume_role_params["ExternalId"] = aws_external_id + + return sts_client.assume_role(**assume_role_params) def _extract_credentials_and_ttl(self, sts_response: dict) -> Tuple[Credentials, Optional[int]]: """Extract credentials and TTL from STS response.""" @@ -558,6 +589,7 @@ class BaseAWSLLM: aws_session_token: Optional[str], aws_role_name: str, aws_session_name: str, + aws_external_id: Optional[str] = None, ) -> Tuple[Credentials, Optional[int]]: """ Authenticate with AWS Role @@ -584,11 +616,11 @@ class BaseAWSLLM: # Check if we need to do cross-account role assumption if aws_role_name != irsa_role_arn: sts_response = self._handle_irsa_cross_account( - irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file + irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file, aws_external_id ) else: sts_response = self._handle_irsa_same_account( - aws_role_name, aws_session_name, region + aws_role_name, aws_session_name, region, aws_external_id ) return self._extract_credentials_and_ttl(sts_response) @@ -619,9 +651,16 @@ class BaseAWSLLM: aws_session_token=aws_session_token, ) - sts_response = sts_client.assume_role( - RoleArn=aws_role_name, RoleSessionName=aws_session_name - ) + assume_role_params = { + "RoleArn": aws_role_name, + "RoleSessionName": aws_session_name + } + + # Add ExternalId parameter if provided + if aws_external_id is not None: + assume_role_params["ExternalId"] = aws_external_id + + sts_response = sts_client.assume_role(**assume_role_params) # Extract the credentials from the response and convert to Session Credentials sts_credentials = sts_response["Credentials"] @@ -800,6 +839,7 @@ class BaseAWSLLM: aws_bedrock_runtime_endpoint = optional_params.pop( "aws_bedrock_runtime_endpoint", None ) # https://bedrock-runtime.{region_name}.amazonaws.com + aws_external_id = optional_params.pop("aws_external_id", None) credentials: Credentials = self.get_credentials( aws_access_key_id=aws_access_key_id, @@ -811,6 +851,7 @@ class 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 Boto3CredentialsInfo( @@ -915,6 +956,7 @@ class BaseAWSLLM: aws_profile_name = optional_params.get("aws_profile_name", None) aws_web_identity_token = optional_params.get("aws_web_identity_token", None) aws_sts_endpoint = optional_params.get("aws_sts_endpoint", None) + aws_external_id = optional_params.get("aws_external_id", None) aws_region_name = self._get_aws_region_name( optional_params=optional_params, model=model ) @@ -929,6 +971,7 @@ class 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, ) sigv4 = SigV4Auth(credentials, service_name, aws_region_name) diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 15a5002f0e4..54c603e5960 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -307,6 +307,7 @@ class BedrockConverseLLM(BaseAWSLLM): ) # https://bedrock-runtime.{region_name}.amazonaws.com aws_web_identity_token = optional_params.pop("aws_web_identity_token", None) aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None) + aws_external_id = optional_params.pop("aws_external_id", None) optional_params.pop("aws_region_name", None) litellm_params[ @@ -323,6 +324,7 @@ class BedrockConverseLLM(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, ) ### SET RUNTIME ENDPOINT ###