From ede6a2bcee2a1ba26a19169b7d6bc65e38fdd74d Mon Sep 17 00:00:00 2001 From: wasimysaid Date: Fri, 19 Jun 2026 17:23:02 +0200 Subject: [PATCH] Avoid Responses stream task-group cleanup --- studio/backend/routes/inference.py | 15 ++++++++- .../tests/test_responses_tool_passthrough.py | 31 +++++++++++++++++++ 2 files changed, 45 insertions(+), 1 deletion(-) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 9c0937907d..f58f3aeaba 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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 = { diff --git a/studio/backend/tests/test_responses_tool_passthrough.py b/studio/backend/tests/test_responses_tool_passthrough.py index 582fcfbe87..9a9cb024d6 100644 --- a/studio/backend/tests/test_responses_tool_passthrough.py +++ b/studio/backend/tests/test_responses_tool_passthrough.py @@ -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": "