fix(router): hold max_parallel_requests slot until streaming response is exhausted or closed (#39859)

* fix(router): hold max_parallel_requests slot until streaming response is exhausted or closed

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(router): normalize deployment_slot once to keep stream_with_fallbacks under the C901 ceiling

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(router): close upstream stream before releasing max_parallel_requests slot

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-05 11:41:35 -07:00 • committed by GitHub
parent f66b3ebe0d
commit e6705510f8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 198 additions and 51 deletions

View file

@ -2597,14 +2597,20 @@ class Router:
model_response: CustomStreamWrapper,
messages: list[dict[str, str]],
initial_kwargs: dict,
deployment_slot: contextlib.AsyncExitStack | None = None,
) -> CustomStreamWrapper:
"""
Helper to iterate over a streaming response.
Catches errors for fallbacks using the router's fallback system
`deployment_slot` holds the deployment's max_parallel_requests semaphore; it is
released when the stream is exhausted, closed, or falls back to another deployment
"""
from litellm.exceptions import MidStreamFallbackError
held_slot: Final = deployment_slot if deployment_slot is not None else contextlib.AsyncExitStack()
class FallbackStreamWrapper(CustomStreamWrapper):
def __init__(self, async_generator: AsyncGenerator):
# Copy attributes from the original model_response
@ -2628,12 +2634,26 @@ class Router:
async def __anext__(self):
return await self._async_generator.__anext__()
async def close_model_response() -> None:
if not hasattr(model_response, "aclose"):
return
try:
await model_response.aclose()
except BaseException as e:
verbose_router_logger.debug(
"stream_with_fallbacks: error closing model_response: %s",
e,
)
async def stream_with_fallbacks():
fallback_response = None # Track for cleanup in finally
try:
async for item in model_response:
yield item
except MidStreamFallbackError as e:
with anyio.CancelScope(shield=True):
await close_model_response()
await held_slot.aclose()
if not e.is_pre_first_chunk and (
e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)
):
@ -2707,14 +2727,8 @@ class Router:
# (e.g. on client disconnect).
# Shield from anyio cancellation so the awaits can complete.
with anyio.CancelScope(shield=True):
if hasattr(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,
)
await close_model_response()
await held_slot.aclose()
if fallback_response is not None and hasattr(fallback_response, "aclose"):
try:
await fallback_response.aclose()
@ -3379,61 +3393,53 @@ class Router:
kwargs=kwargs,
client_type="max_parallel_requests",
)
if rpm_semaphore is not None and isinstance(rpm_semaphore, asyncio.Semaphore):
async with rpm_semaphore:
"""
- Check rpm limits before making the call
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
"""
await self.async_routing_strategy_pre_call_checks(
deployment=deployment,
logging_obj=logging_obj,
parent_otel_span=parent_otel_span,
)
response = await _response
else:
async with contextlib.AsyncExitStack() as deployment_slot:
if isinstance(rpm_semaphore, asyncio.Semaphore):
await deployment_slot.enter_async_context(rpm_semaphore)
await self.async_routing_strategy_pre_call_checks(
deployment=deployment,
logging_obj=logging_obj,
parent_otel_span=parent_otel_span,
)
response = await _response
## CHECK CONTENT FILTER ERROR ##
if isinstance(response, ModelResponse):
_should_raise = self._should_raise_content_policy_error(model=model, response=response, kwargs=kwargs)
if _should_raise:
raise litellm.ContentPolicyViolationError(
message="Response output was blocked.",
model=model,
llm_provider="",
## CHECK CONTENT FILTER ERROR ##
if isinstance(response, ModelResponse):
_should_raise = self._should_raise_content_policy_error(
model=model, response=response, kwargs=kwargs
)
if _should_raise:
raise litellm.ContentPolicyViolationError(
message="Response output was blocked.",
model=model,
llm_provider="",
)
if (
isinstance(response, CustomStreamWrapper)
and response.completion_stream is None
and response.make_call is not None
):
await response.fetch_stream()
if (
isinstance(response, CustomStreamWrapper)
and response.completion_stream is None
and response.make_call is not None
):
await response.fetch_stream()
self.success_calls[model_name] += 1
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
# debug how often this deployment picked
self._track_deployment_metrics(
deployment=deployment,
response=response,
parent_otel_span=parent_otel_span,
)
if isinstance(response, CustomStreamWrapper):
return await self._acompletion_streaming_iterator(
model_response=response,
messages=messages,
initial_kwargs=input_kwargs_for_streaming_fallback,
self.success_calls[model_name] += 1
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
# debug how often this deployment picked
self._track_deployment_metrics(
deployment=deployment,
response=response,
parent_otel_span=parent_otel_span,
)
return response
if isinstance(response, CustomStreamWrapper):
return await self._acompletion_streaming_iterator(
model_response=response,
messages=messages,
initial_kwargs=input_kwargs_for_streaming_fallback,
deployment_slot=deployment_slot.pop_all(),
)
return response
except litellm.Timeout as e:
deployment_request_timeout_param: Final = _timeout_debug_deployment_dict.get("litellm_params", {}).get(
"request_timeout", None

View file

@ -12938,3 +12938,144 @@ async def test_router_retry_policy_controls_upstream_attempt_count(
await router.acompletion(model="gpt-5.6", messages=[{"role": "user", "content": "hi"}])
assert upstream.call_count == expected_upstream_calls
class _InFlightTracker:
def __init__(self) -> None:
self.current = 0
self.peak = 0
def enter(self) -> None:
self.current += 1
self.peak = max(self.peak, self.current)
def exit(self) -> None:
self.current -= 1
_SSE_CHUNKS: Final[tuple[bytes, ...]] = tuple(
b'data: {"id":"c","object":"chat.completion.chunk","created":1,"model":"gpt-5.6",'
b'"choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}\n\n'
for _ in range(5)
)
class _CountingSSEStream(httpx.AsyncByteStream):
def __init__(self, tracker: _InFlightTracker) -> None:
self._tracker = tracker
self._in_flight = False
def _finish(self) -> None:
if self._in_flight:
self._in_flight = False
self._tracker.exit()
async def __aiter__(self):
self._in_flight = True
self._tracker.enter()
try:
for chunk in _SSE_CHUNKS:
await asyncio.sleep(0.02)
yield chunk
finally:
await self.aclose()
yield b"data: [DONE]\n\n"
async def aclose(self) -> None:
await asyncio.sleep(0.02)
self._finish()
def _max_parallel_router(max_parallel_requests: int) -> Router:
return Router(
model_list=[
{
"model_name": "gpt-5.6",
"litellm_params": {
"model": "openai/gpt-5.6",
"api_key": "sk-fake",
"api_base": "https://max-parallel.local/v1",
"max_parallel_requests": max_parallel_requests,
},
}
],
num_retries=0,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [False, True])
async def test_router_max_parallel_requests_bounds_in_flight_upstream_calls(
monkeypatch: pytest.MonkeyPatch, stream: bool
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
tracker: Final = _InFlightTracker()
router: Final = _max_parallel_router(max_parallel_requests=2)
async def upstream(request: httpx.Request) -> httpx.Response:
if stream:
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, stream=_CountingSSEStream(tracker)
)
tracker.enter()
await asyncio.sleep(0.05)
tracker.exit()
return httpx.Response(
200,
json={
"id": "c",
"object": "chat.completion",
"created": 1,
"model": "gpt-5.6",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "x"}, "finish_reason": "stop"}],
},
)
async def one_call() -> None:
response = await router.acompletion(
model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=stream
)
if stream:
async for _ in response:
pass
with respx.mock(assert_all_called=True) as respx_mock:
respx_mock.post("https://max-parallel.local/v1/chat/completions").mock(side_effect=upstream)
await asyncio.wait_for(asyncio.gather(*(one_call() for _ in range(10))), timeout=10)
assert tracker.peak <= 2
assert tracker.current == 0
@pytest.mark.asyncio
async def test_router_max_parallel_requests_slot_released_when_stream_closed_early(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
tracker: Final = _InFlightTracker()
router: Final = _max_parallel_router(max_parallel_requests=1)
with respx.mock() as respx_mock:
respx_mock.post("https://max-parallel.local/v1/chat/completions").mock(
side_effect=lambda request: httpx.Response(
200, headers={"content-type": "text/event-stream"}, stream=_CountingSSEStream(tracker)
)
)
first: Final = await router.acompletion(
model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=True
)
await first.__anext__()
async def second_call() -> None:
second = await router.acompletion(
model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=True
)
async for _ in second:
pass
second_task: Final = asyncio.create_task(second_call())
await asyncio.sleep(0.05)
assert tracker.current == 1
await first.aclose()
await asyncio.wait_for(second_task, timeout=2)
assert tracker.peak == 1
assert tracker.current == 0