mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
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:
commit
39e4f1ae13
6 changed files with 149 additions and 0 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue