mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
propagate execution context into logging tasks
This commit is contained in:
parent
258b674dbb
commit
ee5a9d0aa0
3 changed files with 202 additions and 28 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue