diff --git a/litellm/router.py b/litellm/router.py index 43d3d97f034..281caf2f1a0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -2868,64 +2868,61 @@ class Router: chunks: List = [] aiter = response.__aiter__() - if ttft_timeout is not None: - loop = asyncio.get_running_loop() - deadline = loop.time() + ttft_timeout - first_token_received = False + try: + if ttft_timeout is not None: + loop = asyncio.get_running_loop() + deadline = loop.time() + ttft_timeout + first_token_received = False - while not first_token_received: - remaining = deadline - loop.time() - if remaining <= 0: - verbose_router_logger.warning( - f"ttft_timeout={ttft_timeout}s exceeded for model={response.model}: " - "provider accepted connection but sent no tokens" - ) - raise litellm.Timeout( - message=f"Router ttft_timeout={ttft_timeout}s exceeded: provider accepted connection but sent no tokens", - model=response.model or "", - llm_provider=response.custom_llm_provider or "", - ) - try: - chunk = await asyncio.wait_for(aiter.__anext__(), timeout=remaining) - except asyncio.TimeoutError: - verbose_router_logger.warning( - f"ttft_timeout={ttft_timeout}s exceeded for model={response.model}: " - "provider accepted connection but sent no tokens" - ) - raise litellm.Timeout( - message=f"Router ttft_timeout={ttft_timeout}s exceeded: provider accepted connection but sent no tokens", - model=response.model or "", - llm_provider=response.custom_llm_provider or "", - ) - except StopAsyncIteration: - break - chunks.append(chunk) - delta = chunk.choices[0].delta if chunk.choices else None - if delta and (delta.content or delta.tool_calls): - first_token_received = True + while not first_token_received: + remaining = deadline - loop.time() + try: + if remaining <= 0: + raise asyncio.TimeoutError + chunk = await asyncio.wait_for( + aiter.__anext__(), timeout=remaining + ) + except asyncio.TimeoutError: + verbose_router_logger.warning( + f"ttft_timeout={ttft_timeout}s exceeded for model={response.model}: " + "provider accepted connection but sent no tokens" + ) + raise litellm.Timeout( + message=f"Router ttft_timeout={ttft_timeout}s exceeded: provider accepted connection but sent no tokens", + model=response.model or "", + llm_provider=response.custom_llm_provider or "", + ) + except StopAsyncIteration: + break + chunks.append(chunk) + delta = chunk.choices[0].delta if chunk.choices else None + if delta and (delta.content or delta.tool_calls): + first_token_received = True - if stream_idle_timeout is not None: - while True: - try: - chunk = await asyncio.wait_for( - aiter.__anext__(), timeout=stream_idle_timeout - ) - except asyncio.TimeoutError: - verbose_router_logger.warning( - f"stream_idle_timeout={stream_idle_timeout}s exceeded for model={response.model}: " - "provider stalled mid-stream" - ) - raise litellm.Timeout( - message=f"Router stream_idle_timeout={stream_idle_timeout}s exceeded: provider stalled mid-stream", - model=response.model or "", - llm_provider=response.custom_llm_provider or "", - ) - except StopAsyncIteration: - break - chunks.append(chunk) - else: - async for chunk in aiter: - chunks.append(chunk) + if stream_idle_timeout is not None: + while True: + try: + chunk = await asyncio.wait_for( + aiter.__anext__(), timeout=stream_idle_timeout + ) + except asyncio.TimeoutError: + verbose_router_logger.warning( + f"stream_idle_timeout={stream_idle_timeout}s exceeded for model={response.model}: " + "provider stalled mid-stream" + ) + raise litellm.Timeout( + message=f"Router stream_idle_timeout={stream_idle_timeout}s exceeded: provider stalled mid-stream", + model=response.model or "", + llm_provider=response.custom_llm_provider or "", + ) + except StopAsyncIteration: + break + chunks.append(chunk) + else: + async for chunk in aiter: + chunks.append(chunk) + finally: + await response.aclose() result = stream_chunk_builder(chunks, messages=messages) if result is None: @@ -3039,6 +3036,19 @@ class Router: "litellm_logging_obj", None ) + async def _await_response() -> Union[ModelResponse, CustomStreamWrapper]: + awaited_response = await _response + if _forced_stream_for_ttft and isinstance( + awaited_response, CustomStreamWrapper + ): + return await self._collect_stream_with_ttft_timeout( + response=awaited_response, + messages=messages, + ttft_timeout=_ttft_timeout, + stream_idle_timeout=_stream_idle_timeout, + ) + return awaited_response + rpm_semaphore = self._get_client( deployment=deployment, kwargs=kwargs, @@ -3057,7 +3067,7 @@ class Router: logging_obj=logging_obj, parent_otel_span=parent_otel_span, ) - response = await _response + response = await _await_response() else: await self.async_routing_strategy_pre_call_checks( deployment=deployment, @@ -3065,7 +3075,7 @@ class Router: parent_otel_span=parent_otel_span, ) - response = await _response + response = await _await_response() ## CHECK CONTENT FILTER ERROR ## if isinstance(response, ModelResponse): @@ -3091,22 +3101,6 @@ class Router: ) if isinstance(response, CustomStreamWrapper): - if _forced_stream_for_ttft: - reconstructed = await self._collect_stream_with_ttft_timeout( - response=response, - messages=messages, - ttft_timeout=_ttft_timeout, - stream_idle_timeout=_stream_idle_timeout, - ) - if self._should_raise_content_policy_error( - model=model, response=reconstructed, kwargs=kwargs - ): - raise litellm.ContentPolicyViolationError( - message="Response output was blocked.", - model=model, - llm_provider="", - ) - return reconstructed return await self._acompletion_streaming_iterator( model_response=response, messages=messages, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index c40cf014b17..27787762240 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4778,6 +4778,17 @@ async def _async_chunks(*chunks): yield chunk +def _fake_stream(aiter_factory, model: str = "gpt-4o") -> MagicMock: + from litellm.utils import CustomStreamWrapper + + stream = MagicMock(spec=CustomStreamWrapper) + stream.model = model + stream.custom_llm_provider = "openai" + stream.__aiter__ = lambda self: aiter_factory() + stream.aclose = AsyncMock() + return stream + + @pytest.mark.asyncio async def test_router_ttft_timeout_returns_non_streaming_response(): """Router reconstructs a non-streaming ModelResponse when provider streams normally.""" @@ -4801,10 +4812,7 @@ async def test_router_ttft_timeout_returns_non_streaming_response(): _make_chunk("", finish_reason="stop"), ] - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - fake_stream.__aiter__ = lambda self: _async_chunks(*chunks) + fake_stream = _fake_stream(lambda: _async_chunks(*chunks)) reconstructed = MagicMock(spec=ModelResponse) reconstructed.choices = [MagicMock()] @@ -4819,6 +4827,7 @@ async def test_router_ttft_timeout_returns_non_streaming_response(): ) assert result is reconstructed + fake_stream.aclose.assert_awaited_once() @pytest.mark.asyncio @@ -4841,10 +4850,7 @@ async def test_router_ttft_timeout_raises_on_hung_provider(): return yield - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - fake_stream.__aiter__ = lambda self: hung_stream() + fake_stream = _fake_stream(hung_stream) with pytest.raises(litellm.Timeout) as exc_info: await router._collect_stream_with_ttft_timeout( @@ -4854,6 +4860,7 @@ async def test_router_ttft_timeout_raises_on_hung_provider(): ) assert "ttft_timeout" in str(exc_info.value) + fake_stream.aclose.assert_awaited_once() @pytest.mark.asyncio @@ -4877,10 +4884,7 @@ async def test_router_ttft_timeout_not_reset_by_preamble_chunks(): await asyncio.sleep(0.05) await asyncio.sleep(10) - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - fake_stream.__aiter__ = lambda self: preamble_only_stream() + fake_stream = _fake_stream(preamble_only_stream) with pytest.raises(litellm.Timeout) as exc_info: await router._collect_stream_with_ttft_timeout( @@ -4912,10 +4916,7 @@ async def test_router_ttft_timeout_empty_stream_raises_api_error(): return yield # make it an async generator - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - fake_stream.__aiter__ = lambda self: empty_stream() + fake_stream = _fake_stream(empty_stream) with pytest.raises(litellm.APIError): await router._collect_stream_with_ttft_timeout( @@ -4987,7 +4988,8 @@ async def test_router_ttft_timeout_acompletion_intercept(): @pytest.mark.asyncio async def test_router_stream_idle_timeout_acompletion_intercept(): """When stream_idle_timeout is set and stream=False, _acompletion forces stream=True - and passes stream_idle_timeout (with ttft_timeout=None) to _collect_stream_with_ttft_timeout.""" + and passes stream_idle_timeout (with ttft_timeout=None) to _collect_stream_with_ttft_timeout. + """ from unittest.mock import AsyncMock, patch import litellm @@ -5071,11 +5073,7 @@ async def test_router_stream_idle_timeout_raises_on_stalled_provider(): await asyncio.sleep(10) # stalls; will be killed by stream_idle_timeout yield _make_chunk("", finish_reason="stop") - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - gen = _stalled_after_first() - fake_stream.__aiter__ = lambda self: gen + fake_stream = _fake_stream(_stalled_after_first) with pytest.raises(litellm.Timeout, match="stream_idle_timeout"): await router._collect_stream_with_ttft_timeout( @@ -5085,6 +5083,8 @@ async def test_router_stream_idle_timeout_raises_on_stalled_provider(): stream_idle_timeout=0.05, ) + fake_stream.aclose.assert_awaited_once() + @pytest.mark.asyncio async def test_router_stream_idle_timeout_completes_when_not_stalled(): @@ -5111,10 +5111,7 @@ async def test_router_stream_idle_timeout_completes_when_not_stalled(): _make_chunk("", finish_reason="stop"), ] - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - fake_stream.__aiter__ = lambda self: _async_chunks(*chunks) + fake_stream = _fake_stream(lambda: _async_chunks(*chunks)) reconstructed = MagicMock(spec=ModelResponse) @@ -5153,11 +5150,7 @@ async def test_router_ttft_and_idle_timeout_both_active(): await asyncio.sleep(10) yield _make_chunk("", finish_reason="stop") - fake_stream = MagicMock() - fake_stream.model = "gpt-4o" - fake_stream.custom_llm_provider = "openai" - gen = _stalled_after_first() - fake_stream.__aiter__ = lambda self: gen + fake_stream = _fake_stream(_stalled_after_first) with pytest.raises(litellm.Timeout, match="stream_idle_timeout"): await router._collect_stream_with_ttft_timeout( @@ -5166,3 +5159,151 @@ async def test_router_ttft_and_idle_timeout_both_active(): ttft_timeout=5.0, stream_idle_timeout=0.05, ) + + +@pytest.mark.asyncio +async def test_router_ttft_timeout_closes_stream_on_caller_cancellation(): + """If the caller is cancelled mid-reconstruction, the upstream stream is closed so the + provider connection is released instead of leaked.""" + import asyncio + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + ], + stream_idle_timeout=30.0, + ) + + first_chunk_consumed = asyncio.Event() + + async def slow_stream(): + yield _make_chunk("Hello") + first_chunk_consumed.set() + await asyncio.sleep(10) # caller cancels while we wait here + + fake_stream = _fake_stream(slow_stream) + + task = asyncio.create_task( + router._collect_stream_with_ttft_timeout( + response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + ttft_timeout=None, + stream_idle_timeout=30.0, + ) + ) + await first_chunk_consumed.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + fake_stream.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_router_semaphore_held_through_reconstruction(): + """The max_parallel_requests semaphore must stay held while the promoted stream is drained + and reconstructed; otherwise a stream=False caller can exceed the configured concurrency. + """ + import asyncio + + from litellm import ModelResponse + from litellm.utils import CustomStreamWrapper + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + ], + ttft_timeout=5.0, + ) + + semaphore = asyncio.Semaphore(1) + reconstructed = MagicMock(spec=ModelResponse) + locked_while_reconstructing = {} + + async def fake_collect(**kwargs): + locked_while_reconstructing["value"] = semaphore.locked() + return reconstructed + + fake_stream = MagicMock(spec=CustomStreamWrapper) + fake_stream.model = "gpt-4o" + fake_stream.custom_llm_provider = "openai" + + def fake_get_client(deployment, kwargs, client_type): + return semaphore if client_type == "max_parallel_requests" else None + + with ( + patch.object(router, "_collect_stream_with_ttft_timeout", new=fake_collect), + patch.object(router, "async_get_available_deployment") as mock_dep, + patch.object(router, "_update_kwargs_with_deployment"), + patch.object(router, "_get_client", side_effect=fake_get_client), + patch.object(router, "_track_deployment_metrics"), + patch.object(router, "_should_raise_content_policy_error", return_value=False), + patch("litellm.acompletion", new_callable=AsyncMock, return_value=fake_stream), + ): + mock_dep.return_value = { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + + result = await router._acompletion( + model="test-model", + messages=[{"role": "user", "content": "hi"}], + stream=False, + ) + + assert result is reconstructed + assert locked_while_reconstructing["value"] is True + + +@pytest.mark.asyncio +async def test_router_ttft_timeout_tags_failed_deployment_id(): + """A ttft Timeout must stamp the failed deployment's id on the exception so weighted + failover / cooldown can exclude it on retry instead of re-picking the hung deployment. + """ + import asyncio + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + "model_info": {"id": "deploy-1"}, + } + ], + ttft_timeout=0.05, + ) + + async def hung_stream(): + await asyncio.sleep(10) + return + yield + + fake_stream = _fake_stream(hung_stream) + + with ( + patch.object(router, "async_get_available_deployment") as mock_dep, + patch.object(router, "_update_kwargs_with_deployment"), + patch.object(router, "_get_client", return_value=None), + patch.object(router, "_track_deployment_metrics"), + patch("litellm.acompletion", new_callable=AsyncMock, return_value=fake_stream), + ): + mock_dep.return_value = { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + "model_info": {"id": "deploy-1"}, + } + + with pytest.raises(litellm.Timeout) as exc_info: + await router._acompletion( + model="test-model", + messages=[{"role": "user", "content": "hi"}], + stream=False, + ) + + assert getattr(exc_info.value, "failed_deployment_id", None) == "deploy-1"