feat(s3_v2): support s3_server_side_encryption_kms_key_id to set x-amz-server-side-encryption-aws-kms-key-id

This commit is contained in:
Devin AI 2026-07-17 22:06:58 +00:00
parent c5b4456401
commit b463ddf92c
2 changed files with 112 additions and 0 deletions

View file

@ -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()

View file

@ -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