diff --git a/src/openai/_streaming.py b/src/openai/_streaming.py index bb45a70f6a..45210f259e 100644 --- a/src/openai/_streaming.py +++ b/src/openai/_streaming.py @@ -42,6 +42,7 @@ def __init__( self._client = client self._options = options self._decoder = client._make_sse_decoder() + self._raw_stream = response.iter_bytes() self._iterator = self.__stream__() def __next__(self) -> _T: @@ -53,7 +54,7 @@ def __iter__(self) -> Iterator[_T]: def _iter_events(self) -> Iterator[ServerSentEvent]: try: - yield from self._decoder.iter_bytes(self.response.iter_bytes()) + yield from self._decoder.iter_bytes(self._raw_stream) except timeout_exceptions() as err: raise APITimeoutError(request=self.response.request) from err except request_exceptions() as err: @@ -68,6 +69,16 @@ def __stream__(self) -> Iterator[_T]: try: for sse in iterator: if sse.data.startswith("[DONE]"): + if response.http_version == "HTTP/1.1": + # Finish the HTTP body so the connection can be reused. Resume the + # active byte iterator without decoding discarded SSE data. This + # uses the request's read timeout; HTTP/2 needs no drain for reuse. + try: + for _ in self._raw_stream: + pass + except request_exceptions(): + # Cleanup must not turn a completed stream into a request error. + pass break # we have to special case the Assistants `thread.` events since we won't have an "event" key in the data @@ -156,6 +167,7 @@ def __init__( self._client = client self._options = options self._decoder = client._make_sse_decoder() + self._raw_stream = response.aiter_bytes() self._iterator = self.__stream__() async def __anext__(self) -> _T: @@ -167,7 +179,7 @@ async def __aiter__(self) -> AsyncIterator[_T]: async def _iter_events(self) -> AsyncIterator[ServerSentEvent]: try: - async for sse in self._decoder.aiter_bytes(self.response.aiter_bytes()): + async for sse in self._decoder.aiter_bytes(self._raw_stream): yield sse except timeout_exceptions() as err: raise APITimeoutError(request=self.response.request) from err @@ -183,6 +195,16 @@ async def __stream__(self) -> AsyncIterator[_T]: try: async for sse in iterator: if sse.data.startswith("[DONE]"): + if response.http_version == "HTTP/1.1": + # Finish the HTTP body so the connection can be reused. Resume the + # active byte iterator without decoding discarded SSE data. This + # uses the request's read timeout; HTTP/2 needs no drain for reuse. + try: + async for _ in self._raw_stream: + pass + except request_exceptions(): + # Cleanup must not turn a completed stream into a request error. + pass break # we have to special case the Assistants `thread.` events since we won't have an "event" key in the data diff --git a/tests/test_streaming.py b/tests/test_streaming.py index 4a6ec58b80..27bc68b1de 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -2,9 +2,14 @@ import os import ssl +import socket +import asyncio import importlib +import threading from typing import Any, Iterator, AsyncIterator from contextlib import aclosing, nullcontext +from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler +from typing_extensions import override import anyio import httpx2 @@ -29,6 +34,132 @@ def http_module(request: pytest.FixtureRequest) -> Any: return importlib.import_module(request.param) +@pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) +@pytest.mark.parametrize("ending", ["complete", "truncate", "stall"]) +async def test_done_http11_connection_reuse(sync: bool, ending: str, http_module: Any) -> None: + connections = 0 + requests = 0 + release = threading.Event() + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + @override + def setup(self) -> None: + nonlocal connections + super().setup() + connections += 1 + self.connection.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + + def do_POST(self) -> None: + nonlocal requests + self.rfile.read(int(self.headers["Content-Length"])) + requests += 1 + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Transfer-Encoding", "chunked") + self.end_headers() + payload = ( + b'data: {"id":"synthetic","object":"chat.completion.chunk","created":0,"model":"synthetic",' + b'"choices":[{"index":0,"delta":{"content":"hello"}}]}\n\ndata: [DONE]\n\n' + ) + try: + self.wfile.write(f"{len(payload):x}\r\n".encode() + payload + b"\r\n") + if requests == 1: + if ending == "truncate": + self.close_connection = True + return + if ending == "stall": + release.wait(timeout=5) + self.wfile.write(b"0\r\n\r\n") + except ConnectionError: + self.close_connection = True + + @override + def log_message(self, format: str, *args: object) -> None: # noqa: A002 + pass + + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}) + thread.start() + options: dict[str, Any] = { + "api_key": "synthetic", + "base_url": f"http://127.0.0.1:{server.server_port}/v1", + "max_retries": 0, + "timeout": http_module.Timeout(5, read=0.2), + } + transport_options = {"trust_env": False, "limits": http_module.Limits(max_connections=1)} + try: + if sync: + with OpenAI(**options, http_client=http_module.Client(**transport_options)) as client: + for _ in range(2): + with client.chat.completions.create(model="synthetic", messages=[], stream=True) as stream: + assert [chunk.choices[0].delta.content for chunk in stream] == ["hello"] + else: + async with AsyncOpenAI(**options, http_client=http_module.AsyncClient(**transport_options)) as async_client: + for _ in range(2): + async with await async_client.chat.completions.create( + model="synthetic", messages=[], stream=True + ) as async_stream: + assert [chunk.choices[0].delta.content async for chunk in async_stream] == ["hello"] + finally: + release.set() + server.shutdown() + thread.join(timeout=5) + server.server_close() + + assert not thread.is_alive() + assert requests == 2 + # A failed drain must discard the connection and release the only pool slot. + assert connections == (1 if ending == "complete" else 2) + + +@pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) +@pytest.mark.parametrize("http_version", [b"HTTP/1.1", b"HTTP/2"]) +async def test_done_discards_bytes_only_for_http11( + sync: bool, http_version: bytes, http_module: Any, client: OpenAI, async_client: AsyncOpenAI +) -> None: + read_tail = False + + def body() -> Iterator[bytes]: + nonlocal read_tail + yield b"data: [DONE]\n\n" + read_tail = True + yield b"\xff\n\n" # Trailing data must not be decoded as SSE. + + response = http_module.Response( + 200, content=body() if sync else to_aiter(body()), extensions={"http_version": http_version} + ) + if sync: + assert list(Stream(cast_to=object, response=response, client=client)) == [] + else: + assert [event async for event in AsyncStream(cast_to=object, response=response, client=async_client)] == [] + assert response.is_closed + assert read_tail == (http_version == b"HTTP/1.1") + + +async def test_cancellation_during_done_drain_closes_response(http_module: Any, async_client: AsyncOpenAI) -> None: + draining = asyncio.Event() + + async def body() -> AsyncIterator[bytes]: + yield b"data: [DONE]\n\n" + draining.set() + await asyncio.Event().wait() + + response = http_module.Response(200, content=body()) + stream = AsyncStream(cast_to=object, response=response, client=async_client) + task = asyncio.create_task(stream.__anext__()) + try: + await asyncio.wait_for(draining.wait(), timeout=5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert response.is_closed + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + @pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) @pytest.mark.parametrize("delivered", [False, True], ids=["before-first-event", "after-first-event"]) @pytest.mark.parametrize(