Merge pull request #38727 from BerriAI/litellm_aws_external_id_embed_sagemaker

fix(aws): forward aws_external_id in Bedrock embeddings and SageMaker credential loading
This commit is contained in:
Mateo Wang 2026-08-29 03:30:32 -07:00 committed by GitHub
commit 39e4f1ae13
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 149 additions and 0 deletions

View file

@ -58,6 +58,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -84,6 +85,7 @@ class BedrockEmbedding(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

View file

@ -34,6 +34,7 @@ class SagemakerChatHandler(BaseAWSLLM):
optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -60,6 +61,7 @@ class SagemakerChatHandler(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

View file

@ -57,6 +57,7 @@ class SagemakerLLM(BaseAWSLLM):
optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -83,6 +84,7 @@ class SagemakerLLM(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

View file

@ -985,3 +985,51 @@ def test_bedrock_cohere_embedding_types_wrapped_as_list(
assert "embedding_types" in request_body
assert request_body["embedding_types"] == expected_embedding_types
assert isinstance(request_body["embedding_types"], list)
def test_load_credentials_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied by the deployment's aws_external_id."""
import datetime
import boto3
from botocore.exceptions import ClientError
from litellm.llms.bedrock.embed.embedding import BedrockEmbedding
monkeypatch.delenv("AWS_EXTERNAL_ID", raising=False)
class FakeSTSClient:
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-embed":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIAEMBEDROLEKEY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
optional_params = {
"aws_access_key_id": "AKIAEMBEDCALLERKEY",
"aws_secret_access_key": "pod-caller-secret",
"aws_region_name": "us-east-1",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-embed-role",
"aws_session_name": "litellm-embed-session",
"aws_external_id": "external-id-embed",
}
with patch.object(boto3, "client", return_value=FakeSTSClient()):
credentials, aws_region_name = BedrockEmbedding()._load_credentials(optional_params)
assert credentials.access_key == "ASIAEMBEDROLEKEY"
assert credentials.token == "assumed-session-token"
assert aws_region_name == "us-east-1"
assert "aws_external_id" not in optional_params

View file

@ -0,0 +1,48 @@
import datetime
from unittest.mock import patch
import boto3
from botocore.exceptions import ClientError
from litellm.llms.sagemaker.chat.handler import SagemakerChatHandler
def test_load_credentials_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied by the deployment's aws_external_id."""
monkeypatch.delenv("AWS_EXTERNAL_ID", raising=False)
class FakeSTSClient:
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-sm-chat":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIASMCHATROLEKEY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
optional_params = {
"aws_access_key_id": "AKIASMCHATCALLERKEY",
"aws_secret_access_key": "pod-caller-secret",
"aws_region_name": "us-east-1",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-sm-chat-role",
"aws_session_name": "litellm-sm-chat-session",
"aws_external_id": "external-id-sm-chat",
}
with patch.object(boto3, "client", return_value=FakeSTSClient()):
credentials, aws_region_name = SagemakerChatHandler()._load_credentials(optional_params)
assert credentials.access_key == "ASIASMCHATROLEKEY"
assert credentials.token == "assumed-session-token"
assert aws_region_name == "us-east-1"
assert "aws_external_id" not in optional_params

View file

@ -172,3 +172,50 @@ async def test_async_native_streaming_forwards_each_frame_incrementally():
assert texts == [f"token{i} " for i in range(len(frames))]
assert consumed_at_token == list(range(1, len(frames) + 1))
def test_load_credentials_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied by the deployment's aws_external_id."""
import datetime
import boto3
from botocore.exceptions import ClientError
from unittest.mock import patch
monkeypatch.delenv("AWS_EXTERNAL_ID", raising=False)
class FakeSTSClient:
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-sm-completion":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIASMCOMPROLEKEY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
optional_params = {
"aws_access_key_id": "AKIASMCOMPCALLERKEY",
"aws_secret_access_key": "pod-caller-secret",
"aws_region_name": "us-east-1",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-sm-completion-role",
"aws_session_name": "litellm-sm-completion-session",
"aws_external_id": "external-id-sm-completion",
}
with patch.object(boto3, "client", return_value=FakeSTSClient()):
credentials, aws_region_name = SagemakerLLM()._load_credentials(optional_params)
assert credentials.access_key == "ASIASMCOMPROLEKEY"
assert credentials.token == "assumed-session-token"
assert aws_region_name == "us-east-1"
assert "aws_external_id" not in optional_params