From 75ca44434aa3aff46332e0c032c9c643c1d4d8c7 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Fri, 13 Feb 2026 21:52:04 -0800 Subject: [PATCH] fix: close streaming connections to prevent connection pool exhaustion - Add aclose() to CustomStreamWrapper to delegate to underlying stream - Add finally block in async_data_generator to release HTTP connections - Thread shared_session through async_streaming to reuse connection pool - Set finite default timeout (600s) in _get_openai_client --- .../litellm_core_utils/streaming_handler.py | 6 + litellm/llms/openai/openai.py | 5 +- litellm/proxy/proxy_server.py | 9 ++ .../test_streaming_handler.py | 49 ++++++++- tests/test_litellm/proxy/test_proxy_server.py | 104 ++++++++++++++++++ 5 files changed, 171 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index c6f0f67976f..f3556b2d18f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -155,6 +155,12 @@ class CustomStreamWrapper: def __aiter__(self): return self + async def aclose(self): + if self.completion_stream is not None and hasattr( + self.completion_stream, "aclose" + ): + await self.completion_stream.aclose() + def check_send_stream_usage(self, stream_options: Optional[dict]): return ( stream_options is not None diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index da87852dff5..8f180be8a1d 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -356,7 +356,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_key: Optional[str] = None, api_base: Optional[str] = None, api_version: Optional[str] = None, - timeout: Union[float, httpx.Timeout] = httpx.Timeout(None), + timeout: Union[float, httpx.Timeout] = httpx.Timeout(timeout=600.0, connect=5.0), max_retries: Optional[int] = DEFAULT_MAX_RETRIES, organization: Optional[str] = None, client: Optional[Union[OpenAI, AsyncOpenAI]] = None, @@ -693,6 +693,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, drop_params=drop_params, stream_options=stream_options, + shared_session=shared_session, ) else: return self.acompletion( @@ -1063,6 +1064,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): headers=None, drop_params: Optional[bool] = None, stream_options: Optional[dict] = None, + shared_session: Optional["ClientSession"] = None, ): response = None data = provider_config.transform_request( @@ -1087,6 +1089,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, client=client, + shared_session=shared_session, ) ## LOGGING logging_obj.pre_call( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bc2d32f141d..1421a47dede 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5142,6 +5142,15 @@ async def async_data_generator( ) error_returned = json.dumps({"error": proxy_exception.to_dict()}) yield f"data: {error_returned}\n\n" + finally: + # Close the response stream to release the underlying HTTP connection + # back to the connection pool. This prevents pool exhaustion when + # clients disconnect mid-stream. + if hasattr(response, "aclose"): + try: + await response.aclose() + except Exception: + pass def select_data_generator( diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index ec2f528a35d..73ecaa20a2d 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2,7 +2,7 @@ import json import os import sys import time -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -1185,3 +1185,50 @@ def test_is_chunk_non_empty_with_valid_tool_calls( ) is True ) + + +@pytest.mark.asyncio +async def test_custom_stream_wrapper_aclose(): + """Test that aclose() delegates to the underlying completion_stream's aclose()""" + mock_stream = AsyncMock() + mock_stream.aclose = AsyncMock() + + wrapper = CustomStreamWrapper( + completion_stream=mock_stream, + model=None, + logging_obj=MagicMock(), + custom_llm_provider=None, + ) + + await wrapper.aclose() + mock_stream.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_custom_stream_wrapper_aclose_no_underlying(): + """Test that aclose() is safe when completion_stream has no aclose method""" + mock_stream = MagicMock(spec=[]) # No aclose attribute + + wrapper = CustomStreamWrapper( + completion_stream=mock_stream, + model=None, + logging_obj=MagicMock(), + custom_llm_provider=None, + ) + + # Should not raise + await wrapper.aclose() + + +@pytest.mark.asyncio +async def test_custom_stream_wrapper_aclose_none_stream(): + """Test that aclose() is safe when completion_stream is None""" + wrapper = CustomStreamWrapper( + completion_stream=None, + model=None, + logging_obj=MagicMock(), + custom_llm_provider=None, + ) + + # Should not raise + await wrapper.aclose() diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d65df0087ad..e57523ea058 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3329,3 +3329,107 @@ class TestInvitationEndpoints: # ProxyException handler returns {"error": {...}}, HTTPException returns {"detail": {...}} error_content = body.get("error", body.get("detail", body)) assert "not allowed" in str(error_content).lower() + + +@pytest.mark.asyncio +async def test_async_data_generator_cleanup_on_early_exit(): + """ + Test that async_data_generator calls response.aclose() in the finally block + when the generator is abandoned mid-stream (client disconnect). + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + mock_chunks = [ + {"choices": [{"delta": {"content": "Hello"}}]}, + {"choices": [{"delta": {"content": " world"}}]}, + {"choices": [{"delta": {"content": " more"}}]}, + ] + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + + async def mock_streaming_iterator(*args, **kwargs): + for chunk in mock_chunks: + yield chunk + + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( + mock_streaming_iterator + ) + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( + side_effect=lambda **kwargs: kwargs.get("response") + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + # Create a mock response with aclose + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + # Consume only the first chunk then abandon the generator (simulates client disconnect) + gen = async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ) + first_chunk = await gen.__anext__() + assert first_chunk.startswith("data: ") + + # Close the generator early (simulates what ASGI does on client disconnect) + await gen.aclose() + + # Verify aclose was called on the response to release the HTTP connection + mock_response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_async_data_generator_cleanup_on_normal_completion(): + """ + Test that async_data_generator calls response.aclose() even on normal completion. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + mock_chunks = [ + {"choices": [{"delta": {"content": "Hello"}}]}, + ] + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + + async def mock_streaming_iterator(*args, **kwargs): + for chunk in mock_chunks: + yield chunk + + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( + mock_streaming_iterator + ) + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( + side_effect=lambda **kwargs: kwargs.get("response") + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + yielded_data = [] + async for data in async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ): + yielded_data.append(data) + + # Should have completed normally with [DONE] + assert any("[DONE]" in d for d in yielded_data) + # aclose should still be called via finally block + mock_response.aclose.assert_awaited_once()