Merge pull request #38394 from BerriAI/litellm_lit6253_cache_hit_callback_flush

fix(logging_worker): rescue dequeued logging tasks lost at event loop close
This commit is contained in:
Mateo Wang 2026-08-26 13:50:02 -07:00 • committed by GitHub
commit 3c24f37502
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 134 additions and 12 deletions

View file

@ -4,6 +4,7 @@
import asyncio
import atexit
import contextvars
import inspect
import logging
from collections.abc import Coroutine, Iterator
from typing import Final
@ -53,6 +54,7 @@ class LoggingWorker:
self._queue: asyncio.Queue[LoggingTask] | None = None
self._worker_task: asyncio.Task | None = None
self._running_tasks: set[asyncio.Task] = set()
self._dequeued_tasks: dict[int, LoggingTask] = {} # mutable-ok: refs so flush can rescue never-started tasks
self._sem: asyncio.Semaphore | None = None
self._bound_loop: asyncio.AbstractEventLoop | None = None
self._last_aggressive_clear_time: float = 0.0
@ -61,6 +63,38 @@ class LoggingWorker:
# Register cleanup handler to flush remaining events on exit
atexit.register(self._flush_on_exit)
def _track_dequeued(self, task: LoggingTask) -> None:
self._dequeued_tasks[id(task)] = task
def _untrack_dequeued(self, task: LoggingTask) -> None:
self._dequeued_tasks.pop(id(task), None)
def _unstarted_dequeued_tasks(self) -> tuple[LoggingTask, ...]:
return tuple(
task
for task in self._dequeued_tasks.values()
if inspect.getcoroutinestate(task["coroutine"]) == inspect.CORO_CREATED
)
def _requeue_unstarted_dequeued(self, new_queue: "asyncio.Queue[LoggingTask]") -> int:
revived: Final = self._unstarted_dequeued_tasks()
self._dequeued_tasks.clear()
for index, revived_task in enumerate(revived):
try:
new_queue.put_nowait(revived_task)
except asyncio.QueueFull:
for leftover in revived[index:]:
self._track_dequeued(leftover)
return index
return len(revived)
def _run_coroutine_silently(self, loop: asyncio.AbstractEventLoop, coroutine: Coroutine) -> bool:
try:
loop.run_until_complete(asyncio.wait_for(coroutine, timeout=self.timeout))
except (Exception, asyncio.CancelledError): # noqa: BLE001 # atexit flush must never break the user's program
return False
return True
@staticmethod
def _drain_pending(queue: "asyncio.Queue[LoggingTask]") -> tuple[LoggingTask, ...]:
"""Pop every task still queued, without awaiting them, so they can be moved to another queue."""
@ -90,10 +124,12 @@ class LoggingWorker:
new_queue: Final[asyncio.Queue[LoggingTask]] = asyncio.Queue(maxsize=self.max_queue_size)
for carried_task in carried_over:
new_queue.put_nowait(carried_task)
if carried_over:
revived_count: Final = self._requeue_unstarted_dequeued(new_queue)
if carried_over or revived_count:
verbose_logger.warning(
"LoggingWorker: event loop changed; carried %d pending logging task(s) onto the new loop",
"LoggingWorker: event loop changed; carried %d pending and revived %d dequeued logging task(s) onto the new loop",
len(carried_over),
revived_count,
)
else:
verbose_logger.debug("LoggingWorker: Event loop changed, reinitializing queue and worker")
@ -129,6 +165,7 @@ class LoggingWorker:
except Exception as e:
verbose_logger.exception("LoggingWorker error: %s", e)
finally:
self._untrack_dequeued(task)
self._queue.task_done()
finally:
# Always release semaphore, even if queue is None
@ -146,6 +183,7 @@ class LoggingWorker:
await self._sem.acquire()
try:
task = await self._queue.get()
self._track_dequeued(task)
# Track each spawned coroutine so we can cancel on shutdown.
processing_task = asyncio.create_task(self._process_log_task(task, self._sem))
self._running_tasks.add(processing_task)
@ -298,9 +336,10 @@ class LoggingWorker:
extracted_tasks: Final = []
for _ in range(items_to_extract):
try:
extracted_tasks.append(self._queue.get_nowait())
extracted_tasks.append(extracted := self._queue.get_nowait())
except asyncio.QueueEmpty:
break
self._track_dequeued(extracted)
return extracted_tasks
@ -318,6 +357,7 @@ class LoggingWorker:
# Add new task to extracted tasks to process directly
if new_task is not None:
self._track_dequeued(new_task)
extracted_tasks.append(new_task)
# Process extracted tasks directly
@ -343,6 +383,7 @@ class LoggingWorker:
# Suppress errors during processing to ensure we keep going
pass
finally:
self._untrack_dequeued(task)
self._queue.task_done()
async def _process_extracted_tasks(self, tasks: list[LoggingTask]) -> None:
@ -486,11 +527,12 @@ class LoggingWorker:
self._safe_log("debug", "[LoggingWorker] atexit: No queue initialized")
return
if self._queue.empty():
unstarted_dequeued: Final = self._unstarted_dequeued_tasks()
if self._queue.empty() and not unstarted_dequeued:
self._safe_log("debug", "[LoggingWorker] atexit: Queue is empty")
return
queue_size: Final = self._queue.qsize()
queue_size: Final = self._queue.qsize() + len(unstarted_dequeued)
self._safe_log("info", f"[LoggingWorker] atexit: Flushing {queue_size} remaining events...")
# Create a new event loop since the original is closed
@ -509,6 +551,16 @@ class LoggingWorker:
previous_raise_exceptions: Final = logging.raiseExceptions
logging.raiseExceptions = False
try:
for pending in unstarted_dequeued:
if (
processed >= MAX_ITERATIONS_TO_CLEAR_QUEUE
or loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE
):
break
if self._run_coroutine_silently(loop, pending["coroutine"]):
processed += 1
self._untrack_dequeued(pending)
while not self._queue.empty() and processed < MAX_ITERATIONS_TO_CLEAR_QUEUE:
if loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE:
self._safe_log(
@ -526,11 +578,8 @@ class LoggingWorker:
# Note: We run the coroutine directly, not via create_task,
# since we're in a new event loop context
try:
loop.run_until_complete(task["coroutine"])
processed += 1
except Exception:
# Silent failure to not break user's program
pass
if self._run_coroutine_silently(loop, task["coroutine"]):
processed += 1
finally:
# Clear reference to prevent memory leaks
task = None

View file

@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2918
"limit": 2917
},
"C401": {
"limit": 8
@ -189,7 +189,7 @@
"limit": 0
},
"S110": {
"limit": 218
"limit": 217
},
"S112": {
"limit": 22

View file

@ -97,6 +97,32 @@ class TestLoggingWorker:
logging.raiseExceptions = previous_raise_exceptions
logger.removeHandler(handler)
def test_flush_on_exit_rescues_dequeued_coroutine_never_started(self):
"""
Regression test for cache-hit success callbacks lost in short-lived SDK scripts:
the worker loop dequeues the task, then ``asyncio.run`` cancels the processing
task before it ever runs, so the coroutine leaves the queue without being
awaited and the atexit flush used to find an empty queue and rescue nothing.
"""
worker = LoggingWorker(timeout=1.0, max_queue_size=10)
fired = []
async def marker():
fired.append(True)
async def short_lived_script():
worker.ensure_initialized_and_enqueue(marker())
asyncio.run(short_lived_script())
assert worker._queue is not None
assert worker._queue.qsize() == 0, "precondition: the worker loop dequeued the task before loop close"
assert fired == [], "precondition: the callback never ran before loop close"
worker._flush_on_exit()
assert fired == [True]
def test_flush_on_exit_swallows_errors_and_drains_remaining(self):
"""A failing queued coroutine must not abort the atexit drain of later events."""
worker = LoggingWorker(timeout=1.0, max_queue_size=10)
@ -118,6 +144,53 @@ class TestLoggingWorker:
assert processed == ["ran"]
assert worker._queue.empty()
def test_loop_change_revives_dequeued_coroutine_on_new_loop(self):
"""
A callback dequeued but never started before its loop closed must run on the
next event loop's worker instead of staying stranded until process exit.
"""
worker = LoggingWorker(timeout=1.0, max_queue_size=10)
fired = []
async def marker(name):
fired.append(name)
async def first_script():
worker.ensure_initialized_and_enqueue(marker("first"))
asyncio.run(first_script())
assert fired == [], "precondition: the callback was dequeued but never ran before loop close"
async def second_script():
worker.ensure_initialized_and_enqueue(marker("second"))
assert worker._queue is not None
await asyncio.wait_for(worker._queue.join(), timeout=5)
asyncio.run(second_script())
assert sorted(fired) == ["first", "second"]
def test_flush_on_exit_swallows_cancellation_and_drains_remaining(self):
"""A callback raising CancelledError must not abort the atexit flush of later events."""
worker = LoggingWorker(timeout=1.0, max_queue_size=10)
worker._queue = asyncio.Queue(maxsize=10)
processed = []
async def cancels_during_flush():
raise asyncio.CancelledError()
async def records_during_flush():
processed.append("ran")
worker.enqueue(cancels_during_flush())
worker.enqueue(records_during_flush())
worker._flush_on_exit()
assert processed == ["ran"]
assert worker._queue.empty()
@pytest.mark.asyncio
async def test_worker_handles_cancellation_gracefully(self, logging_worker):
"""Test that the worker handles cancellation without throwing exceptions."""