From b463ddf92c9611ddef89ddd25289a2a2ebf2c2f8 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 17 Jul 2026 22:06:58 +0000 Subject: [PATCH] feat(s3_v2): support s3_server_side_encryption_kms_key_id to set x-amz-server-side-encryption-aws-kms-key-id --- litellm/integrations/s3_v2.py | 17 ++++ tests/test_litellm/integrations/test_s3_v2.py | 95 +++++++++++++++++++ 2 files changed, 112 insertions(+) diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index 5b953035cfd..1b435bd69c8 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -55,6 +55,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, s3_server_side_encryption: Optional[str] = None, + s3_server_side_encryption_kms_key_id: str | None = None, s3_callback_params_override: Optional[dict] = None, **kwargs, ): @@ -94,6 +95,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix=s3_use_key_prefix, s3_use_virtual_hosted_style=s3_use_virtual_hosted_style, s3_server_side_encryption=s3_server_side_encryption, + s3_server_side_encryption_kms_key_id=s3_server_side_encryption_kms_key_id, ) verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}") @@ -148,6 +150,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, s3_server_side_encryption: Optional[str] = None, + s3_server_side_encryption_kms_key_id: str | None = None, params_source: Optional[dict] = None, ): """ @@ -199,6 +202,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption + self.s3_server_side_encryption_kms_key_id = ( + params.get("s3_server_side_encryption_kms_key_id") or s3_server_side_encryption_kms_key_id + ) + return async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -340,6 +347,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): if self.s3_server_side_encryption else {} ), + **( + {"x-amz-server-side-encryption-aws-kms-key-id": self.s3_server_side_encryption_kms_key_id} + if self.s3_server_side_encryption_kms_key_id + else {} + ), } req = requests.Request("PUT", url, data=json_string, headers=headers) prepped = req.prepare() @@ -515,6 +527,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): if self.s3_server_side_encryption else {} ), + **( + {"x-amz-server-side-encryption-aws-kms-key-id": self.s3_server_side_encryption_kms_key_id} + if self.s3_server_side_encryption_kms_key_id + else {} + ), } req = requests.Request("PUT", url, data=json_string, headers=headers) prepped = req.prepare() diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/test_litellm/integrations/test_s3_v2.py index f0a33f2ebfc..344b65eeab8 100644 --- a/tests/test_litellm/integrations/test_s3_v2.py +++ b/tests/test_litellm/integrations/test_s3_v2.py @@ -1388,3 +1388,98 @@ def test_s3_server_side_encryption_read_from_callback_params(): assert logger.s3_server_side_encryption == "aws:kms" finally: litellm.s3_callback_params = original + + +@pytest.mark.asyncio +async def test_async_upload_sets_kms_key_id_header_when_configured(): + """ + When s3_server_side_encryption_kms_key_id is set, the PUT must carry + x-amz-server-side-encryption-aws-kms-key-id so buckets whose IAM policy + enforces an exact customer-managed KMS key stop returning 403. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_server_side_encryption="aws:kms", + s3_server_side_encryption_kms_key_id="arn:aws:kms:us-east-1:123456789012:key/abc-123", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-kms.json", + payload={"test": "kms"}, + s3_object_download_filename="test-kms.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put.return_value = response + + await logger.async_upload_data_to_s3(test_element) + + headers = logger.async_httpx_client.put.call_args.kwargs["headers"] + assert headers["x-amz-server-side-encryption"] == "aws:kms" + assert ( + headers["x-amz-server-side-encryption-aws-kms-key-id"] + == "arn:aws:kms:us-east-1:123456789012:key/abc-123" + ) + + +@pytest.mark.asyncio +async def test_async_upload_omits_kms_key_id_header_when_not_configured(): + """Without s3_server_side_encryption_kms_key_id, the KMS header is absent.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_server_side_encryption="aws:kms", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-no-kms.json", + payload={"test": "no-kms"}, + s3_object_download_filename="test-no-kms.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put.return_value = response + + await logger.async_upload_data_to_s3(test_element) + + headers = logger.async_httpx_client.put.call_args.kwargs["headers"] + assert "x-amz-server-side-encryption-aws-kms-key-id" not in headers + + +def test_s3_server_side_encryption_kms_key_id_read_from_callback_params(): + """s3_server_side_encryption_kms_key_id can be configured via s3_callback_params.""" + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = { + "s3_bucket_name": "from-global", + "s3_server_side_encryption": "aws:kms", + "s3_server_side_encryption_kms_key_id": "arn:aws:kms:us-east-1:123456789012:key/abc-123", + } + try: + logger = S3Logger() + assert ( + logger.s3_server_side_encryption_kms_key_id + == "arn:aws:kms:us-east-1:123456789012:key/abc-123" + ) + finally: + litellm.s3_callback_params = original