mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
ec9353cb69
commit
ee2fda07d9
2 changed files with 279 additions and 67 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue