From ee5a9d0aa0588821bce3bb8575ab1989bec24238 Mon Sep 17 00:00:00 2001 From: Din Date: Thu, 11 Sep 2025 14:53:37 +0100 Subject: [PATCH] propagate execution context into logging tasks --- litellm/litellm_core_utils/logging_worker.py | 81 ++++++++++++------- .../test_litellm_logging.py | 70 ++++++++++++++++ .../litellm_core_utils/test_logging_worker.py | 79 ++++++++++++++++++ 3 files changed, 202 insertions(+), 28 deletions(-) diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 3f83719dd32..16860ad6852 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -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() - diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 48a22dcc8af..fa5164b9c16 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_logging_worker.py b/tests/test_litellm/litellm_core_utils/test_logging_worker.py index 24c77339025..47e626e04af 100644 --- a/tests/test_litellm/litellm_core_utils/test_logging_worker.py +++ b/tests/test_litellm/litellm_core_utils/test_logging_worker.py @@ -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"