diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a7c7786aa91..98d6c08b7fc 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -37,6 +37,7 @@ from litellm._uuid import uuid from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.base_llm.managed_resources.utils import ( resolve_passthrough_managed_id_provider, @@ -1310,8 +1311,8 @@ async def pass_through_request( ## LOG SUCCESS passthrough_logging_payload["response_body"] = response_body end_time = datetime.now() - asyncio.create_task( - pass_through_endpoint_logging.pass_through_async_success_handler( + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( httpx_response=response, response_body=response_body, url_route=str(url), @@ -2153,8 +2154,8 @@ async def websocket_passthrough_request( mock_response = MockWebSocketResponse(target) # Use the same success handler as HTTP passthrough endpoints - asyncio.create_task( - pass_through_endpoint_logging.pass_through_async_success_handler( + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( httpx_response=mock_response, # type: ignore response_body=websocket_messages, # type: ignore url_route=endpoint or "", diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 7a725472dd7..3a5728fd66f 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -1,4 +1,3 @@ -import asyncio from datetime import datetime from typing import List, Optional, Tuple @@ -7,6 +6,7 @@ import httpx import litellm from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.proxy._types import PassThroughEndpointLoggingResultValues from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType @@ -92,8 +92,8 @@ class PassThroughStreamingHandler: if not logging_scheduled and raw_bytes: logging_scheduled = True try: - asyncio.create_task( - PassThroughStreamingHandler._route_streaming_logging_to_handler( + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py index 38990644154..b781190eaef 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.proxy.pass_through_endpoints.streaming_handler import ( PassThroughStreamingHandler, ) @@ -120,6 +121,91 @@ async def test_chunk_processor_does_not_schedule_logging_when_no_chunks(): mock_route.assert_not_called() +@pytest.mark.asyncio +async def test_chunk_processor_routes_logging_through_logging_worker(): + """The spend-log coroutine must be handed to the durable logging worker, which + keeps a strong reference and drains on shutdown, instead of a bare + asyncio.create_task that the event loop only weak-references and can drop + under GC/load, silently losing the SpendLogs row for a successful call.""" + chunks = [b"chunk-1", b"chunk-2"] + response = _make_streaming_response(chunks) + + enqueued = [] + + def _capture(async_coroutine): + enqueued.append(async_coroutine) + async_coroutine.close() + + with ( + patch.object( + PassThroughStreamingHandler, + "_route_streaming_logging_to_handler", + new=AsyncMock(), + ), + patch.object( + GLOBAL_LOGGING_WORKER, + "ensure_initialized_and_enqueue", + side_effect=_capture, + ) as mock_enqueue, + ): + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + ): + received.append(chunk) + + assert received == chunks + mock_enqueue.assert_called_once() + assert asyncio.iscoroutine(enqueued[0]) + + +@pytest.mark.asyncio +async def test_chunk_processor_routes_logging_through_logging_worker_on_disconnect(): + """Even when the client disconnects mid-stream, the partial-usage log must go + through the durable logging worker rather than a droppable bare task.""" + chunks = [b"event-1", b"event-2", b"event-3"] + response = _make_streaming_response(chunks) + + enqueued = [] + + def _capture(async_coroutine): + enqueued.append(async_coroutine) + async_coroutine.close() + + with ( + patch.object( + PassThroughStreamingHandler, + "_route_streaming_logging_to_handler", + new=AsyncMock(), + ), + patch.object( + GLOBAL_LOGGING_WORKER, + "ensure_initialized_and_enqueue", + side_effect=_capture, + ) as mock_enqueue, + ): + gen = PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + ) + await gen.__anext__() + await gen.aclose() + + mock_enqueue.assert_called_once() + assert asyncio.iscoroutine(enqueued[0]) + + def test_convert_raw_bytes_survives_truncated_multibyte_sequence(): """A stream cut mid-multibyte-sequence (client disconnect) must still decode via errors="replace" so the usage events already received are logged, instead