mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
96802e177b
commit
75ca44434a
5 changed files with 171 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue