mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
fix(bedrock): forward aws_external_id in files and batches credential loading
This commit is contained in:
parent
d1320404fe
commit
76839ca9d8
6 changed files with 234 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"}'
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}'
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue