From 765fd0762ee08bd38083568dda2b85f7c4ed5654 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 7 Jul 2026 19:13:50 +0300 Subject: [PATCH] fix(responses): stop scheduling sync success_handler concurrently with async_success_handler (#32239) --- litellm/a2a_protocol/streaming_iterator.py | 14 +- litellm/interactions/streaming_iterator.py | 13 +- .../litellm_core_utils/realtime_streaming.py | 11 +- litellm/responses/streaming_iterator.py | 24 ++- ...t_base_responses_api_streaming_iterator.py | 5 +- .../test_responses_hooks.py | 5 + .../test_a2a_streaming_iterator.py | 102 +++++++++++ .../test_interactions_streaming_iterator.py | 97 +++++++++++ .../test_realtime_streaming.py | 5 +- .../test_responses_streaming_iterator.py | 158 ++++++++++++++++++ .../test_responses_websocket_all_providers.py | 12 +- 11 files changed, 395 insertions(+), 51 deletions(-) create mode 100644 tests/test_litellm/a2a_protocol/test_a2a_streaming_iterator.py create mode 100644 tests/test_litellm/interactions/test_interactions_streaming_iterator.py create mode 100644 tests/test_litellm/responses/test_responses_streaming_iterator.py diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index 529154919f3..1ef174a5eee 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -11,7 +11,6 @@ from litellm._logging import verbose_logger from litellm.a2a_protocol.cost_calculator import A2ACostCalculator from litellm.a2a_protocol.utils import A2ARequestUtils from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.litellm_core_utils.thread_pool_executor import executor if TYPE_CHECKING: from a2a.types import SendStreamingMessageRequest, SendStreamingMessageResponse @@ -128,22 +127,15 @@ class A2AStreamingIterator: # Call success handlers - they will build standard_logging_object asyncio.create_task( - self.logging_obj.async_success_handler( - result=result, + self.logging_obj.dispatch_success_handlers( + result, start_time=self.start_time, end_time=end_time, cache_hit=None, + prefer_async_handlers=True, ) ) - executor.submit( - self.logging_obj.success_handler, - result=result, - cache_hit=None, - start_time=self.start_time, - end_time=end_time, - ) - verbose_logger.info( f"A2A streaming completed: prompt_tokens={prompt_tokens}, " f"completion_tokens={completion_tokens}, total_tokens={total_tokens}, " diff --git a/litellm/interactions/streaming_iterator.py b/litellm/interactions/streaming_iterator.py index 45c5443cfd2..0d9d1b4579c 100644 --- a/litellm/interactions/streaming_iterator.py +++ b/litellm/interactions/streaming_iterator.py @@ -174,22 +174,15 @@ class InteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator): logging_response = copy.deepcopy(self.completed_response) asyncio.create_task( - self.logging_obj.async_success_handler( - result=logging_response, + self.logging_obj.dispatch_success_handlers( + logging_response, start_time=self.start_time, end_time=datetime.now(), cache_hit=None, + prefer_async_handlers=True, ) ) - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=None, - start_time=self.start_time, - end_time=datetime.now(), - ) - class SyncInteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator): """ diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index a1a070eb5b7..220d1caa3d2 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,5 +1,4 @@ import asyncio -import concurrent.futures import json from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, Union, cast @@ -25,9 +24,6 @@ if TYPE_CHECKING: else: CLIENT_CONNECTION_CLASS = Any -# Create a thread pool with a maximum of 10 threads -executor = concurrent.futures.ThreadPoolExecutor(max_workers=10) - class RealtimeEventNormalizer(Protocol): def should_drop(self, event: object) -> bool: ... @@ -315,13 +311,12 @@ class RealTimeStreaming: if self.session_tools or self.tool_calls: self.logging_obj.model_call_details["realtime_tools"] = self.session_tools self.logging_obj.model_call_details["realtime_tool_calls"] = self.tool_calls - ## ASYNC LOGGING # Route through the bounded logging worker (per-coroutine timeout + # concurrency cap) instead of a bare create_task, so a slow callback # can't leave suspended tasks pinning each call's response in memory. - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(self.logging_obj.async_success_handler(self.messages)) - ## SYNC LOGGING - executor.submit(self.logging_obj.success_handler(self.messages)) + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) + ) async def _send_to_backend(self, message: str) -> bool: """Send a message to the backend WebSocket. diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 890b3b636ba..3618331f0f5 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -287,11 +287,12 @@ class BaseResponsesAPIStreamingIterator: end_time = datetime.now() if is_async: asyncio.create_task( - self.logging_obj.async_success_handler( - result=logging_response, + self.logging_obj.dispatch_success_handlers( + logging_response, start_time=self.start_time, end_time=end_time, cache_hit=self._completed_response_cache_hit, + prefer_async_handlers=True, ) ) else: @@ -302,14 +303,13 @@ class BaseResponsesAPIStreamingIterator: end_time=end_time, cache_hit=self._completed_response_cache_hit, ) - - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=self._completed_response_cache_hit, - start_time=self.start_time, - end_time=end_time, - ) + executor.submit( + self.logging_obj.success_handler, + result=logging_response, + cache_hit=self._completed_response_cache_hit, + start_time=self.start_time, + end_time=end_time, + ) self._run_post_success_hooks(end_time=end_time) def _handle_logging_completed_response(self): @@ -1136,7 +1136,6 @@ def _build_synthetic_response_events( # --------------------------------------------------------------------------- from litellm._logging import verbose_logger -from litellm.litellm_core_utils.thread_pool_executor import executor as _ws_executor RESPONSES_WS_LOGGED_EVENT_TYPES = [ "response.created", @@ -1251,8 +1250,7 @@ class ResponsesWebSocketStreaming: if self.input_messages: self.logging_obj.model_call_details["messages"] = self.input_messages if self.messages: - asyncio.create_task(self.logging_obj.async_success_handler(self.messages)) - _ws_executor.submit(self.logging_obj.success_handler, self.messages) + asyncio.create_task(self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)) async def backend_to_client(self) -> None: """Forward events from backend WebSocket to the client.""" diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index 2acced4c679..5388c5aef83 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -647,9 +647,10 @@ class TestBaseResponsesAPIStreamingIterator: assert result.type == ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE assert iterator.completed_response == result - # Success handler should have been called (via _handle_logging_completed_response) + # Success handlers are dispatched as one async task (via _handle_logging_completed_response); + # the sync handler must never be submitted to the executor concurrently (LIT-4210) mock_create_task.assert_called_once() - mock_executor.submit.assert_called_once() + mock_executor.submit.assert_not_called() # Failure handlers should NOT have been called mock_logging_obj.async_failure_handler.assert_not_called() diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py index 3cc5e3984e2..2344a62de4d 100644 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ b/tests/llm_responses_api_testing/test_responses_hooks.py @@ -38,6 +38,11 @@ class _FakeLoggingObj: self.model_call_details = {"litellm_params": {}} # Signature alignment with Logging handlers + async def dispatch_success_handlers(self, *args, **kwargs): + kwargs.pop("prefer_async_handlers", None) + await self.async_success_handler(*args, **kwargs) + self.success_handler(*args, **kwargs) + def success_handler(self, *args, **kwargs): self.success_calls += 1 self.last_success_kwargs = kwargs diff --git a/tests/test_litellm/a2a_protocol/test_a2a_streaming_iterator.py b/tests/test_litellm/a2a_protocol/test_a2a_streaming_iterator.py new file mode 100644 index 00000000000..d86cbb94a91 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/test_a2a_streaming_iterator.py @@ -0,0 +1,102 @@ +""" +Regression test for LIT-4210: completing an A2A stream must not run the sync +success_handler on the thread-pool executor concurrently with +async_success_handler (cross-thread pydantic mutation segfaults pydantic-core). +""" + +import asyncio +import time +from types import SimpleNamespace + +import pytest + +import litellm +from litellm.a2a_protocol import streaming_iterator as a2a_streaming_iterator_module +from litellm.a2a_protocol.streaming_iterator import A2AStreamingIterator +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils import thread_pool_executor as thread_pool_executor_module +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging + + +class RecordingCustomLogger(CustomLogger): + def __init__(self): + super().__init__() + self.async_hook_fired = False + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + self.async_hook_fired = True + + async def async_log_stream_event(self, kwargs, response_obj, start_time, end_time): + self.async_hook_fired = True + + +class RecordingExecutor: + def __init__(self, inner): + self._inner = inner + self.submits: list = [] + + def submit(self, fn, *args, **kwargs): + self.submits.append(fn) + return self._inner.submit(fn, *args, **kwargs) + + def submitted_for(self, logging_obj) -> list: + return [fn for fn in self.submits if getattr(fn, "__self__", None) is logging_obj] + + +@pytest.fixture(autouse=True) +def _isolate_callbacks(): + saved = ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) + yield + ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) = saved + + +@pytest.mark.asyncio +async def test_custom_logger_only_never_submits_sync_success_handler(monkeypatch): + recording_executor = RecordingExecutor(thread_pool_executor_module.executor) + monkeypatch.setattr(thread_pool_executor_module, "executor", recording_executor) + monkeypatch.setattr(a2a_streaming_iterator_module, "executor", recording_executor, raising=False) + + recorder = RecordingCustomLogger() + litellm.success_callback = [recorder] + litellm._async_success_callback = [recorder] + + logging_obj = LitellmLogging( + model="a2a/test-agent", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="a2a_send_message_streaming", + start_time=time.time(), + litellm_call_id="lit-4210-test", + function_id="lit-4210-test", + ) + + async def _empty_stream(): + return + yield + + iterator = A2AStreamingIterator( + stream=_empty_stream(), + request=SimpleNamespace( + params=SimpleNamespace(message={"role": "user", "parts": [{"kind": "text", "text": "hi"}]}) + ), + logging_obj=logging_obj, + agent_name="test-agent", + ) + + await iterator._handle_stream_complete() + await asyncio.sleep(0.5) + + assert recorder.async_hook_fired is True + assert recording_executor.submitted_for(logging_obj) == [] diff --git a/tests/test_litellm/interactions/test_interactions_streaming_iterator.py b/tests/test_litellm/interactions/test_interactions_streaming_iterator.py new file mode 100644 index 00000000000..9f88b2c9611 --- /dev/null +++ b/tests/test_litellm/interactions/test_interactions_streaming_iterator.py @@ -0,0 +1,97 @@ +""" +Regression test for LIT-4210: completing an async Interactions API stream must +not run the sync success_handler on the thread-pool executor concurrently with +async_success_handler (cross-thread pydantic mutation segfaults pydantic-core). +""" + +import asyncio +import time + +import httpx +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.interactions import streaming_iterator as interactions_streaming_iterator_module +from litellm.interactions.streaming_iterator import InteractionsAPIStreamingIterator +from litellm.litellm_core_utils import thread_pool_executor as thread_pool_executor_module +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.types.interactions import InteractionsAPIStreamingResponse + + +class RecordingCustomLogger(CustomLogger): + def __init__(self): + super().__init__() + self.async_hook_fired = False + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + self.async_hook_fired = True + + async def async_log_stream_event(self, kwargs, response_obj, start_time, end_time): + self.async_hook_fired = True + + +class RecordingExecutor: + def __init__(self, inner): + self._inner = inner + self.submits: list = [] + + def submit(self, fn, *args, **kwargs): + self.submits.append(fn) + return self._inner.submit(fn, *args, **kwargs) + + def submitted_for(self, logging_obj) -> list: + return [fn for fn in self.submits if getattr(fn, "__self__", None) is logging_obj] + + +@pytest.fixture(autouse=True) +def _isolate_callbacks(): + saved = ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) + yield + ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) = saved + + +@pytest.mark.asyncio +async def test_custom_logger_only_never_submits_sync_success_handler(monkeypatch): + recording_executor = RecordingExecutor(thread_pool_executor_module.executor) + monkeypatch.setattr(thread_pool_executor_module, "executor", recording_executor) + monkeypatch.setattr(interactions_streaming_iterator_module, "executor", recording_executor) + + recorder = RecordingCustomLogger() + litellm.success_callback = [recorder] + litellm._async_success_callback = [recorder] + + logging_obj = LitellmLogging( + model="gemini/gemini-3-pro-preview", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="ainteraction", + start_time=time.time(), + litellm_call_id="lit-4210-test", + function_id="lit-4210-test", + ) + iterator = InteractionsAPIStreamingIterator( + response=httpx.Response(200), + model="gemini/gemini-3-pro-preview", + interactions_api_config=None, + logging_obj=logging_obj, + ) + iterator.completed_response = InteractionsAPIStreamingResponse() + + iterator._handle_logging_completed_response() + await asyncio.sleep(0.5) + + assert recorder.async_hook_fired is True + assert recording_executor.submitted_for(logging_obj) == [] diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 766befd1a99..dff54515098 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -2961,10 +2961,13 @@ async def test_log_messages_routes_async_logging_through_bounded_worker(): with ( patch("litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER") as mock_worker, patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task, - patch("litellm.litellm_core_utils.realtime_streaming.executor.submit"), ): await streaming.log_messages() mock_worker.ensure_initialized_and_enqueue.assert_called_once() + enqueued = mock_worker.ensure_initialized_and_enqueue.call_args + assert (enqueued.args or tuple(enqueued.kwargs.values()))[0] is logging_obj.dispatch_success_handlers.return_value + logging_obj.dispatch_success_handlers.assert_called_once_with(streaming.messages, prefer_async_handlers=True) + logging_obj.success_handler.assert_not_called() # the bare create_task path must no longer be used for success logging mock_create_task.assert_not_called() diff --git a/tests/test_litellm/responses/test_responses_streaming_iterator.py b/tests/test_litellm/responses/test_responses_streaming_iterator.py new file mode 100644 index 00000000000..9ba7dfcaa80 --- /dev/null +++ b/tests/test_litellm/responses/test_responses_streaming_iterator.py @@ -0,0 +1,158 @@ +""" +Regression tests for LIT-4210: the streaming iterators must never run the sync +success_handler on the thread-pool executor concurrently with +async_success_handler. Concurrent mutation of the shared response object / +model_call_details from two threads segfaults pydantic-core (customer pods +crashed with exit 139 whenever any CustomLogger was registered). +""" + +import asyncio +import time + +import httpx +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils import thread_pool_executor as thread_pool_executor_module +from litellm.responses import streaming_iterator as responses_streaming_iterator_module +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator +from litellm.types.llms.openai import ResponsesAPIResponse + + +class RecordingCustomLogger(CustomLogger): + def __init__(self): + super().__init__() + self.async_hook_started: float | None = None + self.async_hook_finished: float | None = None + + async def _record(self): + self.async_hook_started = time.monotonic() + await asyncio.sleep(0.2) + self.async_hook_finished = time.monotonic() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await self._record() + + async def async_log_stream_event(self, kwargs, response_obj, start_time, end_time): + await self._record() + + +class RecordingExecutor: + def __init__(self, inner): + self._inner = inner + self.submits: list = [] + + def submit(self, fn, *args, **kwargs): + self.submits.append((time.monotonic(), fn)) + return self._inner.submit(fn, *args, **kwargs) + + def submit_times_for(self, logging_obj) -> list: + return [t for t, fn in self.submits if getattr(fn, "__self__", None) is logging_obj] + + +@pytest.fixture(autouse=True) +def _isolate_callbacks(): + saved = ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) + yield + ( + litellm.callbacks, + litellm.success_callback, + litellm._async_success_callback, + litellm.failure_callback, + litellm._async_failure_callback, + ) = saved + + +@pytest.fixture +def recording_executor(monkeypatch): + recording = RecordingExecutor(thread_pool_executor_module.executor) + monkeypatch.setattr(thread_pool_executor_module, "executor", recording) + monkeypatch.setattr(responses_streaming_iterator_module, "executor", recording) + return recording + + +def _make_logging_obj() -> LitellmLogging: + logging_obj = LitellmLogging( + model="gpt-5.4-nano", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="aresponses", + start_time=time.time(), + litellm_call_id="lit-4210-test", + function_id="lit-4210-test", + ) + logging_obj.model_call_details["litellm_params"] = {"aresponses": True} + return logging_obj + + +def _make_iterator(logging_obj: LitellmLogging) -> ResponsesAPIStreamingIterator: + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="gpt-5.4-nano", + responses_api_provider_config=None, + logging_obj=logging_obj, + ) + iterator.completed_response = ResponsesAPIResponse( + id="resp_lit4210", + created_at=1700000000.0, + model="gpt-5.4-nano", + object="response", + output=[], + parallel_tool_calls=True, + tool_choice="auto", + tools=[], + error=None, + incomplete_details=None, + instructions=None, + metadata={}, + temperature=1.0, + top_p=1.0, + ) + return iterator + + +@pytest.mark.asyncio +async def test_custom_logger_only_never_submits_sync_success_handler(recording_executor): + recorder = RecordingCustomLogger() + litellm.success_callback = [recorder] + litellm._async_success_callback = [recorder] + + logging_obj = _make_logging_obj() + iterator = _make_iterator(logging_obj) + + iterator._log_completed_response(is_async=True) + await asyncio.sleep(0.6) + + assert recorder.async_hook_started is not None + assert recording_executor.submit_times_for(logging_obj) == [] + + +@pytest.mark.asyncio +async def test_sync_callbacks_run_only_after_async_handler_completes(recording_executor): + recorder = RecordingCustomLogger() + sync_events: list = [] + + def sync_callback(kwargs, response_obj, start_time, end_time): + sync_events.append(time.monotonic()) + + litellm.success_callback = [recorder, sync_callback] + litellm._async_success_callback = [recorder] + + logging_obj = _make_logging_obj() + iterator = _make_iterator(logging_obj) + + iterator._log_completed_response(is_async=True) + await asyncio.sleep(0.8) + + assert recorder.async_hook_finished is not None + submit_times = recording_executor.submit_times_for(logging_obj) + assert len(submit_times) == 1 + assert submit_times[0] >= recorder.async_hook_finished diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 0d13ff4fd05..4509abc7749 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1122,7 +1122,7 @@ class TestNativeWebSocketGuardrails: client_ws = MagicMock() client_ws.send_text = AsyncMock() logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() delta_event = json.dumps( {"type": "response.output_text.delta", "delta": "alice@example.com"} @@ -1196,7 +1196,7 @@ class TestNativeWebSocketGuardrails: client_ws = MagicMock() client_ws.send_text = AsyncMock() logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() done_events = [ json.dumps( @@ -1895,7 +1895,7 @@ class TestNativeWebSocketGuardrailMasking: ] ) logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() handler = _make_streaming( websocket=websocket, @@ -1951,7 +1951,7 @@ class TestNativeWebSocketGuardrailMasking: ] ) logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() handler = _make_streaming( websocket=websocket, @@ -2014,7 +2014,7 @@ class TestNativeWebSocketGuardrailMasking: ] ) logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() handler = _make_streaming( websocket=websocket, @@ -2077,7 +2077,7 @@ class TestNativeWebSocketGuardrailMasking: ] ) logging_obj = MagicMock() - logging_obj.async_success_handler = AsyncMock() + logging_obj.dispatch_success_handlers = AsyncMock() handler = _make_streaming( websocket=websocket,