mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
d66c3897a9
commit
883a3d19d1
2 changed files with 240 additions and 105 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue