mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #31928 from BerriAI/litellm_s3_v2_content_md5
This commit is contained in:
commit
9b6d0b0e98
2 changed files with 178 additions and 0 deletions
|
|
@ -54,6 +54,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files: bool = False,
|
||||
s3_use_key_prefix: bool = False,
|
||||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: Optional[str] = None,
|
||||
s3_callback_params_override: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -92,6 +93,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files=s3_strip_base64_files,
|
||||
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,
|
||||
)
|
||||
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
|
||||
|
||||
|
|
@ -145,6 +147,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files: bool = False,
|
||||
s3_use_key_prefix: bool = False,
|
||||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: Optional[str] = None,
|
||||
params_source: Optional[dict] = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -194,6 +197,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
|
||||
)
|
||||
|
||||
self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption
|
||||
|
||||
return
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -273,6 +278,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
import requests
|
||||
|
|
@ -317,14 +323,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5 = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**(
|
||||
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
|
||||
if self.s3_server_side_encryption
|
||||
else {}
|
||||
),
|
||||
}
|
||||
req = requests.Request("PUT", url, data=json_string, headers=headers)
|
||||
prepped = req.prepare()
|
||||
|
|
@ -447,6 +462,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
import requests
|
||||
|
|
@ -482,14 +498,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5 = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**(
|
||||
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
|
||||
if self.s3_server_side_encryption
|
||||
else {}
|
||||
),
|
||||
}
|
||||
req = requests.Request("PUT", url, data=json_string, headers=headers)
|
||||
prepped = req.prepare()
|
||||
|
|
|
|||
|
|
@ -1194,3 +1194,156 @@ def test_s3_callback_params_override_empty_dict_is_opt_in():
|
|||
assert logger.s3_bucket_name is None
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
||||
|
||||
def _expected_content_md5(payload: dict) -> str:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
json_string = safe_dumps(payload)
|
||||
return base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
|
||||
def _require_non_security_md5(monkeypatch):
|
||||
import hashlib
|
||||
|
||||
original_md5 = hashlib.md5
|
||||
|
||||
def fips_md5(data=b"", *, usedforsecurity=True):
|
||||
if usedforsecurity:
|
||||
raise ValueError("MD5 blocked for security use")
|
||||
return original_md5(data, usedforsecurity=usedforsecurity)
|
||||
|
||||
monkeypatch.setattr(hashlib, "md5", fips_md5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_upload_sets_content_md5_header(monkeypatch):
|
||||
"""
|
||||
Object Lock buckets reject PUTs without a Content-MD5 header (AWS spec).
|
||||
The async upload must send a base64 md5 of the exact signed body.
|
||||
"""
|
||||
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",
|
||||
)
|
||||
|
||||
payload = {"test": "content-md5"}
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-md5.json",
|
||||
payload=payload,
|
||||
s3_object_download_filename="test-md5.json",
|
||||
)
|
||||
_require_non_security_md5(monkeypatch)
|
||||
|
||||
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["Content-MD5"] == _expected_content_md5(payload)
|
||||
assert "x-amz-server-side-encryption" not in headers
|
||||
|
||||
|
||||
def test_sync_upload_sets_content_md5_header(monkeypatch):
|
||||
"""The sync upload path must also send Content-MD5 for Object Lock buckets."""
|
||||
from unittest.mock import 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",
|
||||
)
|
||||
|
||||
payload = {"test": "sync-content-md5"}
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-sync-md5.json",
|
||||
payload=payload,
|
||||
s3_object_download_filename="test-sync-md5.json",
|
||||
)
|
||||
_require_non_security_md5(monkeypatch)
|
||||
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.raise_for_status = MagicMock()
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.put.return_value = response
|
||||
|
||||
with patch(
|
||||
"litellm.integrations.s3_v2._get_httpx_client",
|
||||
return_value=mock_sync_client,
|
||||
):
|
||||
logger.upload_data_to_s3(test_element)
|
||||
|
||||
headers = mock_sync_client.put.call_args.kwargs["headers"]
|
||||
assert headers["Content-MD5"] == _expected_content_md5(payload)
|
||||
assert "x-amz-server-side-encryption" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_upload_sets_server_side_encryption_header_when_configured():
|
||||
"""
|
||||
When s3_server_side_encryption is set (e.g. buckets with a KMS default
|
||||
encryption policy), the PUT must carry x-amz-server-side-encryption.
|
||||
"""
|
||||
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-sse.json",
|
||||
payload={"test": "sse"},
|
||||
s3_object_download_filename="test-sse.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"
|
||||
|
||||
|
||||
def test_s3_server_side_encryption_read_from_callback_params():
|
||||
"""s3_server_side_encryption 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",
|
||||
}
|
||||
try:
|
||||
logger = S3Logger()
|
||||
assert logger.s3_server_side_encryption == "aws:kms"
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue