mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(s3): await batch uploads before clearing queue
This commit is contained in:
parent
b6fd7f7746
commit
d8f57780f6
2 changed files with 155 additions and 5 deletions
|
|
@ -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,
|
||||
|
|
|
|||
117
tests/test_litellm/integrations/test_s3_v2_batch_flush.py
Normal file
117
tests/test_litellm/integrations/test_s3_v2_batch_flush.py
Normal 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()
|
||||
Loading…
Add table
Reference in a new issue