mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
* ci: run the unit_selection.sh shard files on every event instead of only fork pull requests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: rename fork-flag to unit-flag now that it applies on every event * test: move tests/test_litellm root and small trees into tests/unit Pure renames, no content changes. Follow-up commits in this PR fix references, merge the three files that already existed in tests/unit, keep live-provider tests in tests/test_litellm and wire CI. * test: carry tests/test_litellm conftest isolation into tests/unit Callback lists, routing fallbacks, cached HTTP clients, logger state, AWS, proxy-URL and keychain env, and session-end client cleanup now reset for unit tests too. The environment isolation owns its MonkeyPatch so a test's own monkeypatch is undone before the model-cost teardown runs. * test: merge, split and prune the moved root and small-tree tests Merge batches/test_batch_utils.py and the chat_completions and messages dispatch tests into the files that already existed in tests/unit. Keep the live Gemini interactions tests, the async image-fetch format test and the OpenAI embedding scorer test in tests/test_litellm since they need real network or keys. Put test_router.py under tests/unit/test_router so the existing package no longer shadows it. Delete eight tests the audit found superseded by stronger ones kept in this move. * ci: run the moved root and small-tree tests under their legacy flags Add the misc and responses-caching-types flags to unit_selection.sh and CircleCI, extend enterprise-routing and mcp-integration, and point the legacy GHA shards, Makefile, redis-compat workflow, merge smoke manifest and change classifier at the new paths. * test: make the new tests/unit directories packages tests/unit/test_package_layout.py requires every directory to carry an __init__.py, and without one the moved and retained test_litellm_responses_bridge.py modules collide on import. * test: scope the unit socket block to tests/unit in shared sessions The GHA shards collect the legacy test-path and the unit selection in one pytest session. The unit conftest's loopback-only block leaked into legacy modules that reach the network at import. The legacy conftest now lifts the restriction at collect and setup time, and the unit conftest re-applies it when collecting its own modules. * test: move tests/test_litellm/llms into tests/unit/llms Rename-only. Moves the provider tests and the fine-tuning fixtures they load, mirroring the old paths. Follow-up commits merge, split and wire them. * test: merge, split and prune the moved llms tests Merges the Databricks chat transformation tests into the existing unit file, keeps the tests that need real keys or the network in tests/test_litellm, deletes the audited tests a stronger unit test already covers, and points imports at tests.unit.llms. * ci: run the moved llms tests under their legacy flags The Vertex AI and All Other Providers shards keep their legacy test-path for the retained files and add the llm-vertex-ai and llm-other-providers unit selections. CircleCI gets matching unit jobs. * test: make the tests/unit/llms directories packages Adds __init__.py to the moved dirs and drops the legacy ones whose directories no longer hold tests. * test: drop script runners and path hacks the llms split left dangling The __main__ runners in the split openai_like files and the Databricks e2e runner called tests that now live in the other half of the split or were deleted. The retained legacy halves also no longer need sys.path edits. * test: give the shard-script tests their own GITHUB_OUTPUT They only passed where the runner set it. The CircleCI unit job's env allowlist drops it, so the script's redirect failed there. * test: point the router and module-deletion checks at tests/unit router_code_coverage and code_qa_check_tests only searched tests/test_litellm, so the moved router tests no longer counted. The two silent-experiment tests the audit deleted were the only direct callers of those methods; they are replaced with tests that assert the forwarded shadow request and the recursion guard. * test: move tests/test_litellm integrations and secret_managers into tests/unit Rename-only. Mirrors the old paths, including the directory conftests and the prompt and JSON fixtures. Follow-up commits prune and wire them. * test: prune and repoint the moved integrations tests Deletes the 7 audited tests a stronger test in the same tree already covers, imports the TLS sink helpers from their new conftest path, and restores os.environ after each integrations test. Some presets write OTEL_EXPORTER_OTLP_HEADERS straight into os.environ, and without the legacy tree's test ordering that header leaked into the AgentOps tests. * ci: run the moved integrations tests under their legacy flag The integrations GHA shard and a new CircleCI job run the integrations unit selection. secret_managers joins the misc selection. * docs: point integrations and secret_managers references at tests/unit * test: make the moved integrations directories packages * test: keep the Databricks manual e2e runner and fix the SageMaker Nova run path The Databricks e2e file is a manual script whose main() calls the tests that were pruned, so pruning them broke the documented run. It is back to its main version. The SageMaker Nova docstring now points at the file's real location in tests/local_testing. * test: keep the job's UNIT_FLAG out of the shard-script tests --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2920 lines
104 KiB
Python
2920 lines
104 KiB
Python
import asyncio
|
||
import copy
|
||
import json
|
||
import re
|
||
import sys
|
||
import textwrap
|
||
import uuid
|
||
from collections.abc import Awaitable, Callable
|
||
from contextlib import asynccontextmanager
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||
|
||
import httpx
|
||
import pytest
|
||
import respx
|
||
|
||
from litellm.integrations.s3_v2 import S3Logger
|
||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||
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):
|
||
from unittest.mock import AsyncMock
|
||
|
||
import requests
|
||
|
||
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_400 = MagicMock()
|
||
response_400.status_code = 400
|
||
response_400.raise_for_status = MagicMock(side_effect=Exception("400 Bad Request"))
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = AsyncMock(return_value=response_400)
|
||
|
||
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")
|
||
|
||
|
||
_SIGV4_ACCESS_KEY = re.compile(r"Credential=(AKIA\d+)/")
|
||
|
||
|
||
@pytest.fixture
|
||
def rotating_profile(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> str:
|
||
"""
|
||
A real botocore profile whose credential_process hands out a new key generation on every call and
|
||
expires inside the advisory refresh window, so RefreshableCredentials re-runs it on every property read.
|
||
"""
|
||
counter = tmp_path / "generation"
|
||
script = tmp_path / "rotate_credentials.py"
|
||
script.write_text(
|
||
textwrap.dedent(
|
||
f"""
|
||
import json, sys
|
||
from datetime import datetime, timedelta, timezone
|
||
from pathlib import Path
|
||
|
||
counter = Path({str(counter)!r})
|
||
generation = int(counter.read_text()) if counter.exists() else 0
|
||
counter.write_text(str(generation + 1))
|
||
expiry = (datetime.now(timezone.utc) + timedelta(minutes=12)).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||
json.dump(
|
||
{{
|
||
"Version": 1,
|
||
"AccessKeyId": f"AKIA{{generation}}",
|
||
"SecretAccessKey": f"secret-{{generation}}",
|
||
"SessionToken": f"token-{{generation}}",
|
||
"Expiration": expiry,
|
||
}},
|
||
sys.stdout,
|
||
)
|
||
"""
|
||
)
|
||
)
|
||
profile = f"rotating-{uuid.uuid4().hex}"
|
||
(tmp_path / "config").write_text(f"[profile {profile}]\ncredential_process = {sys.executable} {script}\n")
|
||
monkeypatch.setenv("AWS_CONFIG_FILE", str(tmp_path / "config"))
|
||
return profile
|
||
|
||
|
||
def _generation(request: httpx.Request) -> tuple[str, str]:
|
||
"""(access key generation, session token generation) SigV4 baked into one request."""
|
||
access_key = _SIGV4_ACCESS_KEY.search(request.headers["Authorization"])
|
||
assert access_key is not None
|
||
return access_key.group(1).removeprefix("AKIA"), request.headers["X-Amz-Security-Token"].removeprefix("token-")
|
||
|
||
|
||
@asynccontextmanager
|
||
async def _s3_logger_on_production_handler(profile: str, statuses: list[int]):
|
||
"""
|
||
S3Logger wired to the real AsyncHTTPHandler over an httpx MockTransport that answers with the given
|
||
statuses in order, so the handler's own raise_for_status behaviour is exercised end to end.
|
||
"""
|
||
requests: list[httpx.Request] = []
|
||
replies = iter(statuses)
|
||
|
||
def respond(request: httpx.Request) -> httpx.Response:
|
||
requests.append(request)
|
||
return httpx.Response(next(replies), request=request, text="<Error><Code>SignatureDoesNotMatch</Code></Error>")
|
||
|
||
handler = AsyncHTTPHandler()
|
||
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
|
||
logger = S3Logger(
|
||
s3_bucket_name="test-bucket",
|
||
s3_region_name="us-east-1",
|
||
s3_aws_profile_name=profile,
|
||
s3_flush_interval=3600,
|
||
)
|
||
logger.async_httpx_client = handler
|
||
with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
|
||
yield logger, requests, mock_sleep
|
||
await handler.client.aclose()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_signs_with_one_frozen_credential_snapshot(rotating_profile: str, caplog):
|
||
"""
|
||
RefreshableCredentials refreshes on every property read once inside the advisory window, so signing
|
||
off the live object would mix the access key of one generation with the token of the next.
|
||
"""
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-frozen.json",
|
||
payload={"test": "frozen"},
|
||
s3_object_download_filename="test-frozen.json",
|
||
)
|
||
async with _s3_logger_on_production_handler(rotating_profile, [200]) as (logger, requests, _):
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
assert len(requests) == 1
|
||
access_key_generation, token_generation = _generation(requests[0])
|
||
assert access_key_generation == token_generation
|
||
assert "Error uploading to s3" not in caplog.text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_retries_403_with_fresh_credentials_and_signature(rotating_profile: str, caplog):
|
||
"""
|
||
A 403 (SignatureDoesNotMatch after an IMDS rotation) must be retried, and the retry must fetch
|
||
credentials again and carry a signature computed from that newer generation.
|
||
"""
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-403.json",
|
||
payload={"test": "403"},
|
||
s3_object_download_filename="test-403.json",
|
||
)
|
||
async with _s3_logger_on_production_handler(rotating_profile, [403, 200]) as (logger, requests, mock_sleep):
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
assert len(requests) == 2
|
||
first_key, first_token = _generation(requests[0])
|
||
second_key, second_token = _generation(requests[1])
|
||
assert first_key == first_token
|
||
assert second_key == second_token
|
||
assert int(second_key) > int(first_key)
|
||
assert requests[1].headers["Authorization"] != requests[0].headers["Authorization"]
|
||
mock_sleep.assert_awaited_once_with(1)
|
||
assert "Error uploading to s3" not in caplog.text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_exhausts_403_retries_through_production_http_handler(rotating_profile: str, caplog):
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-403-exhausted.json",
|
||
payload={"test": "403-exhausted"},
|
||
s3_object_download_filename="test-403-exhausted.json",
|
||
)
|
||
async with _s3_logger_on_production_handler(rotating_profile, [403, 403, 403]) as (logger, requests, mock_sleep):
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
assert len(requests) == 3
|
||
assert mock_sleep.await_args_list == [call(1), call(2)]
|
||
assert "Error uploading to s3" in caplog.text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_upload_does_not_retry_404_through_production_http_handler(rotating_profile: str, caplog):
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-404.json",
|
||
payload={"test": "404"},
|
||
s3_object_download_filename="test-404.json",
|
||
)
|
||
async with _s3_logger_on_production_handler(rotating_profile, [404]) as (logger, requests, mock_sleep):
|
||
await logger.async_upload_data_to_s3(test_element)
|
||
|
||
assert len(requests) == 1
|
||
mock_sleep.assert_not_awaited()
|
||
assert "Error uploading to s3" in caplog.text
|
||
|
||
|
||
def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, monkeypatch: pytest.MonkeyPatch):
|
||
monkeypatch.setenv("AWS_PROFILE", rotating_profile)
|
||
logger = S3Logger(s3_bucket_name="test-bucket", s3_region_name="us-east-1", s3_flush_interval=3600)
|
||
test_element = s3BatchLoggingElement(
|
||
s3_object_key="2025-09-14/test-sync-403.json",
|
||
payload={"test": "sync-403"},
|
||
s3_object_download_filename="test-sync-403.json",
|
||
)
|
||
requests: list[httpx.Request] = []
|
||
replies = iter([403, 200])
|
||
|
||
def respond(request: httpx.Request) -> httpx.Response:
|
||
requests.append(request)
|
||
return httpx.Response(next(replies), request=request)
|
||
|
||
handler = HTTPHandler()
|
||
handler.client = httpx.Client(transport=httpx.MockTransport(respond))
|
||
with (
|
||
patch( # test-quality-ok: sync upload builds its HTTPHandler per call, there is no injection seam for it
|
||
"litellm.integrations.s3_v2._get_httpx_client", return_value=handler
|
||
),
|
||
patch("time.sleep") as mock_sleep,
|
||
):
|
||
logger.upload_data_to_s3(test_element)
|
||
|
||
assert len(requests) == 2
|
||
first_key, first_token = _generation(requests[0])
|
||
second_key, second_token = _generation(requests[1])
|
||
assert first_key == first_token
|
||
assert second_key == second_token
|
||
assert int(second_key) > int(first_key)
|
||
mock_sleep.assert_called_once_with(1)
|
||
|
||
|
||
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_and_colons_in_file_name():
|
||
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"
|
||
)
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"response_id",
|
||
[
|
||
"s3://example-batch-bucket/litellm-bedrock-files/input.jsonl",
|
||
"gs://example-batch-bucket/litellm-vertex-files/input.jsonl",
|
||
],
|
||
)
|
||
def test_s3_object_key_has_no_colon_for_cloud_uri_file_ids(response_id: str):
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
key = get_s3_object_key(
|
||
s3_path="",
|
||
prefix="",
|
||
start_time=datetime(2026, 9, 7, 4, 51, 6, 685889),
|
||
s3_file_name=f"time-04-51-06-685889_{response_id}",
|
||
)
|
||
|
||
filename = key.rsplit("/", 1)[-1]
|
||
assert ":" not in filename
|
||
assert filename.endswith("_input.jsonl.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")
|
||
|
||
|
||
# --------------------------------------------------------------
|
||
# object keys bounded to S3's 1024 UTF-8 byte limit
|
||
# --------------------------------------------------------------
|
||
def _oversized_response_id() -> str:
|
||
return "resp_" + "A" * 1100
|
||
|
||
|
||
def test_s3_object_key_at_the_byte_limit_is_left_alone():
|
||
"""A key that still fits is left byte-identical."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
start_time = datetime(2026, 8, 24, 6, 18, 41, 948021)
|
||
fixed_len = len("input/2026-08-24/.json")
|
||
file_name = "x" * (MAX_S3_OBJECT_KEY_BYTES - fixed_len)
|
||
|
||
key = get_s3_object_key(s3_path="input", prefix="", start_time=start_time, s3_file_name=file_name)
|
||
|
||
assert key == f"input/2026-08-24/{file_name}.json"
|
||
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
|
||
|
||
|
||
def test_s3_object_key_is_bounded_for_oversized_response_id():
|
||
"""An oversized Responses API id is shortened to a readable head plus a digest."""
|
||
import hashlib
|
||
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
start_time = datetime(2026, 8, 24, 6, 18, 41, 948021)
|
||
file_name = f"time-06-18-41-948021_{_oversized_response_id()}"
|
||
|
||
key = get_s3_object_key(s3_path="input", prefix="DefaultTeamProd/", start_time=start_time, s3_file_name=file_name)
|
||
|
||
assert len(key.encode("utf-8")) <= MAX_S3_OBJECT_KEY_BYTES
|
||
assert key.startswith("input/DefaultTeamProd/2026-08-24/time-06-18-41-948021_resp_")
|
||
assert key.endswith(f"_{hashlib.sha256(file_name.encode('utf-8')).hexdigest()}.json")
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"s3_path,prefix",
|
||
[
|
||
("input", ""),
|
||
("a" * 900, ""),
|
||
("input", "team-" + "b" * 900 + "/"),
|
||
("c" * 600, "team-" + "d" * 600 + "/key-" + "e" * 600 + "/"),
|
||
# many short segments, so the trim lands exactly on the budget edge
|
||
("", "ssss/" * 200),
|
||
],
|
||
)
|
||
def test_s3_object_key_is_bounded_for_long_paths_and_aliases(s3_path: str, prefix: str):
|
||
"""Long paths, team aliases and key aliases stay within the cap."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
key = get_s3_object_key(
|
||
s3_path=s3_path,
|
||
prefix=prefix,
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name=f"time-06-18-41-948021_{_oversized_response_id()}",
|
||
)
|
||
|
||
assert len(key.encode("utf-8")) <= MAX_S3_OBJECT_KEY_BYTES
|
||
assert key.endswith(".json")
|
||
assert "/2026-08-24/" in key or key.startswith("2026-08-24/")
|
||
assert "/" not in key.rsplit("2026-08-24/", 1)[1]
|
||
|
||
|
||
def test_s3_object_key_trimmed_prefixes_stay_distinct_per_operator():
|
||
"""Prefixes that differ only past the trim point keep separate folders."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
start_time = datetime(2026, 8, 24, 6, 18, 41, 948021)
|
||
keys = [
|
||
get_s3_object_key(
|
||
s3_path="input",
|
||
prefix="team-" + "b" * 1000 + suffix + "/",
|
||
start_time=start_time,
|
||
s3_file_name=f"time-06-18-41-948021_{_oversized_response_id()}",
|
||
)
|
||
for suffix in ("-one", "-two")
|
||
]
|
||
|
||
assert keys[0] != keys[1]
|
||
assert all(key.startswith("input/team-" + "b" * 900) for key in keys)
|
||
assert all(len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES for key in keys)
|
||
|
||
|
||
def test_s3_object_key_bounded_prefix_never_splits_a_multibyte_character():
|
||
"""A multibyte prefix is trimmed on a character boundary."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
s3_path = "\u65e5\u672c\u8a9e" * 200
|
||
|
||
key = get_s3_object_key(
|
||
s3_path=s3_path,
|
||
prefix="\u30c1\u30fc\u30e0" * 200 + "/",
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name=f"time-06-18-41-948021_{_oversized_response_id()}",
|
||
)
|
||
|
||
assert len(key.encode("utf-8")) <= MAX_S3_OBJECT_KEY_BYTES
|
||
assert key.startswith(s3_path[:100])
|
||
assert "\ufffd" not in key
|
||
|
||
|
||
def test_s3_object_key_stays_unique_for_ids_sharing_a_head():
|
||
"""Ids sharing a visible head still get distinct keys."""
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
start_time = datetime(2026, 8, 24, 6, 18, 41, 948021)
|
||
keys = {
|
||
get_s3_object_key(
|
||
s3_path="input",
|
||
prefix="",
|
||
start_time=start_time,
|
||
s3_file_name=f"time-06-18-41-948021_{_oversized_response_id()}{suffix}",
|
||
)
|
||
for suffix in ("first", "second", "third")
|
||
}
|
||
|
||
assert len(keys) == 3
|
||
|
||
|
||
def test_s3_object_key_bounding_matches_the_documented_layout():
|
||
"""The bounded key is `<prefix>/<date>/<head>_<sha256>.json`."""
|
||
import hashlib
|
||
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
file_name = f"time-06-18-41-948021_{_oversized_response_id()}"
|
||
|
||
key = get_s3_object_key(
|
||
s3_path="input",
|
||
prefix="team/",
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name=file_name,
|
||
)
|
||
|
||
digest = hashlib.sha256(file_name.encode("utf-8")).hexdigest()
|
||
assert key == f"input/team/2026-08-24/{file_name[:64]}_{digest}.json"
|
||
|
||
|
||
def test_s3_object_key_keeps_the_configured_prefix_when_only_the_id_overflows():
|
||
"""A 940 byte configured prefix survives whole when only the id overflows."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
prefix = "team-" + "b" * 934 + "/"
|
||
|
||
key = get_s3_object_key(
|
||
s3_path="",
|
||
prefix=prefix,
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name=f"time-06-18-41-948021_{_oversized_response_id()}",
|
||
)
|
||
|
||
assert key.startswith(prefix + "2026-08-24/")
|
||
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
|
||
|
||
|
||
def test_s3_object_key_spends_the_whole_budget_when_the_prefix_must_be_trimmed():
|
||
"""A trimmed prefix keeps every byte the budget allows, not whole segments."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
s3_path = "p" * 400 + "/" + "q" * 600
|
||
|
||
key = get_s3_object_key(
|
||
s3_path=s3_path,
|
||
prefix="",
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name="time-06-18-41-948021_abc",
|
||
)
|
||
|
||
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
|
||
assert key.startswith("p" * 400 + "/" + "q" * 500)
|
||
|
||
|
||
def test_s3_object_key_keeps_a_single_segment_path_as_far_as_it_fits():
|
||
"""A path with no separator is kept as far as it fits, never dropped to the bucket root."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_key
|
||
|
||
key = get_s3_object_key(
|
||
s3_path="a" * 1050,
|
||
prefix="",
|
||
start_time=datetime(2026, 8, 24, 6, 18, 41, 948021),
|
||
s3_file_name="time-06-18-41-948021_chatcmpl-xyz",
|
||
)
|
||
|
||
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
|
||
assert key.startswith("a" * 900)
|
||
|
||
|
||
def test_create_s3_batch_logging_element_bounds_key_and_keeps_full_response_id():
|
||
"""The batch element bounds the key and keeps the full response id in the payload."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
|
||
logger = S3Logger(s3_use_team_prefix=True, s3_use_key_prefix=True)
|
||
response_id = _oversized_response_id()
|
||
payload = StandardLoggingPayload(
|
||
id=response_id,
|
||
metadata={"user_api_key_team_alias": "DefaultTeamProd", "user_api_key_alias": "prod-key"},
|
||
messages=[],
|
||
)
|
||
|
||
result = logger.create_s3_batch_logging_element(datetime(2026, 8, 24, 6, 18, 41, 948021), payload)
|
||
|
||
assert result is not None
|
||
assert len(result.s3_object_key.encode("utf-8")) <= MAX_S3_OBJECT_KEY_BYTES
|
||
assert result.s3_object_key.startswith("DefaultTeamProd/prod-key/2026-08-24/")
|
||
assert result.payload["id"] == response_id
|
||
|
||
|
||
def test_s3_object_download_filename_is_bounded_for_oversized_response_id():
|
||
"""The Content-Disposition filename is bounded too, or the PUT fails with MetadataTooLarge."""
|
||
from litellm.constants import MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES
|
||
from litellm.integrations.s3 import get_s3_object_download_filename
|
||
|
||
file_name = get_s3_object_download_filename(datetime(2026, 8, 24, 6, 18, 41, 948021), _oversized_response_id())
|
||
|
||
assert len(file_name.encode("utf-8")) <= MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES
|
||
assert file_name.startswith("time-2026-08-24T06-18-41-948021_resp_")
|
||
assert file_name.endswith(".json")
|
||
|
||
|
||
def test_s3_object_download_filenames_stay_distinct_when_shortened():
|
||
"""Shortened filenames stay distinct."""
|
||
from litellm.integrations.s3 import get_s3_object_download_filename
|
||
|
||
start_time = datetime(2026, 8, 24, 6, 18, 41, 948021)
|
||
file_names = {
|
||
get_s3_object_download_filename(start_time, _oversized_response_id() + suffix)
|
||
for suffix in ("first", "second", "third")
|
||
}
|
||
|
||
assert len(file_names) == 3
|
||
|
||
|
||
def test_s3_object_download_filename_short_id_is_unchanged():
|
||
"""An ordinary response id keeps the filename it had before."""
|
||
from litellm.integrations.s3 import get_s3_object_download_filename
|
||
|
||
file_name = get_s3_object_download_filename(datetime(2026, 8, 24, 6, 18, 41, 948021), "resp_abc123")
|
||
|
||
assert file_name == "time-2026-08-24T06-18-41-948021_resp_abc123.json"
|
||
|
||
|
||
def test_create_s3_batch_logging_element_bounds_the_download_filename():
|
||
"""The batch element carries a bounded Content-Disposition filename."""
|
||
from litellm.constants import MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES
|
||
|
||
logger = S3Logger()
|
||
payload = StandardLoggingPayload(id=_oversized_response_id(), metadata={}, messages=[])
|
||
|
||
result = logger.create_s3_batch_logging_element(datetime(2026, 8, 24, 6, 18, 41, 948021), payload)
|
||
|
||
assert result is not None
|
||
assert len(result.s3_object_download_filename.encode("utf-8")) <= MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_audit_log_object_key_is_bounded_for_a_long_configured_path():
|
||
"""Audit log keys are bounded by the same builder."""
|
||
from litellm.constants import MAX_S3_OBJECT_KEY_BYTES
|
||
|
||
logger = S3Logger()
|
||
logger.s3_path = "audit-archive/" + "z" * 1100
|
||
|
||
await logger.async_log_audit_log_event({"id": "1a4f7bd0-6f1e-4d0a-9b3c-9f2e1d5a7c88"})
|
||
|
||
assert len(logger.log_queue) == 1
|
||
assert len(logger.log_queue[0].s3_object_key.encode("utf-8")) <= MAX_S3_OBJECT_KEY_BYTES
|
||
assert logger.log_queue[0].s3_object_key.startswith("audit-archive/" + "z" * 900)
|
||
|
||
|
||
def test_s3_object_download_filename_drops_characters_that_break_the_header():
|
||
"""A quote or separator in the response id cannot escape the quoted header value."""
|
||
from litellm.integrations.s3 import get_s3_object_download_filename
|
||
|
||
file_name = get_s3_object_download_filename(datetime(2026, 8, 24, 6, 18, 41, 948021), 'resp_a"b/c')
|
||
|
||
assert file_name == "time-2026-08-24T06-18-41-948021_resp_a_b_c.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"],
|
||
)
|
||
|
||
|
||
def _s3_logger_for_region(region_name: str) -> S3Logger:
|
||
logger = S3Logger.__new__(S3Logger)
|
||
logger.s3_endpoint_url = None
|
||
logger.s3_bucket_name = "my-litellm-audit"
|
||
logger.s3_region_name = region_name
|
||
return logger
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"region_name,expected_url",
|
||
[
|
||
(
|
||
"cn-northwest-1",
|
||
"https://my-litellm-audit.s3.cn-northwest-1.amazonaws.com.cn/2025-01-01/key.json",
|
||
),
|
||
(
|
||
"us-gov-west-1",
|
||
"https://my-litellm-audit.s3.us-gov-west-1.amazonaws.com/2025-01-01/key.json",
|
||
),
|
||
(
|
||
"us-east-1",
|
||
"https://my-litellm-audit.s3.us-east-1.amazonaws.com/2025-01-01/key.json",
|
||
),
|
||
],
|
||
)
|
||
def test_build_object_url_uses_partition_dns_suffix(region_name: str, expected_url: str) -> None:
|
||
assert _s3_logger_for_region(region_name)._build_object_url("2025-01-01/key.json") == expected_url
|
||
|
||
|
||
def _prompts_only_logger(s3_log_prompts_only: bool | None = None) -> S3Logger:
|
||
return 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_log_prompts_only=s3_log_prompts_only,
|
||
)
|
||
|
||
|
||
def _chat_payload() -> StandardLoggingPayload:
|
||
return StandardLoggingPayload(
|
||
id="chatcmpl-prompts-only",
|
||
messages=[{"role": "user", "content": "Reply with exactly the word PINEAPPLE."}],
|
||
response={"choices": [{"message": {"role": "assistant", "content": "PINEAPPLE"}}]},
|
||
metadata={"user_api_key_team_alias": None},
|
||
)
|
||
|
||
|
||
async def _queued_body_via_async_upload(
|
||
logger: S3Logger, log_event: Callable[..., Awaitable[None]]
|
||
) -> dict[str, object]:
|
||
payload = _chat_payload()
|
||
original = copy.deepcopy(payload)
|
||
await log_event(
|
||
kwargs={"standard_logging_object": payload},
|
||
response_obj=None,
|
||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||
end_time=datetime(2026, 7, 30, 12, 0, 1),
|
||
)
|
||
assert payload == original, "the caller's standard_logging_object must not be mutated"
|
||
(element,) = logger.log_queue
|
||
|
||
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)
|
||
return json.loads(logger.async_httpx_client.put.call_args.kwargs["data"])
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.parametrize("event_name", ["async_log_success_event", "async_log_failure_event"])
|
||
async def test_prompts_only_drops_response_but_keeps_messages_in_uploaded_object(
|
||
monkeypatch: pytest.MonkeyPatch, event_name: str
|
||
):
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": True})
|
||
logger = _prompts_only_logger()
|
||
|
||
log_event: Callable[..., Awaitable[None]] = (
|
||
logger.async_log_success_event if event_name == "async_log_success_event" else logger.async_log_failure_event
|
||
)
|
||
body = await _queued_body_via_async_upload(logger, log_event)
|
||
|
||
assert body["messages"] == _chat_payload()["messages"]
|
||
assert body["response"] is None
|
||
assert body["id"] == "chatcmpl-prompts-only"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_prompts_only_default_off_keeps_response_in_uploaded_object(monkeypatch: pytest.MonkeyPatch):
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||
logger = _prompts_only_logger()
|
||
|
||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||
|
||
assert body["response"] == _chat_payload()["response"]
|
||
assert body["messages"] == _chat_payload()["messages"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_prompts_only_explicit_false_in_params_beats_env_var(monkeypatch: pytest.MonkeyPatch):
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": False})
|
||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||
logger = _prompts_only_logger()
|
||
|
||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||
|
||
assert body["response"] == _chat_payload()["response"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_prompts_only_env_var_applies_when_param_unset(monkeypatch: pytest.MonkeyPatch):
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||
logger = _prompts_only_logger()
|
||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||
|
||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||
|
||
assert body["response"] is None
|
||
assert body["messages"] == _chat_payload()["messages"]
|
||
|
||
|
||
@respx.mock
|
||
def test_prompts_only_constructor_kwarg_applies_to_sync_upload(monkeypatch: pytest.MonkeyPatch):
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||
logger = _prompts_only_logger(s3_log_prompts_only=True)
|
||
payload = _chat_payload()
|
||
|
||
element = logger.create_s3_batch_logging_element(
|
||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||
standard_logging_payload=payload,
|
||
)
|
||
assert element is not None
|
||
assert payload["response"] == _chat_payload()["response"]
|
||
|
||
put_route = respx.put(url__regex=r"https://test-bucket\.s3\..*").mock(return_value=httpx.Response(200))
|
||
logger.upload_data_to_s3(element)
|
||
|
||
body = json.loads(put_route.calls.last.request.content)
|
||
assert body["response"] is None
|
||
assert body["messages"] == _chat_payload()["messages"]
|
||
|
||
|
||
@pytest.mark.parametrize("callback_name", ["s3", "s3_v2"])
|
||
def test_prompts_only_toggle_is_exposed_to_admin_ui_for_both_s3_callbacks(callback_name: str):
|
||
from litellm.integrations.custom_logger import CustomLogger
|
||
|
||
assert "S3_LOG_PROMPTS_ONLY" in CustomLogger.get_callback_env_vars(callback_name)
|
||
|
||
|
||
def _element(payload: dict[str, object], key_suffix: str) -> s3BatchLoggingElement:
|
||
return s3BatchLoggingElement(
|
||
s3_object_key=f"2025-09-14/test-{key_suffix}.json",
|
||
payload=payload,
|
||
s3_object_download_filename=f"test-{key_suffix}.json",
|
||
)
|
||
|
||
|
||
def _ok_response() -> MagicMock:
|
||
response = MagicMock()
|
||
response.status_code = 200
|
||
response.raise_for_status = MagicMock()
|
||
return response
|
||
|
||
|
||
class _CountingPut:
|
||
def __init__(self) -> None:
|
||
self.in_flight = 0
|
||
self.peak = 0
|
||
self.calls = 0
|
||
|
||
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
|
||
self.in_flight += 1
|
||
self.peak = max(self.peak, self.in_flight)
|
||
self.calls += 1
|
||
await asyncio.sleep(0.01)
|
||
self.in_flight -= 1
|
||
return _ok_response()
|
||
|
||
|
||
class _RecordingPut:
|
||
def __init__(self) -> None:
|
||
self.calls: tuple[tuple[str, str | None, dict[str, str] | None], ...] = ()
|
||
|
||
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
|
||
self.calls = (*self.calls, (url, data, headers))
|
||
return _ok_response()
|
||
|
||
|
||
class _LateAppendingPut:
|
||
def __init__(self, logger: S3Logger, element: s3BatchLoggingElement, fail_first: bool = False) -> None:
|
||
self.logger = logger
|
||
self.element = element
|
||
self.fail_first = fail_first
|
||
self.appended = False
|
||
|
||
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
|
||
if not self.appended:
|
||
self.appended = True
|
||
self.logger.log_queue.append(self.element)
|
||
if self.fail_first:
|
||
return _failure_response()
|
||
return _ok_response()
|
||
|
||
|
||
class _FailOnSuffixPut:
|
||
def __init__(self, suffixes: tuple[str, ...]) -> None:
|
||
self.failing = True
|
||
self.suffixes = suffixes
|
||
|
||
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
|
||
if self.failing and url.endswith(self.suffixes):
|
||
return _failure_response()
|
||
return _ok_response()
|
||
|
||
|
||
class _FailUntilClearedPut:
|
||
def __init__(self) -> None:
|
||
self.failing = True
|
||
self.calls: tuple[tuple[str, str | None], ...] = ()
|
||
|
||
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
|
||
self.calls = (*self.calls, (url, data))
|
||
if self.failing:
|
||
return _failure_response()
|
||
return _ok_response()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_send_batch_bounds_concurrent_uploads() -> None:
|
||
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_max_concurrent_uploads=4,
|
||
)
|
||
|
||
put = _CountingPut()
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
|
||
logger.log_queue = [_element({"i": i}, f"{i}") for i in range(40)]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
assert put.peak == 4
|
||
assert put.calls == 40
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_send_batch_uploads_single_jsonl_file() -> None:
|
||
import json
|
||
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _RecordingPut()
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
|
||
payloads = [{"id": "req-1"}, {"id": "req-2"}, {"id": "req-3"}]
|
||
logger.log_queue = [_element(payload, f"{i}") for i, payload in enumerate(payloads)]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
assert len(put.calls) == 1
|
||
url, data, headers = put.calls[0]
|
||
assert url.endswith(".jsonl")
|
||
assert data is not None
|
||
assert headers is not None
|
||
assert [json.loads(line) for line in data.splitlines()] == payloads
|
||
assert headers["Content-Type"] == "application/x-ndjson"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_flush_queue_preserves_events_added_during_upload() -> None:
|
||
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",
|
||
)
|
||
|
||
late_element = _element({"id": "late"}, "late")
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = _LateAppendingPut(logger, late_element)
|
||
|
||
logger.log_queue = [_element({"id": "first"}, "first")]
|
||
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == [late_element]
|
||
|
||
|
||
def _override_logger(**overrides: object) -> S3Logger:
|
||
return 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_callback_params_override=overrides,
|
||
)
|
||
|
||
|
||
def test_env_backed_false_string_keeps_per_request_uploads() -> None:
|
||
assert _override_logger(s3_batch_file_upload="false").s3_batch_file_upload is False
|
||
assert _override_logger(s3_batch_file_upload="true").s3_batch_file_upload is True
|
||
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_batch_file_upload=True,
|
||
s3_callback_params_override={"s3_batch_file_upload": "false"},
|
||
)
|
||
assert logger.s3_batch_file_upload is True
|
||
|
||
|
||
@pytest.mark.parametrize("bad", [0, -3, "0", "abc", ""])
|
||
def test_invalid_concurrency_falls_back_to_default(bad: object) -> None:
|
||
from litellm.constants import DEFAULT_S3_MAX_CONCURRENT_UPLOADS
|
||
|
||
logger = _override_logger(s3_max_concurrent_uploads=bad)
|
||
|
||
assert logger.s3_max_concurrent_uploads == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
|
||
assert logger._upload_semaphore._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
|
||
|
||
|
||
def test_env_backed_concurrency_string_is_parsed() -> None:
|
||
logger = _override_logger(s3_max_concurrent_uploads="4")
|
||
|
||
assert logger.s3_max_concurrent_uploads == 4
|
||
assert logger._upload_semaphore._value == 4
|
||
|
||
|
||
@pytest.mark.parametrize("empty", [None, ""])
|
||
def test_empty_config_concurrency_falls_back_to_constructor_value(empty: object) -> None:
|
||
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_max_concurrent_uploads=4,
|
||
s3_callback_params_override={"s3_max_concurrent_uploads": empty},
|
||
)
|
||
|
||
assert logger.s3_max_concurrent_uploads == 4
|
||
assert logger._upload_semaphore._value == 4
|
||
|
||
|
||
def _failure_response() -> MagicMock:
|
||
response = MagicMock()
|
||
response.status_code = 400
|
||
response.raise_for_status = MagicMock(side_effect=Exception("s3 rejected the object"))
|
||
return response
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_failed_uploads_stay_queued_for_next_flush() -> None:
|
||
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",
|
||
)
|
||
|
||
elements = [_element({"i": i}, f"{i}") for i in range(5)]
|
||
put = _FailOnSuffixPut(("test-2.json", "test-4.json"))
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
logger.log_queue = list(elements)
|
||
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == [elements[2], elements[4]]
|
||
|
||
put.failing = False
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == []
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_file_upload_failure_keeps_whole_batch() -> None:
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _FailUntilClearedPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
|
||
elements = [_element({"i": i}, f"{i}") for i in range(3)]
|
||
logger.log_queue = list(elements)
|
||
|
||
await logger.flush_queue()
|
||
|
||
assert len(put.calls) == 1
|
||
assert len(logger.log_queue) == 1
|
||
assert logger.log_queue[0].body == "\n".join(json.dumps(element.payload) for element in elements)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_events_appended_during_failed_flush_survive() -> None:
|
||
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",
|
||
)
|
||
|
||
late = _element({"id": "late"}, "late")
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = _LateAppendingPut(logger, late, fail_first=True)
|
||
|
||
first = _element({"id": "first"}, "first")
|
||
logger.log_queue = [first]
|
||
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == [first, late]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_file_key_shape() -> None:
|
||
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_path="logs",
|
||
s3_batch_file_upload=True,
|
||
)
|
||
|
||
put = _RecordingPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
logger.log_queue = [_element({"id": "req-1"}, "0")]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
((url, _data, headers),) = put.calls
|
||
assert headers is not None
|
||
assert re.search(r".*/2025-09-14/batch_\d{2}-\d{2}-\d{2}_[0-9a-f]{32}\.jsonl$", url)
|
||
assert headers["Content-Disposition"].endswith('.jsonl"')
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_file_groups_raw_elements_by_key_parent() -> None:
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _RecordingPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
|
||
alpha = s3BatchLoggingElement(
|
||
s3_object_key="logs/alpha/2026-01-01/a.json", payload={"id": "a"}, s3_object_download_filename="a.json"
|
||
)
|
||
beta = s3BatchLoggingElement(
|
||
s3_object_key="logs/beta/2026-01-01/b.json", payload={"id": "b"}, s3_object_download_filename="b.json"
|
||
)
|
||
plain = s3BatchLoggingElement(
|
||
s3_object_key="logs/2026-01-01/c.json", payload={"id": "c"}, s3_object_download_filename="c.json"
|
||
)
|
||
root = s3BatchLoggingElement(
|
||
s3_object_key="solo.json", payload={"id": "d"}, s3_object_download_filename="solo.json"
|
||
)
|
||
logger.log_queue = [alpha, beta, plain, root]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
assert len(put.calls) == 4
|
||
by_parent = {
|
||
re.sub(r"(^|/)batch_\d{2}-\d{2}-\d{2}_[0-9a-f]{32}\.jsonl$", "", url.split(".com/", 1)[-1]): (url, data)
|
||
for url, data, _headers in put.calls
|
||
}
|
||
assert sorted(by_parent) == ["", "logs/2026-01-01", "logs/alpha/2026-01-01", "logs/beta/2026-01-01"]
|
||
assert [line for line in by_parent[""][1].splitlines()] == [json.dumps({"id": "d"})]
|
||
assert [line for line in by_parent["logs/alpha/2026-01-01"][1].splitlines()] == [json.dumps({"id": "a"})]
|
||
assert [line for line in by_parent["logs/beta/2026-01-01"][1].splitlines()] == [json.dumps({"id": "b"})]
|
||
assert [line for line in by_parent["logs/2026-01-01"][1].splitlines()] == [json.dumps({"id": "c"})]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None:
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _FailUntilClearedPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
logger.log_queue = [_element({"i": i}, f"{i}") for i in range(3)]
|
||
|
||
await logger.flush_queue()
|
||
|
||
assert len(logger.log_queue) == 1
|
||
assert logger.log_queue[0].body is not None
|
||
assert logger.log_queue[0].s3_object_key.endswith(".jsonl")
|
||
|
||
put.failing = False
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == []
|
||
assert len(put.calls) == 2
|
||
assert put.calls[0] == put.calls[1]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_elements_appended_after_failed_batch_file_get_their_own_file() -> None:
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _FailUntilClearedPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
logger.log_queue = [_element({"id": "first"}, "first")]
|
||
|
||
await logger.flush_queue()
|
||
|
||
late = _element({"id": "late"}, "late")
|
||
logger.log_queue.append(late)
|
||
|
||
put.failing = False
|
||
await logger.flush_queue()
|
||
|
||
assert logger.log_queue == []
|
||
assert len(put.calls) == 3
|
||
assert put.calls[0] == put.calls[1]
|
||
assert put.calls[2][0] != put.calls[0][0]
|
||
assert put.calls[2][1] == json.dumps({"id": "late"})
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_file_mode_disabled_when_s3_v2_is_cold_storage_logger(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
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_batch_file_upload=True,
|
||
)
|
||
|
||
put = _RecordingPut()
|
||
|
||
logger.async_httpx_client = AsyncMock()
|
||
logger.async_httpx_client.put = put
|
||
|
||
import litellm
|
||
|
||
monkeypatch.setattr(litellm, "cold_storage_custom_logger", "s3_v2")
|
||
logger.log_queue = [_element({"id": "req-1"}, "0")]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
assert len(put.calls) == 1
|
||
assert put.calls[0][0].endswith("test-0.json")
|
||
|
||
monkeypatch.setattr(litellm, "cold_storage_custom_logger", None)
|
||
logger.log_queue = [_element({"id": "req-2"}, "1")]
|
||
|
||
await logger.async_send_batch()
|
||
|
||
assert len(put.calls) == 2
|
||
assert put.calls[1][0].endswith(".jsonl")
|