mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Co-authored-by: yucheng <yucheng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
1858 lines
66 KiB
Python
1858 lines
66 KiB
Python
import asyncio
|
||
from datetime import datetime
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
import pytest
|
||
|
||
from litellm.integrations.s3_v2 import S3Logger
|
||
from litellm.types.utils import StandardLoggingPayload
|
||
|
||
|
||
class TestS3V2UnitTests:
|
||
"""Test that S3 v2 integration only uses safe_dumps and not json.dumps"""
|
||
|
||
def test_s3_v2_source_code_analysis(self):
|
||
"""Test that S3 v2 source code only imports and uses safe_dumps"""
|
||
import inspect
|
||
|
||
from litellm.integrations import s3_v2
|
||
|
||
# Get the source code of the s3_v2 module
|
||
source_code = inspect.getsource(s3_v2)
|
||
|
||
# Verify that json.dumps is not used directly in the code
|
||
assert (
|
||
"json.dumps(" not in source_code
|
||
), "S3 v2 should not use json.dumps directly"
|
||
|
||
@patch("asyncio.create_task")
|
||
@patch("litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush")
|
||
def test_s3_v2_endpoint_url(self, mock_periodic_flush, mock_create_task):
|
||
"""testing s3 endpoint url"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
# Mock periodic_flush and create_task to prevent async task creation during init
|
||
mock_periodic_flush.return_value = None
|
||
mock_create_task.return_value = None
|
||
|
||
# Mock response for all tests
|
||
mock_response = MagicMock()
|
||
mock_response.status_code = 200
|
||
mock_response.raise_for_status = MagicMock()
|
||
|
||
# Create a test batch logging element
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-key.json",
|
||
payload={"test": "data"},
|
||
s3_object_download_filename="test-file.json",
|
||
)
|
||
|
||
# Test 1: Custom endpoint URL with bucket name
|
||
s3_logger = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_endpoint_url="https://s3.amazonaws.com",
|
||
s3_aws_access_key_id="test-key",
|
||
s3_aws_secret_access_key="test-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
s3_logger.async_httpx_client = AsyncMock()
|
||
s3_logger.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger.async_upload_data_to_s3(test_element))
|
||
|
||
call_args = s3_logger.async_httpx_client.put.call_args
|
||
assert call_args is not None
|
||
url = call_args[0][0]
|
||
expected_url = "https://s3.amazonaws.com/test-bucket/2025-09-14/test-key.json"
|
||
assert url == expected_url, f"Expected URL {expected_url}, got {url}"
|
||
|
||
# Test 2: MinIO-compatible endpoint
|
||
s3_logger_minio = S3Logger(
|
||
s3_bucket_name="litellm-logs",
|
||
s3_endpoint_url="https://minio.example.com:9000",
|
||
s3_aws_access_key_id="minio-key",
|
||
s3_aws_secret_access_key="minio-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
s3_logger_minio.async_httpx_client = AsyncMock()
|
||
s3_logger_minio.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger_minio.async_upload_data_to_s3(test_element))
|
||
|
||
call_args_minio = s3_logger_minio.async_httpx_client.put.call_args
|
||
assert call_args_minio is not None
|
||
url_minio = call_args_minio[0][0]
|
||
expected_minio_url = (
|
||
"https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url_minio == expected_minio_url
|
||
), f"Expected MinIO URL {expected_minio_url}, got {url_minio}"
|
||
|
||
# Test 3: Custom endpoint without bucket name (should fall back to default)
|
||
s3_logger_no_bucket = S3Logger(
|
||
s3_endpoint_url="https://s3.amazonaws.com",
|
||
s3_aws_access_key_id="test-key",
|
||
s3_aws_secret_access_key="test-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
s3_logger_no_bucket.async_httpx_client = AsyncMock()
|
||
s3_logger_no_bucket.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger_no_bucket.async_upload_data_to_s3(test_element))
|
||
|
||
call_args_no_bucket = s3_logger_no_bucket.async_httpx_client.put.call_args
|
||
assert call_args_no_bucket is not None
|
||
url_no_bucket = call_args_no_bucket[0][0]
|
||
# Should use default S3 URL format when bucket is missing (bucket becomes None in URL)
|
||
assert "s3.us-east-1.amazonaws.com" in url_no_bucket
|
||
assert "https://" in url_no_bucket
|
||
# Should not include the custom endpoint since bucket is missing
|
||
assert "https://s3.amazonaws.com/" not in url_no_bucket
|
||
|
||
# Test 4: Sync upload method with custom endpoint
|
||
s3_logger_sync = S3Logger(
|
||
s3_bucket_name="sync-bucket",
|
||
s3_endpoint_url="https://custom.s3.endpoint.com",
|
||
s3_aws_access_key_id="sync-key",
|
||
s3_aws_secret_access_key="sync-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
mock_sync_client = MagicMock()
|
||
mock_sync_client.put.return_value = mock_response
|
||
|
||
with patch(
|
||
"litellm.integrations.s3_v2._get_httpx_client",
|
||
return_value=mock_sync_client,
|
||
):
|
||
s3_logger_sync.upload_data_to_s3(test_element)
|
||
|
||
call_args_sync = mock_sync_client.put.call_args
|
||
assert call_args_sync is not None
|
||
url_sync = call_args_sync[0][0]
|
||
expected_sync_url = (
|
||
"https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url_sync == expected_sync_url
|
||
), f"Expected sync URL {expected_sync_url}, got {url_sync}"
|
||
|
||
# Test 5: Download method with custom endpoint
|
||
s3_logger_download = S3Logger(
|
||
s3_bucket_name="download-bucket",
|
||
s3_endpoint_url="https://download.s3.endpoint.com",
|
||
s3_aws_access_key_id="download-key",
|
||
s3_aws_secret_access_key="download-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
mock_download_response = MagicMock()
|
||
mock_download_response.status_code = 200
|
||
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
|
||
s3_logger_download.async_httpx_client = AsyncMock()
|
||
s3_logger_download.async_httpx_client.get.return_value = mock_download_response
|
||
|
||
result = asyncio.run(
|
||
s3_logger_download._download_object_from_s3(
|
||
"2025-09-14/download-test-key.json"
|
||
)
|
||
)
|
||
|
||
call_args_download = s3_logger_download.async_httpx_client.get.call_args
|
||
assert call_args_download is not None
|
||
url_download = call_args_download[0][0]
|
||
expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json"
|
||
assert (
|
||
url_download == expected_download_url
|
||
), f"Expected download URL {expected_download_url}, got {url_download}"
|
||
|
||
assert result == {"downloaded": "data"}
|
||
|
||
@patch("asyncio.create_task")
|
||
@patch("litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush")
|
||
def test_s3_v2_virtual_hosted_style(self, mock_periodic_flush, mock_create_task):
|
||
"""Test s3_use_virtual_hosted_style parameter for virtual-hosted-style URLs"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
# Mock periodic_flush and create_task to prevent async task creation during init
|
||
mock_periodic_flush.return_value = None
|
||
mock_create_task.return_value = None
|
||
|
||
# Mock response for all tests
|
||
mock_response = MagicMock()
|
||
mock_response.status_code = 200
|
||
mock_response.raise_for_status = MagicMock()
|
||
|
||
# Create a test batch logging element
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-key.json",
|
||
payload={"test": "data"},
|
||
s3_object_download_filename="test-file.json",
|
||
)
|
||
|
||
# Test 1: Virtual-hosted-style with custom endpoint
|
||
s3_logger_virtual = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_endpoint_url="https://s3.custom-endpoint.com",
|
||
s3_aws_access_key_id="test-key",
|
||
s3_aws_secret_access_key="test-secret",
|
||
s3_region_name="us-east-1",
|
||
s3_use_virtual_hosted_style=True,
|
||
)
|
||
|
||
s3_logger_virtual.async_httpx_client = AsyncMock()
|
||
s3_logger_virtual.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger_virtual.async_upload_data_to_s3(test_element))
|
||
|
||
call_args = s3_logger_virtual.async_httpx_client.put.call_args
|
||
assert call_args is not None
|
||
url = call_args[0][0]
|
||
expected_url = (
|
||
"https://test-bucket.s3.custom-endpoint.com/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url == expected_url
|
||
), f"Expected virtual-hosted-style URL {expected_url}, got {url}"
|
||
|
||
# Test 2: Path-style (default behavior with s3_use_virtual_hosted_style=False)
|
||
s3_logger_path = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_endpoint_url="https://s3.custom-endpoint.com",
|
||
s3_aws_access_key_id="test-key",
|
||
s3_aws_secret_access_key="test-secret",
|
||
s3_region_name="us-east-1",
|
||
s3_use_virtual_hosted_style=False,
|
||
)
|
||
|
||
s3_logger_path.async_httpx_client = AsyncMock()
|
||
s3_logger_path.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger_path.async_upload_data_to_s3(test_element))
|
||
|
||
call_args_path = s3_logger_path.async_httpx_client.put.call_args
|
||
assert call_args_path is not None
|
||
url_path = call_args_path[0][0]
|
||
expected_path_url = (
|
||
"https://s3.custom-endpoint.com/test-bucket/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url_path == expected_path_url
|
||
), f"Expected path-style URL {expected_path_url}, got {url_path}"
|
||
|
||
# Test 3: Virtual-hosted-style with http protocol
|
||
s3_logger_http = S3Logger(
|
||
s3_bucket_name="http-bucket",
|
||
s3_endpoint_url="http://minio.local:9000",
|
||
s3_aws_access_key_id="minio-key",
|
||
s3_aws_secret_access_key="minio-secret",
|
||
s3_region_name="us-east-1",
|
||
s3_use_virtual_hosted_style=True,
|
||
)
|
||
|
||
s3_logger_http.async_httpx_client = AsyncMock()
|
||
s3_logger_http.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger_http.async_upload_data_to_s3(test_element))
|
||
|
||
call_args_http = s3_logger_http.async_httpx_client.put.call_args
|
||
assert call_args_http is not None
|
||
url_http = call_args_http[0][0]
|
||
expected_http_url = (
|
||
"http://http-bucket.minio.local:9000/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url_http == expected_http_url
|
||
), f"Expected virtual-hosted-style URL with http {expected_http_url}, got {url_http}"
|
||
|
||
# Test 4: Sync upload method with virtual-hosted-style
|
||
s3_logger_sync_virtual = S3Logger(
|
||
s3_bucket_name="sync-bucket",
|
||
s3_endpoint_url="https://storage.example.com",
|
||
s3_aws_access_key_id="sync-key",
|
||
s3_aws_secret_access_key="sync-secret",
|
||
s3_region_name="us-east-1",
|
||
s3_use_virtual_hosted_style=True,
|
||
)
|
||
|
||
mock_sync_client = MagicMock()
|
||
mock_sync_client.put.return_value = mock_response
|
||
|
||
with patch(
|
||
"litellm.integrations.s3_v2._get_httpx_client",
|
||
return_value=mock_sync_client,
|
||
):
|
||
s3_logger_sync_virtual.upload_data_to_s3(test_element)
|
||
|
||
call_args_sync = mock_sync_client.put.call_args
|
||
assert call_args_sync is not None
|
||
url_sync = call_args_sync[0][0]
|
||
expected_sync_url = (
|
||
"https://sync-bucket.storage.example.com/2025-09-14/test-key.json"
|
||
)
|
||
assert (
|
||
url_sync == expected_sync_url
|
||
), f"Expected virtual-hosted-style sync URL {expected_sync_url}, got {url_sync}"
|
||
|
||
# Test 5: Download method with virtual-hosted-style
|
||
s3_logger_download_virtual = S3Logger(
|
||
s3_bucket_name="download-bucket",
|
||
s3_endpoint_url="https://download.endpoint.com",
|
||
s3_aws_access_key_id="download-key",
|
||
s3_aws_secret_access_key="download-secret",
|
||
s3_region_name="us-east-1",
|
||
s3_use_virtual_hosted_style=True,
|
||
)
|
||
|
||
mock_download_response = MagicMock()
|
||
mock_download_response.status_code = 200
|
||
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
|
||
s3_logger_download_virtual.async_httpx_client = AsyncMock()
|
||
s3_logger_download_virtual.async_httpx_client.get.return_value = (
|
||
mock_download_response
|
||
)
|
||
|
||
result = asyncio.run(
|
||
s3_logger_download_virtual._download_object_from_s3(
|
||
"2025-09-14/download-test-key.json"
|
||
)
|
||
)
|
||
|
||
call_args_download = s3_logger_download_virtual.async_httpx_client.get.call_args
|
||
assert call_args_download is not None
|
||
url_download = call_args_download[0][0]
|
||
expected_download_url = "https://download-bucket.download.endpoint.com/2025-09-14/download-test-key.json"
|
||
assert (
|
||
url_download == expected_download_url
|
||
), f"Expected virtual-hosted-style download URL {expected_download_url}, got {url_download}"
|
||
|
||
assert result == {"downloaded": "data"}
|
||
|
||
@patch("asyncio.create_task")
|
||
@patch("litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush")
|
||
def test_s3_v2_put_url_encodes_spaces_in_object_key(
|
||
self, mock_periodic_flush, mock_create_task
|
||
):
|
||
import requests
|
||
from unittest.mock import AsyncMock
|
||
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
mock_periodic_flush.return_value = None
|
||
mock_create_task.return_value = None
|
||
|
||
mock_response = MagicMock()
|
||
mock_response.status_code = 200
|
||
mock_response.raise_for_status = MagicMock()
|
||
|
||
s3_object_key = "My Team/2025-09-14/test-key.json"
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key=s3_object_key,
|
||
payload={"test": "data"},
|
||
s3_object_download_filename="test-file.json",
|
||
)
|
||
|
||
s3_logger = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_endpoint_url="https://s3.amazonaws.com",
|
||
s3_aws_access_key_id="test-key",
|
||
s3_aws_secret_access_key="test-secret",
|
||
s3_region_name="us-east-1",
|
||
)
|
||
s3_logger.async_httpx_client = AsyncMock()
|
||
s3_logger.async_httpx_client.put.return_value = mock_response
|
||
|
||
asyncio.run(s3_logger.async_upload_data_to_s3(test_element))
|
||
|
||
call_args = s3_logger.async_httpx_client.put.call_args
|
||
assert call_args is not None
|
||
actual_url = call_args[0][0]
|
||
raw_url = f"https://s3.amazonaws.com/test-bucket/{s3_object_key}"
|
||
expected_url = requests.Request("PUT", raw_url).prepare().url
|
||
assert actual_url == expected_url
|
||
assert " " not in actual_url
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_retries_on_s3_503():
|
||
"""
|
||
Test that async_upload_data_to_s3 retries on transient S3 503 Slow Down
|
||
and succeeds on the second attempt.
|
||
"""
|
||
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",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-retry.json",
|
||
payload={"test": "retry"},
|
||
s3_object_download_filename="test-retry.json",
|
||
)
|
||
|
||
# First call returns 503, second call returns 200
|
||
response_503 = MagicMock()
|
||
response_503.status_code = 503
|
||
response_200 = MagicMock()
|
||
response_200.status_code = 200
|
||
response_200.raise_for_status = MagicMock()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = AsyncMock(side_effect=[response_503, response_200])
|
||
|
||
with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
# Verify PUT was called twice (retry after 503)
|
||
assert logger.async_httpx_client.put.call_count == 2
|
||
# Verify sleep was called with the backoff delay
|
||
mock_sleep.assert_called_once_with(1) # 2**0 = 1s
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_retries_on_s3_500():
|
||
"""
|
||
Test that async_upload_data_to_s3 retries on transient S3 500 errors.
|
||
"""
|
||
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",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-retry-500.json",
|
||
payload={"test": "retry-500"},
|
||
s3_object_download_filename="test-retry-500.json",
|
||
)
|
||
|
||
response_500 = MagicMock()
|
||
response_500.status_code = 500
|
||
response_200 = MagicMock()
|
||
response_200.status_code = 200
|
||
response_200.raise_for_status = MagicMock()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = AsyncMock(side_effect=[response_500, response_200])
|
||
|
||
with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
assert logger.async_httpx_client.put.call_count == 2
|
||
mock_sleep.assert_called_once_with(1)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_exhausts_retries_on_persistent_503():
|
||
"""
|
||
Test that async_upload_data_to_s3 raises after exhausting all retries
|
||
on persistent S3 503.
|
||
"""
|
||
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",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-exhaust.json",
|
||
payload={"test": "exhaust"},
|
||
s3_object_download_filename="test-exhaust.json",
|
||
)
|
||
|
||
# All 3 attempts return 503
|
||
response_503 = MagicMock()
|
||
response_503.status_code = 503
|
||
response_503.raise_for_status = MagicMock(
|
||
side_effect=Exception("503 Service Unavailable")
|
||
)
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = AsyncMock(return_value=response_503)
|
||
|
||
with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
|
||
with patch.object(logger, "handle_callback_failure") as mock_failure:
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
# 3 PUT attempts total
|
||
assert logger.async_httpx_client.put.call_count == 3
|
||
# 2 sleeps (between attempts 1-2 and 2-3)
|
||
assert mock_sleep.call_count == 2
|
||
# Callback failure handler called after exhausting retries
|
||
mock_failure.assert_called_once_with(callback_name="S3Logger")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_no_retry_on_4xx():
|
||
"""
|
||
Test that async_upload_data_to_s3 does NOT retry on 4xx errors (client errors).
|
||
"""
|
||
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",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-no-retry.json",
|
||
payload={"test": "no-retry"},
|
||
s3_object_download_filename="test-no-retry.json",
|
||
)
|
||
|
||
response_403 = MagicMock()
|
||
response_403.status_code = 403
|
||
response_403.raise_for_status = MagicMock(side_effect=Exception("403 Forbidden"))
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = AsyncMock(return_value=response_403)
|
||
|
||
with patch.object(logger, "handle_callback_failure") as mock_failure:
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
# Only 1 attempt — no retry for 4xx
|
||
assert logger.async_httpx_client.put.call_count == 1
|
||
mock_failure.assert_called_once_with(callback_name="S3Logger")
|
||
|
||
|
||
def test_sync_upload_retries_on_s3_503():
|
||
"""
|
||
Test that the sync upload_data_to_s3 retries on transient S3 503.
|
||
"""
|
||
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",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-sync-retry.json",
|
||
payload={"test": "sync-retry"},
|
||
s3_object_download_filename="test-sync-retry.json",
|
||
)
|
||
|
||
response_503 = MagicMock()
|
||
response_503.status_code = 503
|
||
response_200 = MagicMock()
|
||
response_200.status_code = 200
|
||
response_200.raise_for_status = MagicMock()
|
||
|
||
mock_sync_client = MagicMock()
|
||
mock_sync_client.put = MagicMock(side_effect=[response_503, response_200])
|
||
|
||
with patch(
|
||
"litellm.integrations.s3_v2._get_httpx_client",
|
||
return_value=mock_sync_client,
|
||
):
|
||
with patch("time.sleep") as mock_sleep:
|
||
logger.upload_data_to_s3(test_element)
|
||
|
||
assert mock_sync_client.put.call_count == 2
|
||
mock_sleep.assert_called_once_with(1)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_log_event_skips_when_standard_logging_object_missing():
|
||
"""
|
||
Reproduces the bug where _async_log_event_base raises ValueError when
|
||
kwargs has no standard_logging_object (e.g. call_type=afile_delete).
|
||
|
||
The S3 logger should skip gracefully, not raise.
|
||
"""
|
||
logger = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_region_name="us-east-1",
|
||
s3_aws_access_key_id="fake",
|
||
s3_aws_secret_access_key="fake",
|
||
)
|
||
|
||
kwargs_without_slo = {
|
||
"call_type": "afile_delete",
|
||
"model": None,
|
||
"litellm_call_id": "test-call-id",
|
||
}
|
||
|
||
start_time = datetime.utcnow()
|
||
end_time = datetime.utcnow()
|
||
|
||
# Spy on handle_callback_failure — should NOT be called if we skip gracefully.
|
||
# Without the fix, the ValueError is caught by the except block which calls
|
||
# handle_callback_failure. With the fix, we return early and never hit except.
|
||
with patch.object(logger, "handle_callback_failure") as mock_failure:
|
||
await logger._async_log_event_base(
|
||
kwargs=kwargs_without_slo,
|
||
response_obj=None,
|
||
start_time=start_time,
|
||
end_time=end_time,
|
||
)
|
||
|
||
assert not mock_failure.called, (
|
||
"handle_callback_failure should not be called — "
|
||
"missing standard_logging_object should be a graceful skip, not an error"
|
||
)
|
||
|
||
# Nothing should have been queued (catches the case where code falls
|
||
# through without returning and appends None to the queue)
|
||
assert (
|
||
len(logger.log_queue) == 0
|
||
), "log_queue should be empty when standard_logging_object is missing"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_strip_base64_removes_file_and_nontext_entries():
|
||
logger = S3Logger(s3_strip_base64_files=True)
|
||
|
||
payload = {
|
||
"messages": [
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{"type": "text", "text": "Hello world"},
|
||
{
|
||
"type": "image",
|
||
"file": {"file_data": "data:image/png;base64,AAAA"},
|
||
},
|
||
{
|
||
"type": "file",
|
||
"file": {"file_data": "data:application/pdf;base64,BBBB"},
|
||
},
|
||
],
|
||
},
|
||
{
|
||
"role": "assistant",
|
||
"content": [
|
||
{"type": "text", "text": "Response"},
|
||
{
|
||
"type": "audio",
|
||
"file": {"file_data": "data:audio/wav;base64,CCCC"},
|
||
},
|
||
],
|
||
},
|
||
]
|
||
}
|
||
|
||
stripped = await logger._strip_base64_from_messages(payload)
|
||
|
||
# 1️⃣ File/image/audio entries are removed
|
||
assert len(stripped["messages"][0]["content"]) == 1
|
||
assert stripped["messages"][0]["content"][0]["text"] == "Hello world"
|
||
|
||
assert len(stripped["messages"][1]["content"]) == 1
|
||
assert stripped["messages"][1]["content"][0]["text"] == "Response"
|
||
|
||
# 2️⃣ No 'file' keys remain
|
||
for msg in stripped["messages"]:
|
||
for content in msg["content"]:
|
||
assert "file" not in content
|
||
assert content.get("type") == "text"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_strip_base64_keeps_non_file_content():
|
||
logger = S3Logger(s3_strip_base64_files=True)
|
||
|
||
payload = {
|
||
"messages": [
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{"type": "text", "text": "Just text"},
|
||
{"type": "text", "text": "Another message"},
|
||
],
|
||
}
|
||
]
|
||
}
|
||
|
||
stripped = await logger._strip_base64_from_messages(payload)
|
||
|
||
# Should not modify pure text messages
|
||
assert stripped["messages"][0]["content"] == payload["messages"][0]["content"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_strip_base64_handles_empty_or_missing_messages():
|
||
logger = S3Logger(s3_strip_base64_files=True)
|
||
|
||
# Missing messages key
|
||
payload_no_messages = {}
|
||
stripped1 = await logger._strip_base64_from_messages(payload_no_messages)
|
||
assert stripped1 == payload_no_messages
|
||
|
||
# Empty messages list
|
||
payload_empty = {"messages": []}
|
||
stripped2 = await logger._strip_base64_from_messages(payload_empty)
|
||
assert stripped2 == payload_empty
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_strip_base64_mixed_nested_objects():
|
||
"""
|
||
Handles weird/nested content structures gracefully.
|
||
"""
|
||
logger = S3Logger(s3_strip_base64_files=True)
|
||
|
||
payload = {
|
||
"messages": [
|
||
{
|
||
"role": "system",
|
||
"content": [
|
||
{"type": "text", "text": "Keep me"},
|
||
{"type": "custom", "metadata": "ignore but non-text"},
|
||
{"foo": "bar"},
|
||
{"file": {"file_data": "data:application/pdf;base64,XXX"}},
|
||
],
|
||
"extra": {"trace_id": "123"},
|
||
}
|
||
]
|
||
}
|
||
|
||
stripped = await logger._strip_base64_from_messages(payload)
|
||
|
||
# Custom/non-text and file entries removed
|
||
content = stripped["messages"][0]["content"]
|
||
assert len(content) == 2
|
||
assert {"type": "text", "text": "Keep me"} in content
|
||
assert {"foo": "bar"} in content
|
||
# Extra metadata preserved
|
||
assert stripped["messages"][0]["extra"]["trace_id"] == "123"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_s3_verify_false_handling(monkeypatch: pytest.MonkeyPatch):
|
||
"""
|
||
Test that s3_verify=False is properly handled and not treated as None.
|
||
|
||
This is a regression test for the bug where s3_verify=False was being
|
||
ignored because 'False or s3_verify' would evaluate to s3_verify (None).
|
||
"""
|
||
from unittest.mock import AsyncMock, patch
|
||
|
||
import litellm
|
||
|
||
# Set up s3_callback_params with s3_verify=False
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "test-bucket",
|
||
"s3_endpoint_url": "https://localhost:443",
|
||
"s3_aws_access_key_id": "minioadmin",
|
||
"s3_aws_secret_access_key": "minioadmin",
|
||
"s3_region_name": "us-east-1",
|
||
"s3_verify": False, # This should NOT be ignored
|
||
"s3_use_ssl": False, # This should also NOT be ignored
|
||
},
|
||
)
|
||
|
||
with patch("asyncio.create_task"):
|
||
with patch(
|
||
"litellm.integrations.s3_v2.get_async_httpx_client"
|
||
) as mock_get_client:
|
||
mock_client = AsyncMock()
|
||
mock_get_client.return_value = mock_client
|
||
|
||
# Create logger
|
||
logger = S3Logger()
|
||
|
||
# Verify s3_verify is False, not None
|
||
assert (
|
||
logger.s3_verify is False
|
||
), f"Expected s3_verify=False, got {logger.s3_verify}"
|
||
assert (
|
||
logger.s3_use_ssl is False
|
||
), f"Expected s3_use_ssl=False, got {logger.s3_use_ssl}"
|
||
|
||
# Verify that get_async_httpx_client was called with ssl_verify=False
|
||
mock_get_client.assert_called_once()
|
||
call_kwargs = mock_get_client.call_args.kwargs
|
||
assert (
|
||
"params" in call_kwargs
|
||
), "params should be passed to get_async_httpx_client"
|
||
assert call_kwargs["params"] == {
|
||
"ssl_verify": False
|
||
}, f"Expected ssl_verify=False in params, got {call_kwargs.get('params')}"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_s3_verify_none_handling(monkeypatch: pytest.MonkeyPatch):
|
||
"""
|
||
Test that s3_verify=None uses default behavior.
|
||
"""
|
||
from unittest.mock import AsyncMock, patch
|
||
|
||
import litellm
|
||
|
||
# Set up s3_callback_params without s3_verify
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"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",
|
||
},
|
||
)
|
||
|
||
with patch("asyncio.create_task"):
|
||
with patch(
|
||
"litellm.integrations.s3_v2.get_async_httpx_client"
|
||
) as mock_get_client:
|
||
mock_client = AsyncMock()
|
||
mock_get_client.return_value = mock_client
|
||
|
||
# Create logger without explicit s3_verify
|
||
logger = S3Logger()
|
||
|
||
# Verify s3_verify is None (default)
|
||
assert (
|
||
logger.s3_verify is None
|
||
), f"Expected s3_verify=None, got {logger.s3_verify}"
|
||
|
||
# Verify that get_async_httpx_client was called
|
||
mock_get_client.assert_called_once()
|
||
call_kwargs = mock_get_client.call_args.kwargs
|
||
# When s3_verify is None, params={'ssl_verify': None} which is fine - uses default behavior
|
||
# The important thing is it's not False
|
||
if "params" in call_kwargs and call_kwargs["params"] is not None:
|
||
assert call_kwargs["params"].get("ssl_verify") is None
|
||
# Either params is None or params={'ssl_verify': None} is acceptable
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_s3_verify_false_creates_httpx_client_with_verify_false(monkeypatch: pytest.MonkeyPatch):
|
||
"""
|
||
Test that when s3_verify=False, the actual httpx client has verify=False.
|
||
|
||
This validates that ssl_verify=False flows through to the httpx.AsyncClient.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
import litellm
|
||
|
||
# Set up s3_callback_params with s3_verify=False
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "test-bucket",
|
||
"s3_endpoint_url": "https://localhost:443",
|
||
"s3_aws_access_key_id": "minioadmin",
|
||
"s3_aws_secret_access_key": "minioadmin",
|
||
"s3_region_name": "us-east-1",
|
||
"s3_verify": False,
|
||
},
|
||
)
|
||
|
||
with patch("asyncio.create_task"):
|
||
# Create logger - this creates the httpx client
|
||
logger = S3Logger()
|
||
|
||
# Verify the logger has s3_verify=False
|
||
assert logger.s3_verify is False
|
||
|
||
# Check the actual httpx client has verify=False
|
||
# The async_httpx_client.client is the actual httpx.AsyncClient
|
||
if hasattr(logger.async_httpx_client, "client"):
|
||
httpx_client = logger.async_httpx_client.client
|
||
# Check the _verify attribute (httpx internal)
|
||
if hasattr(httpx_client, "_verify"):
|
||
assert (
|
||
httpx_client._verify is False
|
||
), f"Expected httpx client _verify=False, got {httpx_client._verify}"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_s3_verify_false_async_client(monkeypatch: pytest.MonkeyPatch):
|
||
"""
|
||
Test that the async httpx client respects s3_verify=False.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
import litellm
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
# Set up s3_callback_params with s3_verify=False
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "test-bucket",
|
||
"s3_endpoint_url": "https://localhost:443",
|
||
"s3_aws_access_key_id": "minioadmin",
|
||
"s3_aws_secret_access_key": "minioadmin",
|
||
"s3_region_name": "us-east-1",
|
||
"s3_verify": False,
|
||
},
|
||
)
|
||
|
||
with patch("asyncio.create_task"):
|
||
logger = S3Logger()
|
||
|
||
# Verify s3_verify is False
|
||
assert logger.s3_verify is False
|
||
|
||
# Create test element
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-11-03/test-key.json",
|
||
payload={"test": "data"},
|
||
s3_object_download_filename="test-file.json",
|
||
)
|
||
|
||
# Mock the async httpx client's put method
|
||
mock_response = MagicMock()
|
||
mock_response.status_code = 200
|
||
mock_response.raise_for_status = MagicMock()
|
||
logger.async_httpx_client.put = AsyncMock(return_value=mock_response)
|
||
|
||
# Call async upload
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
# Verify put was called
|
||
assert logger.async_httpx_client.put.called
|
||
|
||
# Check that the async httpx client was created with verify=False
|
||
if hasattr(logger.async_httpx_client, "client"):
|
||
httpx_client = logger.async_httpx_client.client
|
||
if hasattr(httpx_client, "_verify"):
|
||
assert (
|
||
httpx_client._verify is False
|
||
), f"Expected async httpx client _verify=False, got {httpx_client._verify}"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_strip_base64_recursive_redaction():
|
||
logger = S3Logger(s3_strip_base64_files=True)
|
||
payload = {
|
||
"messages": [
|
||
{
|
||
"content": [
|
||
{"type": "text", "text": "normal text"},
|
||
{
|
||
"type": "text",
|
||
"text": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg",
|
||
},
|
||
{
|
||
"type": "text",
|
||
"text": "Nested: {'data': 'data:application/pdf;base64,AAA...'}",
|
||
},
|
||
{"file": {"file_data": "data:application/pdf;base64,AAAA"}},
|
||
{"metadata": {"preview": "data:audio/mp3;base64,AAAAA=="}},
|
||
]
|
||
}
|
||
]
|
||
}
|
||
|
||
result = await logger._strip_base64_from_messages(payload)
|
||
content = result["messages"][0]["content"]
|
||
|
||
# Dropped file-type entries
|
||
assert not any("file" in c for c in content)
|
||
|
||
# Base64 redacted globally
|
||
import json
|
||
|
||
for c in content:
|
||
if isinstance(c, dict):
|
||
s = json.dumps(c).lower()
|
||
# "[base64_redacted]" is fine, but raw base64 is not
|
||
assert "base64," not in s, f"Found real base64 blob in: {s}"
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# Shared fixture that silences asyncio.create_task during tests
|
||
# --------------------------------------------------------------
|
||
@pytest.fixture(autouse=True)
|
||
def patch_asyncio_create_task():
|
||
"""Prevent 'no running event loop' errors when S3Logger calls asyncio.create_task()."""
|
||
with patch("asyncio.create_task"):
|
||
yield
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# Parametrized prefix combination test
|
||
# --------------------------------------------------------------
|
||
@pytest.mark.parametrize(
|
||
"use_team_prefix,use_key_prefix,team_alias,key_alias,expected_prefix",
|
||
[
|
||
(False, False, "teamA", "keyA", ""),
|
||
(True, False, "teamA", "keyA", "teamA/"),
|
||
(False, True, "teamA", "keyA", "keyA/"),
|
||
(True, True, "teamA", "keyA", "teamA/keyA/"),
|
||
(True, True, None, "keyA", "keyA/"),
|
||
(True, True, "teamA", None, "teamA/"),
|
||
(True, True, None, None, ""),
|
||
],
|
||
)
|
||
def test_s3_object_key_prefix_combinations(
|
||
use_team_prefix, use_key_prefix, team_alias, key_alias, expected_prefix
|
||
):
|
||
"""
|
||
Validate correct S3 prefix composition for team alias + key alias combinations.
|
||
"""
|
||
with patch("litellm.integrations.s3_v2.get_s3_object_key") as mock_get_key:
|
||
mock_get_key.return_value = "mocked/s3/object/key.json"
|
||
|
||
logger = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_region_name="us-east-1",
|
||
s3_use_team_prefix=use_team_prefix,
|
||
s3_use_key_prefix=use_key_prefix,
|
||
)
|
||
|
||
payload = StandardLoggingPayload(
|
||
id="abc123",
|
||
metadata={
|
||
"user_api_key_team_alias": team_alias,
|
||
"user_api_key_alias": key_alias,
|
||
},
|
||
messages=[{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||
)
|
||
|
||
result = logger.create_s3_batch_logging_element(datetime.utcnow(), payload)
|
||
assert result is not None
|
||
mock_get_key.assert_called_once()
|
||
|
||
prefix_arg = mock_get_key.call_args.kwargs.get("prefix")
|
||
assert prefix_arg == expected_prefix, (
|
||
f"Expected prefix '{expected_prefix}', got '{prefix_arg}' "
|
||
f"for team={team_alias}, key={key_alias}, "
|
||
f"use_team_prefix={use_team_prefix}, use_key_prefix={use_key_prefix}"
|
||
)
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# Test prefix priority and concatenation
|
||
# --------------------------------------------------------------
|
||
def test_prefix_priority_and_path_construction():
|
||
"""
|
||
Validate that prefix components are ordered and joined with '/' only once.
|
||
"""
|
||
with patch("litellm.integrations.s3_v2.get_s3_object_key") as mock_get_key:
|
||
mock_get_key.return_value = "mocked/key"
|
||
|
||
logger = S3Logger(s3_use_team_prefix=True, s3_use_key_prefix=True)
|
||
payload = StandardLoggingPayload(
|
||
id="xyz999",
|
||
metadata={
|
||
"user_api_key_team_alias": "Team-Alpha",
|
||
"user_api_key_alias": "API-12345",
|
||
},
|
||
messages=[],
|
||
)
|
||
|
||
logger.create_s3_batch_logging_element(datetime.utcnow(), payload)
|
||
prefix_arg = mock_get_key.call_args.kwargs.get("prefix", "")
|
||
|
||
assert prefix_arg == "Team-Alpha/API-12345/"
|
||
assert "//" not in prefix_arg
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# Test when prefixes are disabled
|
||
# --------------------------------------------------------------
|
||
def test_prefix_absent_when_flags_disabled():
|
||
"""
|
||
Verify prefix is omitted entirely when prefix flags are False.
|
||
"""
|
||
with patch("litellm.integrations.s3_v2.get_s3_object_key") as mock_get_key:
|
||
mock_get_key.return_value = "mocked/key"
|
||
|
||
logger = S3Logger(s3_use_team_prefix=False, s3_use_key_prefix=False)
|
||
payload = StandardLoggingPayload(
|
||
id="no-prefix",
|
||
metadata={
|
||
"user_api_key_team_alias": "team-x",
|
||
"user_api_key_alias": "key-x",
|
||
},
|
||
messages=[],
|
||
)
|
||
|
||
logger.create_s3_batch_logging_element(datetime.utcnow(), payload)
|
||
prefix_arg = mock_get_key.call_args.kwargs.get("prefix", None)
|
||
assert prefix_arg == "", f"Expected empty prefix, got {prefix_arg}"
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# Integration-style test (asyncio fixture will patch create_task)
|
||
# --------------------------------------------------------------
|
||
@pytest.mark.asyncio
|
||
async def test_combined_prefix_reflects_in_s3_object_key():
|
||
"""
|
||
Integration-style test ensuring final s3_object_key includes both prefixes correctly.
|
||
"""
|
||
logger = S3Logger(s3_use_team_prefix=True, s3_use_key_prefix=True)
|
||
payload = StandardLoggingPayload(
|
||
id="int-test",
|
||
metadata={
|
||
"user_api_key_team_alias": "myteam",
|
||
"user_api_key_alias": "apikey",
|
||
},
|
||
messages=[],
|
||
)
|
||
|
||
result = logger.create_s3_batch_logging_element(datetime.utcnow(), payload)
|
||
key = result.s3_object_key
|
||
assert "myteam/apikey/" in key, f"Expected both prefixes in key: {key}"
|
||
|
||
|
||
def test_s3_object_key_sanitizes_slashes_in_file_name():
|
||
"""Response ids containing slashes (e.g. bedrock batch job ARNs) must not
|
||
create nested S3 folders; only path/prefix/date slashes are separators."""
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
start_time = datetime(2026, 2, 11, 0, 35, 18, 391582)
|
||
file_name = "time-00-35-18-391582_arn:aws:bedrock:us-east-1:123456789012:model-invocation-job/gl18r6skk9yy"
|
||
|
||
key = get_s3_object_key(
|
||
s3_path="LiteLLMAPPLogs",
|
||
prefix="myteam/",
|
||
start_time=start_time,
|
||
s3_file_name=file_name,
|
||
)
|
||
|
||
assert key == (
|
||
"LiteLLMAPPLogs/myteam/2026-02-11/"
|
||
"time-00-35-18-391582_arn:aws:bedrock:us-east-1:123456789012:model-invocation-job_gl18r6skk9yy.json"
|
||
)
|
||
|
||
|
||
def test_create_s3_batch_logging_element_flat_key_for_arn_response_id():
|
||
"""End-to-end through the s3_v2 element builder: an ARN response id must
|
||
yield a flat file directly under the date segment."""
|
||
logger = S3Logger(s3_use_team_prefix=False, s3_use_key_prefix=False)
|
||
payload = StandardLoggingPayload(
|
||
id="arn:aws:bedrock:us-east-1:123456789012:model-invocation-job/gl18r6skk9yy",
|
||
metadata={},
|
||
messages=[],
|
||
)
|
||
|
||
start_time = datetime(2026, 2, 11, 0, 35, 18, 391582)
|
||
result = logger.create_s3_batch_logging_element(start_time, payload)
|
||
|
||
assert result is not None
|
||
date_segment = "2026-02-11/"
|
||
file_segment = result.s3_object_key.split(date_segment, 1)[1]
|
||
assert "/" not in file_segment, f"Expected flat file under date segment, got: {result.s3_object_key}"
|
||
assert file_segment.endswith("model-invocation-job_gl18r6skk9yy.json")
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# params_source / s3_callback_params_override (audit-log decoupling)
|
||
# --------------------------------------------------------------
|
||
def test_s3_callback_params_override_uses_alternate_dict(monkeypatch):
|
||
"""`s3_callback_params_override` makes the logger read its config from
|
||
the override dict instead of `litellm.s3_callback_params`."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "normal-bucket"})
|
||
logger = S3Logger(
|
||
s3_callback_params_override={
|
||
"s3_bucket_name": "audit-bucket",
|
||
"s3_path": "audit-prefix",
|
||
"s3_region_name": "us-west-2",
|
||
}
|
||
)
|
||
assert logger.s3_bucket_name == "audit-bucket"
|
||
assert logger.s3_path == "audit-prefix"
|
||
assert logger.s3_region_name == "us-west-2"
|
||
|
||
|
||
def test_s3_callback_params_override_does_not_mutate_inputs(monkeypatch):
|
||
"""Resolving `os.environ/X` markers must not mutate the override dict
|
||
or `litellm.s3_callback_params`."""
|
||
import litellm
|
||
|
||
monkeypatch.setenv("MY_AUDIT_BUCKET", "resolved-bucket")
|
||
override = {"s3_bucket_name": "os.environ/MY_AUDIT_BUCKET"}
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "os.environ/MY_AUDIT_BUCKET"})
|
||
logger = S3Logger(s3_callback_params_override=override)
|
||
assert logger.s3_bucket_name == "resolved-bucket"
|
||
assert override["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
|
||
assert (
|
||
litellm.s3_callback_params["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
|
||
)
|
||
|
||
|
||
def test_s3_callback_params_override_none_falls_back_to_global(monkeypatch):
|
||
"""No override → behaves exactly as today (reads `litellm.s3_callback_params`)."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "from-global"})
|
||
logger = S3Logger()
|
||
assert logger.s3_bucket_name == "from-global"
|
||
|
||
|
||
def test_s3_callback_params_override_empty_dict_is_opt_in(monkeypatch):
|
||
"""An empty override dict skips the global entirely (env/IAM-only config)."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "from-global"})
|
||
logger = S3Logger(s3_callback_params_override={})
|
||
assert logger.s3_bucket_name is None
|
||
|
||
|
||
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(monkeypatch):
|
||
"""s3_server_side_encryption can be configured via s3_callback_params."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "from-global",
|
||
"s3_server_side_encryption": "aws:kms",
|
||
},
|
||
)
|
||
logger = S3Logger()
|
||
assert logger.s3_server_side_encryption == "aws:kms"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_sets_sse_kms_key_id_header_when_configured():
|
||
"""
|
||
When s3_sse_kms_key_id is set alongside aws:kms, the PUT must carry
|
||
x-amz-server-side-encryption-aws-kms-key-id so objects are encrypted
|
||
with the customer-managed KMS key instead of the bucket default.
|
||
"""
|
||
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_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-sse-kms.json",
|
||
payload={"test": "sse-kms"},
|
||
s3_object_download_filename="test-sse-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:111122223333:key/test-key-id"
|
||
)
|
||
|
||
|
||
def test_sync_upload_sets_sse_kms_key_id_header_when_configured():
|
||
"""The sync upload path must carry the same SSE-KMS headers."""
|
||
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",
|
||
s3_server_side_encryption="aws:kms",
|
||
s3_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-sync-sse-kms.json",
|
||
payload={"test": "sync-sse-kms"},
|
||
s3_object_download_filename="test-sync-sse-kms.json",
|
||
)
|
||
|
||
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["x-amz-server-side-encryption"] == "aws:kms"
|
||
assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == (
|
||
"arn:aws:kms:us-east-1:111122223333:key/test-key-id"
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_omits_kms_key_id_header_when_not_configured():
|
||
"""SSE without a key id must not emit the KMS key id header."""
|
||
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="AES256",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-aes256.json",
|
||
payload={"test": "aes256"},
|
||
s3_object_download_filename="test-aes256.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"] == "AES256"
|
||
assert "x-amz-server-side-encryption-aws-kms-key-id" not in headers
|
||
|
||
|
||
def test_s3_sse_kms_key_id_read_from_callback_params(monkeypatch):
|
||
"""s3_sse_kms_key_id can be configured via s3_callback_params."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "from-global",
|
||
"s3_server_side_encryption": "aws:kms",
|
||
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
},
|
||
)
|
||
logger = S3Logger()
|
||
assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/test-key-id")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_infers_aws_kms_when_only_key_id_set():
|
||
"""
|
||
Setting only s3_sse_kms_key_id must not produce an invalid request
|
||
(S3 rejects a key id without an algorithm); aws:kms is inferred.
|
||
"""
|
||
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_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
)
|
||
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-kms-only.json",
|
||
payload={"test": "kms-only"},
|
||
s3_object_download_filename="test-kms-only.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:111122223333:key/test-key-id"
|
||
)
|
||
|
||
|
||
def test_s3_sse_kms_key_id_read_from_audit_override_params(monkeypatch):
|
||
"""The audit-log override path must honor s3_sse_kms_key_id too."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "normal-logs-bucket"})
|
||
logger = S3Logger(
|
||
s3_callback_params_override={
|
||
"s3_bucket_name": "audit-logs-bucket",
|
||
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/audit-key-id",
|
||
}
|
||
)
|
||
assert logger.s3_bucket_name == "audit-logs-bucket"
|
||
assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/audit-key-id")
|
||
|
||
|
||
def test_kms_key_id_dropped_when_algorithm_is_not_kms(monkeypatch):
|
||
"""
|
||
AES256 plus a KMS key id is an invalid S3 combination; the key id must be
|
||
dropped at init so uploads keep working instead of silently 400ing.
|
||
"""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "from-global",
|
||
"s3_server_side_encryption": "AES256",
|
||
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
},
|
||
)
|
||
logger = S3Logger()
|
||
assert logger.s3_server_side_encryption == "AES256"
|
||
assert logger.s3_sse_kms_key_id is None
|
||
|
||
|
||
def test_non_string_algorithm_is_dropped_and_valid_key_id_is_rescued(monkeypatch):
|
||
"""
|
||
A YAML boolean in s3_server_side_encryption must not crash logger init and
|
||
must not discard the valid key id; aws:kms is inferred from the key id.
|
||
"""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "from-global",
|
||
"s3_server_side_encryption": True,
|
||
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
|
||
},
|
||
)
|
||
logger = S3Logger()
|
||
assert logger.s3_server_side_encryption == "aws:kms"
|
||
assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/test-key-id")
|
||
|
||
|
||
def test_non_string_key_id_is_dropped_and_valid_algorithm_is_kept(monkeypatch):
|
||
"""A mistyped key id (unquoted YAML number) must not disable the valid algorithm."""
|
||
import litellm
|
||
|
||
monkeypatch.setattr(
|
||
litellm,
|
||
"s3_callback_params",
|
||
{
|
||
"s3_bucket_name": "from-global",
|
||
"s3_server_side_encryption": "aws:kms",
|
||
"s3_sse_kms_key_id": 12345,
|
||
},
|
||
)
|
||
logger = S3Logger()
|
||
assert logger.s3_server_side_encryption == "aws:kms"
|
||
assert logger.s3_sse_kms_key_id is None
|
||
|
||
|
||
_ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
|
||
_SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||
_KEY_WITH_SPACE = "LOGS/LLM AI Projects/2026-08-04/time-13-01-00-abc.json"
|
||
|
||
|
||
def _signature_for(signer_cls, url: str, method: str, body: bytes | None, headers: dict[str, str]) -> str:
|
||
from botocore.awsrequest import AWSRequest
|
||
from botocore.credentials import Credentials
|
||
|
||
sent = {name.lower(): value for name, value in headers.items()}
|
||
signed_header_names = sent["authorization"].split("SignedHeaders=")[1].split(", ")[0].split(";")
|
||
request = AWSRequest(
|
||
method=method,
|
||
url=url,
|
||
data=body,
|
||
headers={name: sent[name] for name in signed_header_names if name in sent},
|
||
)
|
||
request.context["timestamp"] = sent["x-amz-date"]
|
||
signer = signer_cls(Credentials(_ACCESS_KEY, _SECRET_KEY), "s3", "us-east-1")
|
||
canonical_request = signer.canonical_request(request)
|
||
return signer.signature(signer.string_to_sign(request, canonical_request), request)
|
||
|
||
|
||
def _as_s3_canonicalizes(url: str) -> str:
|
||
"""
|
||
The path S3 rebuilds from the wire path: percent-encode everything outside the unreserved
|
||
set, without normalizing or double-encoding. `=` becomes `%3D`, `%20` stays `%20`.
|
||
"""
|
||
from urllib.parse import quote, unquote, urlsplit, urlunsplit
|
||
|
||
split = urlsplit(url)
|
||
return urlunsplit(split._replace(path=quote(unquote(split.path), safe="/~")))
|
||
|
||
|
||
def _assert_signed_for_s3_canonicalization(url: str, method: str, body: bytes | None, headers: dict[str, str]) -> None:
|
||
"""
|
||
S3 rebuilds the canonical request from the wire path with single percent-encoding, which
|
||
botocore models as S3SigV4Auth; plain SigV4Auth double-encodes it (%2520 for a space) and S3
|
||
answers 403 SignatureDoesNotMatch. Assert we sent an already-encoded path and signed it the
|
||
way S3 reads it.
|
||
"""
|
||
from botocore.auth import S3SigV4Auth, SigV4Auth
|
||
|
||
assert url == _as_s3_canonicalizes(url)
|
||
sent_signature = headers["Authorization"].split("Signature=")[1].strip()
|
||
assert sent_signature == _signature_for(S3SigV4Auth, url, method, body, headers)
|
||
assert sent_signature != _signature_for(SigV4Auth, url, method, body, headers)
|
||
|
||
|
||
def _logger_for_signing() -> S3Logger:
|
||
return S3Logger(
|
||
s3_bucket_name="logs-bucket",
|
||
s3_aws_access_key_id=_ACCESS_KEY,
|
||
s3_aws_secret_access_key=_SECRET_KEY,
|
||
s3_region_name="us-east-1",
|
||
)
|
||
|
||
|
||
def _element_with_space():
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
return s3BatchLoggingElement(
|
||
s3_object_key=_KEY_WITH_SPACE,
|
||
payload={"test": "sigv4"},
|
||
s3_object_download_filename="log.json",
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_signs_object_key_with_space_the_way_s3_does():
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
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(_element_with_space())
|
||
|
||
call = logger.async_httpx_client.put.call_args
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="PUT",
|
||
body=call.kwargs["data"].encode("utf-8"),
|
||
headers=call.kwargs["headers"],
|
||
)
|
||
|
||
|
||
def test_sync_upload_signs_object_key_with_space_the_way_s3_does():
|
||
from unittest.mock import MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
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(_element_with_space())
|
||
|
||
call = mock_sync_client.put.call_args
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="PUT",
|
||
body=call.kwargs["data"].encode("utf-8"),
|
||
headers=call.kwargs["headers"],
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_download_signs_object_key_with_space_the_way_s3_does():
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
response = MagicMock()
|
||
response.status_code = 200
|
||
response.json = MagicMock(return_value={"downloaded": "data"})
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.get.return_value = response
|
||
|
||
assert await logger._download_object_from_s3(_KEY_WITH_SPACE) == {"downloaded": "data"}
|
||
|
||
call = logger.async_httpx_client.get.call_args
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="GET",
|
||
body=None,
|
||
headers=call.kwargs["headers"],
|
||
)
|
||
|
||
_RESERVED_CHAR_KEYS = (
|
||
"2026-08-21/time-05-29-36_resp_bGl0ZWxsbTpjdXN0b20=.json",
|
||
"session=logs/2026-08-21/time-05-29-36_abc.json",
|
||
"a+b/2026-08-21/time-05-29-36_abc.json",
|
||
"a&b/2026-08-21/time-05-29-36_abc.json",
|
||
"a#b/2026-08-21/time-05-29-36_abc.json",
|
||
"a?b/2026-08-21/time-05-29-36_abc.json",
|
||
"a%b/2026-08-21/time-05-29-36_abc.json",
|
||
_KEY_WITH_SPACE,
|
||
)
|
||
|
||
|
||
def _element_for(s3_object_key: str):
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
|
||
return s3BatchLoggingElement(
|
||
s3_object_key=s3_object_key,
|
||
payload={"test": "sigv4"},
|
||
s3_object_download_filename="log.json",
|
||
)
|
||
|
||
|
||
def _expected_wire_url(s3_object_key: str) -> str:
|
||
"""The URL boto3 itself would put on the wire for this key."""
|
||
from urllib.parse import quote
|
||
|
||
return f"https://logs-bucket.s3.us-east-1.amazonaws.com/{quote(s3_object_key, safe='/')}"
|
||
|
||
|
||
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_percent_encodes_reserved_characters_in_object_key(s3_object_key):
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
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(_element_for(s3_object_key))
|
||
|
||
call = logger.async_httpx_client.put.call_args
|
||
assert call[0][0] == _expected_wire_url(s3_object_key)
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="PUT",
|
||
body=call.kwargs["data"].encode("utf-8"),
|
||
headers=call.kwargs["headers"],
|
||
)
|
||
|
||
|
||
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
|
||
def test_sync_upload_percent_encodes_reserved_characters_in_object_key(s3_object_key):
|
||
from unittest.mock import MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
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(_element_for(s3_object_key))
|
||
|
||
call = mock_sync_client.put.call_args
|
||
assert call[0][0] == _expected_wire_url(s3_object_key)
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="PUT",
|
||
body=call.kwargs["data"].encode("utf-8"),
|
||
headers=call.kwargs["headers"],
|
||
)
|
||
|
||
|
||
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
|
||
@pytest.mark.asyncio
|
||
async def test_download_percent_encodes_reserved_characters_in_object_key(s3_object_key):
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
logger = _logger_for_signing()
|
||
response = MagicMock()
|
||
response.status_code = 200
|
||
response.json = MagicMock(return_value={"downloaded": "data"})
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.get.return_value = response
|
||
|
||
assert await logger._download_object_from_s3(s3_object_key) == {"downloaded": "data"}
|
||
|
||
call = logger.async_httpx_client.get.call_args
|
||
assert call[0][0] == _expected_wire_url(s3_object_key)
|
||
_assert_signed_for_s3_canonicalization(
|
||
url=call[0][0],
|
||
method="GET",
|
||
body=None,
|
||
headers=call.kwargs["headers"],
|
||
)
|