fix(router): hold concurrency semaphore through ttft reconstruction and close stream on exit

The ttft_timeout / stream_idle_timeout path promotes a stream=False call to
streaming and drains it via _collect_stream_with_ttft_timeout. Two correctness
issues are addressed:

- max_parallel_requests semaphore: reconstruction previously ran after the
  'async with rpm_semaphore' block had already exited, so a stream=False caller
  could exceed the configured concurrency for the entire drain. Reconstruction
  now happens inside the semaphore via _await_response, restoring the
  non-streaming concurrency guarantee.

- Connection cleanup: the drain loop now runs under try/finally and calls
  response.aclose() on timeout, cancellation, or normal completion, so a caller
  disconnect mid-reconstruction releases the upstream connection instead of
  leaking it.

The reconstructed ModelResponse now flows through the shared content-policy
check and _track_deployment_metrics, removing the duplicated content-policy
block. The two identical ttft raise sites are collapsed into one.

Tests: semaphore-held-through-reconstruction (fails before the fix),
stream-closed-on-caller-cancellation, and a ttft Timeout tagging
failed_deployment_id so weighted failover / cooldown can exclude the hung
deployment on retry.
This commit is contained in:
Nathan Price 2026-06-15 08:37:42 -05:00
parent d66c3897a9
commit 883a3d19d1
2 changed files with 240 additions and 105 deletions

View file

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

View file

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