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
This commit is contained in:
Ryan Crabbe 2026-02-13 21:52:04 -08:00
parent 96802e177b
commit 75ca44434a
5 changed files with 171 additions and 2 deletions

View file

@ -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

View file

@ -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(

View file

@ -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(

View file

@ -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()

View file

@ -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()