mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
3c24f37502
3 changed files with 134 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue