mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
f66b3ebe0d
commit
e6705510f8
2 changed files with 198 additions and 51 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue