fix: add debug logging to stream cleanup, improve tests

This commit is contained in:
Ryan Crabbe 2026-02-14 17:31:39 -08:00
parent 2253f92f09
commit f5e36066ab
5 changed files with 228 additions and 72 deletions

View file

@ -169,8 +169,11 @@ class CustomStreamWrapper:
result = self.completion_stream.close()
if result is not None:
await result
except BaseException:
pass
except BaseException as e:
verbose_logger.debug(
"CustomStreamWrapper.aclose: error closing completion_stream: %s",
e,
)
def check_send_stream_usage(self, stream_options: Optional[dict]):
return (

View file

@ -5149,8 +5149,10 @@ async def async_data_generator(
if hasattr(response, "aclose"):
try:
await response.aclose()
except Exception:
pass
except Exception as e:
verbose_proxy_logger.debug(
"async_data_generator: error closing response stream: %s", e
)
def select_data_generator(

View file

@ -1603,11 +1603,23 @@ class Router:
# Shield from anyio cancellation so the awaits can complete.
with anyio.CancelScope(shield=True):
if hasattr(model_response, "aclose"):
await model_response.aclose()
try:
await model_response.aclose()
except BaseException as e:
verbose_router_logger.debug(
"stream_with_fallbacks: error closing model_response: %s",
e,
)
if fallback_response is not None and hasattr(
fallback_response, "aclose"
):
await fallback_response.aclose()
try:
await fallback_response.aclose()
except BaseException as e:
verbose_router_logger.debug(
"stream_with_fallbacks: error closing fallback_response: %s",
e,
)
return FallbackStreamWrapper(stream_with_fallbacks())

View file

@ -3433,3 +3433,50 @@ async def test_async_data_generator_cleanup_on_normal_completion():
assert any("[DONE]" in d for d in yielded_data)
# aclose should still be called via finally block
mock_response.aclose.assert_awaited_once()
@pytest.mark.asyncio
async def test_async_data_generator_cleanup_on_midstream_error():
"""
Test that async_data_generator calls response.aclose() via finally block
even when an exception occurs mid-stream.
"""
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_proxy_logging_obj = MagicMock(spec=ProxyLogging)
async def mock_streaming_iterator_with_error(*args, **kwargs):
yield {"choices": [{"delta": {"content": "Hello"}}]}
raise RuntimeError("upstream connection reset")
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = (
mock_streaming_iterator_with_error
)
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 yielded data chunk and then an error chunk
assert len(yielded_data) >= 2
assert any("error" in d for d in yielded_data)
# aclose must still be called via finally block despite the error
mock_response.aclose.assert_awaited_once()

View file

@ -20,6 +20,9 @@ from litellm.llms.custom_httpx.aiohttp_transport import (
)
# ── aiohttp transport layer tests ──────────────────────────────
@pytest.mark.asyncio
async def test_aiohttp_transport_response_uses_stream_not_content():
"""handle_async_request must use stream= so aclose() propagates to AiohttpResponseStream."""
@ -88,6 +91,9 @@ async def test_aiohttp_response_stream_aclose_releases_connection():
assert aexit_called
# ── CustomStreamWrapper.aclose() tests ─────────────────────────
@pytest.mark.asyncio
async def test_aclose_falls_back_to_close():
"""OpenAI's AsyncStream has close() but not aclose(). Must fall back."""
@ -161,27 +167,37 @@ async def test_aclose_completes_under_cancellation():
assert aclose_completed
# ── Router stream_with_fallbacks cleanup tests ──────────────────
@pytest.mark.asyncio
async def test_stream_with_fallbacks_closes_stream_on_generator_close():
"""Closing the generator from async_function_with_fallbacks must aclose() the stream."""
"""Closing the FallbackStreamWrapper must aclose() the underlying model_response
via stream_with_fallbacks' finally block."""
from litellm.router import Router
stream_closed = False
class FakeStream:
class FakeStream(CustomStreamWrapper):
def __init__(self):
self.chunks = ["chunk1", "chunk2", "chunk3"]
self.index = 0
super().__init__(
completion_stream=None,
model="test-model",
logging_obj=MagicMock(),
custom_llm_provider="openai",
)
self._items = ["chunk1", "chunk2", "chunk3"]
self._index = 0
def __aiter__(self):
return self
async def __anext__(self):
if self.index >= len(self.chunks):
if self._index >= len(self._items):
raise StopAsyncIteration
chunk = self.chunks[self.index]
self.index += 1
return chunk
item = self._items[self._index]
self._index += 1
return item
async def aclose(self):
nonlocal stream_closed
@ -201,74 +217,53 @@ async def test_stream_with_fallbacks_closes_stream_on_generator_close():
fake_stream = FakeStream()
with patch.object(router, "acompletion", return_value=fake_stream):
result = await router.async_function_with_fallbacks(
original_function=router.acompletion,
model="test-model",
messages=[{"role": "user", "content": "hi"}],
stream=True,
num_retries=0,
)
# Call _acompletion_streaming_iterator directly so we go through
# stream_with_fallbacks and its finally block
result = await router._acompletion_streaming_iterator(
model_response=fake_stream,
messages=[{"role": "user", "content": "hi"}],
initial_kwargs={"model": "test-model"},
)
async for chunk in result:
break
# Consume one chunk then close (simulates client disconnect)
async for _ in result:
break
await result.aclose()
await result.aclose()
assert stream_closed
assert stream_closed, "model_response stream was not closed by stream_with_fallbacks finally block"
@pytest.mark.asyncio
async def test_stream_with_fallbacks_closes_fallback_response_on_disconnect():
"""When stream_with_fallbacks is closed during fallback iteration,
both model_response and fallback_response must be closed."""
async def test_stream_with_fallbacks_closes_stream_on_normal_completion():
"""stream_with_fallbacks must aclose() model_response even on normal completion."""
from litellm.router import Router
model_closed = False
fallback_closed = False
class FakeModelStream:
"""Simulates a stream that fails mid-stream, triggering fallback."""
stream_closed = False
class FakeStream(CustomStreamWrapper):
def __init__(self):
self.chunks = []
self.model = "test-model"
self.custom_llm_provider = "openai"
self.logging_obj = MagicMock()
super().__init__(
completion_stream=None,
model="test-model",
logging_obj=MagicMock(),
custom_llm_provider="openai",
)
self._items = ["chunk1"]
self._index = 0
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
async def aclose(self):
nonlocal model_closed
model_closed = True
class FakeFallbackStream:
"""Simulates a fallback stream that yields chunks."""
def __init__(self):
self.items = ["fb1", "fb2", "fb3"]
self.index = 0
def __aiter__(self):
return self
async def __anext__(self):
if self.index >= len(self.items):
if self._index >= len(self._items):
raise StopAsyncIteration
item = self.items[self.index]
self.index += 1
item = self._items[self._index]
self._index += 1
return item
async def aclose(self):
nonlocal fallback_closed
fallback_closed = True
# Just verify the finally block closes model_response even on normal completion
fake_model_stream = FakeModelStream()
nonlocal stream_closed
stream_closed = True
router = Router(
model_list=[
@ -282,18 +277,115 @@ async def test_stream_with_fallbacks_closes_fallback_response_on_disconnect():
]
)
with patch.object(router, "acompletion", return_value=fake_model_stream):
result = await router.async_function_with_fallbacks(
original_function=router.acompletion,
model="test-model",
fake_stream = FakeStream()
result = await router._acompletion_streaming_iterator(
model_response=fake_stream,
messages=[{"role": "user", "content": "hi"}],
initial_kwargs={"model": "test-model"},
)
# Exhaust the stream fully
async for _ in result:
pass
await result.aclose()
assert stream_closed, "model_response stream was not closed after normal completion"
@pytest.mark.asyncio
async def test_stream_with_fallbacks_closes_both_on_fallback_disconnect():
"""When a fallback is triggered and the client disconnects during fallback
iteration, both model_response and fallback_response must be closed."""
from litellm.exceptions import MidStreamFallbackError
from litellm.router import Router
model_closed = False
fallback_closed = False
class FakeModelStream(CustomStreamWrapper):
"""Stream that raises MidStreamFallbackError immediately to trigger fallback."""
def __init__(self):
super().__init__(
completion_stream=None,
model="test-model",
logging_obj=MagicMock(),
custom_llm_provider="openai",
)
self.chunks = []
def __aiter__(self):
return self
async def __anext__(self):
raise MidStreamFallbackError(
message="test mid-stream error",
model="test-model",
llm_provider="openai",
generated_content="",
)
async def aclose(self):
nonlocal model_closed
model_closed = True
class FakeFallbackStream:
"""Fallback stream that yields chunks."""
def __init__(self):
self._items = ["fb1", "fb2", "fb3"]
self._index = 0
def __aiter__(self):
return self
async def __anext__(self):
if self._index >= len(self._items):
raise StopAsyncIteration
item = self._items[self._index]
self._index += 1
return item
async def aclose(self):
nonlocal fallback_closed
fallback_closed = True
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "openai/test",
"api_key": "fake",
},
}
]
)
fake_model_stream = FakeModelStream()
fake_fallback_stream = FakeFallbackStream()
# Mock async_function_with_fallbacks_common_utils to return the fallback stream
# instead of actually calling through the full fallback machinery
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=fake_fallback_stream,
):
result = await router._acompletion_streaming_iterator(
model_response=fake_model_stream,
messages=[{"role": "user", "content": "hi"}],
stream=True,
num_retries=0,
initial_kwargs={
"model": "test-model",
"fallbacks": ["other-model"],
},
)
# Exhaust the stream then close
# Consume one fallback chunk then close (simulates client disconnect)
async for _ in result:
pass
break
await result.aclose()
assert model_closed, "model_response stream was not closed"
assert fallback_closed, "fallback_response stream was not closed"