From d56ff13b191ce4ba27d026c9a015e9c490978d2b Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 27 Jul 2026 21:06:25 +0000 Subject: [PATCH] fix(proxy): cover the upstream call with the SSE keepalive The silent window for reasoning models is mostly spent waiting on the provider's response headers, before create_response is reached, so wrap the whole streaming request instead --- litellm/constants.py | 7 - litellm/proxy/common_request_processing.py | 218 +++++++++++------- .../proxy/test_common_request_processing.py | 203 ++++++++-------- 3 files changed, 227 insertions(+), 201 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 8d5166ef129..35247225f35 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -213,13 +213,6 @@ AIOHTTP_SO_KEEPALIVE = os.getenv("AIOHTTP_SO_KEEPALIVE", "False").lower() == "tr AIOHTTP_TCP_KEEPIDLE = int(os.getenv("AIOHTTP_TCP_KEEPIDLE", 60)) AIOHTTP_TCP_KEEPINTVL = int(os.getenv("AIOHTTP_TCP_KEEPINTVL", 30)) AIOHTTP_TCP_KEEPCNT = int(os.getenv("AIOHTTP_TCP_KEEPCNT", 5)) -# Application-level SSE keepalive, in seconds; 0 disables it. While a streaming -# response has not produced its first token the proxy writes nothing to the -# client, so any hop with an idle timeout (AWS ALB and nginx default to 60s) -# reaps a healthy connection during a long time-to-first-token. The socket-level -# knobs above do not help; they never write a byte into the in-flight response. -# When set, the proxy emits an SSE comment frame at this interval until the -# first chunk arrives, which resets those idle watchdogs. SSE_KEEPALIVE_INTERVAL_SECONDS = float(os.getenv("SSE_KEEPALIVE_INTERVAL_SECONDS", 0)) # enable_cleanup_closed is only needed for Python versions with the SSL leak bug # Fixed in Python 3.12.7+ and 3.13.1+ (see https://github.com/python/cpython/pull/118960) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f12eacc3ac1..8be3c82502b 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,4 +1,5 @@ import asyncio +import functools import json import logging import math @@ -9,10 +10,13 @@ from typing import ( TYPE_CHECKING, Any, AsyncGenerator, + Awaitable, Callable, + Concatenate, Dict, Literal, Optional, + ParamSpec, Tuple, Union, ) @@ -86,8 +90,6 @@ from litellm.types.utils import ( _DD_STREAMING_TRACE_ENABLED = not isinstance(tracer, NullTracer) -# SSE comment frame: every compliant SSE client discards it, so it is a safe way -# to put bytes on the wire while a stream has produced no tokens yet. SSE_KEEPALIVE_FRAME = ": litellm-keepalive\n\n" @@ -466,8 +468,7 @@ async def _wait_for_http_disconnect(request: Request) -> None: async def _buffer_first_chunk_honoring_disconnect( generator: AsyncGenerator[str, None], request: Optional[Request], - keepalive_interval_seconds: float = 0.0, -) -> str | asyncio.Task[str]: +) -> str: """Fetch the first streamed chunk, cancelling the upstream LLM call if the client disconnects before it arrives. @@ -478,42 +479,28 @@ async def _buffer_first_chunk_honoring_disconnect( until the request timeout (LIT-3568). Cancelling the fetch propagates into async_streaming_data_generator, whose finally block records the 499 and closes the upstream stream. - - With keepalive_interval_seconds set, buffering is bounded by that interval: - the still-pending fetch is handed back to the caller so it can start the SSE - response and write keepalive frames instead of staying silent on the wire. """ - if request is None and keepalive_interval_seconds <= 0: + if request is None: return await generator.__anext__() chunk_task: asyncio.Task[str] = asyncio.ensure_future(generator.__anext__()) - disconnect_task: asyncio.Task[None] | None = ( - asyncio.ensure_future(_wait_for_http_disconnect(request)) if request is not None else None - ) + disconnect_task: asyncio.Task[None] = asyncio.ensure_future(_wait_for_http_disconnect(request)) try: - await asyncio.wait( - tuple(task for task in (chunk_task, disconnect_task) if task is not None), - timeout=keepalive_interval_seconds if keepalive_interval_seconds > 0 else None, - return_when=asyncio.FIRST_COMPLETED, - ) + await asyncio.wait({chunk_task, disconnect_task}, return_when=asyncio.FIRST_COMPLETED) # A completed disconnect_task has already consumed the http.disconnect # message, so Starlette's later listen_for_disconnect would never see it. # Take the cancellation path whenever a disconnect was observed, even if # the first chunk landed in the same scheduler turn. - disconnect_observed = disconnect_task is not None and disconnect_task.done() + disconnect_observed = disconnect_task.done() finally: - if disconnect_task is not None: - disconnect_task.cancel() - try: - await disconnect_task - except BaseException: # noqa: BLE001 - pass + disconnect_task.cancel() + try: + await disconnect_task + except BaseException: # noqa: BLE001 + pass - if not disconnect_observed: - if chunk_task.done() and not chunk_task.cancelled(): - return chunk_task.result() - if keepalive_interval_seconds > 0: - return chunk_task + if not disconnect_observed and chunk_task.done() and not chunk_task.cancelled(): + return chunk_task.result() chunk_task.cancel() with anyio.CancelScope(shield=True): @@ -551,64 +538,120 @@ def _build_stream_error_payload(e: Exception) -> tuple[int, dict[str, Any]]: return error_status, error_obj -async def _stream_chunks( - first_chunk_value: str | None, - generator: AsyncGenerator[str, None], -) -> AsyncGenerator[str, None]: - if not _DD_STREAMING_TRACE_ENABLED: - # Fast path: no per-chunk span object / context-manager overhead. - if first_chunk_value is not None: - yield first_chunk_value - async for chunk in generator: - yield chunk - return - if first_chunk_value is not None: - with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): - yield first_chunk_value - async for chunk in generator: - with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): - yield chunk +def _sse_error_frame(error_obj: dict[str, Any]) -> str: + return f"data: {json.dumps({'error': error_obj})}\n\ndata: [DONE]\n\n" -async def _keepalive_until_first_chunk( - chunk_task: asyncio.Task[str], - generator: AsyncGenerator[str, None], +def _error_obj_from_response(response: Response) -> dict[str, Any]: + body = getattr(response, "body", b"") + decoded = body.decode("utf-8") if isinstance(body, bytes) else str(body) + fallback: dict[str, Any] = { + "message": decoded, + "type": "None", + "param": "None", + "code": str(response.status_code), + } + try: + parsed = json.loads(decoded) + except json.JSONDecodeError: + return fallback + error_obj = parsed.get("error") if isinstance(parsed, dict) else None + return error_obj if isinstance(error_obj, dict) else fallback + + +async def _close_stream_targets(response: Response | None) -> None: + with anyio.CancelScope(shield=True): + for target in ( + getattr(response, "body_iterator", None), + getattr(response, "_upstream_generator", None), + ): + aclose = getattr(target, "aclose", None) + if aclose is None: + continue + try: + await aclose() + except BaseException as e: # noqa: BLE001 + verbose_proxy_logger.debug("error closing streaming generator after keepalive: %s", e) + + +async def _keepalive_until_response( + response_task: "asyncio.Task[Response]", keepalive_interval_seconds: float, ) -> AsyncGenerator[str, None]: - """Stream SSE comment frames until the first chunk lands, then the response. + """Heartbeat until the request coroutine returns a response, then forward it. - Comment frames are ignored by every SSE client but are real bytes on the - wire, so idle watchdogs between the client and the proxy stop reaping - healthy slow-TTFT streams. The response status and headers are already - committed by the time the first chunk arrives, so an error-only stream is - delivered as an SSE error frame rather than the JSON body create_response - returns when it manages to buffer the first chunk in time. + Providers hold back the upstream response headers while a reasoning model + thinks (OpenAI does this for over a minute on gpt-5.x), so the proxy has + not entered the ASGI response phase yet and nothing further down the + streaming path can put a byte on the wire. Committing the SSE response + early and heartbeating covers that whole window; the tradeoff is that the + per-request ``x-litellm-*`` headers and a non-200 status can no longer be + set, so a failure surfacing after the first heartbeat is delivered as an + SSE error frame. """ + response: Response | None = None try: - while not chunk_task.done(): - await asyncio.wait((chunk_task,), timeout=keepalive_interval_seconds) - if not chunk_task.done(): + while not response_task.done(): + await asyncio.wait((response_task,), timeout=keepalive_interval_seconds) + if not response_task.done(): yield SSE_KEEPALIVE_FRAME try: - first_chunk_value = chunk_task.result() - except StopAsyncIteration: - return + response = response_task.result() except Exception as e: # noqa: BLE001 - verbose_proxy_logger.exception(f"Error consuming first chunk from generator: {e}") + verbose_proxy_logger.exception(f"Error starting stream after keepalive response start: {e}") _, error_obj = _build_stream_error_payload(e) - yield f"data: {json.dumps({'error': error_obj})}\n\n" - yield "data: [DONE]\n\n" + yield _sse_error_frame(error_obj) return - async for chunk in _stream_chunks(first_chunk_value, generator): - yield chunk + if not isinstance(response, StreamingResponse): + yield _sse_error_frame(_error_obj_from_response(response)) + return + async for chunk in response.body_iterator: + yield chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk finally: - if not chunk_task.done(): - chunk_task.cancel() + if not response_task.done(): + response_task.cancel() with anyio.CancelScope(shield=True): try: - await chunk_task + await response_task except BaseException: # noqa: BLE001 pass + await _close_stream_targets(response) + + +_P = ParamSpec("_P") + + +def sse_keepalive_during_slow_ttft( + process_request: Callable[Concatenate["ProxyBaseLLMRequestProcessing", _P], Awaitable[Any]], + *, + interval_seconds: Callable[[], float] = lambda: SSE_KEEPALIVE_INTERVAL_SECONDS, +) -> Callable[Concatenate["ProxyBaseLLMRequestProcessing", _P], Awaitable[Any]]: + """Start the SSE response and heartbeat when a streaming request stays silent. + + Disabled unless ``SSE_KEEPALIVE_INTERVAL_SECONDS`` is positive, in which + case a streaming request that has produced nothing after that many seconds + gets its response committed early so idle watchdogs (AWS ALB and nginx + default to 60s) see bytes instead of reaping a healthy connection. + """ + + @functools.wraps(process_request) + async def wrapper(self: "ProxyBaseLLMRequestProcessing", *args: _P.args, **kwargs: _P.kwargs) -> Any: + keepalive_interval_seconds = interval_seconds() + if keepalive_interval_seconds <= 0 or self.data.get("stream") is not True: + return await process_request(self, *args, **kwargs) + + response_task = asyncio.create_task(process_request(self, *args, **kwargs)) + await asyncio.wait((response_task,), timeout=keepalive_interval_seconds) + if response_task.done(): + return response_task.result() + + return StreamingResponse( + _keepalive_until_response(response_task, keepalive_interval_seconds), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, + ) + + return wrapper async def create_response( @@ -617,7 +660,6 @@ async def create_response( headers: dict, default_status_code: int = status.HTTP_200_OK, request: Optional[Request] = None, - keepalive_interval_seconds: float = SSE_KEEPALIVE_INTERVAL_SECONDS, ) -> Union[StreamingResponse, JSONResponse]: """ Create streaming response, checking if the first chunk is an error. @@ -640,20 +682,7 @@ async def create_response( generator = await generator # Now get the first chunk from the actual generator - buffered = await _buffer_first_chunk_honoring_disconnect(generator, request, keepalive_interval_seconds) - - if isinstance(buffered, asyncio.Task): - # Time-to-first-token exceeded the keepalive interval; start the SSE - # response now and heartbeat until the model produces something. - return _UpstreamClosingStreamingResponse( - _keepalive_until_first_chunk(buffered, generator, keepalive_interval_seconds), - media_type=media_type, - headers=streaming_headers, - status_code=default_status_code, - upstream_generator=generator, - ) - - first_chunk_value = buffered + first_chunk_value = await _buffer_first_chunk_honoring_disconnect(generator, request) if first_chunk_value is not None: try: @@ -715,7 +744,6 @@ async def create_response( # Unexpected error consuming first chunk. verbose_proxy_logger.exception(f"Error consuming first chunk from generator: {e}") - # Preserve status code from HTTPException (e.g., guardrail blocks) error_status, error_obj = _build_stream_error_payload(e) async def error_gen_message() -> AsyncGenerator[str, None]: @@ -729,8 +757,23 @@ async def create_response( status_code=error_status, ) + async def combined_generator() -> AsyncGenerator[str, None]: + if not _DD_STREAMING_TRACE_ENABLED: + # Fast path: no per-chunk span object / context-manager overhead. + if first_chunk_value is not None: + yield first_chunk_value + async for chunk in generator: + yield chunk + return + if first_chunk_value is not None: + with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): + yield first_chunk_value + async for chunk in generator: + with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): + yield chunk + return _UpstreamClosingStreamingResponse( - _stream_chunks(first_chunk_value, generator), + combined_generator(), media_type=media_type, headers=streaming_headers, status_code=final_status_code, @@ -1565,6 +1608,7 @@ class ProxyBaseLLMRequestProcessing: _payload_str, ) + @sse_keepalive_during_slow_ttft async def base_process_llm_request( self, request: Request, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index a2e09b90661..acba57e66d8 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -32,6 +32,7 @@ from litellm.proxy.common_request_processing import ( _should_return_raw_model_name, _UpstreamClosingStreamingResponse, create_response, + sse_keepalive_during_slow_ttft, SSE_KEEPALIVE_FRAME, ) from litellm.proxy.dd_span_tagger import DDSpanTagger @@ -1593,153 +1594,141 @@ class TestCommonRequestProcessingHelpers: @pytest.mark.asyncio -class TestSSEKeepaliveDuringTimeToFirstToken: - """Regression coverage for #34819: with SSE_KEEPALIVE_INTERVAL_SECONDS set, - a stream whose time-to-first-token is long must put bytes on the wire before - an intermediary's idle timeout (AWS ALB and nginx default to 60s) reaps it. +class TestSSEKeepaliveDuringSlowTimeToFirstToken: + """Regression coverage for #34819: with SSE_KEEPALIVE_INTERVAL_SECONDS set, a + streaming request that has produced nothing yet must put bytes on the wire + before an intermediary's idle timeout (AWS ALB and nginx default to 60s) + reaps it. The silent window covers the upstream call itself, which for + reasoning models is where most of it is spent, so the keepalive wraps the + whole request rather than only the first-chunk buffering in create_response. """ @staticmethod - def _connected_request() -> Request: - async def receive(): - await asyncio.Event().wait() + def _processor() -> ProxyBaseLLMRequestProcessing: + return ProxyBaseLLMRequestProcessing(data={"model": "gpt-5.6", "stream": True}) - return Request({"type": "http", "method": "POST", "path": "/", "headers": []}, receive) + @staticmethod + async def _drain(response: StreamingResponse) -> list: + return [chunk async for chunk in response.body_iterator] - async def test_keepalive_frames_are_sent_while_first_token_is_pending(self): - first_token = asyncio.Event() + async def test_heartbeats_until_the_upstream_call_returns_then_streams_it(self): + upstream_ready = asyncio.Event() - async def slow_generator(): - await first_token.wait() - yield 'data: {"content": "hi"}\n\n' - yield "data: [DONE]\n\n" + async def slow_upstream(_self) -> StreamingResponse: + await upstream_ready.wait() + + async def body(): + yield 'data: {"content": "hi"}\n\n' + yield "data: [DONE]\n\n" + + return StreamingResponse(body(), media_type="text/event-stream") + + wrapped = sse_keepalive_during_slow_ttft(slow_upstream, interval_seconds=lambda: 0.01) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) - response = await asyncio.wait_for( - create_response( - slow_generator(), - "text/event-stream", - {}, - request=self._connected_request(), - keepalive_interval_seconds=0.01, - ), - timeout=5, - ) assert isinstance(response, StreamingResponse) - assert response.status_code == status.HTTP_200_OK + assert response.media_type == "text/event-stream" + body_iterator = response.body_iterator.__aiter__() + heartbeats = [await asyncio.wait_for(body_iterator.__anext__(), timeout=5) for _ in range(3)] + assert heartbeats == [SSE_KEEPALIVE_FRAME] * 3 - body = response.body_iterator.__aiter__() - keepalives = [await asyncio.wait_for(body.__anext__(), timeout=5) for _ in range(3)] - assert keepalives == [SSE_KEEPALIVE_FRAME] * 3 + upstream_ready.set() + assert [chunk async for chunk in body_iterator] == ['data: {"content": "hi"}\n\n', "data: [DONE]\n\n"] - first_token.set() - remaining = [chunk async for chunk in body] - assert remaining == ['data: {"content": "hi"}\n\n', "data: [DONE]\n\n"] - - async def test_no_keepalive_frames_when_interval_is_disabled(self): - async def slow_generator(): + async def test_no_heartbeats_when_the_setting_is_disabled(self): + async def slow_upstream(_self) -> StreamingResponse: await asyncio.sleep(0.05) - yield 'data: {"content": "hi"}\n\n' - yield "data: [DONE]\n\n" - response = await asyncio.wait_for( - create_response( - slow_generator(), - "text/event-stream", - {}, - request=self._connected_request(), - ), - timeout=5, - ) - chunks = [chunk async for chunk in response.body_iterator] - assert chunks == ['data: {"content": "hi"}\n\n', "data: [DONE]\n\n"] + async def body(): + yield "data: [DONE]\n\n" - async def test_error_only_stream_still_returns_json_when_first_chunk_is_fast(self): - async def error_generator(): - yield 'data: {"error": {"code": 403, "message": "forbidden"}}\n\n' + return StreamingResponse(body(), media_type="text/event-stream") + + wrapped = sse_keepalive_during_slow_ttft(slow_upstream, interval_seconds=lambda: 0) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) + + assert await self._drain(response) == ["data: [DONE]\n\n"] + + async def test_non_streaming_request_is_untouched(self): + async def slow_upstream(_self) -> JSONResponse: + await asyncio.sleep(0.05) + return JSONResponse(status_code=200, content={"id": "chatcmpl-1"}) + + processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-5.6"}) + wrapped = sse_keepalive_during_slow_ttft(slow_upstream, interval_seconds=lambda: 0.01) + response = await asyncio.wait_for(wrapped(processor), timeout=5) - response = await asyncio.wait_for( - create_response( - error_generator(), - "text/event-stream", - {}, - request=self._connected_request(), - keepalive_interval_seconds=5, - ), - timeout=5, - ) assert isinstance(response, JSONResponse) - assert response.status_code == status.HTTP_403_FORBIDDEN - async def test_error_raised_after_keepalives_is_delivered_as_sse_error_frame(self): - release = asyncio.Event() + async def test_fast_request_keeps_its_status_and_headers(self): + async def fast_error(_self) -> JSONResponse: + return JSONResponse(status_code=429, content={"error": {"message": "rate limited"}}) - async def failing_generator(): - await release.wait() + wrapped = sse_keepalive_during_slow_ttft(fast_error, interval_seconds=lambda: 5) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) + + assert isinstance(response, JSONResponse) + assert response.status_code == 429 + + async def test_failure_after_heartbeats_is_delivered_as_an_sse_error_frame(self): + async def slow_failure(_self) -> StreamingResponse: + await asyncio.sleep(0.05) raise HTTPException(status_code=429, detail="rate limited") - yield "unreachable" - response = await asyncio.wait_for( - create_response( - failing_generator(), - "text/event-stream", - {}, - request=self._connected_request(), - keepalive_interval_seconds=0.01, - ), - timeout=5, - ) - # Status and headers are already committed once keepalives start, so the - # failure can only be reported inside the stream. - assert response.status_code == status.HTTP_200_OK + wrapped = sse_keepalive_during_slow_ttft(slow_failure, interval_seconds=lambda: 0.01) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) - body = response.body_iterator.__aiter__() - assert await asyncio.wait_for(body.__anext__(), timeout=5) == SSE_KEEPALIVE_FRAME - - release.set() - remaining = [chunk async for chunk in body] - assert remaining[-1] == "data: [DONE]\n\n" - error = json.loads(remaining[0][len("data: ") :])["error"] + chunks = await self._drain(response) + assert chunks[0] == SSE_KEEPALIVE_FRAME + error_frame, done_frame = chunks[-1].split("\n\n")[0], chunks[-1].split("\n\n")[1] + error = json.loads(error_frame[len("data: ") :])["error"] assert error["message"] == "rate limited" assert error["code"] == "429" + assert done_frame == "data: [DONE]" - async def test_client_disconnect_during_keepalives_closes_upstream_stream(self): - upstream_closed = asyncio.Event() + async def test_error_response_after_heartbeats_is_delivered_as_an_sse_error_frame(self): + async def slow_error_response(_self) -> JSONResponse: + await asyncio.sleep(0.05) + return JSONResponse(status_code=403, content={"error": {"message": "forbidden", "code": "403"}}) - async def never_first_token(): + wrapped = sse_keepalive_during_slow_ttft(slow_error_response, interval_seconds=lambda: 0.01) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) + + chunks = await self._drain(response) + assert chunks[0] == SSE_KEEPALIVE_FRAME + assert json.loads(chunks[-1].split("\n\n")[0][len("data: ") :])["error"]["message"] == "forbidden" + + async def test_client_disconnect_during_heartbeats_cancels_the_upstream_call(self): + upstream_cancelled = asyncio.Event() + + async def never_returns(_self) -> StreamingResponse: try: await asyncio.Event().wait() - yield "unreachable" - finally: - upstream_closed.set() + except asyncio.CancelledError: + upstream_cancelled.set() + raise + raise AssertionError("unreachable") - response = await asyncio.wait_for( - create_response( - never_first_token(), - "text/event-stream", - {}, - request=self._connected_request(), - keepalive_interval_seconds=0.01, - ), - timeout=5, - ) + wrapped = sse_keepalive_during_slow_ttft(never_returns, interval_seconds=lambda: 0.01) + response = await asyncio.wait_for(wrapped(self._processor()), timeout=5) disconnected = asyncio.Event() - keepalives_sent = 0 + heartbeats_sent = 0 async def receive(): await disconnected.wait() return {"type": "http.disconnect"} async def send(message): - nonlocal keepalives_sent + nonlocal heartbeats_sent if message["type"] == "http.response.body" and message.get("body"): - keepalives_sent += 1 + heartbeats_sent += 1 disconnected.set() await asyncio.wait_for(response({"type": "http"}, receive, send), timeout=5) - assert keepalives_sent >= 1 - assert upstream_closed.is_set() + assert heartbeats_sent >= 1 + assert upstream_cancelled.is_set() class TestExtractErrorFromSSEChunk: