propagate execution context into logging tasks

This commit is contained in:
Din 2025-09-11 14:53:37 +01:00
parent 258b674dbb
commit ee5a9d0aa0
3 changed files with 202 additions and 28 deletions

View file

@ -1,10 +1,21 @@
import asyncio
import contextlib
from typing import Coroutine, Optional
import contextvars
from typing import Coroutine, Optional, TypedDict
from litellm._logging import verbose_logger
class LoggingTask(TypedDict):
"""
A logging task with its associated context to ensure logging is executed in
the original task's context.
"""
coroutine: Coroutine
context: contextvars.Context
class LoggingWorker:
"""
A simple, async logging worker that processes log coroutines in the background.
@ -13,77 +24,84 @@ class LoggingWorker:
This leads to a +200 RPS performance improvement when using LiteLLM Python SDK or Proxy Server.
- Use this to queue coroutine tasks that are not critical to the main flow of the application. e.g Success/Error callbacks, logging, etc.
"""
LOGGING_WORKER_MAX_QUEUE_SIZE = 50_000
LOGGING_WORKER_MAX_TIME_PER_COROUTINE = 20.0
MAX_ITERATIONS_TO_CLEAR_QUEUE = 200
MAX_TIME_TO_CLEAR_QUEUE = 5.0
def __init__(
self,
timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
self,
timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
max_queue_size: int = LOGGING_WORKER_MAX_QUEUE_SIZE,
):
self.timeout = timeout
self.max_queue_size = max_queue_size
self._queue: Optional[asyncio.Queue] = None
self._queue: Optional[asyncio.Queue[LoggingTask]] = None
self._worker_task: Optional[asyncio.Task] = None
def _ensure_queue(self) -> None:
"""Initialize the queue if it doesn't exist."""
if self._queue is None:
self._queue = asyncio.Queue(maxsize=self.max_queue_size)
def start(self) -> None:
"""Start the logging worker. Idempotent - safe to call multiple times."""
self._ensure_queue()
if self._worker_task is None or self._worker_task.done():
self._worker_task = asyncio.create_task(self._worker_loop())
async def _worker_loop(self) -> None:
"""Main worker loop that processes log coroutines sequentially."""
try:
if self._queue is None:
return
while True:
# Process one coroutine at a time to keep event loop load predictable
coroutine = await self._queue.get()
task = await self._queue.get()
try:
await asyncio.wait_for(coroutine, timeout=self.timeout)
# Run the coroutine in its original context
await asyncio.wait_for(
task["context"].run(asyncio.create_task, task["coroutine"]),
timeout=self.timeout,
)
except Exception as e:
verbose_logger.exception(f"LoggingWorker error: {e}")
pass
finally:
self._queue.task_done()
except asyncio.CancelledError:
verbose_logger.debug("LoggingWorker cancelled during shutdown")
# Attempt to clear remaining items to prevent "never awaited" warnings
await self.clear_queue()
def enqueue(self, coroutine: Coroutine) -> None:
"""
Add a coroutine to the logging queue.
Add a coroutine to the logging queue.
Hot path: never blocks, drops logs if queue is full.
"""
if self._queue is None:
return
try:
self._queue.put_nowait(coroutine)
# Capture the current context when enqueueing
task = LoggingTask(coroutine=coroutine, context=contextvars.copy_context())
self._queue.put_nowait(task)
except asyncio.QueueFull as e:
verbose_logger.exception(f"LoggingWorker queue is full: {e}")
# Drop logs on overload to protect request throughput
pass
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine):
"""
Ensure the logging worker is initialized and enqueue the coroutine.
"""
self.start()
self.enqueue(async_coroutine)
async def stop(self) -> None:
"""Stop the logging worker and clean up resources."""
if self._worker_task:
@ -91,34 +109,42 @@ class LoggingWorker:
with contextlib.suppress(Exception):
await self._worker_task
self._worker_task = None
async def flush(self) -> None:
"""Flush the logging queue."""
if self._queue is None:
return
while not self._queue.empty():
await self._queue.join()
async def clear_queue(self):
"""
Clear the queue with a maximum time limit.
"""
if self._queue is None:
return
start_time = asyncio.get_event_loop().time()
for _ in range(self.MAX_ITERATIONS_TO_CLEAR_QUEUE):
# Check if we've exceeded the maximum time
if asyncio.get_event_loop().time() - start_time >= self.MAX_TIME_TO_CLEAR_QUEUE:
verbose_logger.warning(f"clear_queue exceeded max_time of {self.MAX_TIME_TO_CLEAR_QUEUE}s, stopping early")
if (
asyncio.get_event_loop().time() - start_time
>= self.MAX_TIME_TO_CLEAR_QUEUE
):
verbose_logger.warning(
f"clear_queue exceeded max_time of {self.MAX_TIME_TO_CLEAR_QUEUE}s, stopping early"
)
break
try:
coroutine = self._queue.get_nowait()
task = self._queue.get_nowait()
# Await the coroutine to properly execute and avoid "never awaited" warnings
try:
await asyncio.wait_for(coroutine, timeout=self.timeout)
await asyncio.wait_for(
task["context"].run(asyncio.create_task, task["coroutine"]),
timeout=self.timeout,
)
except Exception:
# Suppress errors during cleanup
pass
@ -129,4 +155,3 @@ class LoggingWorker:
# Global instance for backward compatibility
GLOBAL_LOGGING_WORKER = LoggingWorker()

View file

@ -467,3 +467,73 @@ async def test_e2e_generate_cold_storage_object_key_not_configured():
assert result is None
@pytest.mark.asyncio
async def test_logging_opentelemetry_context_propagation():
"""
Test that OpenTelemtry context propagation works with async completion.
"""
import asyncio
import litellm
from litellm.integrations.custom_logger import CustomLogger
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
provider = TracerProvider()
exporter = InMemorySpanExporter()
provider.add_span_processor(SimpleSpanProcessor(exporter))
trace.set_tracer_provider(provider)
tracer = trace.get_tracer(__name__)
class MockOpenTelemetryLogger(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
span = tracer.start_span(start_time=start_time.timestamp() * 1e9, name="async_log_success_event")
span.end(end_time=end_time)
mock_logging_obj = MockOpenTelemetryLogger()
litellm.callbacks = [mock_logging_obj]
with tracer.start_as_current_span("span_1") as span:
span_1_id = span.get_span_context().span_id
await litellm.acompletion(
max_tokens=100,
messages=[{"role": "user", "content": "Hey"}],
model="openai/codex-mini-latest",
mock_response="Hello, world!",
)
with tracer.start_as_current_span("span_2") as span:
span_2_id = span.get_span_context().span_id
await litellm.acompletion(
max_tokens=100,
messages=[{"role": "user", "content": "Hey"}],
model="openai/codex-mini-latest",
mock_response="Hello, world!",
)
await asyncio.sleep(1)
spans = exporter.get_finished_spans()
assert len(spans) == 4
assert span_1_id != span_2_id
sorted_spans = sorted(list(spans), key=lambda x: x.start_time or 0)
assert sorted_spans[0].name == "span_1"
assert sorted_spans[1].name == "async_log_success_event"
assert sorted_spans[2].name == "span_2"
assert sorted_spans[3].name == "async_log_success_event"
first_span_context = sorted_spans[0].get_span_context()
assert first_span_context is not None and first_span_context.span_id == span_1_id
second_span_context = sorted_spans[2].get_span_context()
assert second_span_context is not None and second_span_context.span_id == span_2_id
first_completion_span_parent = sorted_spans[1].parent
assert first_completion_span_parent is not None and first_completion_span_parent.span_id == span_1_id
# This check would fail without the proper context propagation, and span[3] would end up with span_1_id as the parent
second_completion_span_parent = sorted_spans[3].parent
assert second_completion_span_parent is not None and second_completion_span_parent.span_id == span_2_id

View file

@ -2,6 +2,7 @@
Tests for the LoggingWorker class to ensure graceful shutdown handling.
"""
import asyncio
import contextvars
import pytest
from unittest.mock import AsyncMock, patch
@ -139,3 +140,81 @@ class TestLoggingWorker:
# Should have logged queue full exceptions
exception_calls = [call for call in mock_logger.exception.call_args_list if "queue is full" in str(call)]
assert len(exception_calls) > 0
@pytest.mark.asyncio
async def test_context_propagation(self, logging_worker):
"""Test that enqueued tasks execute in their original context."""
# Create a context variable for testing
test_context_var: contextvars.ContextVar[str] = contextvars.ContextVar('test_context_var')
# Track results from multiple tasks
task_results = []
async def test_task(task_id: str):
"""A test coroutine that checks if it can access the context variable."""
# Sleep a bit to simulate real work and ensure context persists
await asyncio.sleep(0.1)
try:
# Try to get the context variable value
value = test_context_var.get()
task_results.append({
'task_id': task_id,
'context_value': value,
'context_accessible': True
})
except LookupError:
# Context variable not found
task_results.append({
'task_id': task_id,
'context_accessible': False,
'context_value': None
})
# Start the logging worker
logging_worker.start()
# Create two separate contexts and enqueue tasks from each
# Context 1: Set context var to "context_1"
ctx1 = contextvars.copy_context()
ctx1.run(test_context_var.set, "context_1")
ctx1.run(logging_worker.enqueue, test_task("task_1"))
# Context 2: Set context var to "context_2"
ctx2 = contextvars.copy_context()
ctx2.run(test_context_var.set, "context_2")
ctx2.run(logging_worker.enqueue, test_task("task_2"))
# Context 3: No context variable set (should get LookupError)
ctx3 = contextvars.copy_context()
ctx3.run(logging_worker.enqueue, test_task("task_3"))
# Wait for all tasks to be processed
await asyncio.sleep(0.5)
# Stop the worker
await logging_worker.stop()
# Sort results by task_id for consistent testing
task_results.sort(key=lambda x: x['task_id'])
# Verify that each task saw its own context
assert len(task_results) == 3, f"Expected 3 results, got {len(task_results)}"
# Task 1 should see "context_1"
task1_result = next((r for r in task_results if r['task_id'] == 'task_1'), None)
assert task1_result is not None, "Task 1 result not found"
assert task1_result['context_accessible'] is True, "Task 1 should have access to context variable"
assert task1_result['context_value'] == "context_1", f"Task 1 should see 'context_1', got: {task1_result['context_value']}"
# Task 2 should see "context_2"
task2_result = next((r for r in task_results if r['task_id'] == 'task_2'), None)
assert task2_result is not None, "Task 2 result not found"
assert task2_result['context_accessible'] is True, "Task 2 should have access to context variable"
assert task2_result['context_value'] == "context_2", f"Task 2 should see 'context_2', got: {task2_result['context_value']}"
# Task 3 should not have access to the context variable
task3_result = next((r for r in task_results if r['task_id'] == 'task_3'), None)
assert task3_result is not None, "Task 3 result not found"
assert task3_result['context_accessible'] is False, "Task 3 should not have access to context variable"