Avoid Responses stream task-group cleanup

This commit is contained in:
wasimysaid 2026-06-19 17:23:02 +02:00
commit ede6a2bcee
2 changed files with 45 additions and 1 deletions

View file

@ -12,6 +12,7 @@ import uuid
from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import StreamingResponse, JSONResponse, Response
from starlette.requests import ClientDisconnect
from typing import Any, List, Optional, Union
import json
import httpx
@ -744,6 +745,18 @@ def _same_task_timeout(timeout_s: float):
return _CompatSameTaskTimeout(timeout_s)
class _SameTaskStreamingResponse(StreamingResponse):
"""StreamingResponse without Starlette's legacy AnyIO task-group wrapper."""
async def __call__(self, scope, receive, send) -> None:
try:
await self.stream_response(send)
except OSError:
raise ClientDisconnect()
if self.background is not None:
await self.background()
async def _preheader_cancelled(cancel_event = None, request: Optional[Request] = None) -> bool:
if cancel_event is not None and cancel_event.is_set():
return True
@ -7839,7 +7852,7 @@ async def _responses_stream(
api_monitor.finish(monitor_id)
yield _sse("response.completed", completed_response)
return StreamingResponse(
return _SameTaskStreamingResponse(
event_generator(),
media_type = "text/event-stream",
headers = {

View file

@ -59,6 +59,7 @@ from models.inference import (
ResponsesUsage,
)
from routes.inference import (
_SameTaskStreamingResponse,
_build_chat_request,
_chat_tool_calls_to_responses_output,
_extract_responses_reasoning,
@ -1075,6 +1076,36 @@ class TestResponsesStreamAdapter:
),
)
def test_stream_response_avoids_legacy_receive_watcher(self, monkeypatch):
self._install_stream_mock(
monkeypatch,
[{"choices": [{"delta": {"content": "33"}}]}],
)
payload = ResponsesRequest(input = "hi", stream = True)
messages = [ChatMessage(role = "user", content = "hi")]
async def run():
response = await _responses_stream(payload, messages, self._Request())
assert isinstance(response, _SameTaskStreamingResponse)
sent = []
async def receive():
raise AssertionError("Responses streams poll disconnects in the generator")
async def send(message):
sent.append(message)
await response({"type": "http", "asgi": {"spec_version": "2.3"}}, receive, send)
return sent
sent = asyncio.run(run())
assert sent[0]["type"] == "http.response.start"
body = b"".join(message.get("body", b"") for message in sent).decode()
assert "response.output_text.delta" in body
assert '"delta":"33"' in body.replace(" ", "")
def test_split_think_markers_stream_as_reasoning_and_visible_text(self, monkeypatch):
chunks = [
{"choices": [{"delta": {"content": "<thi"}}]},