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().
This commit is contained in:
Nathan Price 2026-06-12 21:00:07 -05:00
parent ec9353cb69
commit ee2fda07d9
2 changed files with 279 additions and 67 deletions

View file

@ -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(

View file

@ -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)