diff --git a/litellm/constants.py b/litellm/constants.py index 73fd11fa4d7..a5be2f6568d 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -49,6 +49,7 @@ DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SE DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16")) +DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY: Final = get_env_int("DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY", 200) # https://docs.aws.amazon.com/AmazonS3/latest/userguide/object-keys.html MAX_S3_OBJECT_KEY_BYTES: Final = 1024 S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64 diff --git a/litellm/integrations/adaptive_concurrency.py b/litellm/integrations/adaptive_concurrency.py new file mode 100644 index 00000000000..e0c6b730dda --- /dev/null +++ b/litellm/integrations/adaptive_concurrency.py @@ -0,0 +1,78 @@ +""" +Adaptive in-flight concurrency limiter (AIMD, Vector ARC style). + +Grows the limit additively after `limit` consecutive clean completions and +halves it only on an explicit throttle signal (429, 503, SlowDown, or a +transport error out of the PUT). With floor == ceiling it degenerates to a +fixed-width semaphore. +""" + +import asyncio +from collections import deque +from contextlib import suppress +from dataclasses import dataclass +from typing import Final + + +@dataclass(frozen=True, slots=True) +class PutSample: + throttled: bool + + +class AdaptiveConcurrencyLimiter: + """AIMD in-flight limiter used as `async with limiter:`.""" + + def __init__(self, initial: int, floor: int, ceiling: int) -> None: + if not 1 <= floor <= ceiling: + raise ValueError(f"adaptive limiter bounds must satisfy 1 <= floor <= ceiling, got {floor}..{ceiling}") + self._limit: int = min(max(initial, floor), ceiling) + self._floor: Final[int] = floor + self._ceiling: Final[int] = ceiling + self._clean_streak: int = 0 + self._in_flight: int = 0 + self._waiters: deque[asyncio.Future[None]] = deque() # mutable-ok: waiters queue up behind a full limit + + @property + def limit(self) -> int: + return self._limit + + async def __aenter__(self) -> "AdaptiveConcurrencyLimiter": + if self._in_flight < self._limit: + self._in_flight += 1 + return self + waiter: Final = asyncio.get_running_loop().create_future() + self._waiters.append(waiter) + try: + await waiter + except asyncio.CancelledError: + if waiter.done() and not waiter.cancelled(): + self._in_flight -= 1 + self._grant() + else: + with suppress(ValueError): + self._waiters.remove(waiter) + raise + return self + + def _grant(self) -> None: + while self._in_flight < self._limit and self._waiters: + waiter = self._waiters.popleft() + if waiter.done(): + continue + self._in_flight += 1 + waiter.set_result(None) + + async def __aexit__(self, *_: object) -> None: + self._in_flight -= 1 + self._grant() + + def record(self, sample: PutSample) -> None: + if sample.throttled: + self._limit = max(self._floor, self._limit // 2) + self._clean_streak = 0 + return + self._clean_streak += 1 + if self._clean_streak >= self._limit and self._limit < self._ceiling: + self._limit += 1 + self._clean_streak = 0 + self._grant() diff --git a/litellm/integrations/s3.py b/litellm/integrations/s3.py index 3c8619e82b2..f330ca8e0ac 100644 --- a/litellm/integrations/s3.py +++ b/litellm/integrations/s3.py @@ -36,22 +36,86 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | return True -def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int: +def _resolve_positive_int(setting: str, configured: object, fallback: int, *, reject_bool: bool) -> int: if configured is None or configured == "": return fallback + if reject_bool and isinstance(configured, bool): + verbose_logger.warning( + "s3 logging: %s=%r is a boolean, not an integer, using %s", setting, configured, fallback + ) + return fallback + try: + bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured) + except ValidationError: + verbose_logger.warning("s3 logging: %s=%r is not an integer, using %s", setting, configured, fallback) + return fallback + if bound < 1: + verbose_logger.warning("s3 logging: %s=%r must be at least 1, using %s", setting, configured, fallback) + return fallback + return bound + + +def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int: + return _resolve_positive_int("s3_max_concurrent_uploads", configured, fallback, reject_bool=False) + + +def resolve_s3_max_queue_size(configured: object, fallback: int) -> int: + return _resolve_positive_int("s3_max_queue_size", configured, fallback, reject_bool=True) + + +def resolve_s3_max_retry_age_seconds(configured: object, fallback: int | None) -> int | None: + if configured is None or configured == "": + return None + if isinstance(configured, bool): + verbose_logger.warning( + "s3 logging: s3_max_retry_age_seconds=%r is a boolean, not an integer, falling back to %r", + configured, + fallback, + ) + return fallback try: bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured) except ValidationError: verbose_logger.warning( - "s3 logging: s3_max_concurrent_uploads=%r is not an integer, using %s", configured, fallback + "s3 logging: s3_max_retry_age_seconds=%r is not an integer, falling back to %r", configured, fallback ) return fallback - if bound < 1: + if bound < 0: verbose_logger.warning( - "s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback + "s3 logging: s3_max_retry_age_seconds=%r must be at least 0, falling back to %r", configured, fallback ) return fallback - return bound + return bound or None + + +def resolve_s3_max_adaptive_concurrency(configured: object, fallback: int) -> int: + return _resolve_positive_int("s3_max_adaptive_concurrency", configured, fallback, reject_bool=True) + + +def resolve_s3_drop_on_terminal_error(configured: object) -> bool: + if configured is None or configured == "": + return True + try: + return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured) + except ValidationError: + verbose_logger.warning( + "s3 logging: s3_drop_on_terminal_error=%r is not a boolean, dropping terminal-failed uploads", + configured, + ) + return True + + +def resolve_s3_adaptive_concurrency(configured: object) -> bool: + if configured is None or configured == "": + return False + try: + return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured) + except ValidationError: + verbose_logger.warning( + "s3 logging: s3_adaptive_concurrency=%r is not a boolean, keeping the fixed upload width", + configured, + ) + return False def resolve_s3_batch_file_upload(configured: object) -> bool: diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index dc33fe6c2bd..88d7906cc4b 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -3,14 +3,19 @@ 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; by default each element is uploaded concurrently (bounded by s3_max_concurrent_uploads), or with s3_batch_file_upload the whole flush is written as one .jsonl file +NOTE 1: S3 does not provide a BATCH PUT API endpoint; by default each element is uploaded concurrently with the fixed s3_max_concurrent_uploads bound (or an adaptive bound when s3_adaptive_concurrency is on, backing off only on throttling), or with s3_batch_file_upload the whole flush is written as one .jsonl file """ import asyncio +import contextvars +import logging +import re import time -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass from datetime import datetime, timezone -from typing import TYPE_CHECKING, Final, cast +from functools import partial +from typing import TYPE_CHECKING, Final, Literal, cast from urllib.parse import quote from uuid import uuid4 @@ -21,15 +26,22 @@ from litellm._logging import print_verbose, verbose_logger from litellm.constants import ( DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS, + DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY, DEFAULT_S3_MAX_CONCURRENT_UPLOADS, ) +from litellm.integrations.adaptive_concurrency import AdaptiveConcurrencyLimiter, PutSample from litellm.integrations.s3 import ( get_s3_object_download_filename, get_s3_object_key, prompts_only_payload, + resolve_s3_adaptive_concurrency, resolve_s3_batch_file_upload, + resolve_s3_drop_on_terminal_error, resolve_s3_log_prompts_only, + resolve_s3_max_adaptive_concurrency, resolve_s3_max_concurrent_uploads, + resolve_s3_max_queue_size, + resolve_s3_max_retry_age_seconds, resolve_sse_params, ) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix @@ -50,6 +62,42 @@ if TYPE_CHECKING: from botocore.credentials import Credentials +UploadOutcome = Literal["delivered", "retry", "dropped"] + +_TERMINAL_ERROR_CODES: Final = frozenset( + { + "EntityTooLarge", + "InvalidArgument", + "MalformedXML", + "InvalidDigest", + "KeyTooLongError", + "BadDigest", + "InvalidRequest", + } +) +_BODY_CODED_STATUSES: Final = frozenset({400, 403}) +_RETRYABLE_STATUSES: Final = frozenset({403, 500, 503}) +_S3_ERROR_CODE: Final = re.compile(r"([^<]+)") + + +@dataclass(frozen=True, slots=True) +class _PreparedPut: + json_string: str + headers: Mapping[str, str] + + +def _s3_error_code(response: httpx.Response) -> str | None: + text: Final = response.text + match: Final = _S3_ERROR_CODE.search(text) if isinstance(text, str) else None + return match.group(1) if match else None + + +def _is_terminal(response: httpx.Response) -> bool: + """True only for object-specific, unrecoverable rejections (400/403 with a terminal XML code). + Unknown codes, empty or non-XML bodies, and every other status fail safe toward retry.""" + return response.status_code in _BODY_CODED_STATUSES and _s3_error_code(response) in _TERMINAL_ERROR_CODES + + def _s3_key_parent(s3_object_key: str) -> str: return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else "" @@ -58,11 +106,19 @@ class S3BatchUploadError(Exception): def __init__(self, failed: int, total: int) -> None: self.failed = failed self.total = total - super().__init__(f"{failed} of {total} S3 uploads failed; events kept in queue for the next flush") + super().__init__(f"{failed} of {total} S3 uploads failed; transient failures kept in queue for the next flush") + + +_in_flush: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar("s3_v2_in_flush", default=False) class S3Logger(CustomBatchLogger, BaseAWSLLM): preserve_events_added_during_flush = True + _flush_retries: int = 0 + _requeued_count: int = 0 + _upload_limiter: asyncio.Semaphore | AdaptiveConcurrencyLimiter | None = None + s3_drop_on_terminal_error: bool = True + s3_max_retry_age_seconds: int | None = 3600 def __init__( self, @@ -92,6 +148,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_sse_kms_key_id: str | None = None, s3_log_prompts_only: bool | None = None, s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS, + s3_max_queue_size: int | None = None, + s3_max_retry_age_seconds: int | None = 3600, + s3_drop_on_terminal_error: bool = True, + s3_adaptive_concurrency: bool = False, + s3_max_adaptive_concurrency: int | None = None, s3_batch_file_upload: bool = False, s3_callback_params_override: dict | None = None, **kwargs, @@ -135,9 +196,22 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_sse_kms_key_id=s3_sse_kms_key_id, s3_log_prompts_only=s3_log_prompts_only, s3_max_concurrent_uploads=s3_max_concurrent_uploads, + s3_max_queue_size=s3_max_queue_size, + s3_max_retry_age_seconds=s3_max_retry_age_seconds, + s3_drop_on_terminal_error=s3_drop_on_terminal_error, + s3_adaptive_concurrency=s3_adaptive_concurrency, + s3_max_adaptive_concurrency=s3_max_adaptive_concurrency, s3_batch_file_upload=s3_batch_file_upload, ) - self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads) + self._upload_limiter = ( + AdaptiveConcurrencyLimiter( + initial=self.s3_max_concurrent_uploads, + floor=self.s3_max_concurrent_uploads, + ceiling=max(self.s3_max_concurrent_uploads, self.s3_max_adaptive_concurrency), + ) + if self.s3_adaptive_concurrency + else asyncio.Semaphore(self.s3_max_concurrent_uploads) + ) verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url) # IMPORTANT @@ -158,8 +232,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): flush_lock=self.flush_lock, flush_interval=s3_flush_interval, batch_size=s3_batch_size, + max_queue_size=self.s3_max_queue_size, ) self.log_queue: list[s3BatchLoggingElement] = [] + self._requeued_count = 0 + self._flush_retries = 0 + self._flush_dropped: dict[int, s3BatchLoggingElement] = {} # Call BaseAWSLLM's __init__ BaseAWSLLM.__init__(self) @@ -194,6 +272,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_sse_kms_key_id: str | None = None, s3_log_prompts_only: bool | None = None, s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS, + s3_max_queue_size: int | None = None, + s3_max_retry_age_seconds: int | None = 3600, + s3_drop_on_terminal_error: bool = True, + s3_adaptive_concurrency: bool = False, + s3_max_adaptive_concurrency: int | None = None, s3_batch_file_upload: bool = False, params_source: dict | None = None, ): @@ -259,6 +342,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): DEFAULT_S3_MAX_CONCURRENT_UPLOADS, ) + configured_queue_size: Final = params.get("s3_max_queue_size") + constructor_queue_size: Final = resolve_s3_max_queue_size( + s3_max_queue_size, CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE + ) + self.s3_max_queue_size = resolve_s3_max_queue_size(configured_queue_size, constructor_queue_size) + + configured_retry_age: Final = params.get("s3_max_retry_age_seconds") + constructor_retry_age: Final = resolve_s3_max_retry_age_seconds(s3_max_retry_age_seconds, 3600) + self.s3_max_retry_age_seconds = ( + constructor_retry_age + if configured_retry_age is None or configured_retry_age == "" + else resolve_s3_max_retry_age_seconds(configured_retry_age, constructor_retry_age) + ) + + configured_drop: Final = params.get("s3_drop_on_terminal_error") + self.s3_drop_on_terminal_error = resolve_s3_drop_on_terminal_error( + configured_drop if configured_drop is not None else s3_drop_on_terminal_error + ) + + self.s3_adaptive_concurrency = s3_adaptive_concurrency or resolve_s3_adaptive_concurrency( + params.get("s3_adaptive_concurrency") + ) + + configured_adaptive_ceiling: Final = params.get("s3_max_adaptive_concurrency") + self.s3_max_adaptive_concurrency = resolve_s3_max_adaptive_concurrency( + s3_max_adaptive_concurrency + if configured_adaptive_ceiling is None or configured_adaptive_ceiling == "" + else configured_adaptive_ceiling, + DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY, + ) + self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload( params.get("s3_batch_file_upload") ) @@ -310,6 +424,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): } return {key: value for key, value in candidates.items() if value} + def _prepare_put(self, batch_logging_element: s3BatchLoggingElement) -> _PreparedPut: + try: + import base64 + import hashlib + except ImportError: + raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + + json_string: Final = ( + batch_logging_element.body + if batch_logging_element.body is not None + else safe_dumps(batch_logging_element.payload) + ) + content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest() + content_md5: Final = base64.b64encode( + hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest() + ).decode() + return _PreparedPut( + json_string=json_string, + headers={ + "Content-Type": batch_logging_element.content_type, + "Content-MD5": content_md5, + "x-amz-content-sha256": content_hash, + "Content-Language": "en", + "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"', + "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0", + **self._sse_headers(), + }, + ) + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): await self._async_log_event_base( kwargs=kwargs, @@ -384,12 +527,21 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): verbose_logger.exception("s3 Layer Error - %s", e) self.handle_callback_failure(callback_name="S3Logger") - async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool: - try: - import base64 - import hashlib - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + @property + def _upload_semaphore(self) -> asyncio.Semaphore | AdaptiveConcurrencyLimiter: + limiter: Final = self._upload_limiter + if limiter is None: + raise AttributeError("_upload_semaphore") + return limiter + + @_upload_semaphore.setter + def _upload_semaphore(self, value: asyncio.Semaphore | AdaptiveConcurrencyLimiter) -> None: + self._upload_limiter = value + + async def async_upload_data_to_s3( + self, + batch_logging_element: s3BatchLoggingElement, + ) -> bool: try: from litellm.litellm_core_utils.asyncify import asyncify @@ -400,31 +552,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): url: Final = self._build_object_url(batch_logging_element.s3_object_key) - # Convert JSON to string - json_string: Final = ( - batch_logging_element.body - if batch_logging_element.body is not None - else safe_dumps(batch_logging_element.payload) - ) - - # Calculate SHA256 hash of the content - content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest() - content_md5: Final = base64.b64encode( - hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest() - ).decode() - - # Prepare the request - headers: Final = { - "Content-Type": batch_logging_element.content_type, - "Content-MD5": content_md5, - "x-amz-content-sha256": content_hash, - "Content-Language": "en", - "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"', - "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0", - **self._sse_headers(), - } - - async def signed_put() -> httpx.Response: + async def signed_put(prepared: _PreparedPut) -> httpx.Response: credentials: Final = await asyncified_get_credentials( aws_access_key_id=self.s3_aws_access_key_id, aws_secret_access_key=self.s3_aws_secret_access_key, @@ -436,18 +564,26 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): aws_web_identity_token=self.s3_aws_web_identity_token, aws_sts_endpoint=self.s3_aws_sts_endpoint, ) - signed_headers: Final = await run_aws_signing(self._sign_put, credentials, url, json_string, headers) + signed_headers: Final = await run_aws_signing( + self._sign_put, credentials, url, prepared.json_string, prepared.headers + ) try: - return await self.async_httpx_client.put(url, data=json_string, headers=signed_headers) + return await self.async_httpx_client.put(url, data=prepared.json_string, headers=signed_headers) except httpx.HTTPStatusError as error: return error.response max_retries: Final = 3 + prepared: Final = self._prepare_put(batch_logging_element) for attempt in range(max_retries): - response = await signed_put() - if response.status_code in (403, 500, 503) and attempt < max_retries - 1: + response = await self._recorded_put(partial(signed_put, prepared)) + if ( + response.status_code in _RETRYABLE_STATUSES + and not (self.s3_drop_on_terminal_error and _is_terminal(response)) + and attempt < max_retries - 1 + ): wait_time = 2**attempt # 1s, 2s - verbose_logger.warning( + verbose_logger.log( + logging.DEBUG if _in_flush.get() else logging.WARNING, "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", response.status_code, wait_time, @@ -455,6 +591,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): max_retries, batch_logging_element.s3_object_key, ) + self._flush_retries += 1 await asyncio.sleep(wait_time) continue response.raise_for_status() @@ -462,6 +599,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): except Exception as e: verbose_logger.exception("Error uploading to s3: %s", e) self.handle_callback_failure(callback_name="S3Logger") + if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response): + verbose_logger.warning( + "s3 logging: dropping object %s after terminal status %s", + batch_logging_element.s3_object_key, + e.response.status_code, + ) + self._flush_dropped[id(batch_logging_element)] = batch_logging_element return False return True @@ -483,11 +627,64 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): # see custom_batch_logger.py which triggers the flush ######################################################### uploads: Final = self._batch_file_elements(batch) if self._batch_file_mode_active() else batch - results: Final = await asyncio.gather(*(self._upload_bounded(element) for element in uploads)) - failed: Final = tuple(element for element, ok in zip(uploads, results, strict=True) if not ok) - if not failed: + self._flush_retries = 0 + self._flush_dropped = {} # mutable-ok: per-flush drop marks read back by _upload_bounded + stale: Final = min(self._requeued_count, len(uploads)) if len(uploads) == len(batch) else 0 + order: Final = (*range(stale, len(uploads)), *range(stale)) + ordered: Final = await asyncio.gather(*(self._upload_outcome(uploads[i]) for i in order)) + outcomes: Final = dict(zip(order, ordered, strict=True)) + results: Final = tuple(outcomes[i] for i in range(len(uploads))) + if self._flush_retries: + verbose_logger.warning( + "s3 logging: %s in-call retries across %s uploads this flush", + self._flush_retries, + len(uploads), + ) + delivered: Final = sum(1 for outcome in results if outcome == "delivered") + bucket_wide: Final = delivered == 0 + failed: Final = tuple( + (element, outcome) for element, outcome in zip(uploads, results, strict=True) if outcome != "delivered" + ) + now: Final = time.monotonic() + requeued: Final = ( + tuple(element for element, _ in failed) + if bucket_wide + else tuple( + element + if element.retrying_since is not None or self.s3_max_retry_age_seconds is None + else element.model_copy(update={"retrying_since": now}) + for element, outcome in failed + if outcome != "dropped" + and not ( + self.s3_max_retry_age_seconds is not None + and element.retrying_since is not None + and now - element.retrying_since > self.s3_max_retry_age_seconds + ) + ) + ) + dropped: Final = len(failed) - len(requeued) + if dropped: + verbose_logger.warning( + "s3 logging: %s uploads dropped (terminal or retrying longer than s3_max_retry_age_seconds=%s)", + dropped, + self.s3_max_retry_age_seconds, + ) + if not requeued: + self._requeued_count = 0 return - self.log_queue = [*failed, *self.log_queue[len(batch) :]] + arrivals: Final = self.log_queue[len(batch) :] + overflow: Final = max(0, len(requeued) + len(arrivals) - self.max_queue_size) + if overflow: + verbose_logger.warning( + "s3 logging: queue exceeded max_queue_size=%s after a failed flush, dropped %s oldest events", + self.max_queue_size, + overflow, + ) + self.log_queue = [ # mutable-ok: log_queue is the flush buffer shared with custom_batch_logger + *requeued, + *arrivals, + ][overflow:] + self._requeued_count = max(0, len(requeued) - overflow) raise S3BatchUploadError(failed=len(failed), total=len(uploads)) def _batch_file_mode_active(self) -> bool: @@ -502,8 +699,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): return True async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool: - async with self._upload_semaphore: - return await self.async_upload_data_to_s3(element) + token: Final = _in_flush.set(True) + try: + async with self._upload_semaphore: + return await self.async_upload_data_to_s3(element) + finally: + _in_flush.reset(token) + + async def _upload_outcome(self, element: s3BatchLoggingElement) -> UploadOutcome: + delivered: Final = await self._upload_bounded(element) + if delivered: + return "delivered" + if id(element) in self._flush_dropped: + return "dropped" + return "retry" + + async def _recorded_put(self, signed_put: Callable[[], Awaitable[httpx.Response]]) -> httpx.Response: + limiter: Final = self._upload_limiter + adaptive: Final = limiter if isinstance(limiter, AdaptiveConcurrencyLimiter) else None + try: + response: Final = await signed_put() + except Exception: + if adaptive is not None: + adaptive.record(PutSample(throttled=True)) + raise + if adaptive is not None: + adaptive.record( + PutSample( + throttled=response.status_code in (429, 503) or _s3_error_code(response) == "SlowDown", + ) + ) + return response def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]: now: Final = datetime.now(timezone.utc) @@ -527,6 +753,9 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): content_type="application/x-ndjson", s3_object_key=f"{parent}/{batch_name}.jsonl" if parent else f"{batch_name}.jsonl", s3_object_download_filename=f"{batch_name}.jsonl", + retrying_since=min( + (element.retrying_since for element in elements if element.retrying_since is not None), default=None + ), ) def create_s3_batch_logging_element( @@ -596,58 +825,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): ) def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement): - try: - import base64 - import hashlib - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") try: verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key) url: Final = self._build_object_url(batch_logging_element.s3_object_key) - # Convert JSON to string - json_string: Final = ( - batch_logging_element.body - if batch_logging_element.body is not None - else safe_dumps(batch_logging_element.payload) - ) - - # Calculate SHA256 hash of the content - content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest() - content_md5: Final = base64.b64encode( - hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest() - ).decode() - - # Prepare the request - headers: Final = { - "Content-Type": batch_logging_element.content_type, - "Content-MD5": content_md5, - "x-amz-content-sha256": content_hash, - "Content-Language": "en", - "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"', - "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0", - **self._sse_headers(), - } + prepared: Final = self._prepare_put(batch_logging_element) httpx_client: Final = _get_httpx_client( params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None) ) - def signed_put() -> httpx.Response: + def signed_put(prepared_put: _PreparedPut) -> httpx.Response: credentials: Final = self.get_credentials( aws_access_key_id=self.s3_aws_access_key_id, aws_secret_access_key=self.s3_aws_secret_access_key, aws_session_token=self.s3_aws_session_token, aws_region_name=self.s3_region_name, ) - signed_headers: Final = self._sign_put(credentials, url, json_string, headers) - return httpx_client.put(url, data=json_string, headers=signed_headers) + signed_headers: Final = self._sign_put(credentials, url, prepared_put.json_string, prepared_put.headers) + return httpx_client.put(url, data=prepared_put.json_string, headers=signed_headers) max_retries: Final = 3 for attempt in range(max_retries): - response = signed_put() - if response.status_code in (403, 500, 503) and attempt < max_retries - 1: + response = signed_put(prepared) + if ( + response.status_code in _RETRYABLE_STATUSES + and not (self.s3_drop_on_terminal_error and _is_terminal(response)) + and attempt < max_retries - 1 + ): wait_time = 2**attempt # 1s, 2s verbose_logger.warning( "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", @@ -664,6 +870,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): except Exception as e: verbose_logger.exception("Error uploading to s3: %s", e) self.handle_callback_failure(callback_name="S3Logger") + if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response): + verbose_logger.warning( + "s3 logging: dropping object %s after terminal status %s", + batch_logging_element.s3_object_key, + e.response.status_code, + ) async def _download_object_from_s3(self, s3_object_key: str) -> dict | None: """ diff --git a/litellm/types/integrations/s3_v2.py b/litellm/types/integrations/s3_v2.py index 555b16dc141..3b0dad97e8c 100644 --- a/litellm/types/integrations/s3_v2.py +++ b/litellm/types/integrations/s3_v2.py @@ -11,3 +11,4 @@ class s3BatchLoggingElement(BaseModel): s3_object_download_filename: str body: str | None = None content_type: str = "application/json" + retrying_since: float | None = None diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py index 3652378503e..0713abc86d0 100644 --- a/tests/documentation_tests/test_env_keys.py +++ b/tests/documentation_tests/test_env_keys.py @@ -38,6 +38,7 @@ EXCLUDED_ROLLOUT_FLAGS = { EXCLUDED_INTERNAL_TUNING_VARS = { "ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", + "DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY", } EXCLUDED_TERMINAL_VARS = { diff --git a/tests/integration/observability/_s3_v2_support.py b/tests/integration/observability/_s3_v2_support.py index 104c0eda863..205ac833fcc 100644 --- a/tests/integration/observability/_s3_v2_support.py +++ b/tests/integration/observability/_s3_v2_support.py @@ -27,11 +27,15 @@ class RecordingS3Sink: fail_attempts: int = 0 fail_until: float = 0.0 fail_status: int = 503 + fail_code: str = "SinkFailure" + fail_body: bytes | None = None delay_seconds: float = 0.5 lock: threading.Lock = field(default_factory=threading.Lock) in_flight: int = 0 peak: int = 0 attempts: int = 0 + attempt_log: list[tuple[float, int]] = field(default_factory=list) # mutable-ok: appended under lock per PUT + attempt_counts: dict[str, int] = field(default_factory=dict) # mutable-ok: per-target PUT counts under lock store: dict[str, bytes] = field(default_factory=dict) # mutable-ok: GET reads must see writes from earlier PUTs def respond(self, request: Request) -> Reply: @@ -44,20 +48,31 @@ class RecordingS3Sink: assert request.target.startswith(f"/{BUCKET}/{PREFIX}/"), request.target with self.lock: self.attempts += 1 - if self.attempts <= self.fail_attempts or time.time() < self.fail_until: - return Reply( - status=self.fail_status, - body=b"SinkFailure", - content_type="application/xml", - ) + self.attempt_counts[request.target] = self.attempt_counts.get(request.target, 0) + 1 self.in_flight += 1 self.peak = max(self.peak, self.in_flight) - self.store[request.target] = request.body + self.attempt_log.append((time.time(), self.in_flight)) + failing: Final = self.attempts <= self.fail_attempts or time.time() < self.fail_until + if not failing: + self.store[request.target] = request.body time.sleep(self.delay_seconds) with self.lock: self.in_flight -= 1 + if failing: + return Reply( + status=self.fail_status, + body=self.fail_body + if self.fail_body is not None + else f"{self.fail_code}".encode(), + content_type="application/xml", + ) return Reply() + def peak_between(self, start: float, end: float) -> int: + with self.lock: + samples: Final = tuple(in_flight for when, in_flight in self.attempt_log if start <= when < end) + return max(samples, default=0) + def objects(self) -> Mapping[str, bytes]: with self.lock: return MappingProxyType(dict(self.store)) @@ -240,7 +255,13 @@ SURFACES: Final = ("chat", "chat_stream", "messages", "messages_stream", "respon def call_surface( - candidate: Gateway, surface: str, openai_model: str, anthropic_model: str, key: str, marker: str + candidate: Gateway, + surface: str, + openai_model: str, + anthropic_model: str, + key: str, + marker: str, + no_cache: bool = True, ) -> tuple[str, str | None]: """Drive one request through the given surface; return (client-visible response id, x-litellm-call-id).""" base: Final = str(candidate.client.base_url).rstrip("/") @@ -249,7 +270,7 @@ def call_surface( reply: Final = openai.OpenAI(base_url=f"{base}/v1", api_key=key).chat.completions.create( model=openai_model, messages=[{"role": "user", "content": marker}], - extra_body={"cache": {"no-cache": True}}, + extra_body={"cache": {"no-cache": True}} if no_cache else {}, ) return reply.id, None @@ -258,7 +279,7 @@ def call_surface( model=openai_model, messages=[{"role": "user", "content": marker}], stream=True, - extra_body={"cache": {"no-cache": True}}, + extra_body={"cache": {"no-cache": True}} if no_cache else {}, ) seen = "" async for chunk in stream: @@ -283,7 +304,7 @@ def call_surface( response: Final = candidate.request( "POST", "/v1/responses", - {"model": openai_model, "input": marker, "cache": {"no-cache": True}}, + {"model": openai_model, "input": marker, **({"cache": {"no-cache": True}} if no_cache else {})}, key=key, ) assert response.status_code == 200, response.text @@ -339,6 +360,10 @@ def matched_ids( if payload["id"] in response_ids: landed.append(payload["id"]) continue + uncached: Final = str(payload["id"]).rsplit("_cache_hit", 1)[0] + if uncached in response_ids: + landed.append(str(payload["id"])) + continue assert payload["litellm_call_id"] in call_ids, f"unmatched payload {payload['id']!r}" landed.append(str(payload["id"])) return frozenset(landed) diff --git a/tests/integration/observability/test_s3_v2_flush_surfaces.py b/tests/integration/observability/test_s3_v2_flush_surfaces.py index 2e0b7260a13..5ee117a4b80 100644 --- a/tests/integration/observability/test_s3_v2_flush_surfaces.py +++ b/tests/integration/observability/test_s3_v2_flush_surfaces.py @@ -1,20 +1,24 @@ +import os import re import uuid from pathlib import Path from typing import Final import pytest +from redis import Redis from _s3_v2_support import ( BUCKET, PREFIX, + SURFACES, RecordingS3Sink, + call_surface, collect_payloads, matched_ids, mixed_burst, s3_config, surface_reply, ) -from integration._support.client import Gateway +from integration._support.client import Gateway, eventually from integration._support.process import owned_proxy from integration._support.wire import wire_server @@ -95,3 +99,38 @@ def test_s3_v2_sink_outage_mid_mixed_burst_recovers_every_response_id(gateway: G assert sum(1 for r in provider.drain() if r.method == "POST") == 48 assert matched_ids(payloads, answered) assert len(payloads) == 48, "a stored id was overwritten or duplicated" + + +def test_s3_v2_cache_hit_twins_log_one_object_per_request(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3cache" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(delay_seconds=0.1) + with wire_server(surface_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = s3_config(tmp_path, bucket.url, {}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key" + ) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + cache: Final = Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) + keys_before: Final = cache.dbsize() + warmed: Final = tuple( + call_surface(candidate, surface, openai_model, anthropic_model, key, f"{marker}-{surface}") + for surface in SURFACES + ) + eventually(cache.dbsize, lambda size: size >= keys_before + len(SURFACES), seconds=30) + repeated: Final = tuple( + call_surface(candidate, surface, openai_model, anthropic_model, key, f"{marker}-{surface}", False) + for surface in SURFACES + ) + payloads: Final = collect_payloads(sink, 2 * len(SURFACES)) + assert sum(1 for r in provider.drain() if r.method == "POST") == len(SURFACES), ( + "a repeated request reached the upstream; the six repeats must all be served from cache" + ) + assert len(payloads) == 12 + assert sum(1 for payload in payloads if payload["cache_hit"] is True) == 6 + assert sum(1 for payload in payloads if payload["cache_hit"] is not True) == 6 + assert matched_ids(payloads, warmed + repeated) diff --git a/tests/integration/observability/test_s3_v2_upload_fanout.py b/tests/integration/observability/test_s3_v2_upload_fanout.py index 3ebca152327..b7d101f023f 100644 --- a/tests/integration/observability/test_s3_v2_upload_fanout.py +++ b/tests/integration/observability/test_s3_v2_upload_fanout.py @@ -119,8 +119,8 @@ BATCH_KEY: Final = re.compile( ) -@pytest.mark.covers("other.observability.s3_v2.flush_bounds_concurrent_puts_to_default_and_keeps_every_log") -def test_s3_v2_flush_bounds_concurrent_puts_to_the_default_of_sixteen(gateway: Gateway, tmp_path: Path) -> None: +@pytest.mark.covers("other.observability.s3_v2.flush_bounds_concurrent_puts_to_default_ceiling_and_keeps_every_log") +def test_s3_v2_flush_bounds_concurrent_puts_to_the_default_ceiling(gateway: Gateway, tmp_path: Path) -> None: marker: Final = "s3fan" + uuid.uuid4().hex[:8] sink: Final = S3Sink() with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: @@ -134,7 +134,9 @@ def test_s3_v2_flush_bounds_concurrent_puts_to_the_default_of_sixteen(gateway: G ids: Final = _burst(candidate, model, key, marker) puts: Final = _collect(bucket, count_lines=False, expected=REQUESTS) assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS - assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the default bound for {REQUESTS} queued logs" + assert sink.peak <= 16, ( + f"peak concurrent PUTs {sink.peak} exceeded the default width of 16 for {REQUESTS} queued logs" + ) assert all(PER_REQUEST_KEY.match(put.target) for put in puts), [put.target for put in puts] assert frozenset(json.loads(put.body)["id"] for put in puts) == ids assert len({put.target for put in puts}) == REQUESTS @@ -272,7 +274,7 @@ def test_s3_v2_upstream_failure_events_land_alongside_successes(gateway: Gateway assert all("synthetic upstream rejection" in json.dumps(payload["error_information"]) for payload in failures) -@pytest.mark.covers("other.observability.s3_v2.invalid_or_empty_bound_falls_back_to_sixteen") +@pytest.mark.covers("other.observability.s3_v2.invalid_or_empty_bound_falls_back_to_default_ceiling") @pytest.mark.parametrize( ("bad", "warns"), [ @@ -281,7 +283,7 @@ def test_s3_v2_upstream_failure_events_land_alongside_successes(gateway: Gateway pytest.param("", False, id="empty"), ], ) -def test_s3_v2_invalid_or_empty_bound_falls_back_to_sixteen( +def test_s3_v2_invalid_or_empty_bound_falls_back_to_default_ceiling( gateway: Gateway, tmp_path: Path, bad: JsonValue, warns: bool ) -> None: marker: Final = "s3bound" + uuid.uuid4().hex[:8] @@ -305,14 +307,14 @@ def test_s3_v2_invalid_or_empty_bound_falls_back_to_sixteen( else: assert "s3_max_concurrent_uploads" not in owned.log.read_text() assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS - assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the fallback bound" + assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the fallback width of 16" assert frozenset(payload["id"] for payload in payloads) == ids @pytest.mark.covers("other.observability.s3_v2.sink_rejection_requeues_and_delivers_every_id_once") def test_s3_v2_sink_rejection_requeues_and_delivers_every_id_once(gateway: Gateway, tmp_path: Path) -> None: marker: Final = "s3deny" + uuid.uuid4().hex[:8] - sink: Final = RecordingS3Sink(fail_status=403, delay_seconds=0.2) + sink: Final = RecordingS3Sink(fail_status=503, delay_seconds=0.2) with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: config: Final = _s3_config(tmp_path, bucket.url, {}) with ( @@ -336,6 +338,283 @@ def test_s3_v2_sink_rejection_requeues_and_delivers_every_id_once(gateway: Gatew assert frozenset(payload["id"] for payload in payloads) == ids +@dataclass(slots=True) +class RejectingS3Sink: + """Answers every PUT whose body carries `reject_marker` with `reject_status`, accepts the rest, + and counts the rejected attempts so a test can see whether the proxy keeps re-sending them.""" + + reject_marker: str + reject_status: int + reject_code: str = "AccessDenied" + reject_until: float = float("inf") + lock: threading.Lock = field(default_factory=threading.Lock) + rejected_attempts: int = 0 + rejected_times: list[float] = field(default_factory=list) # mutable-ok: appended under lock per rejected PUT + store: dict[str, bytes] = field(default_factory=dict) # mutable-ok: later PUTs must be visible to earlier polls + + def respond(self, request: Request) -> Reply: + assert request.method == "PUT", request.method + with self.lock: + if self.reject_marker.encode() in request.body and time.time() < self.reject_until: + self.rejected_attempts += 1 + self.rejected_times.append(time.time()) + return Reply(status=self.reject_status, body=f"{self.reject_code}".encode()) + self.store[request.target] = request.body + return Reply() + + def landed_ids(self) -> frozenset[str]: + with self.lock: + bodies: Final = tuple(self.store.values()) + return frozenset(json.loads(line)["id"] for body in bodies for line in body.splitlines()) + + +def _send(candidate: Gateway, model: str, key: str, identity: str) -> None: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": identity}], "cache": {"no-cache": True}}, + key=key, + ) + assert response.status_code == 200, response.text + + +def _send_and_wait_until_landed(candidate: Gateway, model: str, key: str, sink: RejectingS3Sink, identity: str) -> None: + _send(candidate, model, key, identity) + eventually(sink.landed_ids, lambda landed: identity in landed, seconds=60) + + +@pytest.mark.parametrize( + ("status", "code"), + [ + pytest.param(403, "AccessDenied", id="access_denied"), + pytest.param(404, "NoSuchBucket", id="no_such_bucket"), + pytest.param(400, "KMS.DisabledException", id="kms_disabled"), + ], +) +def test_s3_v2_object_rejected_with_a_bucket_wide_code_is_delivered_once_the_fault_clears( + gateway: Gateway, tmp_path: Path, status: int, code: str +) -> None: + marker: Final = "s3fault" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-denied", reject_status=status, reject_code=code) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(candidate, model, key, f"{marker}-denied") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-first-flush") + eventually(lambda: sink.rejected_attempts, lambda attempts: attempts >= 2, seconds=30) + sink.reject_until = time.time() + eventually(sink.landed_ids, lambda landed: f"{marker}-denied" in landed, seconds=60) + readiness: Final = candidate.client.get("/health/readiness") + assert readiness.status_code == 200, readiness.text + assert sum(1 for r in provider.drain() if r.method == "POST") == 2 + assert sink.landed_ids() == {f"{marker}-denied", f"{marker}-first-flush"} + + +def test_s3_v2_terminal_object_is_put_once_and_dropped_by_default(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3toolarge" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-huge", reject_status=400, reject_code="EntityTooLarge") + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(candidate, model, key, f"{marker}-huge") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-sibling") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-second-flush") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-third-flush") + assert sum(1 for r in provider.drain() if r.method == "POST") == 4 + assert sink.rejected_attempts == 1, ( + f"an EntityTooLarge object was PUT {sink.rejected_attempts} times next to delivered siblings; " + "with the default s3_drop_on_terminal_error it must be attempted once and dropped" + ) + + +def test_s3_v2_terminal_object_keeps_retrying_when_opted_out(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3keep" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-huge", reject_status=400, reject_code="EntityTooLarge") + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config( + tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_drop_on_terminal_error": False} + ) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(candidate, model, key, f"{marker}-huge") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-sibling") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-second-flush") + eventually(lambda: sink.rejected_attempts, lambda attempts: attempts >= 2, seconds=30) + assert sum(1 for r in provider.drain() if r.method == "POST") == 3 + + +def test_s3_v2_aged_out_object_is_dropped_next_to_delivered_siblings(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3aged" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-doomed", reject_status=503, reject_code="InternalError") + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_max_retry_age_seconds": 1}) + with ( + owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(owned.gateway, model, key, f"{marker}-doomed") + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-sibling") + _send(owned.gateway, model, key, f"{marker}-trigger") + eventually( + lambda: owned.log.read_text(), + lambda text: "retrying longer than s3_max_retry_age_seconds=1)" in text, + seconds=60, + ) + exhausted: Final = sink.rejected_attempts + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-one-flush-later") + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-two-flushes-later") + assert sum(1 for r in provider.drain() if r.method == "POST") == 5 + assert 3 <= exhausted <= 3 * 3, f"{exhausted} PUTs for an object that aged out after its second flush" + assert sink.rejected_attempts == exhausted, ( + f"a 503 object kept being PUT after ageing out: {exhausted} -> {sink.rejected_attempts}" + ) + + +def test_s3_v2_aged_out_object_stays_queued_while_the_whole_sink_is_down(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3down" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(fail_status=503, delay_seconds=0.1) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_max_retry_age_seconds": 1}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + sink.fail_until = time.time() + 12 + ids: Final = _push(candidate, model, key, marker, 4) + payloads: Final = collect_payloads(sink, 4, seconds=90) + assert sum(1 for r in provider.drain() if r.method == "POST") == 4 + assert frozenset(payload["id"] for payload in payloads) == ids, ( + "a bucket-wide outage longer than the age budget lost events" + ) + + +def test_s3_v2_failing_sink_trims_the_oldest_failed_events_past_the_queue_cap(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3cap" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(fail_status=503, delay_seconds=0.05) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_max_queue_size": 4}) + with ( + owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + sink.fail_until = time.time() + 15 + _send(owned.gateway, model, key, f"{marker}-probe") + eventually(lambda: owned.log.read_text(), lambda text: "S3BatchUploadError" in text, seconds=30) + for index in range(24): + _send(owned.gateway, model, key, f"{marker}-{index}") + eventually( + lambda: owned.log.read_text(), + lambda text: "after a failed flush, dropped" in text, + seconds=30, + ) + payloads: Final = collect_payloads(sink, 4, seconds=90) + landed: Final = frozenset(payload["id"] for payload in payloads) + assert sum(1 for r in provider.drain() if r.method == "POST") == 25 + assert len(landed) == 4, f"{len(landed)} objects landed with s3_max_queue_size=4" + assert f"{marker}-probe" not in landed and f"{marker}-0" not in landed, ( + f"the oldest events survived the cap: {landed}" + ) + assert f"{marker}-23" in landed, f"the newest event was dropped: {landed}" + + +def test_s3_v2_retry_age_zero_keeps_aged_object_queued(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3agezero" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-doomed", reject_status=503, reject_code="SlowDown") + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_max_retry_age_seconds": 0}) + with ( + owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(owned.gateway, model, key, f"{marker}-doomed") + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-sibling") + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-second-flush") + _send_and_wait_until_landed(owned.gateway, model, key, sink, f"{marker}-third-flush") + attempts_before_clear: Final = sink.rejected_attempts + log_text: Final = owned.log.read_text() + assert "uploads dropped" not in log_text, log_text + assert "retrying longer than" not in log_text, log_text + sink.reject_until = time.time() + eventually(sink.landed_ids, lambda landed: f"{marker}-doomed" in landed, seconds=60) + assert sum(1 for r in provider.drain() if r.method == "POST") == 4 + assert attempts_before_clear >= 3, ( + f"only {attempts_before_clear} PUTs for an object that stayed queued through three delivered flushes; " + "with s3_max_retry_age_seconds=0 it must keep retrying longer than any enabled budget" + ) + + +def test_s3_v2_throttled_429_object_is_put_once_per_flush(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3throttle" + uuid.uuid4().hex[:8] + sink: Final = RejectingS3Sink(reject_marker=f"{marker}-throttled", reject_status=429, reject_code="TooManyRequests") + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(candidate, model, key, f"{marker}-throttled") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-sibling") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-second-flush") + _send_and_wait_until_landed(candidate, model, key, sink, f"{marker}-third-flush") + sink.reject_until = time.time() + eventually(sink.landed_ids, lambda landed: f"{marker}-throttled" in landed, seconds=60) + assert sum(1 for r in provider.drain() if r.method == "POST") == 4 + times: Final = tuple(sink.rejected_times) + gaps: Final = tuple(round(later - earlier, 3) for earlier, later in zip(times, times[1:])) + assert len(times) >= 3 and min(gaps) >= 1.5, ( + f"PUTs for a 429 object ran {gaps} apart; the 2 s flush interval allows exactly one attempt per flush " + "because 429 is not an in-call retry status" + ) + + +def test_s3_v2_default_config_retries_access_denied_and_every_event_lands(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3denied" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(fail_status=403, fail_code="AccessDenied", delay_seconds=0.05) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + sink.fail_until = time.time() + 20 + ids: Final = _push(candidate, model, key, marker, 8) + payloads: Final = collect_payloads(sink, 8, seconds=120) + assert sum(1 for r in provider.drain() if r.method == "POST") == 8 + assert frozenset(payload["id"] for payload in payloads) == ids, ( + f"a default-config run lost events through a 20s AccessDenied outage: {len(payloads)} landed" + ) + attempt_totals: Final = tuple(sorted(sink.attempt_counts.values())) + assert len(attempt_totals) == 8 and all(count >= 4 for count in attempt_totals), ( + f"each object must see at least one full 3-PUT retry burst before landing: {attempt_totals}" + ) + + @pytest.mark.covers("other.observability.s3_v2.batch_retry_resends_identical_key_and_body") def test_s3_v2_batch_retry_resends_identical_key_and_body(gateway: Gateway, tmp_path: Path) -> None: marker: Final = "s3retry" + uuid.uuid4().hex[:8] @@ -628,3 +907,187 @@ def test_s3_v2_sigterm_mid_burst_loses_only_inflight_without_duplicates(gateway: ) targets: Final = tuple(sink.objects()) assert len(set(targets)) == len(targets), "the same object was PUT more than once" + + +RAMP_REQUESTS: Final = 400 +RAMP_PUT_DELAY_SECONDS: Final = 1.0 + + +def _push(candidate: Gateway, model: str, key: str, marker: str, count: int) -> frozenset[str]: + ids: Final = tuple(f"{marker}-{index}" for index in range(count)) + + def request(identity: str) -> str: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": identity}], "cache": {"no-cache": True}}, + key=key, + ) + assert response.status_code == 200, response.text + return response.json()["id"] + + with ThreadPoolExecutor(max_workers=64) as pool: + returned: Final = frozenset(pool.map(request, ids)) + assert returned == frozenset(ids) + return returned + + +def test_s3_v2_slow_sink_ramps_concurrency_and_drains_the_backlog(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3ramp" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(delay_seconds=RAMP_PUT_DELAY_SECONDS) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config( + tmp_path, bucket.url, {"s3_batch_file_upload": False, "s3_adaptive_concurrency": True} + ) + with ( + owned_proxy( + gateway, + tmp_path, + {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "1", "DEFAULT_S3_BATCH_SIZE": "5000"}, + config=config, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + ids: Final = _push(candidate, model, key, marker, RAMP_REQUESTS) + drain_started: Final = time.monotonic() + payloads: Final = collect_payloads(sink, RAMP_REQUESTS, seconds=180) + drained_seconds: Final = time.monotonic() - drain_started + fixed_sixteen_estimate: Final = RAMP_REQUESTS * RAMP_PUT_DELAY_SECONDS / 16 + assert sum(1 for r in provider.drain() if r.method == "POST") == RAMP_REQUESTS + assert frozenset(payload["id"] for payload in payloads) == ids + assert sink.peak > 16, f"adaptive limiter never ramped past the old fixed bound: peak {sink.peak}" + assert drained_seconds < 2 * fixed_sixteen_estimate, ( + f"backlog of {RAMP_REQUESTS} drained in {drained_seconds:.1f}s with peak concurrency {sink.peak}; " + f"even a fixed bound of 16 would need only ~{fixed_sixteen_estimate:.1f}s, so the uploads stalled" + ) + + +def test_s3_v2_throttled_sink_halves_in_flight_puts(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3throt" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(fail_status=503, fail_code="SlowDown", delay_seconds=0.3) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config( + tmp_path, + bucket.url, + {"s3_batch_file_upload": False, "s3_adaptive_concurrency": True, "s3_max_concurrent_uploads": 4}, + ) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + healthy_ids: Final = _push(candidate, model, key, f"{marker}-healthy", REQUESTS) + collect_payloads(sink, REQUESTS) + healthy_peak: Final = sink.peak + window_start: Final = time.time() + sink.fail_until = window_start + 60 + throttled_ids: Final = _push(candidate, model, key, f"{marker}-throttled", REQUESTS) + first_fail_at: Final = eventually( + lambda: next((when for when, _ in sink.attempt_log if when >= window_start), None), + lambda when: when is not None, + seconds=30, + ) + window_end: Final = first_fail_at + 8.0 + sink.fail_until = window_end + payloads: Final = collect_payloads(sink, 2 * REQUESTS, seconds=120) + throttled_peak: Final = sink.peak_between(first_fail_at + 5.0, window_end) + throttled_attempts: Final = sum( + 1 for when, _ in sink.attempt_log if first_fail_at + 5.0 <= when < window_end + ) + assert sum(1 for r in provider.drain() if r.method == "POST") == 2 * REQUESTS + assert healthy_peak > 4, ( + f"healthy peak {healthy_peak} never rose above the configured width 4; nothing to back off from" + ) + assert throttled_attempts > 0, ( + "no PUTs observed in the measured SlowDown window; the back-off assertion would be vacuous" + ) + assert throttled_peak < healthy_peak, ( + f"in-flight PUTs during the SlowDown window peaked at {throttled_peak}, not below the healthy peak " + f"{healthy_peak}; the limiter did not back off" + ) + assert frozenset(payload["id"] for payload in payloads) == healthy_ids | throttled_ids + assert len(sink.objects()) == 2 * REQUESTS + + +def test_s3_v2_coded_403_is_transient_and_every_id_lands_once(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3coded" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink(fail_attempts=3, fail_status=403, fail_code="RequestTimeout", delay_seconds=0.1) + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": False}) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + ids: Final = _push(candidate, model, key, marker, 4) + payloads: Final = collect_payloads(sink, 4) + assert frozenset(payload["id"] for payload in payloads) == ids + assert sink.attempts >= 7, ( + f"only {sink.attempts} PUT attempts for 4 objects whose first 3 uploads 403 RequestTimeout; " + "coded 403s must be retried" + ) + + +def test_s3_v2_success_callback_mode_logs_only_successes(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3succ" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink() + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _recording_s3_config( + tmp_path, + bucket.url, + {}, + {"callbacks": [], "success_callback": ["s3_v2"]}, + ) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + ghost: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": f"ghost-{uuid.uuid4().hex}", "messages": [{"role": "user", "content": "hi"}]}, + key=key, + ) + assert ghost.status_code in (400, 403, 404), ghost.text + _send(candidate, model, key, marker) + payloads: Final = collect_payloads(sink, 1) + assert len(payloads) == 1 + assert payloads[0]["id"] == marker + assert payloads[0]["status"] == "success" + + +def test_s3_v2_failure_callback_mode_logs_only_failures(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "s3failcb" + uuid.uuid4().hex[:8] + sink: Final = RecordingS3Sink() + with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket: + config: Final = _recording_s3_config( + tmp_path, + bucket.url, + {}, + {"callbacks": [], "failure_callback": ["s3_v2"]}, + ) + with ( + owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, config=config) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _send(candidate, model, key, marker) + ghost: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": f"ghost-{uuid.uuid4().hex}", "messages": [{"role": "user", "content": "hi"}]}, + key=key, + ) + assert ghost.status_code in (400, 403, 404), ghost.text + payloads: Final = collect_payloads(sink, 1) + assert len(payloads) == 1 + assert payloads[0]["status"] == "failure" + assert payloads[0]["id"] != marker + assert isinstance(payloads[0]["litellm_call_id"], str) and payloads[0]["litellm_call_id"] diff --git a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py index b1d111bf1f9..0dff25965f4 100644 --- a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py +++ b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py @@ -322,6 +322,7 @@ class TestS3LoggerAuditLogEvent: logger.s3_path = "my-prefix" logger.log_queue = [] logger.batch_size = 100 + logger.max_queue_size = 100 audit_log = StandardAuditLogPayload( id="audit-123", @@ -355,6 +356,7 @@ class TestS3LoggerAuditLogEvent: logger.s3_path = None logger.log_queue = [] logger.batch_size = 100 + logger.max_queue_size = 100 audit_log = StandardAuditLogPayload( id="audit-456", diff --git a/tests/unit/integrations/test_adaptive_concurrency.py b/tests/unit/integrations/test_adaptive_concurrency.py new file mode 100644 index 00000000000..15b9d52a42e --- /dev/null +++ b/tests/unit/integrations/test_adaptive_concurrency.py @@ -0,0 +1,179 @@ +import asyncio +from typing import Final + +import pytest + +from litellm.integrations.adaptive_concurrency import AdaptiveConcurrencyLimiter, PutSample + +_real_sleep: Final = asyncio.sleep + + +def _limiter(initial: int = 4, floor: int = 1, ceiling: int = 16) -> AdaptiveConcurrencyLimiter: + return AdaptiveConcurrencyLimiter(initial=initial, floor=floor, ceiling=ceiling) + + +@pytest.mark.asyncio +async def test_limit_grows_after_limit_clean_samples() -> None: + limiter: Final = _limiter(initial=4) + for _ in range(4): + limiter.record(PutSample(throttled=False)) + assert limiter.limit == 5 + + +@pytest.mark.asyncio +async def test_limit_does_not_grow_before_the_streak_completes() -> None: + limiter: Final = _limiter(initial=4) + for _ in range(3): + limiter.record(PutSample(throttled=False)) + assert limiter.limit == 4 + + +@pytest.mark.asyncio +async def test_throttled_sample_halves_the_limit() -> None: + limiter: Final = _limiter(initial=16) + limiter.record(PutSample(throttled=True)) + assert limiter.limit == 8 + + +@pytest.mark.asyncio +async def test_limit_clamps_at_the_floor() -> None: + limiter: Final = _limiter(initial=4, floor=4) + limiter.record(PutSample(throttled=True)) + assert limiter.limit == 4 + + +@pytest.mark.asyncio +async def test_limit_clamps_at_the_ceiling() -> None: + limiter: Final = _limiter(initial=15, ceiling=16) + for _ in range(1000): + limiter.record(PutSample(throttled=False)) + assert limiter.limit == 16 + + +@pytest.mark.asyncio +async def test_throttled_sample_resets_the_clean_streak() -> None: + limiter: Final = _limiter(initial=4, ceiling=32) + for _ in range(3): + limiter.record(PutSample(throttled=False)) + limiter.record(PutSample(throttled=True)) + limiter.record(PutSample(throttled=False)) + assert limiter.limit == 2 + + +@pytest.mark.asyncio +async def test_growing_the_limit_wakes_a_waiting_acquirer() -> None: + limiter: Final = AdaptiveConcurrencyLimiter(initial=1, floor=1, ceiling=4) + acquired: Final[list[str]] = [] # mutable-ok: the waiter task appends to it across the await boundary + released: Final = asyncio.Event() + + async def hold() -> None: + async with limiter: + await released.wait() + + holder: Final = asyncio.create_task(hold()) + + async def waiter() -> None: + async with limiter: + acquired.append("waiter") + + pending: Final = asyncio.create_task(waiter()) + await _real_sleep(0) + assert not acquired + + limiter.record(PutSample(throttled=False)) + await asyncio.wait_for(asyncio.shield(pending), timeout=5) + released.set() + await asyncio.wait_for(holder, timeout=5) + assert tuple(acquired) == ("waiter",) + + +@pytest.mark.asyncio +async def test_releasing_a_slot_wakes_exactly_one_waiter() -> None: + limiter: Final = AdaptiveConcurrencyLimiter(initial=1, floor=1, ceiling=4) + acquired: Final[list[str]] = [] # mutable-ok: the waiter tasks append to it across the await boundary + release: Final = asyncio.Event() + + async def hold() -> None: + async with limiter: + await _real_sleep(0) + + async def waiter(name: str) -> None: + async with limiter: + acquired.append(name) + await release.wait() + + holder: Final = asyncio.create_task(hold()) + waiters: Final = tuple(asyncio.create_task(waiter(f"w{i}")) for i in range(3)) + await _real_sleep(0) + await asyncio.wait_for(holder, timeout=5) + await _real_sleep(0) + assert len(acquired) == 1 + + for _ in range(3): + limiter.record(PutSample(throttled=False)) + await _real_sleep(0) + assert len(acquired) == 3 + release.set() + await asyncio.gather(*waiters) + + +@pytest.mark.asyncio +async def test_double_cancel_during_release_leaves_in_flight_at_zero() -> None: + limiter: Final = AdaptiveConcurrencyLimiter(initial=1, floor=1, ceiling=1) + entered: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def hold() -> None: + async with limiter: + entered.set() + await release.wait() + + holder: Final = asyncio.create_task(hold()) + await entered.wait() + + waiter: Final = asyncio.create_task(hold()) + await _real_sleep(0) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + + release.set() + await asyncio.wait_for(holder, timeout=5) + holder.cancel() + try: + await holder + except asyncio.CancelledError: + pass + + assert limiter._in_flight == 0 + + +@pytest.mark.asyncio +async def test_a_cancelled_waiter_is_skipped_when_a_slot_frees() -> None: + limiter: Final = AdaptiveConcurrencyLimiter(initial=1, floor=1, ceiling=1) + acquired: Final[list[str]] = [] # mutable-ok: waiter tasks append across the await boundary + first_entered: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def hold(name: str, entered: asyncio.Event | None = None) -> None: + async with limiter: + acquired.append(name) + if entered is not None: + entered.set() + await release.wait() + + holder: Final = asyncio.create_task(hold("holder", first_entered)) + await first_entered.wait() + doomed: Final = asyncio.create_task(hold("doomed")) + next_waiter: Final = asyncio.create_task(hold("next")) + await _real_sleep(0) + doomed.cancel() + with pytest.raises(asyncio.CancelledError): + await doomed + release.set() + await asyncio.wait_for(holder, timeout=5) + await asyncio.wait_for(next_waiter, timeout=5) + + assert "doomed" not in acquired + assert "next" in acquired + assert limiter._in_flight == 0 diff --git a/tests/unit/integrations/test_s3_v2.py b/tests/unit/integrations/test_s3_v2.py index c67eaa45112..caab4ff561d 100644 --- a/tests/unit/integrations/test_s3_v2.py +++ b/tests/unit/integrations/test_s3_v2.py @@ -4,22 +4,27 @@ import json import re import sys import textwrap +import time import uuid from collections.abc import Awaitable, Callable from contextlib import asynccontextmanager from datetime import datetime from pathlib import Path +from typing import Final from unittest.mock import AsyncMock, MagicMock, call, patch import httpx import pytest import respx -from litellm.integrations.s3_v2 import S3Logger +from litellm.integrations.s3_v2 import S3BatchUploadError, 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 +_real_sleep: Final = asyncio.sleep +_NOW: Final = 1_000_000.0 + class TestS3V2UnitTests: """Test that S3 v2 integration only uses safe_dumps and not json.dumps""" @@ -387,8 +392,10 @@ async def test_async_upload_retries_on_s3_503(): # First call returns 503, second call returns 200 response_503 = MagicMock() response_503.status_code = 503 + response_503.text = "" response_200 = MagicMock() response_200.status_code = 200 + response_200.text = "" response_200.raise_for_status = MagicMock() logger.async_httpx_client = AsyncMock() @@ -427,8 +434,10 @@ async def test_async_upload_retries_on_s3_500(): response_500 = MagicMock() response_500.status_code = 500 + response_500.text = "" response_200 = MagicMock() response_200.status_code = 200 + response_200.text = "" response_200.raise_for_status = MagicMock() logger.async_httpx_client = AsyncMock() @@ -467,6 +476,7 @@ async def test_async_upload_exhausts_retries_on_persistent_503(): # All 3 attempts return 503 response_503 = MagicMock() response_503.status_code = 503 + response_503.text = "" response_503.raise_for_status = MagicMock(side_effect=Exception("503 Service Unavailable")) logger.async_httpx_client = AsyncMock() @@ -485,9 +495,10 @@ async def test_async_upload_exhausts_retries_on_persistent_503(): @pytest.mark.asyncio -async def test_async_upload_no_retry_on_4xx(): +async def test_async_upload_retries_400_with_an_unknown_error_code(): """ - Test that async_upload_data_to_s3 does NOT retry on 4xx errors (client errors). + A 400 is outside the retry set, so an unknown gets a single PUT and the "retry" outcome + for the flush-level requeue, never an in-call backoff. """ from unittest.mock import AsyncMock, MagicMock @@ -501,24 +512,29 @@ async def test_async_upload_no_retry_on_4xx(): ) 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", + s3_object_key="2025-09-14/test-retry-400.json", + payload={"test": "retry-400"}, + s3_object_download_filename="test-retry-400.json", ) response_400 = MagicMock() response_400.status_code = 400 + response_400.text = "SomethingElse" response_400.raise_for_status = MagicMock(side_effect=Exception("400 Bad Request")) + response_200 = MagicMock() + response_200.status_code = 200 + response_200.text = "" + response_200.raise_for_status = MagicMock() logger.async_httpx_client = AsyncMock() - logger.async_httpx_client.put = AsyncMock(return_value=response_400) + logger.async_httpx_client.put = AsyncMock(side_effect=[response_400, response_200]) - with patch.object(logger, "handle_callback_failure") as mock_failure: - await logger.async_upload_data_to_s3(test_element) + with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + outcome = 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") + mock_sleep.assert_not_awaited() + assert outcome is False _SIGV4_ACCESS_KEY = re.compile(r"Credential=(AKIA\d+)/") @@ -657,21 +673,57 @@ async def test_async_upload_exhausts_403_retries_through_production_http_handler @pytest.mark.asyncio -async def test_async_upload_does_not_retry_404_through_production_http_handler(rotating_profile: str, caplog): +async def test_async_upload_is_single_attempted_on_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) + outcome = await logger.async_upload_data_to_s3(test_element) assert len(requests) == 1 + assert outcome is False mock_sleep.assert_not_awaited() assert "Error uploading to s3" in caplog.text +@pytest.mark.asyncio +async def test_async_upload_access_denied_403_is_retried_and_then_requeued(rotating_profile: str, caplog): + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-403-denied.json", + payload={"test": "403-denied"}, + s3_object_download_filename="test-403-denied.json", + ) + requests: list[httpx.Request] = [] + + def respond(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(403, request=request, text="AccessDenied") + + 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=rotating_profile, + s3_flush_interval=3600, + ) + logger.async_httpx_client = handler + with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + outcome = await logger.async_upload_data_to_s3(test_element) + await handler.client.aclose() + + assert outcome is False + assert len(requests) == 3 + assert mock_sleep.await_args_list == [call(1), call(2)] + assert "Error uploading to s3" in caplog.text + + def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("AWS_ACCESS_KEY_ID", raising=False) + monkeypatch.delenv("AWS_SECRET_ACCESS_KEY", raising=False) + monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) 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( @@ -684,7 +736,7 @@ def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, mon def respond(request: httpx.Request) -> httpx.Response: requests.append(request) - return httpx.Response(next(replies), request=request) + return httpx.Response(next(replies), request=request, text="SignatureDoesNotMatch") handler = HTTPHandler() handler.client = httpx.Client(transport=httpx.MockTransport(respond)) @@ -2481,12 +2533,14 @@ def _element(payload: dict[str, object], key_suffix: str) -> s3BatchLoggingEleme def _ok_response() -> MagicMock: response = MagicMock() response.status_code = 200 + response.text = "" response.raise_for_status = MagicMock() return response class _CountingPut: - def __init__(self) -> None: + def __init__(self, width: int) -> None: + self.width = width self.in_flight = 0 self.peak = 0 self.calls = 0 @@ -2495,7 +2549,10 @@ class _CountingPut: self.in_flight += 1 self.peak = max(self.peak, self.in_flight) self.calls += 1 - await asyncio.sleep(0.01) + for _ in range(50): + if self.in_flight >= self.width: + break + await _real_sleep(0) self.in_flight -= 1 return _ok_response() @@ -2515,36 +2572,56 @@ class _LateAppendingPut: self.element = element self.fail_first = fail_first self.appended = False + self.failed_key: str | None = None 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() + self.failed_key = url + if url == self.failed_key: + return _transient_failure_response() return _ok_response() +class _AppendingFailingPut: + def __init__(self, logger: S3Logger, elements: tuple[s3BatchLoggingElement, ...]) -> None: + self.logger = logger + self.elements = elements + 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 + for element in self.elements: + self.logger.log_queue.append(element) + return _transient_failure_response() + + class _FailOnSuffixPut: def __init__(self, suffixes: tuple[str, ...]) -> None: self.failing = True self.suffixes = suffixes + self.calls: tuple[str, ...] = () async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, url) if self.failing and url.endswith(self.suffixes): - return _failure_response() + return _transient_failure_response() return _ok_response() class _FailUntilClearedPut: - def __init__(self) -> None: + def __init__(self, status: int = 503, code: str | None = "SlowDown", raw_body: str | None = None) -> None: self.failing = True + self.response: Final = _coded_failure_response(status, code, raw_body) 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 self.response return _ok_response() @@ -2558,7 +2635,7 @@ async def test_async_send_batch_bounds_concurrent_uploads() -> None: s3_max_concurrent_uploads=4, ) - put = _CountingPut() + put = _CountingPut(logger.s3_max_concurrent_uploads) logger.async_httpx_client = AsyncMock() logger.async_httpx_client.put = put @@ -2652,14 +2729,14 @@ def test_invalid_concurrency_falls_back_to_default(bad: object) -> None: 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 + assert logger._upload_limiter._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 + assert logger._upload_limiter._value == 4 @pytest.mark.parametrize("empty", [None, ""]) @@ -2674,16 +2751,30 @@ def test_empty_config_concurrency_falls_back_to_constructor_value(empty: object) ) assert logger.s3_max_concurrent_uploads == 4 - assert logger._upload_semaphore._value == 4 + assert logger._upload_limiter._value == 4 -def _failure_response() -> MagicMock: +def _coded_failure_response(status: int, code: str | None, raw_body: str | None = None) -> MagicMock: + body: Final = ( + raw_body if raw_body is not None else (f"{code}" if code is not None else "") + ) response = MagicMock() - response.status_code = 400 - response.raise_for_status = MagicMock(side_effect=Exception("s3 rejected the object")) + response.status_code = status + response.text = body + response.raise_for_status = MagicMock( + side_effect=httpx.HTTPStatusError(str(status), request=MagicMock(), response=response) + ) return response +def _transient_failure_response(status: int = 503) -> MagicMock: + return _coded_failure_response(status, "SlowDown") + + +def _terminal_failure_response() -> MagicMock: + return _coded_failure_response(400, "EntityTooLarge") + + @pytest.mark.asyncio async def test_failed_uploads_stay_queued_for_next_flush() -> None: logger = S3Logger( @@ -2700,12 +2791,16 @@ async def test_failed_uploads_stay_queued_for_next_flush() -> None: logger.async_httpx_client.put = put logger.log_queue = list(elements) - await logger.flush_queue() + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() - assert logger.log_queue == [elements[2], elements[4]] + assert [element.s3_object_key for element in logger.log_queue] == [ + elements[2].s3_object_key, + elements[4].s3_object_key, + ] - put.failing = False - await logger.flush_queue() + put.failing = False + await logger.flush_queue() assert logger.log_queue == [] @@ -2728,9 +2823,10 @@ async def test_batch_file_upload_failure_keeps_whole_batch() -> None: elements = [_element({"i": i}, f"{i}") for i in range(3)] logger.log_queue = list(elements) - await logger.flush_queue() + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() - assert len(put.calls) == 1 + assert len(put.calls) == 3 assert len(logger.log_queue) == 1 assert logger.log_queue[0].body == "\n".join(json.dumps(element.payload) for element in elements) @@ -2752,9 +2848,16 @@ async def test_events_appended_during_failed_flush_survive() -> None: first = _element({"id": "first"}, "first") logger.log_queue = [first] + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == [first.s3_object_key, late.s3_object_key] + assert logger.log_queue[0].retrying_since is None + + logger.async_httpx_client.put.failed_key = None await logger.flush_queue() - assert logger.log_queue == [first, late] + assert logger.log_queue == [] @pytest.mark.asyncio @@ -2841,7 +2944,8 @@ async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None: logger.async_httpx_client.put = put logger.log_queue = [_element({"i": i}, f"{i}") for i in range(3)] - await logger.flush_queue() + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() assert len(logger.log_queue) == 1 assert logger.log_queue[0].body is not None @@ -2851,7 +2955,7 @@ async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None: await logger.flush_queue() assert logger.log_queue == [] - assert len(put.calls) == 2 + assert len(put.calls) == 4 assert put.calls[0] == put.calls[1] @@ -2871,7 +2975,8 @@ async def test_elements_appended_after_failed_batch_file_get_their_own_file() -> logger.async_httpx_client.put = put logger.log_queue = [_element({"id": "first"}, "first")] - await logger.flush_queue() + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() late = _element({"id": "late"}, "late") logger.log_queue.append(late) @@ -2880,10 +2985,13 @@ async def test_elements_appended_after_failed_batch_file_get_their_own_file() -> await logger.flush_queue() assert logger.log_queue == [] - assert len(put.calls) == 3 + assert len(put.calls) == 5 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"}) + second_flush: Final = put.calls[3:] + assert put.calls[0] in second_flush + late_call: Final = next(call for call in second_flush if call != put.calls[0]) + assert late_call[0] != put.calls[0][0] + assert late_call[1] == json.dumps({"id": "late"}) @pytest.mark.asyncio @@ -2918,3 +3026,1769 @@ async def test_batch_file_mode_disabled_when_s3_v2_is_cold_storage_logger(monkey assert len(put.calls) == 2 assert put.calls[1][0].endswith(".jsonl") + + +class _FailOnSuffixCodedPut: + def __init__( + self, suffixes: tuple[str, ...], status: int, code: str | None = None, raw_body: str | None = None + ) -> None: + self.suffixes = suffixes + self.response: Final = _coded_failure_response(status, code, raw_body) + self.calls: tuple[str, ...] = () + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, url) + if url.endswith(self.suffixes): + return self.response + return _ok_response() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "code", "raw_body", "puts_per_element"), + [ + pytest.param(403, "AccessDenied", None, 3, id="access-denied-403"), + pytest.param(403, None, None, 3, id="empty-403"), + pytest.param(403, None, "Forbidden", 3, id="html-403"), + pytest.param(400, "KMS.DisabledException", None, 1, id="kms-disabled-400"), + pytest.param(404, "NoSuchBucket", None, 1, id="no-such-bucket-404"), + ], +) +async def test_non_terminal_failure_is_requeued_and_delivered_on_recovery( + status: int, code: str | None, raw_body: str | None, puts_per_element: int +) -> 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 = _FailUntilClearedPut(status=status, code=code, raw_body=raw_body) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = list(elements) + + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert len(logger.log_queue) == 5 + assert len(put.calls) == 5 * puts_per_element + + put.failing = False + await logger.flush_queue() + + assert logger.log_queue == [] + assert len(put.calls) == 5 * puts_per_element + 5 + landed: Final = frozenset( + element.s3_object_key + for element in elements + if any(call[0].endswith(element.s3_object_key) for call in put.calls[-5:]) + ) + assert landed == frozenset(element.s3_object_key for element in elements) + + +@pytest.mark.asyncio +async def test_persistent_500_stays_queued_through_a_dozen_failed_flushes() -> 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", + ) + + put = _FailUntilClearedPut(status=500, code="InternalError") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)] + + with patch("asyncio.sleep", new_callable=AsyncMock): + for _ in range(12): + await logger.flush_queue() + assert len(logger.log_queue) == 5 + + assert len(put.calls) == 12 * 15 + + put.failing = False + await logger.flush_queue() + + assert logger.log_queue == [] + assert len(put.calls) == 12 * 15 + 5 + + +@pytest.mark.asyncio +async def test_terminal_object_is_dropped_once_next_to_delivered_siblings_when_opted_in() -> 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_drop_on_terminal_error=True, + ) + + elements = [_element({"i": i}, f"{i}") for i in range(5)] + put = _FailOnSuffixCodedPut(("test-1.json",), 400, "EntityTooLarge") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = list(elements) + + await logger.flush_queue() + + assert len(put.calls) == 5 + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_terminal_object_is_requeued_when_opted_out() -> 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_drop_on_terminal_error=False, + ) + + put = _FailOnSuffixCodedPut(("test-1.json",), 400, "EntityTooLarge") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)] + + await logger.flush_queue() + + assert len(logger.log_queue) == 1 + assert logger.log_queue[0].s3_object_key.endswith("test-1.json") + assert sum(call.endswith("test-1.json") for call in put.calls) == 1 + assert len(put.calls) == 5 + + +@pytest.mark.asyncio +async def test_terminal_objects_are_requeued_when_every_upload_in_the_flush_fails() -> 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_drop_on_terminal_error=True, + ) + + put = _FailUntilClearedPut(status=400, code="EntityTooLarge") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(5)] + + await logger.flush_queue() + + assert len(logger.log_queue) == 5 + + +@pytest.mark.asyncio +async def test_retrying_past_the_opted_in_budget_is_dropped_only_next_to_delivered_siblings() -> 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_retry_age_seconds=60, + ) + + aged = _element({"id": "aged"}, "aged").model_copy(update={"retrying_since": _NOW - 120}) + fresh = _element({"id": "fresh"}, "fresh") + put = _FailOnSuffixCodedPut(("test-aged.json",), 503, "SlowDown") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [aged, fresh] + + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=_NOW), + ): + await logger.flush_queue() + + assert logger.log_queue == [] + assert len(put.calls) == 4 + + +@pytest.mark.asyncio +async def test_retrying_past_the_budget_stays_queued_when_the_whole_flush_fails() -> 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_retry_age_seconds=60, + ) + + aged = _element({"id": "aged"}, "aged").model_copy(update={"retrying_since": _NOW - 120}) + put = _FailUntilClearedPut(status=503, code="SlowDown") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [aged] + + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=_NOW), + ): + await logger.flush_queue() + + assert len(logger.log_queue) == 1 + + +@pytest.mark.asyncio +async def test_overflow_after_a_failed_flush_trims_failed_first_and_counts_upload_failures_only( + caplog, +) -> 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_queue_size=4, + ) + + late = tuple(_element({"id": f"late-{index}"}, f"late-{index}") for index in range(3)) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _AppendingFailingPut(logger, late) + logger.log_queue = [_element({"id": "first"}, "first"), _element({"id": "second"}, "second")] + + with ( + patch.object(logger, "handle_callback_failure") as mock_failure, + patch("asyncio.sleep", new_callable=AsyncMock), + pytest.raises(S3BatchUploadError), + ): + await logger.async_send_batch() + + assert [element.payload["id"] for element in logger.log_queue] == ["second", "late-0", "late-1", "late-2"] + failed_uploads: Final = 2 + assert mock_failure.call_count == failed_uploads + mock_failure.assert_called_with(callback_name="S3Logger") + assert "dropped 1 oldest events" in caplog.text + + +@pytest.mark.asyncio +async def test_default_logger_ages_out_elements_retrying_longer_than_an_hour(caplog) -> 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": index}, f"{index}").model_copy(update={"retrying_since": _NOW - 7200}) for index in range(3) + ] + put = _FailOnSuffixPut(("test-1.json", "test-2.json")) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = list(elements) + + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=_NOW), + ): + await logger.flush_queue() + + assert logger.log_queue == [] + assert "uploads dropped" in caplog.text + + +@pytest.mark.asyncio +async def test_opted_out_logger_never_ages_out_long_retrying_elements(caplog) -> 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_retry_age_seconds=0, + ) + + elements = [ + _element({"i": index}, f"{index}").model_copy(update={"retrying_since": _NOW - 7200}) for index in range(3) + ] + put = _FailOnSuffixPut(("test-1.json", "test-2.json")) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = list(elements) + + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=_NOW), + ): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == [ + elements[1].s3_object_key, + elements[2].s3_object_key, + ] + assert "uploads dropped" not in caplog.text + + put.failing = False + await logger.flush_queue() + + assert logger.log_queue == [] + landed: Final = frozenset(call_url.rsplit("/", 1)[-1] for call_url in put.calls) + assert landed == frozenset(f"test-{index}.json" for index in range(3)) + + +@pytest.mark.asyncio +async def test_queue_grows_past_the_cap_while_the_sink_fails_and_everything_lands() -> 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_queue_size=5, + ) + + elements = [_element({"i": index}, f"{index}") for index in range(8)] + put = _FailUntilClearedPut() + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = list(elements[:5]) + + with ( + patch.object(logger, "handle_callback_failure") as mock_failure, + patch("asyncio.sleep", new_callable=AsyncMock), + ): + await logger.flush_queue() + + assert len(logger.log_queue) == 5 + upload_failures: Final = 5 + assert mock_failure.call_count == upload_failures + + for element in elements[5:]: + logger.log_queue.append(element) + + put.failing = False + await logger.flush_queue() + + assert logger.log_queue == [] + landed: Final = frozenset(call[0].rsplit("/", 1)[-1] for call in put.calls[-8:]) + assert landed == frozenset(f"test-{index}.json" for index in range(8)) # calls are (url, data) pairs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "code"), + [ + pytest.param(404, "NoSuchKey", id="404"), + pytest.param(401, None, id="401"), + pytest.param(400, None, id="uncoded-400"), + ], +) +async def test_unlisted_status_gets_one_put_and_stays_queued(status: int, code: str | None) -> 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", + ) + + put = _FailUntilClearedPut(status=status, code=code) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(4)] + + with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + await logger.flush_queue() + + assert len(put.calls) == 4 + mock_sleep.assert_not_awaited() + assert len(logger.log_queue) == 4 + + +class _SyncRecordingClient: + def __init__(self, response: httpx.Response) -> None: + self.response: Final = response + self.put_calls: list = [] # mutable-ok: call log appended once per PUT + + def put(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.put_calls.append(url) + return self.response + + +def test_sync_upload_404_is_single_attempt_without_sleep() -> 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", + ) + + sync_client: Final = _SyncRecordingClient(_coded_failure_response(404, "NoSuchKey")) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client), + patch("time.sleep") as mock_sleep, + ): + logger.upload_data_to_s3(_element({"id": "sync-404"}, "sync-404")) + + assert len(sync_client.put_calls) == 1 + mock_sleep.assert_not_called() + + +@pytest.mark.parametrize( + ("status", "expected_puts", "expected_sleeps"), + [ + pytest.param(429, 1, [], id="429-single"), + pytest.param(408, 1, [], id="408-single"), + pytest.param(502, 1, [], id="502-single"), + pytest.param(504, 1, [], id="504-single"), + pytest.param(503, 3, [call(1), call(2)], id="503-backoff"), + ], +) +def test_sync_upload_retry_set_matches_base(status: int, expected_puts: int, expected_sleeps: list) -> 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", + ) + + sync_client: Final = _SyncRecordingClient(_coded_failure_response(status, "SlowDown")) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client), + patch("time.sleep") as mock_sleep, + ): + logger.upload_data_to_s3(_element({"id": "sync"}, "sync")) + + assert len(sync_client.put_calls) == expected_puts + assert mock_sleep.call_args_list == expected_sleeps + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "code"), + [ + pytest.param(503, "SlowDown", id="503"), + pytest.param(500, "InternalError", id="500"), + pytest.param(403, "AccessDenied", id="access-denied-403"), + ], +) +async def test_retryable_statuses_back_off_three_attempts(status: int, code: str | None) -> 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", + ) + + put = _FailUntilClearedPut(status=status, code=code) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "req"}, "req")] + + with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + await logger.flush_queue() + + assert len(put.calls) == 3 + assert mock_sleep.await_args_list == [call(1), call(2)] + assert len(logger.log_queue) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "code"), + [ + pytest.param(429, "TooManyRequests", id="429"), + pytest.param(408, None, id="408"), + pytest.param(502, None, id="502"), + pytest.param(504, None, id="504"), + ], +) +async def test_non_base_statuses_are_not_retried_in_call(status: int, code: str | None) -> 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", + ) + + put = _FailUntilClearedPut(status=status, code=code) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "req"}, "req")] + + with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + await logger.flush_queue() + + assert len(put.calls) == 1 + assert mock_sleep.await_args_list == [] + assert len(logger.log_queue) == 1 + + +class _FirstFailThenOkPut: + def __init__(self, fail_suffix: str) -> None: + self.fail_suffix = fail_suffix + self.failed_once = False + self.calls: tuple[str, ...] = () + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, url) + if url.endswith(self.fail_suffix) and not self.failed_once: + self.failed_once = True + return _transient_failure_response() + return _ok_response() + + +@pytest.mark.asyncio +async def test_retry_finishes_before_the_next_first_attempt() -> 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=1, + ) + + put = _FirstFailThenOkPut("test-a.json") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "a"}, "a"), _element({"id": "b"}, "b")] + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + await logger.flush_queue() + + assert [call_url.rsplit("/", 1)[-1] for call_url in put.calls] == ["test-a.json", "test-a.json", "test-b.json"] + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_objects_in_backoff_are_bounded_by_the_slot_width() -> 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=2, + ) + + put = _FailUntilClearedPut(status=503) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(20)] + + sleeping: Final[list[int]] = [0] + peak: Final[list[int]] = [0] + + async def counting_sleep(delay: float) -> None: + sleeping[0] += 1 + peak[0] = max(peak[0], sleeping[0]) + for _ in range(10): + await _real_sleep(0) + sleeping[0] -= 1 + + with patch("asyncio.sleep", new=counting_sleep): + await logger.flush_queue() + + assert peak[0] <= 2, f"{peak[0]} objects slept at once, slot width is 2" + assert len(put.calls) == 60 + + +@pytest.mark.asyncio +async def test_subclass_returning_true_drains_the_queue() -> None: + class _TrueUploadLogger(S3Logger): + async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool: + return True + + logger = _TrueUploadLogger( + 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", + ) + logger.async_httpx_client = AsyncMock() + logger.log_queue = [_element({"id": "a"}, "a")] + + await logger.flush_queue() + + assert logger.log_queue == [] + logger.async_httpx_client.put.assert_not_called() + + +@pytest.mark.asyncio +async def test_failed_direct_upload_returns_false() -> 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", + ) + + put = _FailUntilClearedPut(status=500) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + test_element = _element({"id": "x"}, "x") + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + outcome = await logger.async_upload_data_to_s3(test_element) + + assert outcome is False + assert len(put.calls) == 3 + + +@pytest.mark.asyncio +async def test_terminal_drop_of_one_element_does_not_drop_a_sibling_with_the_same_key() -> None: + class _TerminalForMarkerPut: + def __init__(self) -> None: + self.calls: tuple[str | None, ...] = () + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, data) + if data is not None and "terminal-marker" in data: + return _terminal_failure_response() + return _transient_failure_response() + + 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_drop_on_terminal_error=True, + ) + + put = _TerminalForMarkerPut() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + + shared_key = "2025-09-14/shared.json" + dropped = s3BatchLoggingElement( + s3_object_key=shared_key, payload={"m": "terminal-marker"}, s3_object_download_filename="shared.json" + ) + sibling = s3BatchLoggingElement( + s3_object_key=shared_key, payload={"m": "healthy"}, s3_object_download_filename="shared.json" + ) + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await logger._upload_outcome(dropped) == "dropped" + assert await logger._upload_outcome(sibling) == "retry" + + +def test_upload_semaphore_alias_is_the_limiter() -> 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", + ) + + assert logger._upload_semaphore is logger._upload_limiter + + +@pytest.mark.asyncio +async def test_overridden_upload_stays_bounded_by_the_configured_width() -> None: + class _InFlightUploadLogger(S3Logger): + def __init__(self, **kwargs: object) -> None: + super().__init__(**kwargs) + self.in_flight = 0 + self.peak = 0 + + async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool: + self.in_flight += 1 + self.peak = max(self.peak, self.in_flight) + for _ in range(10): + await _real_sleep(0) + self.in_flight -= 1 + return True + + logger = _InFlightUploadLogger( + 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, + ) + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(40)] + + await logger.flush_queue() + + assert logger.peak <= 4 + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_holding_the_semaphore_during_a_direct_upload_does_not_deadlock() -> 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=1, + ) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _RecordingPut() + element = _element({"id": "x"}, "x") + + async def held_upload() -> bool: + async with logger._upload_semaphore: + return await logger.async_upload_data_to_s3(element) + + assert await asyncio.wait_for(held_upload(), timeout=5) is True + + +@pytest.mark.asyncio +async def test_assigning_a_semaphore_changes_the_upload_width() -> 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", + ) + logger._upload_semaphore = asyncio.Semaphore(3) + + put = _CountingPut(width=3) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(30)] + + await logger.flush_queue() + + assert put.peak == 3 + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> 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_drop_on_terminal_error=False, + ) + put = _StatusPut([_coded_failure_response(403, "InvalidRequest"), _ok_response()]) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + failures = AsyncMock() + logger.handle_callback_failure = failures + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True + + assert put.calls == 2 + failures.assert_not_called() + + dropping = 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", + ) + put.calls = 0 + dropping.async_httpx_client = AsyncMock() + dropping.async_httpx_client.put = put + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await dropping.async_upload_data_to_s3(_element({"id": "x"}, "x")) is False + + assert put.calls == 1 + + +def test_sync_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> 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_drop_on_terminal_error=False, + ) + mock_sync_client = MagicMock() + mock_sync_client.put = MagicMock(side_effect=[_coded_failure_response(403, "InvalidRequest"), _ok_response()]) + failures = MagicMock() + logger.handle_callback_failure = failures + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep"), + ): + logger.upload_data_to_s3(_element({"id": "x"}, "x")) + + assert mock_sync_client.put.call_count == 2 + failures.assert_not_called() + + dropping = 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", + ) + mock_sync_client.put = MagicMock(side_effect=[_coded_failure_response(403, "InvalidRequest"), _ok_response()]) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep"), + ): + dropping.upload_data_to_s3(_element({"id": "x"}, "x")) + + assert mock_sync_client.put.call_count == 1 + + +def test_sync_retry_lines_stay_at_warning_level(caplog) -> 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", + ) + mock_sync_client = MagicMock() + mock_sync_client.put = MagicMock( + side_effect=[_transient_failure_response(503), _transient_failure_response(503), _ok_response()] + ) + + with ( + caplog.at_level("WARNING"), + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep"), + ): + logger.upload_data_to_s3(_element({"id": "x"}, "x")) + + assert mock_sync_client.put.call_count == 3 + assert sum(1 for record in caplog.records if "retrying in" in record.getMessage()) == 2 + + +@pytest.mark.asyncio +async def test_direct_async_upload_logs_retry_lines_at_warning_level(caplog) -> 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", + ) + put = _StatusPut([_transient_failure_response(503), _ok_response()]) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + + with caplog.at_level("WARNING"), patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True + + assert put.calls == 2 + assert sum(1 for record in caplog.records if "retrying in" in record.getMessage()) == 1 + + +def _init_bypassed_logger() -> S3Logger: + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + logger = S3Logger.__new__(S3Logger) + logger.iam_cache = BaseAWSLLM._shared_iam_cache + logger.s3_endpoint_url = None + logger.s3_bucket_name = "test-bucket" + logger.s3_region_name = "us-east-1" + logger.s3_use_virtual_hosted_style = False + logger.s3_verify = None + logger.s3_aws_access_key_id = "test-key" + logger.s3_aws_secret_access_key = "test-secret" + logger.s3_aws_session_token = None + logger.s3_aws_session_name = None + logger.s3_aws_profile_name = None + logger.s3_aws_role_name = None + logger.s3_aws_web_identity_token = None + logger.s3_aws_sts_endpoint = None + logger.s3_server_side_encryption = None + logger.s3_sse_kms_key_id = None + logger.s3_log_prompts_only = None + return logger + + +@pytest.mark.asyncio +async def test_init_bypassed_logger_retries_a_503_and_reports_a_404() -> None: + logger = _init_bypassed_logger() + put = _StatusPut([_transient_failure_response(503), _ok_response()]) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await logger.async_upload_data_to_s3(_element({"id": "x"}, "x")) is True + + assert put.calls == 2 + + put.calls = 0 + put.responses = [_coded_failure_response(404, None)] + with patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))): + assert await logger.async_upload_data_to_s3(_element({"id": "y"}, "y")) is False + + assert put.calls == 1 + + +def test_init_bypassed_sync_logger_retries_a_503_and_reports_a_404() -> None: + logger = _init_bypassed_logger() + mock_sync_client = MagicMock() + mock_sync_client.put = MagicMock(side_effect=[_transient_failure_response(503), _ok_response()]) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep"), + ): + logger.upload_data_to_s3(_element({"id": "x"}, "x")) + + assert mock_sync_client.put.call_count == 2 + retried_headers: Final = dict(mock_sync_client.put.call_args.kwargs["headers"]) + assert "X-Amz-Date" in retried_headers + + mock_sync_client.put = MagicMock(return_value=_coded_failure_response(404, None)) + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep"), + ): + logger.upload_data_to_s3(_element({"id": "y"}, "y")) + + assert mock_sync_client.put.call_count == 1 + failed_url: Final = str(mock_sync_client.put.call_args[0][0]) + assert "test-y.json" in failed_url + + +@pytest.mark.asyncio +async def test_subclass_with_base_style_upload_bounded_drains_the_queue() -> None: + class _BaseStyleLogger(S3Logger): + async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool: + return True + + logger = _BaseStyleLogger( + 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", + ) + logger.async_httpx_client = AsyncMock() + logger.log_queue = [_element({"id": "a"}, "a")] + + await logger.flush_queue() + + assert logger.log_queue == [] + logger.async_httpx_client.put.assert_not_called() + + +def test_bool_config_values_fall_back_to_the_default() -> None: + from litellm.integrations.s3 import ( + resolve_s3_max_concurrent_uploads, + resolve_s3_max_queue_size, + resolve_s3_max_retry_age_seconds, + ) + + assert resolve_s3_max_concurrent_uploads(True, 16) == 1 + assert resolve_s3_max_queue_size(True, 50000) == 50000 + assert resolve_s3_max_retry_age_seconds(True, 3600) == 3600 + + +def test_int_env_helper_falls_back_on_non_numeric(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.litellm_core_utils.env_utils import get_env_int + + monkeypatch.setenv("TEST_S3_INT_ENV", "abc") + assert get_env_int("TEST_S3_INT_ENV", 3) == 3 + monkeypatch.setenv("TEST_S3_INT_ENV", "7") + assert get_env_int("TEST_S3_INT_ENV", 3) == 7 + + +class _FailOncePerKeyPut: + def __init__(self) -> None: + self.failed: set[str] = set() + self.calls: tuple[str, ...] = () + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, url) + if url not in self.failed: + self.failed.add(url) + return _transient_failure_response() + return _ok_response() + + +class _SlowFailOncePerKeyPut: + def __init__(self, dumps_count) -> None: + self.failed: set[str] = set() + self.dumps_count = dumps_count + self.first_completed: int | None = None + self.calls: tuple[str, ...] = () + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls = (*self.calls, url) + await _real_sleep(0) + if self.first_completed is None: + self.first_completed = self.dumps_count() + if url not in self.failed: + self.failed.add(url) + return _transient_failure_response() + return _ok_response() + + +@pytest.mark.asyncio +async def test_peak_serialized_bodies_bounded_by_upload_width() -> 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", + ) + + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps as real_safe_dumps + + dumps_calls: list[object] = [] + + def counting_dumps(*args, **kwargs): + dumps_calls.append(args) + return real_safe_dumps(*args, **kwargs) + + put = _SlowFailOncePerKeyPut(lambda: len(dumps_calls)) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(64)] + + with ( + patch("litellm.integrations.s3_v2.safe_dumps", side_effect=counting_dumps), + patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))), + ): + await logger.flush_queue() + + assert put.first_completed is not None + assert put.first_completed <= logger.s3_max_concurrent_uploads + assert len(dumps_calls) == 64 + assert len(put.calls) == 128 + + +@pytest.mark.asyncio +async def test_send_batch_calls_upload_with_one_positional_arg() -> 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", + ) + + uploaded: list[str] = [] # mutable-ok: appended once per upload by the double + + async def mock_upload(batch_logging_element) -> str: + uploaded.append(batch_logging_element.s3_object_key) + return "delivered" + + logger.async_upload_data_to_s3 = mock_upload + logger.log_queue = [_element({"id": "a"}, "a"), _element({"id": "b"}, "b")] + + await logger.flush_queue() + + assert sorted(key.rsplit("/", 1)[-1] for key in uploaded) == ["test-a.json", "test-b.json"] + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_retries_serialize_the_body_once_per_element_per_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", + ) + + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps as real_safe_dumps + + dumps_calls: list[object] = [] + + def counting_dumps(*args, **kwargs): + dumps_calls.append(args) + return real_safe_dumps(*args, **kwargs) + + put = _FailUntilClearedPut(status=503) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(8)] + + with ( + patch("litellm.integrations.s3_v2.safe_dumps", side_effect=counting_dumps), + patch("asyncio.sleep", new=AsyncMock(side_effect=lambda delay: _real_sleep(0))), + ): + await logger.flush_queue() + + assert len(dumps_calls) == 8 + assert len(put.calls) == 24 + assert len(logger.log_queue) == 8 + + +@pytest.mark.asyncio +async def test_async_flush_logs_one_retry_warning(caplog) -> 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", + ) + + put = _FailOncePerKeyPut() + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"i": index}, f"{index}") for index in range(8)] + + with caplog.at_level("WARNING"), patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert logger.log_queue == [] + assert sum(1 for record in caplog.records if "in-call retries" in record.getMessage()) == 1 + assert all("retrying in" not in record.getMessage() for record in caplog.records) + + +class _AppendingSuffixFailingPut: + def __init__(self, logger: S3Logger, element: s3BatchLoggingElement, fail_suffixes: tuple[str, ...]) -> None: + self.logger = logger + self.element = element + self.fail_suffixes = fail_suffixes + 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 url.endswith(self.fail_suffixes): + return _transient_failure_response() + return _ok_response() + + +@pytest.mark.asyncio +async def test_failed_elements_stay_oldest_first_when_requeued() -> 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_retry_age_seconds=3600, + ) + + late = _element({"id": "late"}, "late") + failed = _element({"id": "f2"}, "f2") + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",)) + logger.log_queue = [_element({"id": "f1"}, "f1"), failed] + + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == [failed.s3_object_key, late.s3_object_key] + assert logger.log_queue[0].retrying_since is not None + + +@pytest.mark.asyncio +async def test_overflow_prefers_arrivals_over_failed_elements_without_counting_the_trim(caplog) -> 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_queue_size=1, + ) + + late = _element({"id": "late"}, "late") + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",)) + logger.log_queue = [_element({"id": "f1"}, "f1"), _element({"id": "f2"}, "f2")] + + with ( + patch.object(logger, "handle_callback_failure") as mock_failure, + patch("asyncio.sleep", new_callable=AsyncMock), + ): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == [late.s3_object_key] + failed_uploads: Final = 1 + assert mock_failure.call_count == failed_uploads + assert "dropped 1 oldest events" in caplog.text + + +@pytest.mark.asyncio +async def test_fresh_elements_upload_before_stale_retries_after_a_failed_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", + s3_max_concurrent_uploads=1, + ) + + late = _element({"id": "late"}, "late") + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _AppendingSuffixFailingPut(logger, late, ("test-f2.json",)) + logger.log_queue = [_element({"id": "f1"}, "f1"), _element({"id": "f2"}, "f2")] + + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + recovered = _FailOnSuffixPut(("never-matches",)) + logger.async_httpx_client.put = recovered + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert [call_url.rsplit("/", 1)[-1] for call_url in recovered.calls] == ["test-late.json", "test-f2.json"] + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_repeated_overflow_trims_oldest_across_failed_flushes(caplog) -> 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_queue_size=3, + ) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _AppendingFailingPut(logger, (_element({"id": "d"}, "d"),)) + logger.log_queue = [_element({"id": name}, name) for name in ("a", "b", "c")] + + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert [element.payload["id"] for element in logger.log_queue] == ["b", "c", "d"] + + logger.async_httpx_client.put = _AppendingFailingPut(logger, (_element({"id": "e"}, "e"),)) + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert [element.payload["id"] for element in logger.log_queue] == ["c", "d", "e"] + assert caplog.text.count("dropped 1 oldest events") == 2 + + +@pytest.mark.asyncio +async def test_retry_age_budget_drops_after_the_clock_set_by_a_partial_failure(caplog) -> 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_retry_age_seconds=1, + ) + + put = _FailOnSuffixCodedPut(("test-poison.json",), 503, "SlowDown") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "poison"}, "poison"), _element({"id": "good"}, "good")] + + t0: Final = _NOW + with patch.object(logger, "handle_callback_failure") as mock_failure: + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=t0), + ): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == ["2025-09-14/test-poison.json"] + assert logger.log_queue[0].retrying_since == t0 + + logger.log_queue.append(_element({"id": "good-2"}, "good-2")) + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=t0 + 2), + ): + await logger.flush_queue() + + assert logger.log_queue == [] + assert "retrying longer than s3_max_retry_age_seconds=1" in caplog.text + poison_puts: Final = sum(1 for call_url in put.calls if call_url.endswith("test-poison.json")) + assert poison_puts == 6 + upload_failures: Final = 2 + assert mock_failure.call_count == upload_failures + + +@pytest.mark.asyncio +async def test_the_retry_clock_starts_at_the_first_partial_failure_not_first_seen() -> 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_retry_age_seconds=1, + ) + + put = _FailUntilClearedPut(status=503, code="SlowDown") + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "poison"}, "poison")] + + t0: Final = _NOW + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=t0), + ): + await logger.flush_queue() + + assert len(logger.log_queue) == 1 + assert logger.log_queue[0].retrying_since is None + + logger.log_queue.append(_element({"id": "good"}, "good")) + put.failing = False + failing_poison: Final = _FailOnSuffixCodedPut(("test-poison.json",), 503, "SlowDown") + logger.async_httpx_client.put = failing_poison + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=t0 + 500), + ): + await logger.flush_queue() + + assert [element.s3_object_key for element in logger.log_queue] == ["2025-09-14/test-poison.json"] + assert logger.log_queue[0].retrying_since == t0 + 500 + + +def test_sync_upload_retries_access_denied_403(caplog): + 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-403.json", + payload={"test": "sync-403"}, + s3_object_download_filename="test-sync-403.json", + ) + + mock_sync_client = MagicMock() + mock_sync_client.put = MagicMock(return_value=_coded_failure_response(403, "AccessDenied")) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep") as mock_sleep, + ): + logger.upload_data_to_s3(test_element) + + assert mock_sync_client.put.call_count == 3 + assert mock_sleep.call_args_list == [call(1), call(2)] + assert "dropping object" not in caplog.text + + +def test_sync_upload_drops_terminal_object_once_and_logs_it_only_when_opted_in(caplog): + 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_drop_on_terminal_error=True, + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-sync-terminal.json", + payload={"test": "sync-terminal"}, + s3_object_download_filename="test-sync-terminal.json", + ) + + mock_sync_client = MagicMock() + mock_sync_client.put = MagicMock(return_value=_coded_failure_response(400, "EntityTooLarge")) + + with ( + patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("time.sleep") as mock_sleep, + ): + logger.upload_data_to_s3(test_element) + + assert mock_sync_client.put.call_count == 1 + mock_sleep.assert_not_called() + assert "dropping object" in caplog.text + + +@pytest.mark.asyncio +async def test_requeued_batch_file_keeps_the_earliest_member_retrying_since() -> 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 + + stale: Final = _NOW - 30 + retried = _element({"id": "retried"}, "retried").model_copy(update={"retrying_since": stale}) + fresh = _element({"id": "fresh"}, "fresh") + logger.log_queue = [retried, fresh] + + with ( + patch("asyncio.sleep", new_callable=AsyncMock), + patch("time.monotonic", return_value=_NOW), + ): + await logger.flush_queue() + + assert len(logger.log_queue) == 1 + assert logger.log_queue[0].s3_object_key.endswith(".jsonl") + assert logger.log_queue[0].retrying_since == stale + + +@pytest.mark.asyncio +async def test_an_unlisted_5xx_is_requeued_without_an_extra_attempt() -> 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", + ) + + put = _FailUntilClearedPut(status=507) + + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = put + logger.log_queue = [_element({"id": "req-507"}, "507")] + + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert len(logger.log_queue) == 1 + assert len(put.calls) == 1 + + +@pytest.mark.parametrize("configured", [0, "0", None, ""]) +def test_retry_age_resolution_disables_the_budget(configured: object) -> None: + from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds + + assert resolve_s3_max_retry_age_seconds(configured, 3600) is None + + +@pytest.mark.parametrize("configured", ["abc", -5, True]) +def test_invalid_retry_age_resolution_falls_back_with_a_warning(configured: object, caplog) -> None: + from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds + + assert resolve_s3_max_retry_age_seconds(configured, 3600) == 3600 + assert "s3_max_retry_age_seconds" in caplog.text + + +def test_retry_age_resolution_accepts_a_positive_int() -> None: + from litellm.integrations.s3 import resolve_s3_max_retry_age_seconds + + assert resolve_s3_max_retry_age_seconds(30, 3600) == 30 + + +def test_default_logger_sets_a_one_hour_retry_age_budget() -> 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", + ) + + assert logger.s3_max_retry_age_seconds == 3600 + + +def test_constructor_zero_disables_the_retry_age_budget() -> 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_retry_age_seconds=0, + ) + + assert logger.s3_max_retry_age_seconds is None + + +def test_invalid_callback_params_retry_age_falls_back_to_the_default() -> 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_callback_params_override={"s3_max_retry_age_seconds": "abc"}, + ) + + assert logger.s3_max_retry_age_seconds == 3600 + + +def test_callback_params_retry_age_wins_over_constructor() -> 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_retry_age_seconds=30, + s3_callback_params_override={"s3_max_retry_age_seconds": 60}, + ) + + assert logger.s3_max_retry_age_seconds == 60 + + +def test_callback_params_drop_terminal_error_wins_over_constructor() -> 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_drop_on_terminal_error=False, + s3_callback_params_override={"s3_drop_on_terminal_error": True}, + ) + + assert logger.s3_drop_on_terminal_error is True + + +def test_invalid_callback_params_drop_terminal_error_falls_back_to_constructor_value() -> 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_drop_on_terminal_error=True, + s3_callback_params_override={"s3_drop_on_terminal_error": "banana"}, + ) + + assert logger.s3_drop_on_terminal_error is True + + +def test_callback_params_adaptive_concurrency_wins_over_constructor() -> 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_adaptive_concurrency=False, + s3_callback_params_override={"s3_adaptive_concurrency": "true"}, + ) + + assert logger.s3_adaptive_concurrency is True + assert logger._upload_limiter._ceiling > logger._upload_limiter.limit + + +def test_invalid_callback_params_max_adaptive_concurrency_falls_back_to_default() -> None: + from litellm.constants import DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY + + logger = _override_logger(s3_adaptive_concurrency=True, s3_max_adaptive_concurrency="abc") + + assert logger.s3_max_adaptive_concurrency == DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY + assert logger._upload_limiter._ceiling == DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY + + +def test_callback_params_queue_size_wins_over_constructor() -> 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_queue_size=7, + s3_callback_params_override={"s3_max_queue_size": 4}, + ) + + assert logger.s3_max_queue_size == 4 + assert logger.max_queue_size == 4 + + +def test_invalid_callback_params_queue_size_falls_back_to_constructor_value() -> 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_queue_size=7, + s3_callback_params_override={"s3_max_queue_size": "abc"}, + ) + + assert logger.s3_max_queue_size == 7 + + +def test_invalid_constructor_queue_size_falls_back_to_default() -> None: + from litellm.integrations.custom_batch_logger import CustomBatchLogger + + 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_queue_size="abc", + ) + + assert logger.s3_max_queue_size == CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE + assert logger.max_queue_size == CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE + + +class _StatusPut: + def __init__(self, responses: "list[MagicMock | Exception]") -> None: + self.responses = responses + self.calls = 0 + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls += 1 + outcome = self.responses[min(self.calls - 1, len(self.responses) - 1)] + if isinstance(outcome, Exception): + raise outcome + return outcome + + +def _slow_down_response(status: int = 200) -> MagicMock: + response = _ok_response() if status == 200 else _transient_failure_response(status) + response.text = "SlowDown" + return response + + +@pytest.mark.asyncio +async def test_503_response_lowers_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut([_transient_failure_response(503), _ok_response()]) + + logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth + with patch("asyncio.sleep", new_callable=AsyncMock): + logger.log_queue = [_element({"i": 0}, "0")] + await logger.flush_queue() + + assert logger._upload_limiter.limit == 32 + + +@pytest.mark.asyncio +async def test_429_response_lowers_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut([_transient_failure_response(429), _ok_response()]) + + logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth + with patch("asyncio.sleep", new_callable=AsyncMock): + logger.log_queue = [_element({"i": 0}, "0")] + await logger.flush_queue() + + assert logger._upload_limiter.limit == 32 + + +@pytest.mark.asyncio +async def test_slow_down_body_code_lowers_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut([_slow_down_response()]) + + logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth + logger.log_queue = [_element({"i": 0}, "0")] + await logger.async_send_batch() + + assert logger._upload_limiter.limit == 32 + + +@pytest.mark.asyncio +async def test_transport_error_lowers_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut( + [httpx.ConnectError("connect refused", request=MagicMock()), _ok_response()] + ) + + logger._upload_limiter._limit = 64 # mutable-ok: seed the AIMD state above the floor without replaying growth + with patch("asyncio.sleep", new_callable=AsyncMock): + logger.log_queue = [_element({"i": 0}, "0")] + await logger.flush_queue() + + assert logger._upload_limiter.limit == 32 + + +@pytest.mark.asyncio +async def test_fast_uploads_raise_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _RecordingPut() + + before: Final = logger._upload_limiter.limit + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(before)] + await logger.async_send_batch() + + assert logger._upload_limiter.limit > before + + +@pytest.mark.asyncio +async def test_configured_concurrency_is_the_fixed_limit_when_adaptive_is_off() -> 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=64, + ) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut([_transient_failure_response(503), _ok_response()]) + + assert logger._upload_limiter._value == 64 + + logger.log_queue = [_element({"i": 0}, "0")] + with patch("asyncio.sleep", new_callable=AsyncMock): + await logger.flush_queue() + + assert logger._upload_limiter._value == 64 + + +@pytest.mark.asyncio +async def test_the_limit_never_falls_below_the_configured_width() -> None: + logger = _override_logger(s3_adaptive_concurrency=True, s3_max_concurrent_uploads=8) + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _StatusPut( + [_transient_failure_response(503), _transient_failure_response(503), _transient_failure_response(503)] + ) + + with patch("asyncio.sleep", new_callable=AsyncMock): + logger.log_queue = [_element({"i": 0}, "0")] + await logger.flush_queue() + + assert logger._upload_limiter.limit == 8 + + +def test_default_upload_width_is_16() -> None: + from litellm.constants import DEFAULT_S3_MAX_CONCURRENT_UPLOADS + + logger = _override_logger() + + assert logger._upload_semaphore._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS + assert DEFAULT_S3_MAX_CONCURRENT_UPLOADS == 16 + + +@pytest.mark.asyncio +async def test_a_slow_put_does_not_lower_the_adaptive_limit() -> None: + logger = _override_logger(s3_adaptive_concurrency=True) + logger.async_httpx_client = AsyncMock() + + async def slow_put(url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + await _real_sleep(0) + return _ok_response() + + logger.async_httpx_client.put = slow_put + + before: Final = logger._upload_limiter.limit + logger.log_queue = [_element({"i": 0}, "0")] + await logger.async_send_batch() + + assert logger._upload_limiter.limit >= before + + +class _FastOkPut: + def __init__(self) -> None: + self.calls = 0 + + async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock: + self.calls += 1 + await _real_sleep(0) + return _ok_response() + + +async def _timed_send_batch(size: int) -> float: + logger = _override_logger() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put = _FastOkPut() + logger.log_queue = [_element({"i": i}, f"{i}") for i in range(size)] + started = time.perf_counter() + await logger.async_send_batch() + return time.perf_counter() - started + + +@pytest.mark.asyncio +async def test_send_batch_time_grows_linearly_with_the_batch() -> None: + baseline: Final = await _timed_send_batch(2_000) + quadrupled: Final = await _timed_send_batch(8_000) + + assert quadrupled / baseline < 8, f"2k took {baseline:.3f}s, 8k took {quadrupled:.3f}s"