From d8f57780f62e1c20ac01ce33320edb373c251129 Mon Sep 17 00:00:00 2001 From: Ganni Galea Curmi Date: Tue, 26 May 2026 20:20:28 -0400 Subject: [PATCH] fix(s3): await batch uploads before clearing queue --- litellm/integrations/s3_v2.py | 43 ++++++- .../integrations/test_s3_v2_batch_flush.py | 117 ++++++++++++++++++ 2 files changed, 155 insertions(+), 5 deletions(-) create mode 100644 tests/test_litellm/integrations/test_s3_v2_batch_flush.py diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index 4ed8a809a13..daf84e40c06 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -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, diff --git a/tests/test_litellm/integrations/test_s3_v2_batch_flush.py b/tests/test_litellm/integrations/test_s3_v2_batch_flush.py new file mode 100644 index 00000000000..db08c55f759 --- /dev/null +++ b/tests/test_litellm/integrations/test_s3_v2_batch_flush.py @@ -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()