Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 24 additions & 2 deletions src/openai/_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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
Expand Down
131 changes: 131 additions & 0 deletions tests/test_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand Down
Loading