mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
c5b4456401
commit
b463ddf92c
2 changed files with 112 additions and 0 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue