mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
test: update streaming cleanup tests for review findings
This commit is contained in:
parent
81be07aaeb
commit
2253f92f09
1 changed files with 80 additions and 7 deletions
|
|
@ -5,7 +5,7 @@ Regression tests for streaming connection pool leak fix.
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
|
|
@ -143,7 +143,7 @@ async def test_aclose_completes_under_cancellation():
|
|||
|
||||
class SlowCloseStream:
|
||||
async def aclose(self):
|
||||
await asyncio.sleep(0)
|
||||
await anyio.sleep(0)
|
||||
nonlocal aclose_completed
|
||||
aclose_completed = True
|
||||
|
||||
|
|
@ -218,9 +218,82 @@ async def test_stream_with_fallbacks_closes_stream_on_generator_close():
|
|||
assert stream_closed
|
||||
|
||||
|
||||
def test_streaming_response_monkey_patch_applied():
|
||||
"""_disconnect_aware_call must be applied to StreamingResponse.__call__."""
|
||||
from litellm.proxy.common_request_processing import _disconnect_aware_call
|
||||
from starlette.responses import StreamingResponse
|
||||
@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."""
|
||||
from litellm.router import Router
|
||||
|
||||
assert StreamingResponse.__call__ is _disconnect_aware_call
|
||||
model_closed = False
|
||||
fallback_closed = False
|
||||
|
||||
class FakeModelStream:
|
||||
"""Simulates a stream that fails mid-stream, triggering fallback."""
|
||||
|
||||
def __init__(self):
|
||||
self.chunks = []
|
||||
self.model = "test-model"
|
||||
self.custom_llm_provider = "openai"
|
||||
self.logging_obj = MagicMock()
|
||||
|
||||
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):
|
||||
raise StopAsyncIteration
|
||||
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()
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/test",
|
||||
"api_key": "fake",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with patch.object(router, "acompletion", return_value=fake_model_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,
|
||||
)
|
||||
|
||||
# Exhaust the stream then close
|
||||
async for _ in result:
|
||||
pass
|
||||
await result.aclose()
|
||||
|
||||
assert model_closed, "model_response stream was not closed"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue