diff --git a/.release-please-manifest.json b/.release-please-manifest.json index ca66de7f50..d11c8fc3a0 100644 --- a/.release-please-manifest.json +++ b/.release-please-manifest.json @@ -1,3 +1,3 @@ { - ".": "3.19.2" + ".": "3.20.0" } \ No newline at end of file diff --git a/CHANGELOG.md b/CHANGELOG.md index c0b5b98ae8..478a8f81b9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,35 @@ # Changelog +## [3.20.0](https://github.com/openai/openai-python/compare/v3.19.2...v3.20.0) (2026-09-28) + + +### Features + +* **api:** add Agents credential and session options ([#3967](https://github.com/openai/openai-python/issues/3967)) ([bb68198](https://github.com/openai/openai-python/commit/bb68198a4bb4a161cbb32a28c3888a51540eaf0f)) +* **api:** add Cyber access programs to Responses ([#3956](https://github.com/openai/openai-python/issues/3956)) ([09c5b6f](https://github.com/openai/openai-python/commit/09c5b6f13f716ad4e417fd7ba9209a8a057e0383)) +* **responses:** opt in to incremental WebSocket text and tool snapshots ([#3973](https://github.com/openai/openai-python/issues/3973)) ([d0207b4](https://github.com/openai/openai-python/commit/d0207b48c043741ff483d7a8c945f8003d1ee714)) +* **responses:** preserve detailed WebSocket accumulator snapshots ([#3981](https://github.com/openai/openai-python/issues/3981)) ([a380cf2](https://github.com/openai/openai-python/commit/a380cf256abc2aa51cbf5d49437ae039b9aa9b12)) + + +### Bug Fixes + +* **client:** retry unmapped TLS transport errors ([#3982](https://github.com/openai/openai-python/issues/3982)) ([0d35a26](https://github.com/openai/openai-python/commit/0d35a2640458b64ed2ff4a7c9b1f1a1e5705336e)) +* **live:** avoid hangs at fractional transcript grouping deadlines ([#3970](https://github.com/openai/openai-python/issues/3970)) ([4ef4129](https://github.com/openai/openai-python/commit/4ef4129e85e81a285905755e57b392aa9e77f9bd)) +* **live:** keep query parameters out of WebSocket endpoint paths ([#3972](https://github.com/openai/openai-python/issues/3972)) ([f9c458b](https://github.com/openai/openai-python/commit/f9c458b759e3756979fceebf50bffa78d9cf0b97)) +* **live:** preserve caller queues and prevent uncertain WebSocket replay ([#3980](https://github.com/openai/openai-python/issues/3980)) ([80e9686](https://github.com/openai/openai-python/commit/80e96860b5ccbfcf4a1c9dd9947fc31cd021037a)) +* **realtime:** preserve base URL queries in WebSocket upgrades ([#3971](https://github.com/openai/openai-python/issues/3971)) ([f7bd4a7](https://github.com/openai/openai-python/commit/f7bd4a703cad904e4f5d91ec9d7abac70d8c0bd6)) +* **realtime:** retain configured queues without replaying attempted sends ([#3978](https://github.com/openai/openai-python/issues/3978)) ([a52805c](https://github.com/openai/openai-python/commit/a52805c2537422ad602eba2bf66072880ea3fd22)) + + +### Chores + +* **api:** clarify documented API error responses ([#3965](https://github.com/openai/openai-python/issues/3965)) ([384fee3](https://github.com/openai/openai-python/commit/384fee3252e2336a85450ff38f9bae41f9199616)) +* **api:** document batch error responses ([#3961](https://github.com/openai/openai-python/issues/3961)) ([6e4a79c](https://github.com/openai/openai-python/commit/6e4a79cc8c7e640e7ac4be710db32fe20b1020f2)) +* **api:** document files and uploads error responses ([#3960](https://github.com/openai/openai-python/issues/3960)) ([a9d727f](https://github.com/openai/openai-python/commit/a9d727ff9c4a38fc5dca8d31bd3dd76f473cfc27)) +* **api:** document fine-tuning and model errors ([#3964](https://github.com/openai/openai-python/issues/3964)) ([5d4003c](https://github.com/openai/openai-python/commit/5d4003c12d5df5a5faed35971cac3bcb711fdf7d)) +* **api:** document Responses not-found errors ([#3959](https://github.com/openai/openai-python/issues/3959)) ([63099e7](https://github.com/openai/openai-python/commit/63099e739d25cb18f90cae93648471bf037bc46f)) +* **api:** document stored chat completion errors ([#3963](https://github.com/openai/openai-python/issues/3963)) ([a73fe0c](https://github.com/openai/openai-python/commit/a73fe0c3d404335a342d1251bd32709b8e3d76f2)) + ## [3.19.2](https://github.com/openai/openai-python/compare/v3.19.1...v3.19.2) (2026-09-23) diff --git a/pyproject.toml b/pyproject.toml index fb70a50977..8d5cca8ac9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "openai" -version = "3.19.2" +version = "3.20.0" description = "The official Python library for the openai API" dynamic = ["readme"] license = "Apache-2.0" diff --git a/src/openai/_httpx2.py b/src/openai/_httpx2.py index 8944928d84..48222a10ae 100644 --- a/src/openai/_httpx2.py +++ b/src/openai/_httpx2.py @@ -1,8 +1,10 @@ from __future__ import annotations +import ssl import sys from typing import Any, Protocol, cast +import anyio import httpx2 from ._constants import DEFAULT_TIMEOUT, DEFAULT_CONNECTION_LIMITS @@ -100,9 +102,11 @@ def timeout_exceptions() -> tuple[type[httpx2.TimeoutException], ...]: return (httpx2.TimeoutException,) if module is None else (httpx2.TimeoutException, module.TimeoutException) -def request_exceptions() -> tuple[type[httpx2.RequestError], ...]: +def request_exceptions() -> tuple[type[Exception], ...]: module = _loaded_legacy_httpx() - return (httpx2.RequestError,) if module is None else (httpx2.RequestError, module.RequestError) + # Shim unmapped AnyIO TLS failures pending https://github.com/pydantic/httpx2/issues/854. + errors = (httpx2.RequestError, ssl.SSLError, anyio.EndOfStream) + return errors if module is None else (*errors, module.RequestError) def status_exceptions() -> tuple[type[httpx2.HTTPStatusError], ...]: diff --git a/src/openai/_send_queue.py b/src/openai/_send_queue.py index 56d7824422..b3eb812b0e 100644 --- a/src/openai/_send_queue.py +++ b/src/openai/_send_queue.py @@ -37,11 +37,12 @@ def enqueue(self, data: str) -> None: self._queue.append((data, byte_length)) self._bytes += byte_length - def flush_sync(self, send: typing.Callable[[str], object]) -> None: + def flush_sync(self, send: typing.Callable[[str], object], *, requeue_failed: bool = True) -> None: """Send every queued message via *send*. If *send* raises, the failing message and all subsequent messages - are re-queued and the error is re-raised. + are re-queued and the error is re-raised. When `requeue_failed` is + false, release the attempted message even on failure or interruption. """ while isinstance(pending := self._begin_flush(), threading.Event): pending.wait() @@ -49,14 +50,21 @@ def flush_sync(self, send: typing.Callable[[str], object]) -> None: try: while pending: data, byte_length = pending[0] - send(data) - with self._lock: - pending.popleft() - self._bytes -= byte_length + sent = False + try: + send(data) + sent = True + finally: + if sent or not requeue_failed: + with self._lock: + pending.popleft() + self._bytes -= byte_length finally: self._end_flush(pending) - async def flush_async(self, send: typing.Callable[[str], typing.Awaitable[object]]) -> None: + async def flush_async( + self, send: typing.Callable[[str], typing.Awaitable[object]], *, requeue_failed: bool = True + ) -> None: """Async variant of :meth:`flush_sync`.""" while isinstance(pending := self._begin_flush(), threading.Event): # Waiting in a worker keeps the event loop responsive. Cancellation @@ -66,10 +74,15 @@ async def flush_async(self, send: typing.Callable[[str], typing.Awaitable[object try: while pending: data, byte_length = pending[0] - await send(data) - with self._lock: - pending.popleft() - self._bytes -= byte_length + sent = False + try: + await send(data) + sent = True + finally: + if sent or not requeue_failed: + with self._lock: + pending.popleft() + self._bytes -= byte_length finally: self._end_flush(pending) diff --git a/src/openai/_version.py b/src/openai/_version.py index e7c432efa8..cf3e8564f3 100644 --- a/src/openai/_version.py +++ b/src/openai/_version.py @@ -1,2 +1,2 @@ __title__ = "openai" -__version__ = "3.19.2" # x-release-please-version +__version__ = "3.20.0" # x-release-please-version diff --git a/src/openai/lib/live/README.md b/src/openai/lib/live/README.md new file mode 100644 index 0000000000..55394b91ea --- /dev/null +++ b/src/openai/lib/live/README.md @@ -0,0 +1,33 @@ +# Live WebSocket lifecycle + +`client.live.connect()` and `client.live.forks.connect(session_id=...)` leave +startup to the caller: send `connection.session.start(...)` and wait for +`session.started`. A sideband connection created by +`client.live.sideband.connect(session_id=...)` attaches to an existing session. +It does not send a start or require a new `session.started`; an attachment's +short replay is not a complete session snapshot. + +Direct `recv()` and `recv_bytes()` calls report transport errors to their +caller. Iterator reconnection is available only when an `on_reconnecting` +callback was explicitly supplied. The callback controls retry and can update +credentials or query parameters. A new socket is not proof that the previous +Live session, recording, or application state was restored. Applications own +their recovery decision and any necessary startup or state reconstruction. +Don't use multiple physical readers: a dispatcher owns the read loop while it +runs. Detaching or closing one transcript grouper doesn't cancel another +observer or the connection. + +All modes preserve the manager's caller-configured `max_queue_size`, including +if the queue was empty at connection time. Only messages that haven't been +attempted on a socket remain eligible for flushing after a retry. A direct send +exception is raised to its caller and cannot prove the server didn't receive +that message. A failure while flushing an already queued message still logs a +warning, as before; it cannot be raised at the original queueing call. Neither +failed attempt is automatically retried. Unattempted pre-open messages and +messages explicitly queued during recovery remain in order within their +existing budget. No automatic session reopening or restoration is implied. + +Treat each wire error as an event to handle, not as a successful session result. +In particular, preserve `session_storage_failed` if `session.closed` follows; +the close doesn't mean that recording storage succeeded. Unknown events and +fields are preserved for callers that need them. diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index 41c3dc7f71..ef046787bf 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -108,8 +108,47 @@ a changed nonempty item ID starts fresh at its index. A supplied response output list overrides the projected items, including an explicit empty list; omitted or null output retains only the helper's earlier projections. +`detailed_snapshot()` returns a separate mutable view when you need the fields +beyond selected text and tool inputs. It contains `stream_id`, `response_id`, +`terminal_type`, `response` and `output`. `response` is the last observed +lifecycle response metadata (excluding `output`), or `None` if none arrived. +`output` is a list of `{"output_index": index, "item": metadata, "content": rows}`. +Each content row is `{"content_index": index, "part": observed_fields}`. The +part's `annotations`, when present as a list, use +`{"annotation_index": index, "annotation": observed_fields}` rows. Indices may +be sparse; list position is not the API index. +For a message that has no projected content, the row's `content` is omitted, +null or empty according to what was actually received. + +For example, after collecting an item or terminal as above: + +```python +details = accumulator.detailed_snapshot() +for output in details["output"]: + for content in output.get("content") or []: + part = content["part"] + # Fields exist only if received: a WS delta may have no part/item type. + text = part.get("text") + citations = part.get("annotations") + token_scores = part.get("logprobs") +``` + +Logprobs accumulate with text deltas and a supplied `output_text.done.logprobs` +replaces them, including empty or null values. Content/item/lifecycle replacements +also replace their corresponding annotations and other metadata; text-only done +events do not erase citations. Refusal, tool/MCP and unknown item/part fields are +retained as observed, with unset and null distinct. Unknown standalone events still +pass through unchanged and are not accumulated. Annotation events enrich matching +known items on the accumulator's lane; annotations received before any matching +item/text remain available in the original event and do not start or replace a turn. +A partial field or unknown type +is provisional, not a fabricated validated response or a successful tool result. +You may mutate this returned view without changing the accumulator, events, or +earlier snapshots. The original `snapshot()` remains immutable and hashable. + `snapshot()` materializes the entire current projection and joins retained -fragments. It is proportional to the accumulated output, so requesting it after +fragments. `detailed_snapshot()` also materializes and copies its whole projection. +Both are proportional to the accumulated output, so requesting either after every delta or completed item repeatedly rebuilds growing prefixes. Use the original event for progress, including the final item itself on `response.output_item.done`. Read the full snapshot at a terminal or on explicit diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index c9cecd0963..cba9f5f82c 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -1,10 +1,12 @@ from __future__ import annotations -from typing import cast +from copy import deepcopy +from typing import Any, cast from dataclasses import field, dataclass from ._session import ResponsesWebSocketError, _field -from ..._compat import model_copy +from ..._compat import PYDANTIC_V1, model_copy +from ..._models import BaseModel from ...types.responses import Response from ...types.responses.responses_server_event import ResponsesServerEvent @@ -47,6 +49,9 @@ class _Output: arguments: list[str] = field(default_factory=list[str]) input: list[str] = field(default_factory=list[str]) text: dict[str, list[str]] = field(default_factory=dict[str, list[str]]) + data: dict[str, object] = field(default_factory=dict[str, object]) + content: dict[str, dict[str, object]] = field(default_factory=dict[str, dict[str, object]]) + annotations: dict[str, dict[str, object]] = field(default_factory=dict[str, dict[str, object]]) class ResponsesWebSocketAccumulator: @@ -67,6 +72,7 @@ def __init__(self) -> None: # String keys get Python's randomized hash. Hex preserves arbitrary-size # non-negative indices without decimal string conversion limits. self._output: dict[str, _Output] = {} + self._response: dict[str, object] | None = None self._final: Response | None = None self._error: Exception | None = None @@ -75,6 +81,7 @@ def reset(self) -> None: self._stream_id = self._response_id = self._terminal_type = None self._bound = False self._output.clear() + self._response = None self._final = self._error = None def snapshot(self) -> ResponsesWebSocketSnapshot: @@ -105,6 +112,61 @@ def snapshot(self) -> ResponsesWebSocketSnapshot: ), ) + def detailed_snapshot(self) -> dict[str, Any]: + """Return an independent mutable projection of observed response, item and part data. + + Output and content are lists of indexed rows, even when wire indices are + sparse. Part annotations use the same indexed-row form. Unset and null + fields stay distinct; missing scaffolding never invents a Response or + an item/part type. Cost is proportional to the full accumulated data. + """ + output: list[dict[str, object]] = [] + for index, item in sorted(self._output.items(), key=lambda pair: int(pair[0], 16)): + data = deepcopy(item.data) + if item.item_id: + data["id"] = item.item_id + if item.arguments: + data["arguments"] = "".join(item.arguments) + if item.input: + data["input"] = "".join(item.input) + content: list[dict[str, object]] = [] + for pos in sorted( + item.content.keys() | item.text.keys() | item.annotations.keys(), key=lambda k: int(k, 16) + ): + part = deepcopy(item.content.get(pos, {})) + if pos in item.text: + part["text"] = "".join(item.text[pos]) + original = part.get("annotations") + if isinstance(original, list) or pos in item.annotations: + annotations = ( + {hex(i): value for i, value in enumerate(cast("list[object]", original))} + if isinstance(original, list) + else {} + ) + annotations.update(deepcopy(item.annotations.get(pos, {}))) + part["annotations"] = [ + {"annotation_index": int(i, 16), "annotation": annotation} + for i, annotation in sorted(annotations.items(), key=lambda pair: int(pair[0], 16)) + ] + content.append({"content_index": int(pos, 16), "part": part}) + row: dict[str, object] = {"output_index": int(index, 16), "item": data} + if item.type == "message": + # Message content is extracted into indexed rows. Retain the + # presence marker until a part/delta supplies projected content. + if "content" in data or content: + original_content = data.pop("content", None) + row["content"] = content if content or isinstance(original_content, list) else None + else: + row["content"] = content + output.append(row) + return { + "stream_id": self._stream_id, + "response_id": self._response_id, + "terminal_type": self._terminal_type, + "response": deepcopy(self._response), + "output": output, + } + def get_final_response(self) -> Response: """Return an independent copy of the received terminal response, never a partial success.""" error = self._error @@ -134,6 +196,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.content_part.done", "response.output_text.delta", "response.output_text.done", + "response.output_text.annotation.added", "response.function_call_arguments.delta", "response.function_call_arguments.done", "response.mcp_call_arguments.delta", @@ -144,6 +207,31 @@ def add_event(self, event: ResponsesServerEvent) -> None: }: return stream_id = _field(event, "stream_id") + # Annotations were historically ignored. They may enrich the matching + # known item, but never bind a lane/turn, replace or retire an item, or + # change the errors observed by existing accumulator callers. + if kind == "response.output_text.annotation.added": + output_pos = _field(event, "output_index") + content_pos = _field(event, "content_index") + annotation_pos = _field(event, "annotation_index") + if ( + self._bound + and self._stream_id == stream_id + and self._terminal_type is None + and self._error is None + and all( + isinstance(value, int) and not isinstance(value, bool) and value >= 0 + for value in (output_pos, content_pos, annotation_pos) + ) + ): + existing = self._output.get(hex(output_pos)) + if existing is not None and existing.item_id == _field(event, "item_id"): + value_data = _data(event, include={"annotation"}) + if "annotation" in value_data: + existing.annotations.setdefault(hex(content_pos), {})[hex(annotation_pos)] = value_data[ + "annotation" + ] + return if stream_id is not None and not isinstance(stream_id, str): raise ValueError("WebSocket stream_id must be a string or null") if self._bound and self._stream_id != stream_id: @@ -186,6 +274,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: self._error = error raise self._output = replacement + self._response = _data(response, exclude={"output"}) self._response_id = response_id or self._response_id self._bound, self._stream_id = True, stream_id if terminal: @@ -248,16 +337,28 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.output_text.delta", "response.output_text.done", }: + position = hex(pos) if kind == "response.output_text.delta": - item.text.setdefault(hex(pos), []).append(value) + item.text.setdefault(position, []).append(value) + data = item.content.setdefault(position, {}) + prob_data = _data(event, include={"logprobs"}) + previous = data.get("logprobs") + incoming = prob_data.get("logprobs") + if isinstance(previous, list) and isinstance(incoming, list): + cast("list[object]", previous).extend(cast("list[object]", incoming)) + else: + data.update(prob_data) elif kind == "response.output_text.done": - item.text[hex(pos)] = [value] + item.text[position] = [value] + item.content.setdefault(position, {}).update(_data(event, include={"logprobs"})) else: part = _field(event, "part") + item.content[position] = _data(part) + item.annotations.pop(position, None) if _field(part, "type") == "output_text": - item.text[hex(pos)] = [value] + item.text[position] = [value] else: - item.text.pop(hex(pos), None) + item.text.pop(position, None) elif kind in {"response.function_call_arguments.delta", "response.mcp_call_arguments.delta"}: item.arguments.append(value) elif kind in {"response.function_call_arguments.done", "response.mcp_call_arguments.done"}: @@ -276,21 +377,32 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None: type=_text_field(source, "type"), name=_optional_text_field(source, "name"), call_id=_optional_text_field(source, "call_id"), + data=_data(source, exclude={"content"} if _field(source, "type") == "message" else None), ) if item.type == "message": content = _field(source, "content") if content is not None: if not isinstance(content, list): raise ValueError("WebSocket message content must be a list or null") + # Marker only; keep actual parts indexed once, never copied twice. + item.data["content"] = [] for pos, part in enumerate(cast("list[object]", content)): + if part is not None: + item.content[hex(pos)] = _data(part) if part is not None and _text_field(part, "type") == "output_text": text = _optional_text_field(part, "text") if text is not None: item.text[hex(pos)] = [text] + else: + item.data.update(_data(source, include={"content"})) elif item.type in {"function_call", "mcp_call", "mcp_approval_request"}: - item.arguments = [_optional_text_field(source, "arguments") or ""] + arguments = _optional_text_field(source, "arguments") + if arguments is not None: + item.arguments = [arguments] elif item.type == "custom_tool_call": - item.input = [_optional_text_field(source, "input") or ""] + input_text = _optional_text_field(source, "input") + if input_text is not None: + item.input = [input_text] key = hex(index) previous = output.get(key) if previous is not None: @@ -302,6 +414,24 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None: output[key] = item +def _data(source: object, *, include: set[str] | None = None, exclude: set[str] | None = None) -> dict[str, object]: + if isinstance(source, BaseModel): + return deepcopy( + source.model_dump( + mode="python", by_alias=True, exclude_unset=True, include=include, exclude=exclude, warnings=PYDANTIC_V1 + ) + ) + if isinstance(source, dict): + return deepcopy( + { + key: value + for key, value in cast("dict[str, object]", source).items() + if (include is None or key in include) and (exclude is None or key not in exclude) + } + ) + return {} + + def _text_field(value: object, name: str) -> str: text = _field(value, name) if not isinstance(text, str): diff --git a/src/openai/resources/live/forks.py b/src/openai/resources/live/forks.py index e27846290e..8deb5fbab5 100644 --- a/src/openai/resources/live/forks.py +++ b/src/openai/resources/live/forks.py @@ -139,7 +139,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=False) self.session = AsyncForksSessionResource(self) @@ -215,11 +215,7 @@ async def send(self, event: ForkClientEvent | ForkClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - await self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + await self._connection.send(data) async def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -324,7 +320,7 @@ async def _send(data: str) -> None: await self._connection.send(data) try: - await self._send_queue.flush_async(_send) + await self._send_queue.flush_async(_send, requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") @@ -638,7 +634,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=True) self.session = ForksSessionResource(self) @@ -714,11 +710,7 @@ def send(self, event: ForkClientEvent | ForkClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + self._connection.send(data) def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -817,7 +809,7 @@ def _reconnect(self, exc: Exception) -> bool: def _flush_send_queue(self) -> None: """Send all queued messages over the current connection.""" try: - self._send_queue.flush_sync(lambda data: self._connection.send(data)) + self._send_queue.flush_sync(lambda data: self._connection.send(data), requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") diff --git a/src/openai/resources/live/live.py b/src/openai/resources/live/live.py index 1af46d4e3f..19a78791e1 100644 --- a/src/openai/resources/live/live.py +++ b/src/openai/resources/live/live.py @@ -360,7 +360,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=False) self.session = AsyncLiveSessionResource(self) @@ -436,11 +436,7 @@ async def send(self, event: ClientEvent | ClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - await self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + await self._connection.send(data) async def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -545,7 +541,7 @@ async def _send(data: str) -> None: await self._connection.send(data) try: - await self._send_queue.flush_async(_send) + await self._send_queue.flush_async(_send, requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") @@ -856,7 +852,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=True) self.session = LiveSessionResource(self) @@ -932,11 +928,7 @@ def send(self, event: ClientEvent | ClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + self._connection.send(data) def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -1035,7 +1027,7 @@ def _reconnect(self, exc: Exception) -> bool: def _flush_send_queue(self) -> None: """Send all queued messages over the current connection.""" try: - self._send_queue.flush_sync(lambda data: self._connection.send(data)) + self._send_queue.flush_sync(lambda data: self._connection.send(data), requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") diff --git a/src/openai/resources/live/sideband.py b/src/openai/resources/live/sideband.py index 043a7233e0..83256c2c1e 100644 --- a/src/openai/resources/live/sideband.py +++ b/src/openai/resources/live/sideband.py @@ -142,7 +142,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=False) self.session = AsyncSidebandSessionResource(self) @@ -218,11 +218,7 @@ async def send(self, event: ConnectClientEvent | ConnectClientEventParam) -> Non if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - await self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + await self._connection.send(data) async def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -329,7 +325,7 @@ async def _send(data: str) -> None: await self._connection.send(data) try: - await self._send_queue.flush_async(_send) + await self._send_queue.flush_async(_send, requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") @@ -646,7 +642,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=True) self.session = SidebandSessionResource(self) @@ -722,11 +718,7 @@ def send(self, event: ConnectClientEvent | ConnectClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + self._connection.send(data) def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -827,7 +819,7 @@ def _reconnect(self, exc: Exception) -> bool: def _flush_send_queue(self) -> None: """Send all queued messages over the current connection.""" try: - self._send_queue.flush_sync(lambda data: self._connection.send(data)) + self._send_queue.flush_sync(lambda data: self._connection.send(data), requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") diff --git a/src/openai/resources/realtime/realtime.py b/src/openai/resources/realtime/realtime.py index aef7757cb0..f541bdd886 100644 --- a/src/openai/resources/realtime/realtime.py +++ b/src/openai/resources/realtime/realtime.py @@ -294,7 +294,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=False) self.session = AsyncRealtimeSessionResource(self) @@ -373,11 +373,7 @@ async def send(self, event: RealtimeClientEvent | RealtimeClientEventParam) -> N if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - await self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + await self._connection.send(data) async def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -484,7 +480,7 @@ async def _send(data: str) -> None: await self._connection.send(data) try: - await self._send_queue.flush_async(_send) + await self._send_queue.flush_async(_send, requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") @@ -818,7 +814,7 @@ def __init__( self._extra_headers = extra_headers self._intentionally_closed = False self._is_reconnecting = False - self._send_queue = send_queue or SendQueue() + self._send_queue = send_queue if send_queue is not None else SendQueue() self._event_handler_registry = EventHandlerRegistry(use_lock=True) self.session = RealtimeSessionResource(self) @@ -897,11 +893,7 @@ def send(self, event: RealtimeClientEvent | RealtimeClientEventParam) -> None: if self._is_reconnecting: self._send_queue.enqueue(data) return - try: - self._connection.send(data) - except Exception: - self._send_queue.enqueue(data) - raise + self._connection.send(data) def send_raw(self, data: bytes | str) -> None: if self._is_reconnecting: @@ -1002,7 +994,7 @@ def _reconnect(self, exc: Exception) -> bool: def _flush_send_queue(self) -> None: """Send all queued messages over the current connection.""" try: - self._send_queue.flush_sync(lambda data: self._connection.send(data)) + self._send_queue.flush_sync(lambda data: self._connection.send(data), requeue_failed=False) except Exception: log.warning("Failed to flush send queue after reconnect") diff --git a/tests/lib/live/test_websocket_roles.py b/tests/lib/live/test_websocket_roles.py index e3b0ae5cdc..3b62f73e04 100644 --- a/tests/lib/live/test_websocket_roles.py +++ b/tests/lib/live/test_websocket_roles.py @@ -6,13 +6,25 @@ import httpx2 import pytest +from websockets.sync.client import ClientConnection from websockets.sync.server import ServerConnection +from websockets.asyncio.client import ClientConnection as AsyncClientConnection from openai import OpenAI, AsyncOpenAI -from openai.types.live import SessionStartedEvent, SessionUpdatedEvent, OutputTranscriptDeltaEvent +from openai.types.live import ( + ErrorEvent, + ServerEvent, + SessionClosedEvent, + SessionStartedEvent, + SessionUpdatedEvent, + OutputTranscriptDeltaEvent, +) +from openai._exceptions import WebSocketQueueFullError from openai.resources.live.live import LiveConnection, AsyncLiveConnection from openai.resources.live.forks import ForksConnection, AsyncForksConnection +from openai.types.websocket_reconnection import ReconnectingEvent +from .helpers import FakeClock, Recording, ClockGrouper, AsyncClockGrouper from ..responses.test_websocket_session import script_server @@ -22,6 +34,169 @@ def bypass_loopback_proxy(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("no_proxy", "127.0.0.1") +@pytest.mark.parametrize("role", ["primary", "fork", "sideband"]) +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_empty_live_manager_queue_keeps_budget_and_caller_messages(role: str, mode: str) -> None: + opened = 0 + + def script(socket: ServerConnection) -> None: + nonlocal opened + opened += 1 + if opened == 1: + socket.close(code=1011, reason="synthetic restart") + return + # Read first event without blocking on a lost manager queue. + socket.send('{"type": "session.output_transcript.delta", "event_id": "ready", "delta": "go"}') + assert json.loads(socket.recv(timeout=5)) == {"type": "session.update", "event_id": "queued", "session": {}} + + with script_server(script, expected_connections=2) as url: + if mode == "sync": + with OpenAI(api_key="fake-live-key", base_url=url, http_client=httpx2.Client(trust_env=False)) as client: + + def on_retry(_event: ReconnectingEvent) -> None: + manager.send({"type": "session.update", "event_id": "queued", "session": {}}) + with pytest.raises(WebSocketQueueFullError): + manager.send({"type": "session.update", "event_id": "large" * 150, "session": {}}) + + if role == "primary": + manager = client.live.connect(max_queue_size=100, on_reconnecting=on_retry, initial_delay=0) + elif role == "fork": + manager = client.live.forks.connect( + session_id="stored", max_queue_size=100, on_reconnecting=on_retry, initial_delay=0 + ) + else: + manager = client.live.sideband.connect( + session_id="stored", max_queue_size=100, on_reconnecting=on_retry, initial_delay=0 + ) + with manager as connection: + assert next(iter(connection)).type == "session.output_transcript.delta" + else: + async with AsyncOpenAI( + api_key="fake-live-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client: + + def on_async_retry(_event: ReconnectingEvent) -> None: + async_manager.send({"type": "session.update", "event_id": "queued", "session": {}}) + with pytest.raises(WebSocketQueueFullError): + async_manager.send({"type": "session.update", "event_id": "large" * 150, "session": {}}) + + if role == "primary": + async_manager = async_client.live.connect( + max_queue_size=100, on_reconnecting=on_async_retry, initial_delay=0 + ) + elif role == "fork": + async_manager = async_client.live.forks.connect( + session_id="stored", max_queue_size=100, on_reconnecting=on_async_retry, initial_delay=0 + ) + else: + async_manager = async_client.live.sideband.connect( + session_id="stored", max_queue_size=100, on_reconnecting=on_async_retry, initial_delay=0 + ) + async with async_manager as async_connection: + assert ( + await asyncio.wait_for(anext(aiter(async_connection)), 5) + ).type == "session.output_transcript.delta" + + +@pytest.mark.parametrize("role", ["primary", "fork", "sideband"]) +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("preopen", [True, False], ids=["preopen-flush", "direct"]) +async def test_live_recovery_never_replays_an_attempted_command( + role: str, mode: str, preopen: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + opened = 0 + recorded: list[tuple[int, str]] = [] + + def script(socket: ServerConnection) -> None: + nonlocal opened + opened += 1 + current = opened + if current == 1: + for _ in range(2 if preopen else 1): + recorded.append((current, json.loads(socket.recv(timeout=5))["event_id"])) + socket.close(code=1011, reason="synthetic restart") + return + socket.send('{"type": "session.output_transcript.delta", "event_id": "ready", "delta": "go"}') + for _ in range(2 if preopen else 1): + recorded.append((current, json.loads(socket.recv(timeout=5))["event_id"])) + + # Complete the actual wire send and then raise, so delivery cannot be inferred + # from the SDK-visible error. The real server checks what each socket received. + original = ClientConnection.send + original_async = AsyncClientConnection.send + + def send_then_interrupt(self: ClientConnection, message: object, *args: object, **kwargs: object) -> None: + original(self, message, *args, **kwargs) # type: ignore[arg-type] + if json.loads(message)["event_id"] == "attempted": # type: ignore[arg-type] + raise OSError("synthetic interruption after send") + + async def async_send_then_interrupt( + self: AsyncClientConnection, message: object, *args: object, **kwargs: object + ) -> None: + await original_async(self, message, *args, **kwargs) # type: ignore[arg-type] + if json.loads(message)["event_id"] == "attempted": # type: ignore[arg-type] + raise OSError("synthetic interruption after send") + + if mode == "sync": + monkeypatch.setattr(ClientConnection, "send", send_then_interrupt) + else: + monkeypatch.setattr(AsyncClientConnection, "send", async_send_then_interrupt) + with script_server(script, expected_connections=2) as url: + if mode == "sync": + with OpenAI(api_key="fake-live-key", base_url=url, http_client=httpx2.Client(trust_env=False)) as client: + + def on_retry(_event: ReconnectingEvent) -> None: + manager.send({"type": "session.update", "event_id": "during-recovery", "session": {}}) + + if role == "primary": + manager = client.live.connect(on_reconnecting=on_retry, initial_delay=0) + elif role == "fork": + manager = client.live.forks.connect(session_id="stored", on_reconnecting=on_retry, initial_delay=0) + else: + manager = client.live.sideband.connect( + session_id="stored", on_reconnecting=on_retry, initial_delay=0 + ) + if preopen: + for identity in ("first", "attempted", "unattempted"): + manager.send({"type": "session.update", "event_id": identity, "session": {}}) + with manager as connection: + if not preopen: + with pytest.raises(OSError, match="synthetic interruption after send"): + connection.send({"type": "session.update", "event_id": "attempted", "session": {}}) + assert next(iter(connection)).type == "session.output_transcript.delta" + else: + async with AsyncOpenAI( + api_key="fake-live-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client: + + def on_async_retry(_event: ReconnectingEvent) -> None: + async_manager.send({"type": "session.update", "event_id": "during-recovery", "session": {}}) + + if role == "primary": + async_manager = async_client.live.connect(on_reconnecting=on_async_retry, initial_delay=0) + elif role == "fork": + async_manager = async_client.live.forks.connect( + session_id="stored", on_reconnecting=on_async_retry, initial_delay=0 + ) + else: + async_manager = async_client.live.sideband.connect( + session_id="stored", on_reconnecting=on_async_retry, initial_delay=0 + ) + if preopen: + for identity in ("first", "attempted", "unattempted"): + async_manager.send({"type": "session.update", "event_id": identity, "session": {}}) + async with async_manager as async_connection: + if not preopen: + with pytest.raises(OSError, match="synthetic interruption after send"): + await async_connection.send( + {"type": "session.update", "event_id": "attempted", "session": {}} + ) + event = await asyncio.wait_for(anext(aiter(async_connection)), 5) + assert event.type == "session.output_transcript.delta" + expected = [(1, "first"), (1, "attempted"), (2, "unattempted"), (2, "during-recovery")] + assert recorded == (expected if preopen else [(1, "attempted"), (2, "during-recovery")]) + + @pytest.mark.parametrize("role", ["primary", "fork", "sideband"]) @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("base_query", ["", "?tenant=sample"], ids=["custom-path", "custom-path-and-query"]) @@ -157,3 +332,238 @@ def script(socket: ServerConnection) -> None: assert isinstance(async_updated, SessionUpdatedEvent) assert async_updated.client_event_id == "caller-update" assert async_updated.session.id == "live_fixture" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("role", ["primary", "fork", "sideband"]) +async def test_disposed_grouper_leaves_dispatcher_other_observers_and_socket_usable(mode: str, role: str) -> None: + session = {"id": "live_fixture", "model": "gpt-live-1", "status": "active", "expires_at": 123} + first_clock, second_clock = FakeClock(), FakeClock() + transcripts = [ + { + "type": "session.output_transcript.delta", + "event_id": f"part-{index}", + "delta": value, + "start_ms": index * 200, + "end_ms": (index + 1) * 200, + } + for index, value in enumerate(["One", " two", " three", " four"]) + ] + observed: list[ServerEvent] = [] + typed: list[OutputTranscriptDeltaEvent] = [] + + def script(socket: ServerConnection) -> None: + if role != "sideband": + assert json.loads(socket.recv(timeout=5)) == { + "type": "session.start", + "event_id": "caller-start", + "session": {"model": "gpt-live-1"} if role == "primary" else {}, + } + socket.send(json.dumps({"type": "session.started", "event_id": "started", "session": session})) + for event in transcripts[:2]: + socket.send(json.dumps(event)) + # Issued only after the caller detaches and closes its first grouper. + assert json.loads(socket.recv(timeout=5)) == { + "type": "session.update", + "event_id": "after-dispose", + "session": {}, + } + socket.send( + json.dumps( + { + "type": "session.updated", + "event_id": "updated", + "client_event_id": "after-dispose", + "session": session, + "future_metadata": {"explicit_null": None, "nested": [1, "retained"]}, + } + ) + ) + for event in transcripts[2:]: + socket.send(json.dumps(event)) + assert json.loads(socket.recv(timeout=5)) == {"type": "session.close", "event_id": "caller-finish"} + socket.send(json.dumps({"type": "session.closed", "event_id": "closed", "reason": "client_close"})) + # The shared fixture rejects extra writes/reconnects and waits for caller close. + + with script_server(script) as url: + if mode == "sync": + first = ClockGrouper(first_clock) + second = ClockGrouper(second_clock) + first_record, second_record = Recording(first, first_clock), Recording(second, second_clock) + with OpenAI(api_key="ek_fake_live", base_url=url, http_client=httpx2.Client(trust_env=False)) as client: + if role == "primary": + manager = client.live.connect() + elif role == "fork": + manager = client.live.forks.connect(session_id="stored") + else: + manager = client.live.sideband.connect(session_id="stored") + with manager as connection: + connection.on("session.output_transcript.delta", first.push) + connection.on("session.output_transcript.delta", second.push) + connection.on("session.output_transcript.delta", typed.append) + connection.on("event", observed.append) + + def manage(event: OutputTranscriptDeltaEvent) -> None: + if event.event_id == "part-1": + assert first_clock.pending + connection.off("session.output_transcript.delta", first.push) + first.close() + first.close() + assert not first_clock.pending + connection.session.update(session={}, event_id="after-dispose") + elif event.event_id == "part-3": + connection.session.close(event_id="caller-finish") + + def close_connection(_event: SessionClosedEvent) -> None: + connection.close() + + connection.on("session.output_transcript.delta", manage) + connection.on("session.closed", close_connection) + if isinstance(connection, LiveConnection): + connection.session.start(session={"model": "gpt-live-1"}, event_id="caller-start") + elif isinstance(connection, ForksConnection): + connection.session.start(session={}, event_id="caller-start") + connection.dispatch_events() + second.close() + else: + async_first = AsyncClockGrouper(first_clock) + async_second = AsyncClockGrouper(second_clock) + first_record, second_record = ( + Recording(async_first, first_clock), + Recording(async_second, second_clock), + ) + async with AsyncOpenAI( + api_key="ek_fake_live", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client: + if role == "primary": + async_manager = async_client.live.connect() + elif role == "fork": + async_manager = async_client.live.forks.connect(session_id="stored") + else: + async_manager = async_client.live.sideband.connect(session_id="stored") + async with async_manager as async_connection: + async_connection.on("session.output_transcript.delta", async_first.push) + async_connection.on("session.output_transcript.delta", async_second.push) + async_connection.on("session.output_transcript.delta", typed.append) + async_connection.on("event", observed.append) + + async def async_manage(event: OutputTranscriptDeltaEvent) -> None: + if event.event_id == "part-1": + assert first_clock.pending + async_connection.off("session.output_transcript.delta", async_first.push) + await async_first.close() + await async_first.close() + assert not first_clock.pending + await async_connection.session.update(session={}, event_id="after-dispose") + elif event.event_id == "part-3": + await async_connection.session.close(event_id="caller-finish") + + async def close_async_connection(_event: SessionClosedEvent) -> None: + await async_connection.close() + + async_connection.on("session.output_transcript.delta", async_manage) + async_connection.on("session.closed", close_async_connection) + if isinstance(async_connection, AsyncLiveConnection): + await async_connection.session.start(session={"model": "gpt-live-1"}, event_id="caller-start") + elif isinstance(async_connection, AsyncForksConnection): + await async_connection.session.start(session={}, event_id="caller-start") + await asyncio.wait_for(async_connection.dispatch_events(), timeout=5) + await async_second.close() + + ids = ["part-0", "part-1", "updated", "part-2", "part-3", "closed"] + assert [event.to_dict().get("event_id") for event in observed] == (ids if role == "sideband" else ["started", *ids]) + assert [event.to_dict() for event in typed] == transcripts + updated = observed[2 if role == "sideband" else 3] + assert isinstance(updated, SessionUpdatedEvent) + assert updated.to_dict()["future_metadata"] == {"explicit_null": None, "nested": [1, "retained"]} + assert [(event.segment.text, event.reason) for event in first_record.closed] == [("One two", "manual")] + assert [(event.segment.text, event.reason) for event in second_record.closed] == [("One two three four", "manual")] + assert not first_clock.pending and not second_clock.pending + + +@pytest.mark.parametrize("role", ["primary", "fork", "sideband"]) +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_live_wire_preserves_unknown_events_and_storage_failure_before_closed(role: str, mode: str) -> None: + wire_events = [ + { + "type": "future.live.event", + "event_id": "future", + "nested": {"explicit_null": None, "values": ["東京🙂", 1]}, + }, + { + "type": "error", + "event_id": "storage", + "client_event_id": "caller-finish", + "error": { + "type": "server_error", + "code": "session_storage_failed", + "message": "synthetic recording failure", + "param": None, + "future_detail": {"retry_allowed": False}, + }, + }, + { + "type": "session.closed", + "event_id": "closed", + "reason": "close_requested", + "session": {"id": "live_fixture", "model": "gpt-live-1", "status": "closed", "expires_at": 123}, + "usage": {"future_counter": 0}, + }, + ] + + def script(socket: ServerConnection) -> None: + if role != "sideband": + start = json.loads(socket.recv(timeout=5)) + assert start == { + "type": "session.start", + "event_id": "caller-start", + "session": {"model": "gpt-live-1"} if role == "primary" else {}, + } + socket.send('{"type": "session.started", "event_id": "started"}') + assert json.loads(socket.recv(timeout=5)) == {"type": "session.close", "event_id": "caller-finish"} + for event in wire_events: + socket.send(json.dumps(event)) + + with script_server(script) as url: + if mode == "sync": + with OpenAI(api_key="fake-live-key", base_url=url, http_client=httpx2.Client(trust_env=False)) as client: + if role == "primary": + manager = client.live.connect() + elif role == "fork": + manager = client.live.forks.connect(session_id="stored") + else: + manager = client.live.sideband.connect(session_id="stored") + with manager as connection: + if isinstance(connection, LiveConnection): + connection.session.start(session={"model": "gpt-live-1"}, event_id="caller-start") + assert isinstance(connection.recv(), SessionStartedEvent) + elif isinstance(connection, ForksConnection): + connection.session.start(session={}, event_id="caller-start") + assert isinstance(connection.recv(), SessionStartedEvent) + connection.session.close(event_id="caller-finish") + events = [connection.recv() for _ in wire_events] + else: + async with AsyncOpenAI( + api_key="fake-live-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client: + if role == "primary": + async_manager = async_client.live.connect() + elif role == "fork": + async_manager = async_client.live.forks.connect(session_id="stored") + else: + async_manager = async_client.live.sideband.connect(session_id="stored") + async with async_manager as async_connection: + if isinstance(async_connection, AsyncLiveConnection): + await async_connection.session.start(session={"model": "gpt-live-1"}, event_id="caller-start") + assert isinstance(await asyncio.wait_for(async_connection.recv(), 5), SessionStartedEvent) + elif isinstance(async_connection, AsyncForksConnection): + await async_connection.session.start(session={}, event_id="caller-start") + assert isinstance(await asyncio.wait_for(async_connection.recv(), 5), SessionStartedEvent) + await async_connection.session.close(event_id="caller-finish") + events = [await asyncio.wait_for(async_connection.recv(), 5) for _ in wire_events] + + # The final close must never replace the actionable recording failure. + assert [event.to_dict(exclude_unset=True) for event in events] == wire_events + assert isinstance(events[1], ErrorEvent) + assert events[1].error.code == "session_storage_failed" + assert isinstance(events[2], SessionClosedEvent) diff --git a/tests/lib/responses/test_websocket_accumulator_details.py b/tests/lib/responses/test_websocket_accumulator_details.py new file mode 100644 index 0000000000..7ffe55bc87 --- /dev/null +++ b/tests/lib/responses/test_websocket_accumulator_details.py @@ -0,0 +1,347 @@ +from __future__ import annotations + +import sys +import json +from copy import deepcopy + +import pytest +from websockets.sync.server import ServerConnection + +from openai.lib.responses_websocket import ResponsesWebSocketAccumulator + +from .test_websocket_session import session_for, script_server, response_event + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_partial_details_over_the_wire_preserve_metadata_and_typed_logprobs(mode: str) -> None: + # WS delta logprobs legitimately have no 'bytes'. Never manufacture a + # ResponseOutputText (its final Logprob requires bytes/top_logprobs). + pos = sys.hash_info.modulus * 3 + delta = { + "type": "response.output_text.delta", + "output_index": pos, + "content_index": pos + 1, + "item_id": "msg_fixture", + "delta": "h", + "logprobs": [{"token": "h", "logprob": -0.3, "extra_fixture": [1]}], + } + frames = [ + response_event("created", output=None, metadata={"fixture": "before"}, model="fixture-model"), + delta, + {**delta, "delta": "i", "logprobs": [{"token": "i", "logprob": -0.4, "top_logprobs": None}]}, + { + "type": "response.output_text.annotation.added", + "output_index": pos, + "content_index": pos + 1, + "item_id": "msg_fixture", + "annotation_index": pos + 2, + "annotation": { + "type": "url_citation", + "title": "fixture", + "url": "https://example.com/fixture", + "start_index": 0, + "end_index": 2, + "extra_fixture": {"known": False}, + }, + }, + { + "type": "response.output_text.done", + "output_index": pos, + "content_index": pos + 1, + "item_id": "msg_fixture", + "text": "corrected", + "logprobs": [{"token": "corrected", "logprob": -0.2}], + }, + response_event("incomplete", output=None, metadata={"fixture": "after"}), + ] + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + for frame in frames: + socket.send(json.dumps(frame)) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + for _ in range(4): + event = await driver.call(lane, "recv") + original = deepcopy(event.to_dict()) + acc.add_event(event) + assert event.to_dict() == original + frozen = acc.snapshot() + frozen_hash = hash(frozen) + saved = acc.detailed_snapshot() + row = saved["output"][0] + assert saved["response"]["metadata"] == {"fixture": "before"} + assert saved["response"]["model"] == "fixture-model" + assert "output" not in saved["response"] + assert row["output_index"] == pos and row["item"]["id"] == "msg_fixture" + assert "type" not in row["item"] # Missing setup never fabricates a type. + assert row["content"][0]["content_index"] == pos + 1 + part = row["content"][0]["part"] + assert part["text"] == "hi" + assert part["logprobs"] == [ + {"token": "h", "logprob": -0.3, "extra_fixture": [1]}, + {"token": "i", "logprob": -0.4, "top_logprobs": None}, + ] + assert part["annotations"] == [{"annotation_index": pos + 2, "annotation": frames[3]["annotation"]}] + untouched = deepcopy(saved) + # Both received events and returned mutable snapshots may be freely + # changed without damaging later progress or old frozen snapshots. + saved["response"]["metadata"]["fixture"] = "caller" + part["logprobs"][0]["extra_fixture"].append(99) + part["annotations"][0]["annotation"]["extra_fixture"]["known"] = True + assert acc.detailed_snapshot() == untouched + acc.add_event(await driver.call(lane, "recv")) + detailed = acc.detailed_snapshot() + part = detailed["output"][0]["content"][0]["part"] + assert part["text"] == "corrected" + assert part["logprobs"] == [{"token": "corrected", "logprob": -0.2}] + assert part["annotations"] == untouched["output"][0]["content"][0]["part"]["annotations"] + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.detailed_snapshot()["response"]["metadata"] == {"fixture": "after"} + assert acc.detailed_snapshot()["terminal_type"] == "response.incomplete" + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + assert acc.snapshot().output_text == "corrected" + assert frozen.output_text == "hi" and hash(frozen) == frozen_hash + acc.reset() + assert not acc.detailed_snapshot()["output"] + assert acc.detailed_snapshot()["response"] is None + assert untouched["response"]["metadata"] == {"fixture": "before"} + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_details_replaced_at_part_item_and_response_boundaries(mode: str) -> None: + pos = sys.hash_info.modulus + fixture_parts: list[dict[str, object]] = [ + {"type": "output_text", "text": "first", "annotations": [], "logprobs": None}, + {"type": "refusal", "refusal": "fixture refusal"}, + {"type": "future_fixture_part", "raw_fixture": [True]}, + ] + message: dict[str, object] = { + "id": "fixture", + "type": "message", + "role": "assistant", + "status": "in_progress", + "source_note": {"fixture": [1]}, + "content": fixture_parts, + } + replacement = { + "type": "mcp_call", + "id": "fixture_new", + "name": "search", + "arguments": "{}", + "server_label": "fixture_server", + "call_id": "fixture_call", + "output": None, + "error": None, + } + frames: list[dict[str, object]] = [ + {"type": "response.output_item.added", "output_index": pos, "item": message}, + { + "type": "response.output_text.annotation.added", + "output_index": pos, + "content_index": 0, + "item_id": "fixture", + "annotation_index": pos, + "annotation": {"type": "future_fixture_citation", "raw_fixture": [1]}, + }, + { + "type": "response.content_part.done", + "output_index": pos, + "content_index": 0, + "item_id": "fixture", + "part": {"type": "output_text", "text": "final", "annotations": [], "logprobs": []}, + }, + {"type": "response.output_item.done", "output_index": pos, "item": replacement}, + { + "type": "response.output_text.annotation.added", + "output_index": pos, + "content_index": 0, + "item_id": "fixture", + "annotation_index": 0, + "annotation": None, + }, + response_event("completed", output=[]), + ] + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + for frame in frames: + socket.send(json.dumps(frame)) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + first = acc.detailed_snapshot() + out = first["output"][0] + assert out["item"] == {k: v for k, v in message.items() if k != "content"} + assert [row["part"] for row in out["content"]] == message["content"] + assert out["content"][0]["part"]["logprobs"] is None + acc.add_event(await driver.call(lane, "recv")) + assert len(acc.detailed_snapshot()["output"][0]["content"][0]["part"]["annotations"]) == 1 + acc.add_event(await driver.call(lane, "recv")) + parts = acc.detailed_snapshot()["output"][0]["content"] + assert [row["part"] for row in parts] == [frames[2]["part"], *fixture_parts[1:]] + assert first["output"][0]["content"][0]["part"]["text"] == "first" + acc.add_event(await driver.call(lane, "recv")) + saved = acc.detailed_snapshot() + assert saved["output"] == [{"output_index": pos, "item": replacement, "content": []}] + # A retired item's late annotation cannot replace the new item. + acc.add_event(await driver.call(lane, "recv")) + assert acc.detailed_snapshot() == saved + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.detailed_snapshot()["output"] == [] + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("bad_position", [None, -1, False, "fixture-invalid"]) +async def test_malformed_new_annotation_does_not_break_old_projection_or_retire_item( + mode: str, bad_position: object +) -> None: + delta = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "fixture", + "delta": "ok", + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + for event in [ + delta, + { + "type": "response.output_text.annotation.added", + "output_index": 0, + "content_index": 0, + "item_id": "replacement", + "annotation_index": bad_position, + "annotation": {"type": "future_fixture_annotation"}, + }, + {**delta, "delta": " next"}, + response_event("completed", output=None), + ]: + socket.send(json.dumps(event)) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.detailed_snapshot() + event = await driver.call(lane, "recv") + # Even if Pydantic normalizes false to 0, this other item's + # annotation cannot replace the original or end its progress. + acc.add_event(event) + assert acc.detailed_snapshot() == prior + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot().output_text == "ok next" + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "completed" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_unknown_full_items_keep_their_content_field_and_never_become_final(mode: str) -> None: + future = { + "type": "future_fixture_item", + "id": "fixture_future", + "content": {"shape": ["not", "a", "message"], "value": None}, + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps({"type": "response.output_item.done", "output_index": 3, "item": future})) + socket.close() + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + assert acc.detailed_snapshot()["output"] == [{"output_index": 3, "item": future, "content": []}] + with pytest.raises(EOFError): + await driver.call(lane, "recv") + with pytest.raises(RuntimeError, match="No terminal"): + acc.get_final_response() + assert acc.detailed_snapshot()["terminal_type"] is None + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("annotation_first", [False, True]) +async def test_annotation_cannot_start_a_turn_or_retire_an_unrelated_item(mode: str, annotation_first: bool) -> None: + delta = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "fixture_original", + "delta": "kept", + } + annotation: dict[str, object] = { + "type": "response.output_text.annotation.added", + "output_index": 0, + "content_index": 0, + "annotation_index": 0, + "item_id": "fixture_other", + } + if annotation_first: + annotation["stream_id"] = "unregistered" + annotation["annotation"] = {"type": "future_fixture"} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + for event in [annotation, delta] if annotation_first else [delta, annotation]: + socket.send(json.dumps(event)) + socket.send(json.dumps({**delta, "delta": " more"})) + socket.send(json.dumps(response_event("completed", output=None))) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + for _ in range(3): + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot().output_text == "kept more" + assert acc.detailed_snapshot()["output"][0]["item"]["id"] == "fixture_original" + assert "annotations" not in acc.detailed_snapshot()["output"][0]["content"][0]["part"] + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "completed" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("content_case", ["missing", "null", "empty"]) +async def test_message_content_presence_preserved_before_terminal(mode: str, content_case: str) -> None: + item: dict[str, object] = {"type": "message", "id": "fixture", "role": "assistant"} + if content_case != "missing": + item["content"] = None if content_case == "null" else [] + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps({"type": "response.output_item.done", "output_index": 1, "item": item})) + socket.send(json.dumps(response_event("completed", output=None))) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane = driver.session.default + await driver.call(lane, "send", {"type": "response.create", "input": "synthetic fixture"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + row = acc.detailed_snapshot()["output"][0] + assert row["item"] == {"type": "message", "id": "fixture", "role": "assistant"} + if content_case == "missing": + assert "content" not in row + else: + assert row["content"] == item["content"] + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "completed" diff --git a/tests/lib/test_realtime_websocket_contract.py b/tests/lib/test_realtime_websocket_contract.py index 5dddcd2707..d1dde5f00c 100644 --- a/tests/lib/test_realtime_websocket_contract.py +++ b/tests/lib/test_realtime_websocket_contract.py @@ -6,9 +6,13 @@ import httpx2 import pytest +from websockets.sync.client import ClientConnection from websockets.sync.server import ServerConnection +from websockets.asyncio.client import ClientConnection as AsyncClientConnection from openai import OpenAI, AsyncOpenAI +from openai._exceptions import WebSocketQueueFullError +from openai.types.websocket_reconnection import ReconnectingEvent from openai.types.realtime.realtime_error_event import RealtimeErrorEvent from openai.types.realtime.input_audio_buffer_cleared_event import InputAudioBufferClearedEvent @@ -133,3 +137,136 @@ def script(socket: ServerConnection) -> None: "event_id": "future-1", "text": "東京🙂", } + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_realtime_empty_manager_queue_keeps_budget_after_open(mode: str) -> None: + opened = 0 + + def script(socket: ServerConnection) -> None: + nonlocal opened + opened += 1 + if opened == 1: + socket.close(code=1011, reason="synthetic restart") + return + # Let the iterator resume even if its connection lost the manager's + # queue. Server-side inspection still proves exactly what was sent. + socket.send('{"type": "input_audio_buffer.cleared", "event_id": "ready"}') + assert json.loads(socket.recv(timeout=5)) == {"type": "input_audio_buffer.clear", "event_id": "queued"} + + with script_server(script, expected_connections=2) as url: + if mode == "sync": + with OpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client: + + def on_retry(_event: ReconnectingEvent) -> None: + manager.send({"type": "input_audio_buffer.clear", "event_id": "queued"}) + with pytest.raises(WebSocketQueueFullError): + manager.send({"type": "input_audio_buffer.append", "audio": "AAAA" * 100}) + + manager = client.realtime.connect( + model="gpt-realtime", max_queue_size=96, on_reconnecting=on_retry, initial_delay=0 + ) + with manager as connection: + event = next(iter(connection)) + assert event.type == "input_audio_buffer.cleared" + else: + async with AsyncOpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as client: + + def on_async_retry(_event: ReconnectingEvent) -> None: + async_manager.send({"type": "input_audio_buffer.clear", "event_id": "queued"}) + with pytest.raises(WebSocketQueueFullError): + async_manager.send({"type": "input_audio_buffer.append", "audio": "AAAA" * 100}) + + async_manager = client.realtime.connect( + model="gpt-realtime", max_queue_size=96, on_reconnecting=on_async_retry, initial_delay=0 + ) + async with async_manager as connection: + event = await asyncio.wait_for(anext(aiter(connection)), 5) + assert event.type == "input_audio_buffer.cleared" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("preopen", [True, False], ids=["preopen-flush", "direct"]) +async def test_realtime_recovery_never_replays_an_attempted_command( + mode: str, preopen: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + opened = 0 + recorded: list[tuple[int, str]] = [] + + def script(socket: ServerConnection) -> None: + nonlocal opened + opened += 1 + current = opened + if current == 1: + for _ in range(2 if preopen else 1): + recorded.append((current, json.loads(socket.recv(timeout=5))["event_id"])) + socket.close(code=1011, reason="synthetic restart") + return + socket.send('{"type": "input_audio_buffer.cleared", "event_id": "ready"}') + for _ in range(2 if preopen else 1): + recorded.append((current, json.loads(socket.recv(timeout=5))["event_id"])) + + # Fail after the real wire send: an exception cannot establish non-delivery. + original = ClientConnection.send + original_async = AsyncClientConnection.send + + def send_then_interrupt(self: ClientConnection, message: object, *args: object, **kwargs: object) -> None: + original(self, message, *args, **kwargs) # type: ignore[arg-type] + if json.loads(message)["event_id"] == "attempted": # type: ignore[arg-type] + raise OSError("synthetic interruption after send") + + async def async_send_then_interrupt( + self: AsyncClientConnection, message: object, *args: object, **kwargs: object + ) -> None: + await original_async(self, message, *args, **kwargs) # type: ignore[arg-type] + if json.loads(message)["event_id"] == "attempted": # type: ignore[arg-type] + raise OSError("synthetic interruption after send") + + if mode == "sync": + monkeypatch.setattr(ClientConnection, "send", send_then_interrupt) + else: + monkeypatch.setattr(AsyncClientConnection, "send", async_send_then_interrupt) + with script_server(script, expected_connections=2) as url: + if mode == "sync": + with OpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client: + + def on_retry(_event: ReconnectingEvent) -> None: + manager.send({"type": "input_audio_buffer.clear", "event_id": "during-recovery"}) + + manager = client.realtime.connect(model="gpt-realtime", on_reconnecting=on_retry, initial_delay=0) + if preopen: + for identity in ("first", "attempted", "unattempted"): + manager.send({"type": "input_audio_buffer.clear", "event_id": identity}) + with manager as connection: + if not preopen: + with pytest.raises(OSError, match="synthetic interruption after send"): + connection.send({"type": "input_audio_buffer.clear", "event_id": "attempted"}) + assert next(iter(connection)).type == "input_audio_buffer.cleared" + else: + async with AsyncOpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as client: + + def on_async_retry(_event: ReconnectingEvent) -> None: + async_manager.send({"type": "input_audio_buffer.clear", "event_id": "during-recovery"}) + + async_manager = client.realtime.connect( + model="gpt-realtime", on_reconnecting=on_async_retry, initial_delay=0 + ) + if preopen: + for identity in ("first", "attempted", "unattempted"): + async_manager.send({"type": "input_audio_buffer.clear", "event_id": identity}) + async with async_manager as connection: + if not preopen: + with pytest.raises(OSError, match="synthetic interruption after send"): + await connection.send({"type": "input_audio_buffer.clear", "event_id": "attempted"}) + event = await asyncio.wait_for(anext(aiter(connection)), 5) + assert event.type == "input_audio_buffer.cleared" + expected = [(1, "first"), (1, "attempted"), (2, "unattempted"), (2, "during-recovery")] + assert recorded == (expected if preopen else [(1, "attempted"), (2, "during-recovery")]) diff --git a/tests/lib/test_realtime_websocket_recovery.py b/tests/lib/test_realtime_websocket_recovery.py new file mode 100644 index 0000000000..b8334dce9a --- /dev/null +++ b/tests/lib/test_realtime_websocket_recovery.py @@ -0,0 +1,252 @@ +from __future__ import annotations + +import json +import asyncio +import threading +from urllib.parse import parse_qs, urlsplit + +import httpx2 +import pytest +from websockets.exceptions import ConnectionClosedError +from websockets.sync.server import ServerConnection + +from openai import OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI +from openai.lib.azure import API_KEY_SENTINEL +from openai.types.websocket_reconnection import ReconnectingEvent, ReconnectingOverrides +from openai.types.realtime.realtime_error_event import RealtimeErrorEvent +from openai.types.realtime.input_audio_buffer_cleared_event import InputAudioBufferClearedEvent + +from .responses.test_websocket_session import script_server + + +@pytest.fixture(autouse=True) +def bypass_loopback_proxy(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NO_PROXY", "127.0.0.1") + monkeypatch.setenv("no_proxy", "127.0.0.1") + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("profile", ["direct-recv", "no-callback", "caller-abort", "nonrecoverable", "clean"]) +async def test_realtime_recovery_stops_without_an_extra_upgrade(mode: str, profile: str) -> None: + attempts: list[int] = [] + code = {"nonrecoverable": 1008, "clean": 1000}.get(profile, 1011) + + def script(socket: ServerConnection) -> None: + socket.close(code, "synthetic close") + + def on_retry(event: ReconnectingEvent) -> ReconnectingOverrides: + attempts.append(event.attempt) + return {"abort": True} + + with script_server(script) as url: + callback = None if profile == "no-callback" else on_retry + if mode == "sync": + with OpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client: + with client.realtime.connect(model="gpt-realtime", on_reconnecting=callback, initial_delay=0) as conn: + if profile == "clean": + assert list(conn) == [] + else: + with pytest.raises(ConnectionClosedError) as error: + if profile == "direct-recv": + conn.recv() + else: + next(iter(conn)) + assert error.value.rcvd is not None and error.value.rcvd.code == code + else: + async with AsyncOpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client: + async with async_client.realtime.connect( + model="gpt-realtime", on_reconnecting=callback, initial_delay=0 + ) as async_conn: + if profile == "clean": + assert [e async for e in async_conn] == [] + else: + with pytest.raises(ConnectionClosedError) as async_error: + if profile == "direct-recv": + await asyncio.wait_for(async_conn.recv(), timeout=5) + else: + await asyncio.wait_for(async_conn.__aiter__().__anext__(), timeout=5) + assert async_error.value.rcvd is not None and async_error.value.rcvd.code == code + assert attempts == ([1] if profile == "caller-abort" else []) + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_realtime_admission_errors_do_not_reset_the_retry_budget(mode: str) -> None: + attempts: list[int] = [] + seen: list[RealtimeErrorEvent] = [] + + def script(socket: ServerConnection) -> None: + socket.send( + json.dumps( + { + "type": "error", + "event_id": "fake-error", + "error": {"type": "server_error", "code": "busy", "message": "Synthetic busy"}, + } + ) + ) + socket.close(1011, "synthetic close") + + def on_retry(event: ReconnectingEvent) -> None: + assert event.close_code == 1011 + assert event.max_attempts == 2 + attempts.append(event.attempt) + + with script_server(script, expected_connections=3) as url: + if mode == "sync": + with OpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client: + with client.realtime.connect( + model="gpt-realtime", on_reconnecting=on_retry, max_retries=2, initial_delay=0 + ) as conn: + with pytest.raises(ConnectionClosedError): + for event in conn: + assert isinstance(event, RealtimeErrorEvent) + seen.append(event) + assert len(seen) <= 3 + else: + async with AsyncOpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as client_async: + async with client_async.realtime.connect( + model="gpt-realtime", on_reconnecting=on_retry, max_retries=2, initial_delay=0 + ) as async_conn: + with pytest.raises(ConnectionClosedError): + async for event in async_conn: + assert isinstance(event, RealtimeErrorEvent) + seen.append(event) + assert len(seen) <= 3 + assert attempts == [1, 2] + assert [e.error.code for e in seen] == ["busy", "busy", "busy"] + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("provider", ["openai", "azure"]) +async def test_realtime_replacement_refreshes_provider_and_preserves_connection_options( + mode: str, provider: str +) -> None: + credential = "fake-realtime-before" + upgrades: list[str] = [] + + def script(socket: ServerConnection) -> None: + assert socket.request is not None + target = urlsplit(socket.request.path) + upgrades.append(socket.request.path) + assert target.path.endswith("/customer/realtime") + query = parse_qs(target.query) + assert query["contract"] == ["socket"] + assert query["tenant"] == ["sample"] + assert socket.request.headers.get_all("Authorization") == [f"Bearer {credential}"] + assert socket.request.headers.get("api-key") is None + assert socket.request.headers["X-Realtime-Test"] == "connection" + assert socket.request.headers.get("Sec-WebSocket-Extensions") is None + if len(upgrades) == 1: + socket.close(1011, "synthetic close") + else: + socket.send('{"type":"input_audio_buffer.cleared","event_id":"recovered"}') + assert json.loads(socket.recv(timeout=5)) == {"type": "input_audio_buffer.clear", "event_id": "next"} + + def on_retry(event: ReconnectingEvent) -> None: + nonlocal credential + assert event.attempt == 1 + credential = "fake-realtime-after" + + def get_token() -> str: + return credential + + async def get_async_token() -> str: + return credential + + with script_server(script, expected_connections=2) as url: + if mode == "sync": + client = ( + AzureOpenAI( + api_key=API_KEY_SENTINEL, + azure_ad_token_provider=get_token, + azure_endpoint="https://origin.test", + websocket_base_url=f"{url.replace('http://', 'ws://')}/customer", + api_version="2024-01-01", + http_client=httpx2.Client(trust_env=False), + ) + if provider == "azure" + else OpenAI( + api_key=get_token, + base_url=f"{url}/customer?tenant=sample", + http_client=httpx2.Client(trust_env=False), + ) + ) + with client: + with client.realtime.connect( + model="gpt-realtime", + on_reconnecting=on_retry, + initial_delay=0, + extra_query={"contract": "socket", **({"tenant": "sample"} if provider == "azure" else {})}, + extra_headers={"X-Realtime-Test": "connection"}, + websocket_connection_options={"compression": None}, + ) as conn: + event = next(iter(conn)) + assert isinstance(event, InputAudioBufferClearedEvent) + assert event.event_id == "recovered" + conn.input_audio_buffer.clear(event_id="next") + else: + async_client = ( + AsyncAzureOpenAI( + api_key=API_KEY_SENTINEL, + azure_ad_token_provider=get_token, + azure_endpoint="https://origin.test", + websocket_base_url=f"{url.replace('http://', 'ws://')}/customer", + api_version="2024-01-01", + http_client=httpx2.AsyncClient(trust_env=False), + ) + if provider == "azure" + else AsyncOpenAI( + api_key=get_async_token, + base_url=f"{url}/customer?tenant=sample", + http_client=httpx2.AsyncClient(trust_env=False), + ) + ) + async with async_client: + async with async_client.realtime.connect( + model="gpt-realtime", + on_reconnecting=on_retry, + initial_delay=0, + extra_query={"contract": "socket", **({"tenant": "sample"} if provider == "azure" else {})}, + extra_headers={"X-Realtime-Test": "connection"}, + websocket_connection_options={"compression": None}, + ) as async_conn: + async_event = await asyncio.wait_for(async_conn.__aiter__().__anext__(), timeout=5) + assert isinstance(async_event, InputAudioBufferClearedEvent) + assert async_event.event_id == "recovered" + await async_conn.input_audio_buffer.clear(event_id="next") + assert len(upgrades) == 2 and upgrades[0] == upgrades[1] + + +async def test_realtime_cancelled_receive_preserves_next_event_and_socket() -> None: + release = threading.Event() + + def script(socket: ServerConnection) -> None: + assert release.wait(timeout=5) + socket.send('{"type":"input_audio_buffer.cleared","event_id":"after-cancel"}') + assert json.loads(socket.recv(timeout=5)) == {"type": "input_audio_buffer.clear", "event_id": "after-cancel"} + + with script_server(script) as url: + async with AsyncOpenAI( + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as client: + async with client.realtime.connect(model="gpt-realtime") as conn: + receiving = asyncio.create_task(conn.recv()) + await asyncio.sleep(0) + receiving.cancel() + try: + with pytest.raises(asyncio.CancelledError): + await receiving + finally: + release.set() + event = await asyncio.wait_for(conn.recv(), timeout=5) + assert isinstance(event, InputAudioBufferClearedEvent) + assert event.event_id == "after-cancel" + await conn.input_audio_buffer.clear(event_id="after-cancel") diff --git a/tests/test_client.py b/tests/test_client.py index 6e8b4d2e72..006fe24d0a 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -3,6 +3,7 @@ import gc import io import os +import ssl import sys import json import math @@ -14,6 +15,7 @@ from unittest import mock from typing_extensions import Literal, AsyncIterator, override +import anyio import httpx2 import pytest from pydantic import ValidationError @@ -3647,3 +3649,46 @@ async def test_invalid_request_retry_limit(is_async: bool, value: Any) -> None: await client.close() else: client.close() + + +@pytest.mark.parametrize("is_async", [False, True]) +@pytest.mark.parametrize("error_type", [ssl.SSLError, anyio.EndOfStream]) +@pytest.mark.parametrize("max_retries,failures", [(0, 1), (2, 3), (2, 1)]) +async def test_unmapped_transport_errors( + is_async: bool, error_type: type[Exception], max_retries: int, failures: int +) -> None: + requests: list[httpx2.Request] = [] + error = error_type("synthetic transport failure") + + def handler(request: httpx2.Request) -> httpx2.Response: + requests.append(request) + if len(requests) <= failures: + raise error + return httpx2.Response(200, json={"recovered": True}) + + transport = httpx2.MockTransport(handler) + client = ( + AsyncOpenAI(api_key="fake-key", max_retries=max_retries, http_client=httpx2.AsyncClient(transport=transport)) + if is_async + else OpenAI(api_key="fake-key", max_retries=max_retries, http_client=httpx2.Client(transport=transport)) + ) + try: + with mock.patch("openai._base_client.BaseClient._calculate_retry_timeout", return_value=0): + try: + if isinstance(client, AsyncOpenAI): + result = await client.get("/test", cast_to=object) + else: + result = client.get("/test", cast_to=object) + except APIConnectionError as exc: + assert failures > max_retries + assert exc.__cause__ is error + assert exc.request is requests[-1] + else: + assert failures <= max_retries + assert result == {"recovered": True} + assert len(requests) == min(failures + 1, max_retries + 1) + finally: + if isinstance(client, AsyncOpenAI): + await client.close() + else: + client.close() diff --git a/tests/test_send_queue_reconnect.py b/tests/test_send_queue_reconnect.py index fec12479ab..5d48ed5097 100644 --- a/tests/test_send_queue_reconnect.py +++ b/tests/test_send_queue_reconnect.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -22,15 +23,17 @@ def test_reconnect_retries_bounded_send_queue( q.enqueue("aaa") ws = MagicMock() attempts = 0 + realtime = connection_type is RealtimeConnection def failing_send(data: str) -> None: nonlocal attempts - assert data == "aaa" + assert data == ("b" if realtime and attempts == 1 else "aaa") if attempts == 0: q.enqueue("b") attempts += 1 - with pytest.raises(WebSocketQueueFullError): - q.enqueue("c") + if not realtime or attempts == 1: + with pytest.raises(WebSocketQueueFullError): + q.enqueue("c") raise RuntimeError("fake send failure") ws.send.side_effect = failing_send @@ -44,10 +47,10 @@ def failing_send(data: str) -> None: ) for expected_attempts in range(1, 4): assert connection._reconnect(RuntimeError("fake disconnect")) - assert q._bytes == 4 - assert attempts == expected_attempts + assert q._bytes == ((1 if expected_attempts == 1 else 0) if realtime else 4) + assert attempts == (min(expected_attempts, 2) if realtime else expected_attempts) assert not connection._reconnect(RuntimeError("fake disconnect")) - assert attempts == 3 + assert attempts == (2 if realtime else 3) # A healthy application event resets the budget; an upgrade alone does not. ws.recv.return_value = '{"type": "response.created"}' @@ -56,7 +59,7 @@ def failing_send(data: str) -> None: sent: list[str] = [] ws.send.side_effect = sent.append assert connection._reconnect(RuntimeError("fake disconnect")) - assert sent == ["aaa", "b"] + assert sent == ([] if realtime else ["aaa", "b"]) assert q._bytes == 0 for _ in range(2): assert connection._reconnect(RuntimeError("fake disconnect")) @@ -76,15 +79,17 @@ async def test_async_reconnect_retries_bounded_send_queue( q.enqueue("aaa") ws = MagicMock() attempts = 0 + realtime = connection_type is AsyncRealtimeConnection async def failing_send(data: str) -> None: nonlocal attempts - assert data == "aaa" + assert data == ("b" if realtime and attempts == 1 else "aaa") if attempts == 0: q.enqueue("b") attempts += 1 - with pytest.raises(WebSocketQueueFullError): - q.enqueue("c") + if not realtime or attempts == 1: + with pytest.raises(WebSocketQueueFullError): + q.enqueue("c") raise RuntimeError("fake send failure") ws.send = AsyncMock(side_effect=failing_send) @@ -98,10 +103,10 @@ async def failing_send(data: str) -> None: ) for expected_attempts in range(1, 4): assert await connection._reconnect(RuntimeError("fake disconnect")) - assert q._bytes == 4 - assert attempts == expected_attempts + assert q._bytes == ((1 if expected_attempts == 1 else 0) if realtime else 4) + assert attempts == (min(expected_attempts, 2) if realtime else expected_attempts) assert not await connection._reconnect(RuntimeError("fake disconnect")) - assert attempts == 3 + assert attempts == (2 if realtime else 3) # A healthy application event resets the budget; an upgrade alone does not. ws.recv = AsyncMock(return_value='{"type": "response.created"}') @@ -110,8 +115,44 @@ async def failing_send(data: str) -> None: sent: list[str] = [] ws.send.side_effect = sent.append assert await connection._reconnect(RuntimeError("fake disconnect")) - assert sent == ["aaa", "b"] + assert sent == ([] if realtime else ["aaa", "b"]) assert q._bytes == 0 for _ in range(2): assert await connection._reconnect(RuntimeError("fake disconnect")) assert not await connection._reconnect(RuntimeError("retry budget exhausted")) + + +@pytest.mark.asyncio +async def test_cancelled_realtime_flush_does_not_replay_active_send() -> None: + q = SendQueue(max_bytes=4) + q.enqueue("é") + q.enqueue("b") + active = asyncio.Event() + socket = MagicMock() + + async def suspended_send(data: str) -> None: + assert data == "é" + active.set() + await asyncio.Event().wait() + + socket.send = AsyncMock(side_effect=suspended_send) + connection = AsyncRealtimeConnection(socket, send_queue=q) + task = asyncio.create_task(connection._flush_send_queue()) + try: + await asyncio.wait_for(active.wait(), 5) + q.enqueue("c") + with pytest.raises(WebSocketQueueFullError): + q.enqueue("d") + finally: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + # Only the active send was attempted; its bytes are released, and the + # queued tail and concurrent enqueue remain in order for the next socket. + q.enqueue("dd") + sent: list[str] = [] + socket.send.side_effect = sent.append + await connection._flush_send_queue() + assert sent == ["b", "c", "dd"] + q.enqueue("1234") diff --git a/tests/test_streaming.py b/tests/test_streaming.py index fe74fcc558..4a6ec58b80 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -1,10 +1,12 @@ from __future__ import annotations import os +import ssl import importlib from typing import Any, Iterator, AsyncIterator from contextlib import aclosing, nullcontext +import anyio import httpx2 import pytest @@ -35,12 +37,19 @@ def http_module(request: pytest.FixtureRequest) -> Any: ("ReadTimeout", APITimeoutError), ("RemoteProtocolError", APIConnectionError), ("DecodingError", APIConnectionError), + (ssl.SSLError, APIConnectionError), + (anyio.EndOfStream, APIConnectionError), ], ) async def test_request_errors_are_wrapped( - sync: bool, delivered: bool, error_name: str, expected_error: type[APIConnectionError], http_module: Any + sync: bool, + delivered: bool, + error_name: str | type[Exception], + expected_error: type[APIConnectionError], + http_module: Any, ) -> None: - error = getattr(http_module, error_name)("synthetic stream failure") + error_type = getattr(http_module, error_name) if isinstance(error_name, str) else error_name + error = error_type("synthetic stream failure") requests: list[Any] = [] first = ( b'data: {"id":"synthetic","object":"chat.completion.chunk","created":0,"model":"synthetic",' diff --git a/uv.lock b/uv.lock index 0e69bdb0f9..9e256efd0d 100644 --- a/uv.lock +++ b/uv.lock @@ -1535,7 +1535,7 @@ wheels = [ [[package]] name = "openai" -version = "3.19.2" # x-release-please-version +version = "3.20.0" # x-release-please-version source = { editable = "." } dependencies = [ { name = "anyio" },