fix(s3): await batch uploads before clearing queue

This commit is contained in:
Ganni Galea Curmi 2026-05-26 20:20:28 -04:00
parent b6fd7f7746
commit d8f57780f6
2 changed files with 155 additions and 5 deletions

View file

@ -3,7 +3,9 @@ s3 Bucket Logging Integration
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to upload each element individually
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so each queued
element is uploaded individually and the batch flush awaits all uploads
before clearing the queue
"""
import asyncio
@ -30,6 +32,8 @@ from .custom_batch_logger import CustomBatchLogger
class S3Logger(CustomBatchLogger, BaseAWSLLM):
preserve_events_added_during_flush = True
def __init__(
self,
s3_bucket_name: Optional[str] = None,
@ -318,7 +322,9 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
self.handle_callback_failure(callback_name="S3Logger")
async def async_upload_data_to_s3(
self, batch_logging_element: s3BatchLoggingElement
self,
batch_logging_element: s3BatchLoggingElement,
raise_on_error: bool = False,
):
try:
import hashlib
@ -427,6 +433,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
response.raise_for_status()
break
except Exception as e:
if raise_on_error:
raise
verbose_logger.exception(f"Error uploading to s3: {str(e)}")
self.handle_callback_failure(callback_name="S3Logger")
@ -437,19 +445,44 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
Returns: None
Raises: Does not raise an exception, will only verbose_logger.exception()
Raises: RuntimeError when one or more uploads fail. The parent
CustomBatchLogger preserves the failed events for retry.
"""
verbose_logger.debug(f"s3_v2 logger - sending batch of {len(self.log_queue)}")
if not self.log_queue:
return
log_queue_snapshot = list(self.log_queue)
#########################################################
# Flush the log queue to s3
# the log queue can be bounded by DEFAULT_S3_BATCH_SIZE
# see custom_batch_logger.py which triggers the flush
#########################################################
for payload in self.log_queue:
asyncio.create_task(self.async_upload_data_to_s3(payload))
results = await asyncio.gather(
*(
self.async_upload_data_to_s3(payload, raise_on_error=True)
for payload in log_queue_snapshot
),
return_exceptions=True,
)
failed_payloads = [
payload
for payload, result in zip(log_queue_snapshot, results)
if isinstance(result, BaseException)
]
if failed_payloads:
# The base flush handler preserves the current queue when
# async_send_batch raises. Replace the flushed snapshot with only
# the failed payloads so successful uploads are not duplicated on
# the next retry, while events appended during this flush remain.
self.log_queue[:] = (
failed_payloads + self.log_queue[len(log_queue_snapshot) :]
)
raise RuntimeError(
f"S3 batch upload failed for {len(failed_payloads)} of "
f"{len(log_queue_snapshot)} queued events"
)
def create_s3_batch_logging_element(
self,

View file

@ -0,0 +1,117 @@
import asyncio
from unittest.mock import Mock, patch
import pytest
from litellm.integrations.s3_v2 import S3Logger
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
def _s3_batch_element(key: str) -> s3BatchLoggingElement:
return s3BatchLoggingElement(
s3_object_key=f"2025-09-14/{key}.json",
payload={"test": key},
s3_object_download_filename=f"{key}.json",
)
def _s3_logger() -> S3Logger:
with (
patch("asyncio.create_task"),
patch(
"litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush",
new=lambda self: None,
),
):
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",
)
@pytest.mark.asyncio
async def test_async_send_batch_awaits_all_s3_uploads():
logger = _s3_logger()
first = _s3_batch_element("first")
second = _s3_batch_element("second")
logger.log_queue = [first, second]
completed_uploads = []
async def upload(payload, raise_on_error=False):
assert raise_on_error is True
await asyncio.sleep(0)
completed_uploads.append(payload.s3_object_key)
logger.async_upload_data_to_s3 = upload
await logger.async_send_batch()
assert set(completed_uploads) == {first.s3_object_key, second.s3_object_key}
@pytest.mark.asyncio
async def test_async_send_batch_noops_when_queue_is_empty():
logger = _s3_logger()
logger.async_upload_data_to_s3 = Mock()
await logger.async_send_batch()
logger.async_upload_data_to_s3.assert_not_called()
@pytest.mark.asyncio
async def test_flush_queue_preserves_failed_s3_uploads_for_retry():
logger = _s3_logger()
first = _s3_batch_element("first")
failed = _s3_batch_element("failed")
added_during_flush = _s3_batch_element("added-during-flush")
logger.log_queue = [first, failed]
async def upload(payload, raise_on_error=False):
assert raise_on_error is True
if payload is failed:
logger.log_queue.append(added_during_flush)
raise RuntimeError("s3 upload failed")
logger.async_upload_data_to_s3 = upload
await logger.flush_queue()
assert logger.log_queue == [failed, added_during_flush]
@pytest.mark.asyncio
async def test_flush_queue_preserves_s3_events_added_during_successful_flush():
logger = _s3_logger()
first = _s3_batch_element("first")
second = _s3_batch_element("second")
added_during_flush = _s3_batch_element("added-during-flush")
logger.log_queue = [first, second]
async def upload(payload, raise_on_error=False):
assert raise_on_error is True
if payload is first:
logger.log_queue.append(added_during_flush)
logger.async_upload_data_to_s3 = upload
await logger.flush_queue()
assert logger.log_queue == [added_during_flush]
@pytest.mark.asyncio
async def test_async_upload_data_to_s3_reraises_without_callback_failure():
logger = _s3_logger()
logger.get_credentials = Mock(side_effect=RuntimeError("credential failure"))
logger.handle_callback_failure = Mock()
with pytest.raises(RuntimeError, match="credential failure"):
await logger.async_upload_data_to_s3(
_s3_batch_element("failed"),
raise_on_error=True,
)
logger.handle_callback_failure.assert_not_called()