Avoid Responses stream task-group cleanup
This commit is contained in:
parent
475ff786d8
commit
ede6a2bcee
2 changed files with 45 additions and 1 deletions
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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"}}]},
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue