fix(bedrock): forward aws_external_id in files and batches credential loading

This commit is contained in:
mateo-berri 2026-08-31 21:18:36 -07:00
parent d1320404fe
commit 76839ca9d8
6 changed files with 234 additions and 0 deletions

View file

@ -1487,6 +1487,7 @@ class CommonBatchFilesUtils:
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Prepare the request data

View file

@ -113,6 +113,7 @@ class BedrockFilesHandler(BaseAWSLLM):
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Create S3 client

View file

@ -146,6 +146,7 @@ class _BedrockS3RequestParams(BaseModel):
aws_role_name: str | None = None
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
s3_region_name: str | None = None
s3_endpoint_url: str | None = None
@ -1029,6 +1030,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Calculate SHA256 hash of the content (REQUIRED for S3)
@ -1290,6 +1292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
aws_role_name=request_params.aws_role_name,
aws_web_identity_token=request_params.aws_web_identity_token,
aws_sts_endpoint=request_params.aws_sts_endpoint,
aws_external_id=request_params.aws_external_id,
)
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()

View file

@ -204,3 +204,71 @@ def test_should_forward_trusted_model_credentials_to_retrieve_provider_config():
assert response is mock_response
litellm_params = mock_retrieve_file.call_args.kwargs["litellm_params"]
assert litellm_params["_litellm_internal_model_credentials"] is trusted_credentials
@pytest.mark.asyncio
async def test_afile_content_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
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-files-download":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIAFILESDOWNLOADROLE",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
class FakeS3Body:
def read(self):
return b'{"custom_id": "req-1"}'
class FakeS3Client:
def get_object(self, Bucket, Key):
return {"Body": FakeS3Body()}
s3_client_kwargs = {}
def fake_boto3_client(service_name, **kwargs):
if service_name == "sts":
return FakeSTSClient()
s3_client_kwargs.update(kwargs)
return FakeS3Client()
optional_params = {
"_litellm_internal_model_credentials": MappingProxyType({"s3_bucket_name": "safe-bucket"}),
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIAFILESDOWNLOADCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-files-download-role",
"aws_session_name": "litellm-files-download-session",
"aws_external_id": "external-id-files-download",
}
with patch.object(boto3, "client", side_effect=fake_boto3_client):
response = await BedrockFilesHandler().afile_content(
file_content_request={"file_id": "s3://safe-bucket/litellm-bedrock-files-model-id-abc.jsonl"},
optional_params=optional_params,
timeout=10.0,
max_retries=None,
)
assert s3_client_kwargs["aws_access_key_id"] == "ASIAFILESDOWNLOADROLE"
assert s3_client_kwargs["aws_session_token"] == "assumed-session-token"
assert response.content == b'{"custom_id": "req-1"}'

View file

@ -2404,3 +2404,111 @@ class TestBedrockFilesS3SignatureEncoding:
body=None,
headers=litellm_params[S3_SIGNED_GET_HEADERS_PARAM],
)
def test_sign_s3_request_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied when signing the S3 upload request."""
import datetime
from unittest.mock import patch
import boto3
from botocore.exceptions import ClientError
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
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-files-put":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIAFILESPUTROLE",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
optional_params = {
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIAFILESPUTCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-files-put-role",
"aws_session_name": "litellm-files-put-session",
"aws_external_id": "external-id-files-put",
}
with patch.object(boto3, "client", return_value=FakeSTSClient()):
signed_headers, _signed_body = BedrockFilesConfig()._sign_s3_request(
content='{"custom_id": "req-1"}',
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
optional_params=optional_params,
)
authorization = {key.lower(): value for key, value in signed_headers.items()}["authorization"]
assert "ASIAFILESPUTROLE" in authorization
def test_sign_s3_get_request_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied when signing the S3 download request."""
import datetime
from unittest.mock import patch
import boto3
from botocore.exceptions import ClientError
from litellm.llms.bedrock.files.transformation import (
BedrockFilesConfig,
_BedrockS3RequestParams,
)
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-files-get":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIAFILESGETROLE",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
request_params = _BedrockS3RequestParams.model_validate(
{
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIAFILESGETCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-files-get-role",
"aws_session_name": "litellm-files-get-session",
"aws_external_id": "external-id-files-get",
}
)
assert request_params.aws_external_id == "external-id-files-get"
with patch.object(boto3, "client", return_value=FakeSTSClient()):
signed_headers = BedrockFilesConfig()._sign_s3_get_request(
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
aws_region_name="us-east-1",
request_params=request_params,
)
authorization = {key.lower(): value for key, value in signed_headers.items()}["authorization"]
assert "ASIAFILESGETROLE" in authorization

View file

@ -520,3 +520,56 @@ def test_merge_bedrock_aws_request_params_keeps_caller_credentials_without_stati
assert merged["aws_secret_access_key"] == "caller-secret"
assert merged["aws_session_token"] == "caller-token"
assert merged["aws_region_name"] == "us-west-2"
def test_sign_aws_request_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied when signing batch API requests."""
import datetime
from unittest.mock import patch
import boto3
from botocore.exceptions import ClientError
from litellm.llms.bedrock.common_utils import CommonBatchFilesUtils
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-batch-sign":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIABATCHSIGNROLE",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
optional_params = {
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIABATCHSIGNCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-batch-sign-role",
"aws_session_name": "litellm-batch-sign-session",
"aws_external_id": "external-id-batch-sign",
}
with patch.object(boto3, "client", return_value=FakeSTSClient()):
signed_headers, signed_data = CommonBatchFilesUtils().sign_aws_request(
service_name="bedrock",
data={"jobName": "litellm-batch-job"},
endpoint_url="https://bedrock.us-east-1.amazonaws.com/model-invocation-job",
optional_params=optional_params,
)
authorization = {key.lower(): value for key, value in signed_headers.items()}["authorization"]
assert "ASIABATCHSIGNROLE" in authorization
assert signed_data == b'{"jobName": "litellm-batch-job"}'