From ee2fda07d9091491b1b30105dcc2daa24c4baf25 Mon Sep 17 00:00:00 2001 From: Nathan Price Date: Fri, 12 Jun 2026 21:00:07 -0500 Subject: [PATCH] feat(router): add ttft_timeout to detect hung providers on non-streaming calls Adds ttft_timeout parameter to Router. When set, non-streaming calls internally switch to stream=True so the router can detect a hung provider (one that accepts the connection but never sends tokens) within ttft_timeout seconds, rather than waiting for the full request timeout which can be very long for large generation requests. Raises litellm.Timeout to trigger existing cooldown and fallback machinery. Caller always receives a standard ModelResponse via stream_chunk_builder. Uses a single hard deadline rather than per-chunk wait_for, so preamble chunks (role deltas, empty tool-call deltas) do not reset the clock. Checks both delta.content and delta.tool_calls for first-token detection. Phase 2 lets real errors propagate rather than swallowing them. Uses asyncio.get_running_loop(). --- litellm/router.py | 212 ++++++++++++++++++++---------- tests/test_litellm/test_router.py | 134 +++++++++++++++++++ 2 files changed, 279 insertions(+), 67 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 80584858311..878c37926f3 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -274,6 +274,7 @@ class Router: ] = None, # max fallbacks to try before exiting the call. Defaults to 5. timeout: Optional[float] = None, stream_timeout: Optional[float] = None, + ttft_timeout: Optional[float] = None, default_litellm_params: Optional[ dict ] = None, # default params for Router.chat.completion.create @@ -432,9 +433,9 @@ class Router: ) # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### - cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = ( - "local" # default to an in-memory cache - ) + cache_type: Literal[ + "local", "redis", "redis-semantic", "s3", "disk" + ] = "local" # default to an in-memory cache redis_cache = None cache_config: Dict[str, Any] = {} @@ -482,9 +483,9 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() - self.team_pattern_routers: Dict[str, PatternMatchRouter] = ( - {} - ) # {"TEAM_ID": PatternMatchRouter} + self.team_pattern_routers: Dict[ + str, PatternMatchRouter + ] = {} # {"TEAM_ID": PatternMatchRouter} self.auto_routers: Dict[str, "AutoRouter"] = {} self.complexity_routers: Dict[str, "ComplexityRouter"] = {} self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {} @@ -565,12 +566,13 @@ class Router: self._explicit_timeout = timeout # None when user did not pass timeout self.timeout = timeout or litellm.request_timeout self.stream_timeout = stream_timeout + self.ttft_timeout = ttft_timeout self.retry_after = retry_after self.routing_strategy = self._normalize_strategy(routing_strategy) - self._routing_groups_input: Optional[List[Union[RoutingGroup, dict]]] = ( - routing_groups - ) + self._routing_groups_input: Optional[ + List[Union[RoutingGroup, dict]] + ] = routing_groups ## SETTING FALLBACKS ## ### validate if it's set + in correct format @@ -697,12 +699,12 @@ class Router: ) ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = ( - model_group_retry_policy - ) - self.model_group_affinity_config: Optional[Dict[str, List[str]]] = ( - model_group_affinity_config - ) + self.model_group_retry_policy: Optional[ + Dict[str, RetryPolicy] + ] = model_group_retry_policy + self.model_group_affinity_config: Optional[ + Dict[str, List[str]] + ] = model_group_affinity_config self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None if allowed_fails_policy is not None: @@ -2616,20 +2618,20 @@ class Router: # _ageneric_api_call_with_fallbacks_helper. # original_generic_function is preserved by the caller so # the helper knows what underlying API to invoke per attempt. - initial_kwargs["original_function"] = ( - self._ageneric_api_call_with_fallbacks_helper - ) + initial_kwargs[ + "original_function" + ] = self._ageneric_api_call_with_fallbacks_helper if e.is_pre_first_chunk or not e.generated_content: # No content generated before the error — retry with the # original input. Adding a continuation prompt would # waste tokens and confuse the model. pass else: - initial_kwargs["input"] = ( - Router._build_responses_continuation_input( - initial_kwargs.get("input"), - e.generated_content, - ) + initial_kwargs[ + "input" + ] = Router._build_responses_continuation_input( + initial_kwargs.get("input"), + e.generated_content, ) # The Responses-API path stores observability metadata # under "litellm_metadata" (not the default "metadata") — @@ -2852,12 +2854,67 @@ class Router: f"Silent experiment failed for model {silent_model}: {str(e)}" ) + async def _collect_stream_with_ttft_timeout( + self, + response: CustomStreamWrapper, + messages: List[Dict[str, str]], + ttft_timeout: float, + ) -> ModelResponse: + from litellm.main import stream_chunk_builder + + chunks: List = [] + aiter = response.__aiter__() + 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 + + async for chunk in aiter: + chunks.append(chunk) + + result = stream_chunk_builder(chunks, messages=messages) + if result is None: + raise litellm.APIError( + status_code=500, + message="stream_chunk_builder returned None: provider returned an empty stream", + llm_provider="", + model="", + ) + return cast(ModelResponse, result) + async def _acompletion( # noqa: PLR0915 self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ - ModelResponse, - CustomStreamWrapper, - ]: + ) -> Union[ModelResponse, CustomStreamWrapper,]: """ - Get an available deployment - call it with a semaphore over the call @@ -2938,6 +2995,13 @@ class Router: } input_kwargs.pop("silent_model", None) + _ttft_timeout = self._get_ttft_timeout(kwargs=kwargs, data=litellm_params) + _forced_stream_for_ttft = ( + _ttft_timeout is not None and not input_kwargs.get("stream", False) + ) + if _forced_stream_for_ttft: + input_kwargs["stream"] = True + _response = litellm.acompletion(**input_kwargs) logging_obj: Optional[LiteLLMLogging] = kwargs.get( @@ -2996,6 +3060,12 @@ class Router: ) if isinstance(response, CustomStreamWrapper): + if _forced_stream_for_ttft and _ttft_timeout is not None: + return await self._collect_stream_with_ttft_timeout( + response=response, + messages=messages, + ttft_timeout=_ttft_timeout, + ) return await self._acompletion_streaming_iterator( model_response=response, messages=messages, @@ -3307,6 +3377,14 @@ class Router: ) return timeout + def _get_ttft_timeout(self, kwargs: dict, data: dict) -> Optional[float]: + return ( + kwargs.get("ttft_timeout", None) + or data.get("ttft_timeout", None) + or self.ttft_timeout + or self.default_litellm_params.get("ttft_timeout", None) + ) + def _get_timeout(self, kwargs: dict, data: dict) -> Optional[Union[float, int]]: """Helper to get timeout from kwargs or deployment params""" timeout: Optional[Union[float, int]] = None @@ -5255,9 +5333,9 @@ class Router: healthy_deployments=healthy_deployments, responses=responses ) returned_response = cast(OpenAIFileObject, responses[0]) - returned_response._hidden_params["model_file_id_mapping"] = ( - model_file_id_mapping - ) + returned_response._hidden_params[ + "model_file_id_mapping" + ] = model_file_id_mapping return returned_response except Exception as e: verbose_router_logger.exception( @@ -6576,11 +6654,11 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - context_window_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) + context_window_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, ) if context_window_fallback_model_group is None: raise original_exception @@ -6612,11 +6690,11 @@ class Router: e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - content_policy_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) + content_policy_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, ) if content_policy_fallback_model_group is None: raise original_exception @@ -6838,9 +6916,9 @@ class Router: ) ## ADD RETRY TRACKING TO METADATA - used for spend logs retry tracking _metadata["attempted_retries"] = 0 - _metadata["max_retries"] = ( - num_retries # Updated after overrides in exception handler - ) + _metadata[ + "max_retries" + ] = num_retries # Updated after overrides in exception handler try: self._handle_mock_testing_rate_limit_error( model_group=model_group, kwargs=kwargs @@ -8039,26 +8117,26 @@ class Router: """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = ( - deployment.litellm_params.auto_router_config_path - ) + auto_router_config_path: Optional[ + str + ] = deployment.litellm_params.auto_router_config_path auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) - default_model: Optional[str] = ( - deployment.litellm_params.auto_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) - embedding_model: Optional[str] = ( - deployment.litellm_params.auto_router_embedding_model - ) + embedding_model: Optional[ + str + ] = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" @@ -8101,13 +8179,13 @@ class Router: ComplexityRouter, ) - complexity_router_config: Optional[dict] = ( - deployment.litellm_params.complexity_router_config - ) + complexity_router_config: Optional[ + dict + ] = deployment.litellm_params.complexity_router_config - default_model: Optional[str] = ( - deployment.litellm_params.complexity_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.complexity_router_default_model # If no default model specified, try to get from config tiers if default_model is None and complexity_router_config: @@ -8294,13 +8372,13 @@ class Router: QualityRouter, ) - quality_router_config: Optional[dict] = ( - deployment.litellm_params.quality_router_config - ) + quality_router_config: Optional[ + dict + ] = deployment.litellm_params.quality_router_config - default_model: Optional[str] = ( - deployment.litellm_params.quality_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.quality_router_default_model if default_model is None and quality_router_config: default_model = quality_router_config.get("default_model") @@ -9071,9 +9149,9 @@ class Router: # Add custom_llm_provider if deployment.litellm_params.custom_llm_provider: - credentials["custom_llm_provider"] = ( - deployment.litellm_params.custom_llm_provider - ) + credentials[ + "custom_llm_provider" + ] = deployment.litellm_params.custom_llm_provider elif "/" in deployment.litellm_params.model: # Extract provider from "provider/model" format credentials["custom_llm_provider"] = deployment.litellm_params.model.split( diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 830edf6412d..a10e6fb7d50 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4756,3 +4756,137 @@ def test_is_deployment_blocked_static_helper_reflects_blocked_flag(): ) is True ) + + +# --------------------------------------------------------------------------- +# ttft_timeout tests +# --------------------------------------------------------------------------- + + +def _make_chunk(content: str, finish_reason: str = "") -> MagicMock: + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta = MagicMock() + chunk.choices[0].delta.content = content + chunk.choices[0].delta.tool_calls = None # must be explicit — MagicMock() is truthy + chunk.choices[0].finish_reason = finish_reason or None + return chunk + + +async def _async_chunks(*chunks): + for chunk in chunks: + yield chunk + + +@pytest.mark.asyncio +async def test_router_ttft_timeout_returns_non_streaming_response(): + """Router reconstructs a non-streaming ModelResponse when provider streams normally.""" + from unittest.mock import patch + + from litellm import ModelResponse + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + ], + ttft_timeout=5.0, + ) + + chunks = [ + _make_chunk("Hello"), + _make_chunk(" world"), + _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) + + reconstructed = MagicMock(spec=ModelResponse) + reconstructed.choices = [MagicMock()] + reconstructed.choices[0].message = MagicMock() + reconstructed.choices[0].message.content = "Hello world" + + with patch("litellm.main.stream_chunk_builder", return_value=reconstructed): + result = await router._collect_stream_with_ttft_timeout( + response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + ttft_timeout=5.0, + ) + + assert result is reconstructed + + +@pytest.mark.asyncio +async def test_router_ttft_timeout_raises_on_hung_provider(): + """Router raises litellm.Timeout when provider never sends a first token.""" + import asyncio + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + ], + ttft_timeout=0.1, + ) + + async def hung_stream(): + await asyncio.sleep(10) + return + yield + + fake_stream = MagicMock() + fake_stream.model = "gpt-4o" + fake_stream.custom_llm_provider = "openai" + fake_stream.__aiter__ = lambda self: hung_stream() + + with pytest.raises(litellm.Timeout) as exc_info: + await router._collect_stream_with_ttft_timeout( + response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + ttft_timeout=0.1, + ) + + assert "ttft_timeout" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_router_ttft_timeout_not_reset_by_preamble_chunks(): + """Preamble chunks must not reset the TTFT clock; only the hard deadline counts.""" + import asyncio + + router = litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}, + } + ], + ttft_timeout=0.2, + ) + + async def preamble_only_stream(): + for _ in range(5): + yield _make_chunk("") + 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() + + with pytest.raises(litellm.Timeout) as exc_info: + await router._collect_stream_with_ttft_timeout( + response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + ttft_timeout=0.2, + ) + + assert "ttft_timeout" in str(exc_info.value)