From b8b966efa2597905bac0871ca372b8361ee7844f Mon Sep 17 00:00:00 2001 From: Marcus Wood Date: Thu, 8 Oct 2026 18:08:22 +0000 Subject: [PATCH 1/2] fix(streaming): reuse HTTP/1.1 connections after DONE --- src/openai/_streaming.py | 26 ++++- tests/test_streaming.py | 149 ++++++++++++++++++++++++- tests/test_streaming_http11.py | 192 +++++++++++++++++++++++++++++++++ 3 files changed, 364 insertions(+), 3 deletions(-) create mode 100644 tests/test_streaming_http11.py 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..544236e34e 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -2,6 +2,7 @@ import os import ssl +import asyncio import importlib from typing import Any, Iterator, AsyncIterator from contextlib import aclosing, nullcontext @@ -10,7 +11,7 @@ import httpx2 import pytest -from openai import OpenAI, AsyncOpenAI, APITimeoutError, APIConnectionError +from openai import OpenAI, APIError, AsyncOpenAI, APITimeoutError, APIConnectionError from openai._streaming import Stream, AsyncStream, ServerSentEvent @@ -29,6 +30,152 @@ 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("http_version", [b"HTTP/1.1", b"HTTP/2"]) +@pytest.mark.parametrize("ending", ["eof", "timeout", "protocol-error", "unexpected-error"]) +async def test_done_drains_bytes_without_decoding_trailing_events( + sync: bool, http_version: bytes, ending: str, http_module: Any +) -> None: + reached_eof = False + error = { + "timeout": http_module.ReadTimeout("synthetic timeout"), + "protocol-error": http_module.RemoteProtocolError("synthetic incomplete body"), + "unexpected-error": ValueError("synthetic application error"), + }.get(ending) + + def body() -> Iterator[bytes]: + nonlocal reached_eof + yield b'data: {"foo":true}\n\ndata: [DONE]\n\n' + # Invalid UTF-8 and an unterminated line must be discarded without SSE decoding. + yield b"\xff" * 65536 + if error is not None: + raise error + reached_eof = True + + async def async_body() -> AsyncIterator[bytes]: + for chunk in body(): + yield chunk + + def handler(_request: Any) -> Any: + return http_module.Response( + 200, + content=body() if sync else async_body(), + extensions={"http_version": http_version}, + ) + + context = ( + pytest.raises(ValueError, match="synthetic application error") + if ending == "unexpected-error" and http_version == b"HTTP/1.1" + else nullcontext() + ) + received: list[object] = [] + if sync: + with OpenAI( + api_key="synthetic", + http_client=http_module.Client(transport=http_module.MockTransport(handler), trust_env=False), + ) as client: + stream = client.post("/synthetic", cast_to=object, stream=True, stream_cls=Stream[object]) + with context: + received.extend(stream) + assert stream.response.is_closed + else: + async with AsyncOpenAI( + api_key="synthetic", + http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), + ) as async_client: + async_stream = await async_client.post( + "/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object] + ) + with context: + async for event in async_stream: + received.append(event) + assert async_stream.response.is_closed + + assert received == [{"foo": True}] + assert reached_eof == (ending == "eof" and http_version == b"HTTP/1.1") + + +@pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) +@pytest.mark.parametrize("exit_kind", ["early-close", "api-error", "invalid-json"]) +async def test_incomplete_stream_does_not_drain(sync: bool, exit_kind: str, http_module: Any) -> None: + read_tail = False + + def body() -> Iterator[bytes]: + nonlocal read_tail + if exit_kind == "api-error": + yield b'data: {"error":{"message":"synthetic API error"}}\n\n' + elif exit_kind == "invalid-json": + yield b"data: invalid json\n\n" + else: + yield b'data: {"foo":true}\n\n' + read_tail = True + yield b"data: [DONE]\n\n" + + async def async_body() -> AsyncIterator[bytes]: + for chunk in body(): + yield chunk + + def handler(_request: Any) -> Any: + return http_module.Response(200, content=body() if sync else async_body()) + + context = ( + nullcontext() + if exit_kind == "early-close" + else pytest.raises(APIError if exit_kind == "api-error" else ValueError) + ) + if sync: + with OpenAI( + api_key="synthetic", + http_client=http_module.Client(transport=http_module.MockTransport(handler), trust_env=False), + ) as client: + stream = client.post("/synthetic", cast_to=object, stream=True, stream_cls=Stream[object]) + with context, stream: + next(stream) + assert stream.response.is_closed + else: + async with AsyncOpenAI( + api_key="synthetic", + http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), + ) as async_client: + async_stream = await async_client.post( + "/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object] + ) + with context: + async with async_stream: + await async_stream.__anext__() + assert async_stream.response.is_closed + + assert not read_tail + + +async def test_cancellation_during_done_drain_closes_response(http_module: Any) -> None: + draining = asyncio.Event() + + async def body() -> AsyncIterator[bytes]: + yield b"data: [DONE]\n\n" + draining.set() + await asyncio.Event().wait() + + def handler(_request: Any) -> Any: + return http_module.Response(200, content=body()) + + async with AsyncOpenAI( + api_key="synthetic", + http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), + ) as client: + stream = await client.post("/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object]) + 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 stream.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( diff --git a/tests/test_streaming_http11.py b/tests/test_streaming_http11.py new file mode 100644 index 0000000000..1910db52bb --- /dev/null +++ b/tests/test_streaming_http11.py @@ -0,0 +1,192 @@ +from __future__ import annotations + +import os +import json +import socket +import importlib +import threading +from typing import Any, Iterator +from dataclasses import field, dataclass +from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler +from typing_extensions import override + +import pytest + +from openai import OpenAI, AsyncOpenAI + + +@dataclass +class StreamingServer: + url: str = "" + connections: int = 0 + requests: int = 0 + ending: str = "complete" + release: threading.Event = field(default_factory=threading.Event) + + +@pytest.fixture +def streaming_server() -> Iterator[StreamingServer]: + state = StreamingServer() + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + @override + def setup(self) -> None: + super().setup() + state.connections += 1 + self.connection.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + + def do_POST(self) -> None: + self.rfile.read(int(self.headers["Content-Length"])) + state.requests += 1 + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Transfer-Encoding", "chunked") + self.end_headers() + + if self.path.endswith("/responses"): + event: dict[str, object] = { + "type": "response.completed", + "sequence_number": 0, + "response": { + "id": "resp_synthetic", + "object": "response", + "created_at": 0, + "status": "completed", + "model": "synthetic", + "output": [], + }, + } + payload = f"event: response.completed\ndata: {json.dumps(event)}\n\n".encode() + else: + event = { + "id": "chatcmpl-synthetic", + "object": "chat.completion.chunk", + "created": 0, + "model": "synthetic", + "choices": [{"index": 0, "delta": {"content": "hello"}, "finish_reason": "stop"}], + } + payload = f"data: {json.dumps(event)}\n\ndata: [DONE]\n\n".encode() + + try: + self.wfile.write(f"{len(payload):x}\r\n".encode() + payload + b"\r\n") + if state.ending == "stall": + state.release.wait(timeout=5) + elif state.ending == "truncate": + self.close_connection = True + return + 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) + state.url = f"http://127.0.0.1:{server.server_port}/v1" + thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}) + thread.start() + try: + yield state + finally: + state.release.set() + server.shutdown() + thread.join(timeout=5) + server.server_close() + assert not thread.is_alive() + + +@pytest.fixture( + params=[ + "httpx2", + pytest.param( + "httpx", + marks=pytest.mark.skipif( + os.environ.get("OPENAI_TEST_LEGACY_HTTPX") != "1", reason="requires the legacy HTTPX compatibility lane" + ), + ), + ] +) +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("api", ["chat", "responses"]) +async def test_fully_consumed_stream_reuses_http11_connection( + sync: bool, api: str, http_module: Any, streaming_server: StreamingServer +) -> None: + if sync: + with OpenAI( + api_key="synthetic", + base_url=streaming_server.url, + max_retries=0, + http_client=http_module.Client(trust_env=False, limits=http_module.Limits(max_connections=1)), + ) as client: + for _ in range(5): + if api == "chat": + with client.chat.completions.create(model="synthetic", messages=[], stream=True) as stream: + assert [chunk.choices[0].delta.content for chunk in stream] == ["hello"] + else: + with client.responses.create(model="synthetic", input="hello", stream=True) as responses_stream: + assert [event.type for event in responses_stream] == ["response.completed"] + else: + async with AsyncOpenAI( + api_key="synthetic", + base_url=streaming_server.url, + max_retries=0, + http_client=http_module.AsyncClient(trust_env=False, limits=http_module.Limits(max_connections=1)), + ) as async_client: + for _ in range(5): + if api == "chat": + 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"] + else: + async with await async_client.responses.create( + model="synthetic", input="hello", stream=True + ) as async_responses_stream: + assert [event.type async for event in async_responses_stream] == ["response.completed"] + + assert streaming_server.requests == 5 + assert streaming_server.connections == 1 + + +@pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) +@pytest.mark.parametrize("ending", ["truncate", "stall"]) +async def test_bad_http_ending_preserves_completed_output_and_releases_pool_slot( + sync: bool, ending: str, http_module: Any, streaming_server: StreamingServer +) -> None: + streaming_server.ending = ending + if sync: + with OpenAI( + api_key="synthetic", + base_url=streaming_server.url, + max_retries=0, + timeout=http_module.Timeout(5, read=0.05), + http_client=http_module.Client(trust_env=False, limits=http_module.Limits(max_connections=1)), + ) 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"] + streaming_server.ending = "complete" + else: + async with AsyncOpenAI( + api_key="synthetic", + base_url=streaming_server.url, + max_retries=0, + timeout=http_module.Timeout(5, read=0.05), + http_client=http_module.AsyncClient(trust_env=False, limits=http_module.Limits(max_connections=1)), + ) 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"] + streaming_server.ending = "complete" + + assert streaming_server.requests == 2 + assert streaming_server.connections == 2 From 83833ea92b0da18e51e32f1d39650aeff691430f Mon Sep 17 00:00:00 2001 From: Marcus Wood Date: Thu, 8 Oct 2026 18:13:53 +0000 Subject: [PATCH 2/2] test(streaming): focus coverage on connection reuse and cleanup --- tests/test_streaming.py | 228 +++++++++++++++------------------ tests/test_streaming_http11.py | 192 --------------------------- 2 files changed, 106 insertions(+), 314 deletions(-) delete mode 100644 tests/test_streaming_http11.py diff --git a/tests/test_streaming.py b/tests/test_streaming.py index 544236e34e..27bc68b1de 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -2,16 +2,20 @@ 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 import pytest -from openai import OpenAI, APIError, AsyncOpenAI, APITimeoutError, APIConnectionError +from openai import OpenAI, AsyncOpenAI, APITimeoutError, APIConnectionError from openai._streaming import Stream, AsyncStream, ServerSentEvent @@ -31,124 +35,110 @@ def http_module(request: pytest.FixtureRequest) -> Any: @pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) -@pytest.mark.parametrize("http_version", [b"HTTP/1.1", b"HTTP/2"]) -@pytest.mark.parametrize("ending", ["eof", "timeout", "protocol-error", "unexpected-error"]) -async def test_done_drains_bytes_without_decoding_trailing_events( - sync: bool, http_version: bytes, ending: str, http_module: Any -) -> None: - reached_eof = False - error = { - "timeout": http_module.ReadTimeout("synthetic timeout"), - "protocol-error": http_module.RemoteProtocolError("synthetic incomplete body"), - "unexpected-error": ValueError("synthetic application error"), - }.get(ending) - - def body() -> Iterator[bytes]: - nonlocal reached_eof - yield b'data: {"foo":true}\n\ndata: [DONE]\n\n' - # Invalid UTF-8 and an unterminated line must be discarded without SSE decoding. - yield b"\xff" * 65536 - if error is not None: - raise error - reached_eof = True - - async def async_body() -> AsyncIterator[bytes]: - for chunk in body(): - yield chunk - - def handler(_request: Any) -> Any: - return http_module.Response( - 200, - content=body() if sync else async_body(), - extensions={"http_version": http_version}, - ) - - context = ( - pytest.raises(ValueError, match="synthetic application error") - if ending == "unexpected-error" and http_version == b"HTTP/1.1" - else nullcontext() - ) - received: list[object] = [] - if sync: - with OpenAI( - api_key="synthetic", - http_client=http_module.Client(transport=http_module.MockTransport(handler), trust_env=False), - ) as client: - stream = client.post("/synthetic", cast_to=object, stream=True, stream_cls=Stream[object]) - with context: - received.extend(stream) - assert stream.response.is_closed - else: - async with AsyncOpenAI( - api_key="synthetic", - http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), - ) as async_client: - async_stream = await async_client.post( - "/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object] +@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' ) - with context: - async for event in async_stream: - received.append(event) - assert async_stream.response.is_closed - - assert received == [{"foo": True}] - assert reached_eof == (ending == "eof" and http_version == b"HTTP/1.1") + 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("exit_kind", ["early-close", "api-error", "invalid-json"]) -async def test_incomplete_stream_does_not_drain(sync: bool, exit_kind: str, http_module: Any) -> None: +@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 - if exit_kind == "api-error": - yield b'data: {"error":{"message":"synthetic API error"}}\n\n' - elif exit_kind == "invalid-json": - yield b"data: invalid json\n\n" - else: - yield b'data: {"foo":true}\n\n' - read_tail = True yield b"data: [DONE]\n\n" + read_tail = True + yield b"\xff\n\n" # Trailing data must not be decoded as SSE. - async def async_body() -> AsyncIterator[bytes]: - for chunk in body(): - yield chunk - - def handler(_request: Any) -> Any: - return http_module.Response(200, content=body() if sync else async_body()) - - context = ( - nullcontext() - if exit_kind == "early-close" - else pytest.raises(APIError if exit_kind == "api-error" else ValueError) + response = http_module.Response( + 200, content=body() if sync else to_aiter(body()), extensions={"http_version": http_version} ) if sync: - with OpenAI( - api_key="synthetic", - http_client=http_module.Client(transport=http_module.MockTransport(handler), trust_env=False), - ) as client: - stream = client.post("/synthetic", cast_to=object, stream=True, stream_cls=Stream[object]) - with context, stream: - next(stream) - assert stream.response.is_closed + assert list(Stream(cast_to=object, response=response, client=client)) == [] else: - async with AsyncOpenAI( - api_key="synthetic", - http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), - ) as async_client: - async_stream = await async_client.post( - "/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object] - ) - with context: - async with async_stream: - await async_stream.__anext__() - assert async_stream.response.is_closed - - assert not read_tail + 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) -> None: +async def test_cancellation_during_done_drain_closes_response(http_module: Any, async_client: AsyncOpenAI) -> None: draining = asyncio.Event() async def body() -> AsyncIterator[bytes]: @@ -156,24 +146,18 @@ async def body() -> AsyncIterator[bytes]: draining.set() await asyncio.Event().wait() - def handler(_request: Any) -> Any: - return http_module.Response(200, content=body()) - - async with AsyncOpenAI( - api_key="synthetic", - http_client=http_module.AsyncClient(transport=http_module.MockTransport(handler), trust_env=False), - ) as client: - stream = await client.post("/synthetic", cast_to=object, stream=True, stream_cls=AsyncStream[object]) - 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 stream.response.is_closed - finally: - task.cancel() - await asyncio.gather(task, return_exceptions=True) + 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"]) diff --git a/tests/test_streaming_http11.py b/tests/test_streaming_http11.py deleted file mode 100644 index 1910db52bb..0000000000 --- a/tests/test_streaming_http11.py +++ /dev/null @@ -1,192 +0,0 @@ -from __future__ import annotations - -import os -import json -import socket -import importlib -import threading -from typing import Any, Iterator -from dataclasses import field, dataclass -from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler -from typing_extensions import override - -import pytest - -from openai import OpenAI, AsyncOpenAI - - -@dataclass -class StreamingServer: - url: str = "" - connections: int = 0 - requests: int = 0 - ending: str = "complete" - release: threading.Event = field(default_factory=threading.Event) - - -@pytest.fixture -def streaming_server() -> Iterator[StreamingServer]: - state = StreamingServer() - - class Handler(BaseHTTPRequestHandler): - protocol_version = "HTTP/1.1" - - @override - def setup(self) -> None: - super().setup() - state.connections += 1 - self.connection.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - - def do_POST(self) -> None: - self.rfile.read(int(self.headers["Content-Length"])) - state.requests += 1 - self.send_response(200) - self.send_header("Content-Type", "text/event-stream") - self.send_header("Transfer-Encoding", "chunked") - self.end_headers() - - if self.path.endswith("/responses"): - event: dict[str, object] = { - "type": "response.completed", - "sequence_number": 0, - "response": { - "id": "resp_synthetic", - "object": "response", - "created_at": 0, - "status": "completed", - "model": "synthetic", - "output": [], - }, - } - payload = f"event: response.completed\ndata: {json.dumps(event)}\n\n".encode() - else: - event = { - "id": "chatcmpl-synthetic", - "object": "chat.completion.chunk", - "created": 0, - "model": "synthetic", - "choices": [{"index": 0, "delta": {"content": "hello"}, "finish_reason": "stop"}], - } - payload = f"data: {json.dumps(event)}\n\ndata: [DONE]\n\n".encode() - - try: - self.wfile.write(f"{len(payload):x}\r\n".encode() + payload + b"\r\n") - if state.ending == "stall": - state.release.wait(timeout=5) - elif state.ending == "truncate": - self.close_connection = True - return - 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) - state.url = f"http://127.0.0.1:{server.server_port}/v1" - thread = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}) - thread.start() - try: - yield state - finally: - state.release.set() - server.shutdown() - thread.join(timeout=5) - server.server_close() - assert not thread.is_alive() - - -@pytest.fixture( - params=[ - "httpx2", - pytest.param( - "httpx", - marks=pytest.mark.skipif( - os.environ.get("OPENAI_TEST_LEGACY_HTTPX") != "1", reason="requires the legacy HTTPX compatibility lane" - ), - ), - ] -) -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("api", ["chat", "responses"]) -async def test_fully_consumed_stream_reuses_http11_connection( - sync: bool, api: str, http_module: Any, streaming_server: StreamingServer -) -> None: - if sync: - with OpenAI( - api_key="synthetic", - base_url=streaming_server.url, - max_retries=0, - http_client=http_module.Client(trust_env=False, limits=http_module.Limits(max_connections=1)), - ) as client: - for _ in range(5): - if api == "chat": - with client.chat.completions.create(model="synthetic", messages=[], stream=True) as stream: - assert [chunk.choices[0].delta.content for chunk in stream] == ["hello"] - else: - with client.responses.create(model="synthetic", input="hello", stream=True) as responses_stream: - assert [event.type for event in responses_stream] == ["response.completed"] - else: - async with AsyncOpenAI( - api_key="synthetic", - base_url=streaming_server.url, - max_retries=0, - http_client=http_module.AsyncClient(trust_env=False, limits=http_module.Limits(max_connections=1)), - ) as async_client: - for _ in range(5): - if api == "chat": - 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"] - else: - async with await async_client.responses.create( - model="synthetic", input="hello", stream=True - ) as async_responses_stream: - assert [event.type async for event in async_responses_stream] == ["response.completed"] - - assert streaming_server.requests == 5 - assert streaming_server.connections == 1 - - -@pytest.mark.parametrize("sync", [True, False], ids=["sync", "async"]) -@pytest.mark.parametrize("ending", ["truncate", "stall"]) -async def test_bad_http_ending_preserves_completed_output_and_releases_pool_slot( - sync: bool, ending: str, http_module: Any, streaming_server: StreamingServer -) -> None: - streaming_server.ending = ending - if sync: - with OpenAI( - api_key="synthetic", - base_url=streaming_server.url, - max_retries=0, - timeout=http_module.Timeout(5, read=0.05), - http_client=http_module.Client(trust_env=False, limits=http_module.Limits(max_connections=1)), - ) 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"] - streaming_server.ending = "complete" - else: - async with AsyncOpenAI( - api_key="synthetic", - base_url=streaming_server.url, - max_retries=0, - timeout=http_module.Timeout(5, read=0.05), - http_client=http_module.AsyncClient(trust_env=False, limits=http_module.Limits(max_connections=1)), - ) 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"] - streaming_server.ending = "complete" - - assert streaming_server.requests == 2 - assert streaming_server.connections == 2