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)