From 08f9e730a108c1151ca1e8a817bec1268ed1e4b8 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 18:53:39 +0000 Subject: [PATCH 01/15] feat(responses): add opt-in websocket accumulation --- src/openai/lib/responses_websocket/README.md | 34 +++ .../lib/responses_websocket/__init__.py | 5 + .../lib/responses_websocket/_accumulator.py | 217 +++++++++++++++ .../responses/test_websocket_accumulator.py | 255 ++++++++++++++++++ 4 files changed, 511 insertions(+) create mode 100644 src/openai/lib/responses_websocket/_accumulator.py create mode 100644 tests/lib/responses/test_websocket_accumulator.py diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index 9a94fec82f..3c0a5a18d4 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -68,6 +68,40 @@ the next response. It does not wait for a hypothetical future successor. Use `recv()` to observe that boundary when coordinating steering; accepted steering alone does not prove that a successor has started. +For incremental text and tool-input snapshots, opt in with a separate helper +fed by the events you already receive: + +```python +from openai.lib.responses_websocket import ResponsesWebSocketAccumulator + +accumulator = ResponsesWebSocketAccumulator() +while True: + event = await lane.recv() # Use lane.recv(timeout=...) in a sync session. + accumulator.add_event(event) + # Original event fields, including unknown variants/fields, remain available. + snapshot = accumulator.snapshot() + print("Current text:", snapshot.output_text) + if snapshot.terminal_type is not None: + response = accumulator.get_final_response() + break +accumulator.reset() # Does not close the lane or connection. +``` + +Use one helper per lane and reset it before the next turn. Snapshots are +immutable and contain selected text, function arguments, and custom tool input, +grouped by output/content index and item ID. They never execute tools. Done +events replace provisional fields. Full item replacements discard old fields; +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. + +The helper's `get_final_response()` returns a copy of the exact received server +response, including its original missing/null/empty output, on completed, +failed or incomplete. It never substitutes the projection for that response. +Before a valid terminal it raises; protocol errors retain their original event. +This caller-fed helper owns no socket, reader or timers and can also observe +events from a connection when there is no session. + Cancel an async `recv` or `get_final_response` wait with normal asyncio cancellation. It leaves queued events and accumulated state available for a later wait, and does not close the lane or connection. Synchronous waits accept diff --git a/src/openai/lib/responses_websocket/__init__.py b/src/openai/lib/responses_websocket/__init__.py index eb827d7186..9b89a900b1 100644 --- a/src/openai/lib/responses_websocket/__init__.py +++ b/src/openai/lib/responses_websocket/__init__.py @@ -7,3 +7,8 @@ ResponsesWebSocketBufferError as ResponsesWebSocketBufferError, AsyncResponsesWebSocketSession as AsyncResponsesWebSocketSession, ) +from ._accumulator import ( + ResponsesWebSocketOutput as ResponsesWebSocketOutput, + ResponsesWebSocketSnapshot as ResponsesWebSocketSnapshot, + ResponsesWebSocketAccumulator as ResponsesWebSocketAccumulator, +) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py new file mode 100644 index 0000000000..644de0ad60 --- /dev/null +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -0,0 +1,217 @@ +from __future__ import annotations + +from typing import cast +from dataclasses import field, dataclass + +from ._session import ResponsesWebSocketError, _field +from ..._compat import model_copy +from ...types.responses import Response +from ...types.responses.responses_server_event import ResponsesServerEvent + + +@dataclass(frozen=True) +class ResponsesWebSocketOutput: + """Selected fields of one output item. Text pairs contain (content_index, text).""" + + output_index: int + item_id: str | None + type: str | None + name: str | None + call_id: str | None + arguments: str + input: str + text: tuple[tuple[int, str], ...] + + +@dataclass(frozen=True) +class ResponsesWebSocketSnapshot: + """Immutable partial projection; terminal_type is None until a valid terminal arrives.""" + + stream_id: str | None + response_id: str | None + terminal_type: str | None + output: tuple[ResponsesWebSocketOutput, ...] + + @property + def output_text(self) -> str: + return "".join(text for item in self.output for _, text in item.text) + + +@dataclass +class _Output: + item_id: str | None + type: str | None = None + name: str | None = None + call_id: str | None = None + arguments: list[str] = field(default_factory=list[str]) + input: list[str] = field(default_factory=list[str]) + text: dict[int, list[str]] = field(default_factory=dict[int, list[str]]) + + +class ResponsesWebSocketAccumulator: + """Opt-in, caller-fed collection of text, function arguments and custom tool input. + + Feed typed events from a connection or lane's recv(). Use one accumulator per + lane and call reset() before another turn. The helper never reads, sends, + closes the connection, or executes tools. It is not safe for concurrent use. + Snapshot fields are projections; get_final_response() retains the exact + server response, including failed/incomplete and missing/null/empty output. + """ + + def __init__(self) -> None: + self._stream_id: str | None = None + self._response_id: str | None = None + self._terminal_type: str | None = None + self._bound = False + self._output: dict[int, _Output] = {} + self._final: Response | None = None + self._error: Exception | None = None + + def reset(self) -> None: + """Release this helper's state only; prior snapshots and the socket remain usable.""" + self._stream_id = self._response_id = self._terminal_type = None + self._bound = False + self._output.clear() + self._final = self._error = None + + def snapshot(self) -> ResponsesWebSocketSnapshot: + return ResponsesWebSocketSnapshot( + stream_id=self._stream_id, + response_id=self._response_id, + terminal_type=self._terminal_type, + output=tuple( + ResponsesWebSocketOutput( + output_index=index, + item_id=item.item_id, + type=item.type, + name=item.name, + call_id=item.call_id, + arguments="".join(item.arguments), + input="".join(item.input), + text=tuple((pos, "".join(parts)) for pos, parts in sorted(item.text.items())), + ) + for index, item in sorted(self._output.items()) + ), + ) + + def get_final_response(self) -> Response: + """Return an independent copy of the received terminal response, never a partial success.""" + if self._error is not None: + raise self._error + if self._final is None: + raise RuntimeError("No terminal response has been received") + return model_copy(self._final, deep=True) + + def add_event(self, event: ResponsesServerEvent) -> None: + """Observe an event without consuming or changing it; unknown events are ignored. + + Final fields replace deltas. A supplied output list, even [], replaces + the projection; null or omitted output retains prior projected fields. + Original event fields remain accessible to the caller. + """ + kind = _field(event, "type") + if kind not in { + "response.created", + "response.in_progress", + "response.completed", + "response.failed", + "response.incomplete", + "response.output_item.added", + "response.output_item.done", + "response.content_part.added", + "response.content_part.done", + "response.output_text.delta", + "response.output_text.done", + "response.function_call_arguments.delta", + "response.function_call_arguments.done", + "response.custom_tool_call_input.delta", + "response.custom_tool_call_input.done", + "error", + }: + return + stream_id = _field(event, "stream_id") + if self._bound and self._stream_id != stream_id: + raise ValueError("Event belongs to another WebSocket lane") + if self._terminal_type is not None or self._error is not None: + raise RuntimeError("Reset the accumulator before adding another turn") + if kind == "error": + self._error = ResponsesWebSocketError(event) + raise self._error + terminal = kind in {"response.completed", "response.failed", "response.incomplete"} + if kind in {"response.created", "response.in_progress"} or terminal: + response = _field(event, "response") + if not isinstance(response, Response): + error = ValueError("WebSocket event is missing a valid response") + if terminal: + self._error = error + raise error + response_id = _field(response, "id") + if self._response_id is not None and response_id is not None and self._response_id != response_id: + raise ValueError("Event belongs to another response") + self._response_id = response_id or self._response_id + self._bound, self._stream_id = True, stream_id + output = _field(response, "output") + if isinstance(output, list): + self._output.clear() + for index, item in enumerate(cast("list[object]", output)): + self._add_item(index, item) + if terminal: + self._final = model_copy(response, deep=True) + self._terminal_type = kind + return + self._bound, self._stream_id = True, stream_id + index = _field(event, "output_index") + if not isinstance(index, int): + raise ValueError("WebSocket output event is missing output_index") + if kind in {"response.output_item.added", "response.output_item.done"}: + self._add_item(index, _field(event, "item")) + return + item_id = _field(event, "item_id") + item = self._output.get(index) + if item is None or (item_id and item.item_id and item_id != item.item_id): + item = _Output(item_id=item_id) + self._output[index] = item + if item_id: + item.item_id = item_id + if kind in { + "response.content_part.added", + "response.content_part.done", + "response.output_text.delta", + "response.output_text.done", + }: + pos = _field(event, "content_index") + if not isinstance(pos, int): + raise ValueError("WebSocket text event is missing content_index") + if kind == "response.output_text.delta": + item.text.setdefault(pos, []).append(_field(event, "delta")) + elif kind == "response.output_text.done": + item.text[pos] = [_field(event, "text")] + else: + part = _field(event, "part") + if _field(part, "type") == "output_text": + item.text[pos] = [_field(part, "text")] + elif kind == "response.function_call_arguments.delta": + item.arguments.append(_field(event, "delta")) + elif kind == "response.function_call_arguments.done": + item.arguments = [_field(event, "arguments")] + elif kind == "response.custom_tool_call_input.delta": + item.input.append(_field(event, "delta")) + elif kind == "response.custom_tool_call_input.done": + item.input = [_field(event, "input")] + + def _add_item(self, index: int, source: object) -> None: + item = _Output( + item_id=_field(source, "id"), + type=_field(source, "type"), + name=_field(source, "name"), + call_id=_field(source, "call_id"), + ) + self._output[index] = item + if item.type == "message": + for pos, part in enumerate(_field(source, "content") or []): + if _field(part, "type") == "output_text": + item.text[pos] = [_field(part, "text")] + elif item.type == "function_call": + item.arguments = [_field(source, "arguments") or ""] + elif item.type == "custom_tool_call": + item.input = [_field(source, "input") or ""] diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py new file mode 100644 index 0000000000..9d8e63f280 --- /dev/null +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +import json +from typing import cast + +import pytest +from websockets.sync.server import ServerConnection + +from openai import omit +from openai._compat import model_copy +from openai.types.responses import ResponseStreamEvent +from openai.lib.responses_websocket import ResponsesWebSocketError, ResponsesWebSocketAccumulator +from openai.lib.streaming.responses import ResponseStreamState + +from .test_websocket_session import session_for, script_server, response_event + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("terminal", ["completed", "failed", "incomplete"]) +@pytest.mark.parametrize("output", ["missing", "null", "empty"]) +async def test_opt_in_accumulation_exact_terminals(mode: str, terminal: str, output: str) -> None: + created = response_event("created") + frames = [ + {"type": "response.output_text.delta", "output_index": 4, "content_index": 1, "item_id": "msg", "delta": "pre"}, + {"type": "response.output_text.delta", "output_index": 4, "content_index": 1, "item_id": "msg", "delta": "fix"}, + {"type": "response.function_call_arguments.delta", "output_index": 0, "item_id": "fc", "delta": "{"}, + { + "type": "response.function_call_arguments.done", + "output_index": 0, + "item_id": "fc", + "arguments": '{"fixed":true}', + }, + {"type": "response.custom_tool_call_input.delta", "output_index": 2, "item_id": "ct", "delta": "draft"}, + {"type": "response.custom_tool_call_input.done", "output_index": 2, "item_id": "ct", "input": "tool data"}, + { + "type": "response.output_text.done", + "output_index": 4, + "content_index": 1, + "item_id": "msg", + "text": "corrected", + }, + {"type": "response.future", "unmodeled": {"keep": True}}, + ] + final = response_event(terminal) + if output == "missing": + del final["response"]["output"] + elif output == "null": + final["response"]["output"] = None + + def script(socket: ServerConnection) -> None: + assert json.loads(socket.recv(timeout=5))["type"] == "response.create" + for event in [created, *frames, final]: + socket.send(json.dumps(event)) + # Reset does not release the socket or lane; a next turn can use it. + assert json.loads(socket.recv(timeout=5))["input"] == "next" + socket.send(json.dumps(response_event("completed", id="resp_next"))) + + 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", "model": "test", "input": "first"}) + acc = ResponsesWebSocketAccumulator() + initial = await driver.call(lane, "recv") + acc.add_event(initial) + # The existing SSE state needs item + part scaffolding. Prove it + # cannot collect the same valid WS text event, then retain that + # untouched typed event for the caller-fed helper. + sse: ResponseStreamState[object] = ResponseStreamState(input_tools=omit, text_format=omit) + sse.handle_event(cast(ResponseStreamEvent, model_copy(initial, deep=True))) + first = await driver.call(lane, "recv") + before = first.to_dict() + with pytest.raises(RuntimeError, match="before receiving its output item"): + sse.handle_event(cast(ResponseStreamEvent, model_copy(first, deep=True))) + acc.add_event(first) + saved = acc.snapshot() + assert saved.output_text == "pre" + assert first.to_dict() == before + with pytest.raises(RuntimeError, match="No terminal"): + acc.get_final_response() + for _ in frames[1:]: + event = await driver.call(lane, "recv") + before = event.to_dict() + acc.add_event(event) + assert event.to_dict() == before + projected = acc.snapshot() + assert projected.output_text == "corrected" + assert projected.terminal_type is None + assert projected.output[0].arguments == '{"fixed":true}' + assert projected.output[1].input == "tool data" + assert saved.output_text == "pre" + received = await driver.call(lane, "recv") + original = received.response.to_dict() + acc.add_event(received) + result = acc.get_final_response() + assert result.to_dict() == original + assert result.status == terminal + result.id = "caller mutation" + assert acc.get_final_response().to_dict() == original + assert received.response.to_dict() == original + assert acc.snapshot().terminal_type == "response." + terminal + assert acc.snapshot().output_text == ("" if output == "empty" else "corrected") + # Existing get_final_response remains a terminal/final-item collector. + # It must not silently gain delta reconstruction. + legacy = await driver.call(lane, "get_final_response") + assert not legacy.output + acc.reset() + assert not acc.snapshot().output + assert saved.output_text == "pre" + await driver.call(lane, "send", {"type": "response.create", "model": "test", "input": "next"}) + received = await driver.call(lane, "recv") + acc.add_event(received) + assert acc.get_final_response().id == "resp_next" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_opt_in_accumulator_lane_and_item_replacements(mode: str) -> None: + a_events = [ + response_event("created", "a"), + { + "type": "response.output_item.added", + "stream_id": "a", + "output_index": 0, + "item": { + "type": "message", + "id": "msg", + "content": [{"type": "output_text", "text": "first"}, {"type": "output_text", "text": "stale"}], + }, + }, + { + "type": "response.output_item.done", + "stream_id": "a", + "output_index": 0, + "item": {"type": "message", "id": "msg", "content": [{"type": "output_text", "text": "fixed"}]}, + }, + { + "type": "response.output_item.done", + "stream_id": "a", + "output_index": 0, + "item": {"type": "function_call", "id": "msg", "name": "lookup", "call_id": "call", "arguments": "{}"}, + }, + { + "type": "response.function_call_arguments.delta", + "stream_id": "a", + "output_index": 0, + "item_id": "fc_new", + "delta": '{"fresh":', + }, + ] + + def script(socket: ServerConnection) -> None: + for _ in range(2): + socket.recv(timeout=5) + for a in a_events: + socket.send(json.dumps(a)) + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "stream_id": "b", + "output_index": 9, + "content_index": 4, + "item_id": "msg_b", + "delta": "b", + } + ) + ) + socket.send(json.dumps(response_event("incomplete", "a", output=None))) + socket.send(json.dumps(response_event("failed", "b", output=None))) + + with script_server(script) as url: + async with session_for(mode, url) as driver: + lane_a, lane_b = driver.session.lane("a"), driver.session.lane("b") + for lane in (lane_a, lane_b): + await driver.call(lane, "send", {"type": "response.create", "model": "test", "input": "text"}) + a, b = ResponsesWebSocketAccumulator(), ResponsesWebSocketAccumulator() + saved = a.snapshot() + for i in range(len(a_events)): + ea, eb = await driver.call(lane_a, "recv"), await driver.call(lane_b, "recv") + a.add_event(ea) + with pytest.raises(ValueError, match="another WebSocket lane"): + a.add_event(eb) + b.add_event(eb) + if i == 1: + saved = a.snapshot() + assert saved.output_text == "firststale" + elif i == 2: + assert a.snapshot().output_text == "fixed" + assert saved.output_text == "firststale" + elif i == 3: + assert a.snapshot().output_text == "" + assert a.snapshot().output[0].arguments == "{}" + elif i == 4: + projected = a.snapshot().output[0] + assert projected.item_id == "fc_new" and projected.arguments == '{"fresh":' + assert projected.name is None and projected.call_id is None + a.add_event(await driver.call(lane_a, "recv")) + b.add_event(await driver.call(lane_b, "recv")) + assert a.get_final_response().id == "resp_a" + assert b.get_final_response().id == "resp_b" + assert a.snapshot().output_text == "" + assert b.snapshot().output_text == "b" * len(a_events) + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("ending", ["EOF", "error", "missing", "null"]) +async def test_opt_in_accumulator_never_invents_a_final(mode: str, ending: str) -> None: + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "item_id": "msg", + "output_index": 0, + "content_index": 0, + "delta": "partial", + } + ) + ) + if ending == "EOF": + socket.close() + elif ending == "error": + socket.send( + json.dumps( + { + "type": "error", + "error": {"type": "invalid_request_error", "code": "test", "message": "fixture error"}, + } + ) + ) + else: + event = {"type": "response.completed"} + if ending == "null": + event["response"] = None # type: ignore[assignment] + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + if ending == "EOF": + with pytest.raises(EOFError): + await driver.call(lane, "recv") + expected = RuntimeError + else: + expected = ResponsesWebSocketError if ending == "error" else ValueError + received = await driver.call(lane, "recv") + with pytest.raises(expected): + acc.add_event(received) + with pytest.raises(expected): + acc.get_final_response() + assert acc.snapshot().output_text == "partial" + assert acc.snapshot().terminal_type is None From 12ac7ed8bed0fa32a26ad2c427fee5f9b54e9504 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 19:05:16 +0000 Subject: [PATCH 02/15] fix(responses): reject malformed accumulation data before updating snapshots --- .../lib/responses_websocket/_accumulator.py | 50 +++++++++++++++---- .../responses/test_websocket_accumulator.py | 42 ++++++++++++++++ 2 files changed, 81 insertions(+), 11 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 644de0ad60..66bbec8377 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -166,6 +166,25 @@ def add_event(self, event: ResponsesServerEvent) -> None: if kind in {"response.output_item.added", "response.output_item.done"}: self._add_item(index, _field(event, "item")) return + # Validate consumed values before replacing an item's retained state. + # The wire decoder preserves known events even when fields are malformed. + value = "" + if kind in { + "response.output_text.delta", + "response.function_call_arguments.delta", + "response.custom_tool_call_input.delta", + }: + value = _text_field(event, "delta") + elif kind == "response.output_text.done": + value = _text_field(event, "text") + elif kind == "response.function_call_arguments.done": + value = _text_field(event, "arguments") + elif kind == "response.custom_tool_call_input.done": + value = _text_field(event, "input") + elif kind in {"response.content_part.added", "response.content_part.done"}: + part = _field(event, "part") + if _field(part, "type") == "output_text": + value = _text_field(part, "text") item_id = _field(event, "item_id") item = self._output.get(index) if item is None or (item_id and item.item_id and item_id != item.item_id): @@ -183,21 +202,21 @@ def add_event(self, event: ResponsesServerEvent) -> None: if not isinstance(pos, int): raise ValueError("WebSocket text event is missing content_index") if kind == "response.output_text.delta": - item.text.setdefault(pos, []).append(_field(event, "delta")) + item.text.setdefault(pos, []).append(value) elif kind == "response.output_text.done": - item.text[pos] = [_field(event, "text")] + item.text[pos] = [value] else: part = _field(event, "part") if _field(part, "type") == "output_text": - item.text[pos] = [_field(part, "text")] + item.text[pos] = [value] elif kind == "response.function_call_arguments.delta": - item.arguments.append(_field(event, "delta")) + item.arguments.append(value) elif kind == "response.function_call_arguments.done": - item.arguments = [_field(event, "arguments")] + item.arguments = [value] elif kind == "response.custom_tool_call_input.delta": - item.input.append(_field(event, "delta")) + item.input.append(value) elif kind == "response.custom_tool_call_input.done": - item.input = [_field(event, "input")] + item.input = [value] def _add_item(self, index: int, source: object) -> None: item = _Output( @@ -206,12 +225,21 @@ def _add_item(self, index: int, source: object) -> None: name=_field(source, "name"), call_id=_field(source, "call_id"), ) - self._output[index] = item if item.type == "message": for pos, part in enumerate(_field(source, "content") or []): if _field(part, "type") == "output_text": - item.text[pos] = [_field(part, "text")] + item.text[pos] = [_text_field(part, "text")] elif item.type == "function_call": - item.arguments = [_field(source, "arguments") or ""] + value = _field(source, "arguments") + item.arguments = ["" if value is None else _text_field(source, "arguments")] elif item.type == "custom_tool_call": - item.input = [_field(source, "input") or ""] + value = _field(source, "input") + item.input = ["" if value is None else _text_field(source, "input")] + self._output[index] = item + + +def _text_field(value: object, name: str) -> str: + text = _field(value, name) + if not isinstance(text, str): + raise ValueError(f"WebSocket output {name} must be a string") + return text diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 9d8e63f280..5cbf59e28c 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -253,3 +253,45 @@ def script(socket: ServerConnection) -> None: acc.get_final_response() assert acc.snapshot().output_text == "partial" assert acc.snapshot().terminal_type is None + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize( + "invalid", + [ + {"type": "response.output_text.delta", "content_index": 0, "delta": None}, + {"type": "response.function_call_arguments.delta", "delta": {"unexpected": True}}, + {"type": "response.custom_tool_call_input.done", "input": 22}, + ], +) +async def test_rejects_invalid_fields_without_poisoning_prior_snapshot(mode: str, invalid: dict[str, object]) -> None: + good = { + "type": "response.output_text.delta", + "item_id": "msg", + "output_index": 0, + "content_index": 0, + "delta": "saved", + } + bad = {"output_index": 0, "item_id": "msg", **invalid} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(good)) + socket.send(json.dumps(bad)) + socket.send(json.dumps(response_event("incomplete", 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match="must be a string"): + acc.add_event(received) + assert acc.snapshot() == prior + assert received == original + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "incomplete" From a9a072e9ac5b6d7f23bb4d5b52c0d109695a440a Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 19:18:28 +0000 Subject: [PATCH 03/15] fix(responses): retain valid projections and avoid per-delta snapshots --- src/openai/lib/responses_websocket/README.md | 17 +++++-- .../lib/responses_websocket/_accumulator.py | 29 +++++++---- .../responses/test_websocket_accumulator.py | 48 +++++++++++++++++++ 3 files changed, 83 insertions(+), 11 deletions(-) diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index 3c0a5a18d4..f372be6402 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -79,9 +79,12 @@ while True: event = await lane.recv() # Use lane.recv(timeout=...) in a sync session. accumulator.add_event(event) # Original event fields, including unknown variants/fields, remain available. - snapshot = accumulator.snapshot() - print("Current text:", snapshot.output_text) - if snapshot.terminal_type is not None: + if event.type == "response.output_text.delta": + print(event.delta, end="", flush=True) # Provisional progress log. + elif event.type == "response.output_item.done": + snapshot = accumulator.snapshot() + print("\nCurrent projected text:", snapshot.output_text) + elif event.type in {"response.completed", "response.failed", "response.incomplete"}: response = accumulator.get_final_response() break accumulator.reset() # Does not close the lane or connection. @@ -95,6 +98,14 @@ 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. +`snapshot()` materializes the entire current projection and joins retained +fragments. It is proportional to the accumulated output, so requesting it after +every small delta repeatedly rebuilds growing prefixes. Use the original event +for per-delta progress and request a full snapshot only when needed, such as +after an item is done. The progress log above is provisional; done events can +shorten or correct earlier text. A UI should replace its displayed projection +at those boundaries, not append the cumulative `snapshot.output_text`. + The helper's `get_final_response()` returns a copy of the exact received server response, including its original missing/null/empty output, on completed, failed or incomplete. It never substitutes the projection for that response. diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 66bbec8377..12826ab76d 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -75,6 +75,11 @@ def reset(self) -> None: self._final = self._error = None def snapshot(self) -> ResponsesWebSocketSnapshot: + """Materialize the entire current immutable projection. + + This joins retained fragments. Use original events for per-delta progress + and request snapshots intentionally (for example, on output_item.done). + """ return ResponsesWebSocketSnapshot( stream_id=self._stream_id, response_id=self._response_id, @@ -96,8 +101,9 @@ def snapshot(self) -> ResponsesWebSocketSnapshot: def get_final_response(self) -> Response: """Return an independent copy of the received terminal response, never a partial success.""" - if self._error is not None: - raise self._error + error = self._error + if error is not None: + raise error if self._final is None: raise RuntimeError("No terminal response has been received") return model_copy(self._final, deep=True) @@ -135,8 +141,9 @@ def add_event(self, event: ResponsesServerEvent) -> None: if self._terminal_type is not None or self._error is not None: raise RuntimeError("Reset the accumulator before adding another turn") if kind == "error": - self._error = ResponsesWebSocketError(event) - raise self._error + protocol_error = ResponsesWebSocketError(event) + self._error = protocol_error + raise protocol_error terminal = kind in {"response.completed", "response.failed", "response.incomplete"} if kind in {"response.created", "response.in_progress"} or terminal: response = _field(event, "response") @@ -159,12 +166,12 @@ def add_event(self, event: ResponsesServerEvent) -> None: self._final = model_copy(response, deep=True) self._terminal_type = kind return - self._bound, self._stream_id = True, stream_id index = _field(event, "output_index") if not isinstance(index, int): raise ValueError("WebSocket output event is missing output_index") if kind in {"response.output_item.added", "response.output_item.done"}: self._add_item(index, _field(event, "item")) + self._bound, self._stream_id = True, stream_id return # Validate consumed values before replacing an item's retained state. # The wire decoder preserves known events even when fields are malformed. @@ -185,6 +192,15 @@ def add_event(self, event: ResponsesServerEvent) -> None: part = _field(event, "part") if _field(part, "type") == "output_text": value = _text_field(part, "text") + pos = _field(event, "content_index") + if kind in { + "response.content_part.added", + "response.content_part.done", + "response.output_text.delta", + "response.output_text.done", + } and not isinstance(pos, int): + raise ValueError("WebSocket text event is missing content_index") + self._bound, self._stream_id = True, stream_id item_id = _field(event, "item_id") item = self._output.get(index) if item is None or (item_id and item.item_id and item_id != item.item_id): @@ -198,9 +214,6 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.output_text.delta", "response.output_text.done", }: - pos = _field(event, "content_index") - if not isinstance(pos, int): - raise ValueError("WebSocket text event is missing content_index") if kind == "response.output_text.delta": item.text.setdefault(pos, []).append(value) elif kind == "response.output_text.done": diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 5cbf59e28c..c0a9158a6a 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -295,3 +295,51 @@ def script(socket: ServerConnection) -> None: assert received == original acc.add_event(await driver.call(lane, "recv")) assert acc.get_final_response().status == "incomplete" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("invalid_index", [{}, {"content_index": "wrong"}, {"content_index": None}]) +@pytest.mark.parametrize( + "fields", + [ + {"type": "response.output_text.delta", "delta": "bad"}, + {"type": "response.output_text.done", "text": "bad"}, + {"type": "response.content_part.added", "part": {"type": "output_text", "text": "bad"}}, + {"type": "response.content_part.done", "part": {"type": "output_text", "text": "bad"}}, + ], +) +async def test_invalid_text_position_does_not_replace_previous_item( + mode: str, fields: dict[str, object], invalid_index: dict[str, object] +) -> None: + initial = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "original", + "delta": "saved", + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(initial)) + socket.send(json.dumps({"output_index": 0, "item_id": "replacement", **fields, **invalid_index})) + socket.send(json.dumps({**initial, "delta": " after"})) + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + invalid = await driver.call(lane, "recv") + untouched = model_copy(invalid, deep=True) + with pytest.raises(ValueError, match="content_index"): + acc.add_event(invalid) + assert acc.snapshot() == prior + assert invalid == untouched + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot().output_text == "saved after" + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "completed" From a67024a2602b7432f590933f5a2a92cfb38b1c32 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 19:29:31 +0000 Subject: [PATCH 04/15] fix(responses): preserve sparse websocket output and document default lane filtering --- src/openai/lib/responses_websocket/README.md | 14 +- .../lib/responses_websocket/_accumulator.py | 4 +- .../responses/test_websocket_accumulator.py | 124 ++++++++++++++++++ 3 files changed, 138 insertions(+), 4 deletions(-) diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index f372be6402..b9e3024d7a 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -75,10 +75,14 @@ fed by the events you already receive: from openai.lib.responses_websocket import ResponsesWebSocketAccumulator accumulator = ResponsesWebSocketAccumulator() +expected_stream_id = "conversation" # Same ID passed to session.lane(); None for session.default. while True: event = await lane.recv() # Use lane.recv(timeout=...) in a sync session. - accumulator.add_event(event) # Original event fields, including unknown variants/fields, remain available. + # Inspect or log all raw events here before filtering for this response. + if getattr(event, "stream_id", None) != expected_stream_id: + continue + accumulator.add_event(event) if event.type == "response.output_text.delta": print(event.delta, end="", flush=True) # Provisional progress log. elif event.type == "response.output_item.done": @@ -90,8 +94,12 @@ while True: accumulator.reset() # Does not close the lane or connection. ``` -Use one helper per lane and reset it before the next turn. Snapshots are -immutable and contain selected text, function arguments, and custom tool input, +Use one helper per lane and reset it before the next turn. The default lane +also receives raw events for unregistered or detached lanes. Observe those +events before filtering, then feed only the response you are collecting. +The same check applies when consuming directly from a connection. + +Snapshots are immutable and contain selected text, function arguments, and custom tool input, grouped by output/content index and item ID. They never execute tools. Done events replace provisional fields. Full item replacements discard old fields; a changed nonempty item ID starts fresh at its index. A supplied response diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 12826ab76d..1a7ecc1e5f 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -232,6 +232,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: item.input = [value] def _add_item(self, index: int, source: object) -> None: + if source is None: + return item = _Output( item_id=_field(source, "id"), type=_field(source, "type"), @@ -240,7 +242,7 @@ def _add_item(self, index: int, source: object) -> None: ) if item.type == "message": for pos, part in enumerate(_field(source, "content") or []): - if _field(part, "type") == "output_text": + if _field(part, "type") == "output_text" and _field(part, "text") is not None: item.text[pos] = [_text_field(part, "text")] elif item.type == "function_call": value = _field(source, "arguments") diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index c0a9158a6a..e671596e36 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -343,3 +343,127 @@ def script(socket: ServerConnection) -> None: assert acc.snapshot().output_text == "saved after" acc.add_event(await driver.call(lane, "recv")) assert acc.get_final_response().status == "completed" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("item_field", [{}, {"item": None}]) +async def test_empty_output_item_keeps_previous_projection(mode: str, item_field: dict[str, object]) -> None: + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send( + json.dumps( + {"type": "response.function_call_arguments.delta", "output_index": 0, "item_id": "fc", "delta": "{"} + ) + ) + for phase in ("added", "done"): + socket.send(json.dumps({"type": "response.output_item." + phase, "output_index": 0, **item_field})) + socket.send( + json.dumps( + {"type": "response.function_call_arguments.delta", "output_index": 0, "item_id": "fc", "delta": "}"} + ) + ) + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + for _ in range(2): + received = await driver.call(lane, "recv") + before = model_copy(received, deep=True) + acc.add_event(received) + assert acc.snapshot() == prior + assert received == before + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot().output[0].arguments == "{}" + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().status == "completed" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("nullable", [{}, {"text": None}]) +async def test_nullable_finalized_text_does_not_block_exact_terminal(mode: str, nullable: dict[str, object]) -> None: + item: dict[str, object] = { + "id": "msg", + "type": "message", + "role": "assistant", + "content": [ + {"type": "output_text", "text": "hello", "annotations": []}, + {"type": "output_text", "annotations": [], **nullable}, + {"type": "output_text", "text": " world", "annotations": []}, + ], + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "item_id": "msg", + "output_index": 0, + "content_index": 0, + "delta": "old", + } + ) + ) + socket.send(json.dumps({"type": "response.output_item.done", "output_index": 0, "item": item})) + socket.send(json.dumps(response_event("incomplete", output=[item]))) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + for _ in range(2): + received = await driver.call(lane, "recv") + before = model_copy(received, deep=True) + acc.add_event(received) + assert acc.snapshot().output_text == "hello world" + assert received == before + if received.type == "response.incomplete": + assert acc.get_final_response().to_dict() == received.response.to_dict() + assert acc.get_final_response().output_text == "hello world" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("foreign_first", [True, False]) +async def test_default_lane_foreign_events_can_be_inspected_without_feeding_the_projection( + mode: str, foreign_first: bool +) -> None: + own = { + "type": "response.output_text.delta", + "item_id": "msg", + "output_index": 0, + "content_index": 0, + "delta": "default", + } + foreign = {**own, "stream_id": "unregistered", "delta": "foreign"} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + for frame in [foreign, own] if foreign_first else [own, foreign]: + socket.send(json.dumps(frame)) + 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": "test"}) + accumulator = ResponsesWebSocketAccumulator() + expected_stream_id = None + observed: list[str] = [] + for _ in range(3): + event = await driver.call(lane, "recv") + if event.type == "response.output_text.delta": + observed.append(event.delta) + if getattr(event, "stream_id", None) != expected_stream_id: + continue + accumulator.add_event(event) + assert set(observed) == {"default", "foreign"} + assert accumulator.snapshot().output_text == "default" + assert accumulator.get_final_response().status == "completed" From 3c86dea9cbe4e74d629f1e7dde345c9ad608eff1 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 20:00:05 +0000 Subject: [PATCH 05/15] fix(responses): validate WebSocket item IDs and project MCP arguments --- .../lib/responses_websocket/_accumulator.py | 15 ++++++---- .../responses/test_websocket_accumulator.py | 30 ++++++++++++++++++- 2 files changed, 38 insertions(+), 7 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 1a7ecc1e5f..24eb7f3339 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -49,7 +49,7 @@ class _Output: class ResponsesWebSocketAccumulator: - """Opt-in, caller-fed collection of text, function arguments and custom tool input. + """Opt-in, caller-fed collection of text, function/MCP arguments and custom tool input. Feed typed events from a connection or lane's recv(). Use one accumulator per lane and call reset() before another turn. The helper never reads, sends, @@ -130,6 +130,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.output_text.done", "response.function_call_arguments.delta", "response.function_call_arguments.done", + "response.mcp_call_arguments.delta", + "response.mcp_call_arguments.done", "response.custom_tool_call_input.delta", "response.custom_tool_call_input.done", "error", @@ -179,12 +181,13 @@ def add_event(self, event: ResponsesServerEvent) -> None: if kind in { "response.output_text.delta", "response.function_call_arguments.delta", + "response.mcp_call_arguments.delta", "response.custom_tool_call_input.delta", }: value = _text_field(event, "delta") elif kind == "response.output_text.done": value = _text_field(event, "text") - elif kind == "response.function_call_arguments.done": + elif kind in {"response.function_call_arguments.done", "response.mcp_call_arguments.done"}: value = _text_field(event, "arguments") elif kind == "response.custom_tool_call_input.done": value = _text_field(event, "input") @@ -200,8 +203,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.output_text.done", } and not isinstance(pos, int): raise ValueError("WebSocket text event is missing content_index") + item_id = _text_field(event, "item_id") self._bound, self._stream_id = True, stream_id - item_id = _field(event, "item_id") item = self._output.get(index) if item is None or (item_id and item.item_id and item_id != item.item_id): item = _Output(item_id=item_id) @@ -222,9 +225,9 @@ def add_event(self, event: ResponsesServerEvent) -> None: part = _field(event, "part") if _field(part, "type") == "output_text": item.text[pos] = [value] - elif kind == "response.function_call_arguments.delta": + elif kind in {"response.function_call_arguments.delta", "response.mcp_call_arguments.delta"}: item.arguments.append(value) - elif kind == "response.function_call_arguments.done": + elif kind in {"response.function_call_arguments.done", "response.mcp_call_arguments.done"}: item.arguments = [value] elif kind == "response.custom_tool_call_input.delta": item.input.append(value) @@ -244,7 +247,7 @@ def _add_item(self, index: int, source: object) -> None: for pos, part in enumerate(_field(source, "content") or []): if _field(part, "type") == "output_text" and _field(part, "text") is not None: item.text[pos] = [_text_field(part, "text")] - elif item.type == "function_call": + elif item.type in {"function_call", "mcp_call"}: value = _field(source, "arguments") item.arguments = ["" if value is None else _text_field(source, "arguments")] elif item.type == "custom_tool_call": diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index e671596e36..1ee8232bea 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -39,6 +39,19 @@ async def test_opt_in_accumulation_exact_terminals(mode: str, terminal: str, out "item_id": "msg", "text": "corrected", }, + {"type": "response.mcp_call_arguments.delta", "output_index": 6, "item_id": "mcp", "delta": '{"search":'}, + {"type": "response.mcp_call_arguments.done", "output_index": 6, "item_id": "mcp", "arguments": '{"search":1}'}, + { + "type": "response.output_item.done", + "output_index": 6, + "item": { + "type": "mcp_call", + "id": "mcp", + "arguments": '{"search":2}', + "server_label": "fixture", + "name": "search", + }, + }, {"type": "response.future", "unmodeled": {"keep": True}}, ] final = response_event(terminal) @@ -82,11 +95,16 @@ def script(socket: ServerConnection) -> None: before = event.to_dict() acc.add_event(event) assert event.to_dict() == before + if event.type == "response.mcp_call_arguments.delta": + assert acc.snapshot().output[-1].arguments == '{"search":' + elif event.type == "response.mcp_call_arguments.done": + assert acc.snapshot().output[-1].arguments == '{"search":1}' projected = acc.snapshot() assert projected.output_text == "corrected" assert projected.terminal_type is None assert projected.output[0].arguments == '{"fixed":true}' assert projected.output[1].input == "tool data" + assert projected.output[-1].arguments == '{"search":2}' assert saved.output_text == "pre" received = await driver.call(lane, "recv") original = received.response.to_dict() @@ -102,7 +120,14 @@ def script(socket: ServerConnection) -> None: # Existing get_final_response remains a terminal/final-item collector. # It must not silently gain delta reconstruction. legacy = await driver.call(lane, "get_final_response") - assert not legacy.output + if output == "empty": + assert not legacy.output + else: + # Only the full MCP item can be recovered by the existing + # final-item collector; all provisional text/tools stay absent. + assert legacy.output is not None + assert len(legacy.output) == 1 + assert legacy.output[0].type == "mcp_call" acc.reset() assert not acc.snapshot().output assert saved.output_text == "pre" @@ -262,6 +287,9 @@ def script(socket: ServerConnection) -> None: {"type": "response.output_text.delta", "content_index": 0, "delta": None}, {"type": "response.function_call_arguments.delta", "delta": {"unexpected": True}}, {"type": "response.custom_tool_call_input.done", "input": 22}, + {"type": "response.output_text.delta", "content_index": 0, "delta": "corrupt", "item_id": None}, + {"type": "response.function_call_arguments.delta", "delta": "corrupt", "item_id": 0}, + {"type": "response.custom_tool_call_input.done", "input": "corrupt", "item_id": {"bad": True}}, ], ) async def test_rejects_invalid_fields_without_poisoning_prior_snapshot(mode: str, invalid: dict[str, object]) -> None: From a4446ef16ebd1babe8434c71132ca7007b21e2b2 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 20:09:09 +0000 Subject: [PATCH 06/15] fix(responses): reject invalid indices and atomically replace lifecycle output --- .../lib/responses_websocket/_accumulator.py | 29 +++--- .../responses/test_websocket_accumulator.py | 89 ++++++++++++++++++- 2 files changed, 106 insertions(+), 12 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 24eb7f3339..e60cebeb85 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -157,22 +157,28 @@ def add_event(self, event: ResponsesServerEvent) -> None: response_id = _field(response, "id") if self._response_id is not None and response_id is not None and self._response_id != response_id: raise ValueError("Event belongs to another response") - self._response_id = response_id or self._response_id - self._bound, self._stream_id = True, stream_id output = _field(response, "output") if isinstance(output, list): - self._output.clear() - for index, item in enumerate(cast("list[object]", output)): - self._add_item(index, item) + replacement: dict[int, _Output] = {} + try: + for index, item in enumerate(cast("list[object]", output)): + self._add_item(replacement, index, item) + except ValueError as error: + if terminal: + self._error = error + raise + self._output = replacement + self._response_id = response_id or self._response_id + self._bound, self._stream_id = True, stream_id if terminal: self._final = model_copy(response, deep=True) self._terminal_type = kind return index = _field(event, "output_index") - if not isinstance(index, int): + if not isinstance(index, int) or isinstance(index, bool): raise ValueError("WebSocket output event is missing output_index") if kind in {"response.output_item.added", "response.output_item.done"}: - self._add_item(index, _field(event, "item")) + self._add_item(self._output, index, _field(event, "item")) self._bound, self._stream_id = True, stream_id return # Validate consumed values before replacing an item's retained state. @@ -201,7 +207,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.content_part.done", "response.output_text.delta", "response.output_text.done", - } and not isinstance(pos, int): + } and (not isinstance(pos, int) or isinstance(pos, bool)): raise ValueError("WebSocket text event is missing content_index") item_id = _text_field(event, "item_id") self._bound, self._stream_id = True, stream_id @@ -234,7 +240,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: elif kind == "response.custom_tool_call_input.done": item.input = [value] - def _add_item(self, index: int, source: object) -> None: + @staticmethod + def _add_item(output: dict[int, _Output], index: int, source: object) -> None: if source is None: return item = _Output( @@ -247,13 +254,13 @@ def _add_item(self, index: int, source: object) -> None: for pos, part in enumerate(_field(source, "content") or []): if _field(part, "type") == "output_text" and _field(part, "text") is not None: item.text[pos] = [_text_field(part, "text")] - elif item.type in {"function_call", "mcp_call"}: + elif item.type in {"function_call", "mcp_call", "mcp_approval_request"}: value = _field(source, "arguments") item.arguments = ["" if value is None else _text_field(source, "arguments")] elif item.type == "custom_tool_call": value = _field(source, "input") item.input = ["" if value is None else _text_field(source, "input")] - self._output[index] = item + output[index] = item def _text_field(value: object, name: str) -> str: diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 1ee8232bea..f49fb2a473 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -326,7 +326,9 @@ def script(socket: ServerConnection) -> None: @pytest.mark.parametrize("mode", ["sync", "async"]) -@pytest.mark.parametrize("invalid_index", [{}, {"content_index": "wrong"}, {"content_index": None}]) +@pytest.mark.parametrize( + "invalid_index", [{}, {"content_index": "wrong"}, {"content_index": None}, {"content_index": False}] +) @pytest.mark.parametrize( "fields", [ @@ -373,6 +375,91 @@ def script(socket: ServerConnection) -> None: assert acc.get_final_response().status == "completed" +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("phase", ["created", "in_progress", "completed", "failed", "incomplete"]) +async def test_invalid_lifecycle_replacement_is_atomic(mode: str, phase: str) -> None: + first = { + "type": "response.output_text.delta", + "item_id": "m", + "output_index": 0, + "content_index": 0, + "delta": "safe", + } + malformed = [ + { + "type": "mcp_approval_request", + "id": "approval", + "name": "search", + "server_label": "fixture", + "arguments": '{"key":1}', + }, + {"type": "function_call", "id": "f", "name": "bad", "arguments": {"bad": True}, "call_id": "c"}, + ] + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(first)) + socket.send(json.dumps(response_event(phase, output=malformed))) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + before = acc.snapshot() + received = await driver.call(lane, "recv") + untouched = model_copy(received, deep=True) + with pytest.raises(ValueError, match="arguments"): + acc.add_event(received) + assert acc.snapshot() == before + assert received == untouched + error = ValueError if phase in {"completed", "failed", "incomplete"} else RuntimeError + with pytest.raises(error): + acc.get_final_response() + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_mcp_approval_items_keep_pending_arguments_and_reject_bool_output_index(mode: str) -> None: + approval = { + "type": "mcp_approval_request", + "id": "approval", + "name": "search", + "server_label": "fixture", + "arguments": '{"key":1}', + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps({"type": "response.output_item.added", "output_index": 0, "item": approval})) + socket.send( + json.dumps( + { + "type": "response.mcp_call_arguments.delta", + "output_index": False, + "item_id": "approval", + "delta": "bad", + } + ) + ) + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + before = acc.snapshot() + assert before.output[0].arguments == '{"key":1}' + with pytest.raises(ValueError, match="output_index"): + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot() == before + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("item_field", [{}, {"item": None}]) async def test_empty_output_item_keeps_previous_projection(mode: str, item_field: dict[str, object]) -> None: From a86c8978d0b57574dc27ffbc9ee6a795fa0681c7 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 20:23:54 +0000 Subject: [PATCH 07/15] fix(responses): reject negative indexes in opt-in WebSocket accumulation --- .../lib/responses_websocket/_accumulator.py | 8 ++--- .../responses/test_websocket_accumulator.py | 30 +++++++++++-------- 2 files changed, 22 insertions(+), 16 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index e60cebeb85..472cbeb145 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -175,8 +175,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: self._terminal_type = kind return index = _field(event, "output_index") - if not isinstance(index, int) or isinstance(index, bool): - raise ValueError("WebSocket output event is missing output_index") + if not isinstance(index, int) or isinstance(index, bool) or index < 0: + raise ValueError("WebSocket output_index must be a non-negative integer") if kind in {"response.output_item.added", "response.output_item.done"}: self._add_item(self._output, index, _field(event, "item")) self._bound, self._stream_id = True, stream_id @@ -207,8 +207,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.content_part.done", "response.output_text.delta", "response.output_text.done", - } and (not isinstance(pos, int) or isinstance(pos, bool)): - raise ValueError("WebSocket text event is missing content_index") + } and (not isinstance(pos, int) or isinstance(pos, bool) or pos < 0): + raise ValueError("WebSocket content_index must be a non-negative integer") item_id = _text_field(event, "item_id") self._bound, self._stream_id = True, stream_id item = self._output.get(index) diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index f49fb2a473..50421c25ab 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -327,7 +327,8 @@ def script(socket: ServerConnection) -> None: @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize( - "invalid_index", [{}, {"content_index": "wrong"}, {"content_index": None}, {"content_index": False}] + "invalid_index", + [{}, {"content_index": "wrong"}, {"content_index": None}, {"content_index": False}, {"content_index": -1}], ) @pytest.mark.parametrize( "fields", @@ -420,7 +421,21 @@ def script(socket: ServerConnection) -> None: @pytest.mark.parametrize("mode", ["sync", "async"]) -async def test_mcp_approval_items_keep_pending_arguments_and_reject_bool_output_index(mode: str) -> None: +@pytest.mark.parametrize( + "invalid", + [ + {"type": "response.mcp_call_arguments.delta", "output_index": False, "item_id": "approval", "delta": "bad"}, + {"type": "response.mcp_call_arguments.delta", "output_index": -1, "item_id": "approval", "delta": "bad"}, + { + "type": "response.output_item.done", + "output_index": -1, + "item": {"type": "message", "id": "bad", "content": [{"type": "output_text", "text": "bad"}]}, + }, + ], +) +async def test_mcp_approval_items_keep_pending_arguments_and_reject_invalid_output_index( + mode: str, invalid: dict[str, object] +) -> None: approval = { "type": "mcp_approval_request", "id": "approval", @@ -432,16 +447,7 @@ async def test_mcp_approval_items_keep_pending_arguments_and_reject_bool_output_ def script(socket: ServerConnection) -> None: socket.recv(timeout=5) socket.send(json.dumps({"type": "response.output_item.added", "output_index": 0, "item": approval})) - socket.send( - json.dumps( - { - "type": "response.mcp_call_arguments.delta", - "output_index": False, - "item_id": "approval", - "delta": "bad", - } - ) - ) + socket.send(json.dumps(invalid)) socket.send(json.dumps(response_event("completed", output=None))) with script_server(script) as url: From 80973d71a5781dc746a79a6639f46f60086c6c98 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 20:37:07 +0000 Subject: [PATCH 08/15] fix(responses): keep sparse WebSocket index collection scalable --- .../lib/responses_websocket/_accumulator.py | 34 +++++++++++-------- .../responses/test_websocket_accumulator.py | 33 ++++++++++++++++++ 2 files changed, 53 insertions(+), 14 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 472cbeb145..cec474c5ce 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -45,7 +45,7 @@ class _Output: call_id: str | None = None arguments: list[str] = field(default_factory=list[str]) input: list[str] = field(default_factory=list[str]) - text: dict[int, list[str]] = field(default_factory=dict[int, list[str]]) + text: dict[str, list[str]] = field(default_factory=dict[str, list[str]]) class ResponsesWebSocketAccumulator: @@ -63,7 +63,9 @@ def __init__(self) -> None: self._response_id: str | None = None self._terminal_type: str | None = None self._bound = False - self._output: dict[int, _Output] = {} + # 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._final: Response | None = None self._error: Exception | None = None @@ -86,16 +88,19 @@ def snapshot(self) -> ResponsesWebSocketSnapshot: terminal_type=self._terminal_type, output=tuple( ResponsesWebSocketOutput( - output_index=index, + output_index=int(index, 16), item_id=item.item_id, type=item.type, name=item.name, call_id=item.call_id, arguments="".join(item.arguments), input="".join(item.input), - text=tuple((pos, "".join(parts)) for pos, parts in sorted(item.text.items())), + text=tuple( + (int(pos, 16), "".join(parts)) + for pos, parts in sorted(item.text.items(), key=lambda pair: int(pair[0], 16)) + ), ) - for index, item in sorted(self._output.items()) + for index, item in sorted(self._output.items(), key=lambda pair: int(pair[0], 16)) ), ) @@ -159,7 +164,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: raise ValueError("Event belongs to another response") output = _field(response, "output") if isinstance(output, list): - replacement: dict[int, _Output] = {} + replacement: dict[str, _Output] = {} try: for index, item in enumerate(cast("list[object]", output)): self._add_item(replacement, index, item) @@ -211,10 +216,11 @@ def add_event(self, event: ResponsesServerEvent) -> None: raise ValueError("WebSocket content_index must be a non-negative integer") item_id = _text_field(event, "item_id") self._bound, self._stream_id = True, stream_id - item = self._output.get(index) + key = hex(index) + item = self._output.get(key) if item is None or (item_id and item.item_id and item_id != item.item_id): item = _Output(item_id=item_id) - self._output[index] = item + self._output[key] = item if item_id: item.item_id = item_id if kind in { @@ -224,13 +230,13 @@ def add_event(self, event: ResponsesServerEvent) -> None: "response.output_text.done", }: if kind == "response.output_text.delta": - item.text.setdefault(pos, []).append(value) + item.text.setdefault(hex(pos), []).append(value) elif kind == "response.output_text.done": - item.text[pos] = [value] + item.text[hex(pos)] = [value] else: part = _field(event, "part") if _field(part, "type") == "output_text": - item.text[pos] = [value] + item.text[hex(pos)] = [value] 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"}: @@ -241,7 +247,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: item.input = [value] @staticmethod - def _add_item(output: dict[int, _Output], index: int, source: object) -> None: + def _add_item(output: dict[str, _Output], index: int, source: object) -> None: if source is None: return item = _Output( @@ -253,14 +259,14 @@ def _add_item(output: dict[int, _Output], index: int, source: object) -> None: if item.type == "message": for pos, part in enumerate(_field(source, "content") or []): if _field(part, "type") == "output_text" and _field(part, "text") is not None: - item.text[pos] = [_text_field(part, "text")] + item.text[hex(pos)] = [_text_field(part, "text")] elif item.type in {"function_call", "mcp_call", "mcp_approval_request"}: value = _field(source, "arguments") item.arguments = ["" if value is None else _text_field(source, "arguments")] elif item.type == "custom_tool_call": value = _field(source, "input") item.input = ["" if value is None else _text_field(source, "input")] - output[index] = item + output[hex(index)] = item def _text_field(value: object, name: str) -> str: diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 50421c25ab..854b915e77 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -1,5 +1,6 @@ from __future__ import annotations +import sys import json from typing import cast @@ -8,13 +9,45 @@ from openai import omit from openai._compat import model_copy +from openai._models import construct_type_unchecked from openai.types.responses import ResponseStreamEvent from openai.lib.responses_websocket import ResponsesWebSocketError, ResponsesWebSocketAccumulator from openai.lib.streaming.responses import ResponseStreamState +from openai.types.responses.responses_server_event import ResponsesServerEvent from .test_websocket_session import session_for, script_server, response_event +@pytest.mark.parametrize("field", ["output_index", "content_index"]) +def test_large_sparse_indices_keep_numeric_order_and_prior_snapshots(field: str) -> None: + acc = ResponsesWebSocketAccumulator() + size = 2048 + step = sys.hash_info.modulus + for i in reversed(range(size)): + acc.add_event( + construct_type_unchecked( + type_=ResponsesServerEvent, + value={ + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "msg", + "delta": str(i) + ",", + field: i * step, + }, + ) + ) + prior = acc.snapshot() + assert prior.output_text == "".join(str(i) + "," for i in range(size)) + if field == "output_index": + assert [item.output_index for item in prior.output] == [i * step for i in range(size)] + else: + assert [index for index, _ in prior.output[0].text] == [i * step for i in range(size)] + acc.reset() + assert not acc.snapshot().output + assert prior.output_text == "".join(str(i) + "," for i in range(size)) + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("terminal", ["completed", "failed", "incomplete"]) @pytest.mark.parametrize("output", ["missing", "null", "empty"]) From fd3319b8bdda644d4606e872c42b8509f3a68744 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 20:49:19 +0000 Subject: [PATCH 09/15] fix(responses): retain valid projections across item and part corrections --- src/openai/lib/responses_websocket/README.md | 16 +- .../lib/responses_websocket/_accumulator.py | 34 +++- .../responses/test_websocket_accumulator.py | 177 ++++++++++++++++-- 3 files changed, 202 insertions(+), 25 deletions(-) diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index b9e3024d7a..2134f7857d 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -86,9 +86,10 @@ while True: if event.type == "response.output_text.delta": print(event.delta, end="", flush=True) # Provisional progress log. elif event.type == "response.output_item.done": - snapshot = accumulator.snapshot() - print("\nCurrent projected text:", snapshot.output_text) + print("\nCompleted item:", event.item) # Replaces the provisional item. elif event.type in {"response.completed", "response.failed", "response.incomplete"}: + snapshot = accumulator.snapshot() # One full projection of all items. + print("\nProjected text:", snapshot.output_text) response = accumulator.get_final_response() break accumulator.reset() # Does not close the lane or connection. @@ -108,11 +109,12 @@ omitted or null output retains only the helper's earlier projections. `snapshot()` materializes the entire current projection and joins retained fragments. It is proportional to the accumulated output, so requesting it after -every small delta repeatedly rebuilds growing prefixes. Use the original event -for per-delta progress and request a full snapshot only when needed, such as -after an item is done. The progress log above is provisional; done events can -shorten or correct earlier text. A UI should replace its displayed projection -at those boundaries, not append the cumulative `snapshot.output_text`. +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 +user demand. The progress log above is provisional; done events can shorten or +correct earlier text. A UI can use the received item to replace the item at +`event.output_index`, and use the terminal snapshot for the whole projection. The helper's `get_final_response()` returns a copy of the exact received server response, including its original missing/null/empty output, on completed, diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index cec474c5ce..a0cbb4b813 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -40,6 +40,7 @@ def output_text(self) -> str: @dataclass class _Output: item_id: str | None + retired_ids: set[str] = field(default_factory=set[str]) type: str | None = None name: str | None = None call_id: str | None = None @@ -79,8 +80,8 @@ def reset(self) -> None: def snapshot(self) -> ResponsesWebSocketSnapshot: """Materialize the entire current immutable projection. - This joins retained fragments. Use original events for per-delta progress - and request snapshots intentionally (for example, on output_item.done). + This joins retained fragments. Use original events for delta and item + progress; read a full snapshot at a terminal or on explicit user demand. """ return ResponsesWebSocketSnapshot( stream_id=self._stream_id, @@ -163,6 +164,11 @@ def add_event(self, event: ResponsesServerEvent) -> None: if self._response_id is not None and response_id is not None and self._response_id != response_id: raise ValueError("Event belongs to another response") output = _field(response, "output") + if output is not None and not isinstance(output, list): + error = ValueError("WebSocket response output must be a list or null") + if terminal: + self._error = error + raise error if isinstance(output, list): replacement: dict[str, _Output] = {} try: @@ -218,9 +224,14 @@ def add_event(self, event: ResponsesServerEvent) -> None: self._bound, self._stream_id = True, stream_id key = hex(index) item = self._output.get(key) - if item is None or (item_id and item.item_id and item_id != item.item_id): + if item is not None and item_id in item.retired_ids: + return + if item is None: item = _Output(item_id=item_id) self._output[key] = item + elif item_id and item.item_id and item_id != item.item_id: + item = _Output(item_id=item_id, retired_ids=item.retired_ids | {item.item_id}) + self._output[key] = item if item_id: item.item_id = item_id if kind in { @@ -237,6 +248,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: part = _field(event, "part") if _field(part, "type") == "output_text": item.text[hex(pos)] = [value] + elif isinstance(_field(part, "type"), str): + item.text.pop(hex(pos), 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"}: @@ -250,9 +263,12 @@ def add_event(self, event: ResponsesServerEvent) -> None: def _add_item(output: dict[str, _Output], index: int, source: object) -> None: if source is None: return + item_id = _field(source, "id") + if item_id is not None: + item_id = _text_field(source, "id") item = _Output( - item_id=_field(source, "id"), - type=_field(source, "type"), + item_id=item_id, + type=_text_field(source, "type"), name=_field(source, "name"), call_id=_field(source, "call_id"), ) @@ -266,7 +282,13 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None: elif item.type == "custom_tool_call": value = _field(source, "input") item.input = ["" if value is None else _text_field(source, "input")] - output[hex(index)] = item + key = hex(index) + previous = output.get(key) + if previous is not None: + item.retired_ids = previous.retired_ids.copy() + if previous.item_id and item.item_id and previous.item_id != item.item_id: + item.retired_ids.add(previous.item_id) + output[key] = item def _text_field(value: object, name: str) -> str: diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 854b915e77..0cadd72753 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -9,11 +9,10 @@ from openai import omit from openai._compat import model_copy -from openai._models import construct_type_unchecked from openai.types.responses import ResponseStreamEvent from openai.lib.responses_websocket import ResponsesWebSocketError, ResponsesWebSocketAccumulator from openai.lib.streaming.responses import ResponseStreamState -from openai.types.responses.responses_server_event import ResponsesServerEvent +from openai.types.responses.responses_server_event import ResponseTextWsDelta from .test_websocket_session import session_for, script_server, response_event @@ -25,16 +24,14 @@ def test_large_sparse_indices_keep_numeric_order_and_prior_snapshots(field: str) step = sys.hash_info.modulus for i in reversed(range(size)): acc.add_event( - construct_type_unchecked( - type_=ResponsesServerEvent, - value={ - "type": "response.output_text.delta", - "output_index": 0, - "content_index": 0, - "item_id": "msg", - "delta": str(i) + ",", - field: i * step, - }, + ResponseTextWsDelta( + type="response.output_text.delta", + output_index=i * step if field == "output_index" else 0, + content_index=i * step if field == "content_index" else 0, + item_id="msg", + delta=str(i) + ",", + logprobs=[], + sequence_number=size - i, ) ) prior = acc.snapshot() @@ -537,6 +534,162 @@ def script(socket: ServerConnection) -> None: assert acc.get_final_response().status == "completed" +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("bad_item", [0, True, "not an item", ["bad"], {"missing": "type"}]) +async def test_invalid_item_shape_preserves_retained_projection(mode: str, bad_item: object) -> None: + delta = {"type": "response.function_call_arguments.delta", "output_index": 0, "item_id": "fc", "delta": "saved"} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(delta)) + for phase in ("added", "done"): + socket.send(json.dumps({"type": "response.output_item." + phase, "output_index": 0, "item": bad_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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + for _ in range(2): + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match="type"): + acc.add_event(received) + assert acc.snapshot() == prior + assert received == original + acc.add_event(await driver.call(lane, "recv")) + assert acc.snapshot().output[0].arguments == "saved" + assert acc.get_final_response().status == "completed" + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +async def test_retired_items_do_not_revert_full_or_incremental_replacements(mode: str) -> None: + delta = {"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "item_id": "a", "delta": "A"} + replacement = { + "type": "response.output_item.done", + "output_index": 0, + "item": {"type": "message", "id": "b", "role": "assistant", "content": [{"type": "output_text", "text": "B"}]}, + } + frames = [ + delta, + replacement, + {**delta, "delta": "stale A"}, + {**delta, "item_id": "c", "delta": "C"}, + {**delta, "item_id": "b", "delta": "stale B"}, + {**delta, "delta": "stale A"}, + response_event("completed", output=None), + ] + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + for expected in ("A", "B", "B", "C", "C", "C"): + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + acc.add_event(received) + assert acc.snapshot().output_text == expected + assert received == original + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.snapshot().output_text == "C" + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("phase", ["created", "in_progress", "completed", "failed", "incomplete"]) +@pytest.mark.parametrize("bad_output", [{}, 42, "bad"]) +async def test_nonlist_lifecycle_output_is_never_a_valid_final(mode: str, phase: str, bad_output: object) -> None: + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "saved", + } + ) + ) + socket.send(json.dumps(response_event(phase, output=bad_output))) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match="output"): + acc.add_event(received) + assert acc.snapshot() == prior + assert received == original + error = ValueError if phase in {"completed", "failed", "incomplete"} else RuntimeError + with pytest.raises(error): + acc.get_final_response() + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("part_phase", ["added", "done"]) +async def test_nontext_part_correction_clears_only_its_own_text_position(mode: str, part_phase: str) -> None: + delta = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "saved", + } + frames = [ + delta, + {**delta, "content_index": 1, "delta": "neighbor"}, + { + "type": "response.content_part." + part_phase, + "output_index": 0, + "content_index": 0, + "item_id": "a", + "part": {"type": "refusal", "refusal": "corrected"}, + }, + response_event("completed", output=None), + ] + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + for _ in range(2): + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + correction = await driver.call(lane, "recv") + original = model_copy(correction, deep=True) + acc.add_event(correction) + assert correction == original + assert prior.output_text == "savedneighbor" + assert acc.snapshot().output[0].text == ((1, "neighbor"),) + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.snapshot().output_text == "neighbor" + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("nullable", [{}, {"text": None}]) async def test_nullable_finalized_text_does_not_block_exact_terminal(mode: str, nullable: dict[str, object]) -> None: From 27bd5b4097b63f40589ccbc63b1ac1e613b0d085 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:03:02 +0000 Subject: [PATCH 10/15] fix(responses): validate WebSocket parts before replacing state and retain item IDs linearly --- .../lib/responses_websocket/_accumulator.py | 9 ++-- .../responses/test_websocket_accumulator.py | 54 +++++++++++++++++++ 2 files changed, 59 insertions(+), 4 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index a0cbb4b813..cffd1a71c1 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -210,7 +210,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: value = _text_field(event, "input") elif kind in {"response.content_part.added", "response.content_part.done"}: part = _field(event, "part") - if _field(part, "type") == "output_text": + if _text_field(part, "type") == "output_text": value = _text_field(part, "text") pos = _field(event, "content_index") if kind in { @@ -230,7 +230,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: item = _Output(item_id=item_id) self._output[key] = item elif item_id and item.item_id and item_id != item.item_id: - item = _Output(item_id=item_id, retired_ids=item.retired_ids | {item.item_id}) + item.retired_ids.add(item.item_id) + item = _Output(item_id=item_id, retired_ids=item.retired_ids) self._output[key] = item if item_id: item.item_id = item_id @@ -248,7 +249,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: part = _field(event, "part") if _field(part, "type") == "output_text": item.text[hex(pos)] = [value] - elif isinstance(_field(part, "type"), str): + else: item.text.pop(hex(pos), None) elif kind in {"response.function_call_arguments.delta", "response.mcp_call_arguments.delta"}: item.arguments.append(value) @@ -285,7 +286,7 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None: key = hex(index) previous = output.get(key) if previous is not None: - item.retired_ids = previous.retired_ids.copy() + item.retired_ids = previous.retired_ids if previous.item_id and item.item_id and previous.item_id != item.item_id: item.retired_ids.add(previous.item_id) output[key] = item diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 0cadd72753..48501e6b3d 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -690,6 +690,60 @@ def script(socket: ServerConnection) -> None: assert acc.get_final_response().to_dict() == terminal.response.to_dict() +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize( + "bad_part", + [{}, {"part": None}, {"part": "bad"}, {"part": 42}, {"part": []}, {"part": {}}, {"part": {"type": 7}}], +) +async def test_invalid_part_never_replaces_current_item(mode: str, bad_part: dict[str, object]) -> None: + delta = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "kept", + } + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(delta)) + for phase in ("added", "done"): + socket.send( + json.dumps( + { + "type": "response.content_part." + phase, + "output_index": 0, + "content_index": 0, + "item_id": "b", + **bad_part, + } + ) + ) + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + before = acc.snapshot() + for _ in range(2): + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match="type"): + acc.add_event(received) + assert acc.snapshot() == before + assert received == original + acc.add_event(await driver.call(lane, "recv")) + assert before.output_text == "kept" + assert acc.snapshot().output_text == "kept more" + terminal = await driver.call(lane, "recv") + acc.add_event(terminal) + assert acc.get_final_response().to_dict() == terminal.response.to_dict() + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("nullable", [{}, {"text": None}]) async def test_nullable_finalized_text_does_not_block_exact_terminal(mode: str, nullable: dict[str, object]) -> None: From 9075dab4236450f3d94e053eb8d5b988686452bf Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:15:19 +0000 Subject: [PATCH 11/15] fix: preserve response projection on malformed full items --- src/openai/lib/responses_websocket/README.md | 10 +- .../lib/responses_websocket/_accumulator.py | 33 ++++--- .../responses/test_websocket_accumulator.py | 93 +++++++++++++++++++ 3 files changed, 118 insertions(+), 18 deletions(-) diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index 2134f7857d..32d87db646 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -79,17 +79,17 @@ expected_stream_id = "conversation" # Same ID passed to session.lane(); None fo while True: event = await lane.recv() # Use lane.recv(timeout=...) in a sync session. # Original event fields, including unknown variants/fields, remain available. - # Inspect or log all raw events here before filtering for this response. + # Inspect events in memory before filtering. Never log raw events: they can + # contain customer data, tool arguments, credentials, or error details. if getattr(event, "stream_id", None) != expected_stream_id: continue accumulator.add_event(event) if event.type == "response.output_text.delta": - print(event.delta, end="", flush=True) # Provisional progress log. + pass # Update the application's UI with event.delta (provisional). elif event.type == "response.output_item.done": - print("\nCompleted item:", event.item) # Replaces the provisional item. + pass # Replace the provisional item at event.output_index with event.item. elif event.type in {"response.completed", "response.failed", "response.incomplete"}: snapshot = accumulator.snapshot() # One full projection of all items. - print("\nProjected text:", snapshot.output_text) response = accumulator.get_final_response() break accumulator.reset() # Does not close the lane or connection. @@ -112,7 +112,7 @@ fragments. It is proportional to the accumulated output, so requesting it 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 -user demand. The progress log above is provisional; done events can shorten or +user demand. A per-delta progress display is provisional; done events can shorten or correct earlier text. A UI can use the received item to replace the item at `event.output_index`, and use the terminal snapshot for the whole projection. diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index cffd1a71c1..c789313abb 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -264,25 +264,26 @@ def add_event(self, event: ResponsesServerEvent) -> None: def _add_item(output: dict[str, _Output], index: int, source: object) -> None: if source is None: return - item_id = _field(source, "id") - if item_id is not None: - item_id = _text_field(source, "id") item = _Output( - item_id=item_id, + item_id=_optional_text_field(source, "id"), type=_text_field(source, "type"), - name=_field(source, "name"), - call_id=_field(source, "call_id"), + name=_optional_text_field(source, "name"), + call_id=_optional_text_field(source, "call_id"), ) if item.type == "message": - for pos, part in enumerate(_field(source, "content") or []): - if _field(part, "type") == "output_text" and _field(part, "text") is not None: - item.text[hex(pos)] = [_text_field(part, "text")] + 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") + for pos, part in enumerate(cast("list[object]", content)): + 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] elif item.type in {"function_call", "mcp_call", "mcp_approval_request"}: - value = _field(source, "arguments") - item.arguments = ["" if value is None else _text_field(source, "arguments")] + item.arguments = [_optional_text_field(source, "arguments") or ""] elif item.type == "custom_tool_call": - value = _field(source, "input") - item.input = ["" if value is None else _text_field(source, "input")] + item.input = [_optional_text_field(source, "input") or ""] key = hex(index) previous = output.get(key) if previous is not None: @@ -297,3 +298,9 @@ def _text_field(value: object, name: str) -> str: if not isinstance(text, str): raise ValueError(f"WebSocket output {name} must be a string") return text + + +def _optional_text_field(value: object, name: str) -> str | None: + if _field(value, name) is None: + return None + return _text_field(value, name) diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 48501e6b3d..6559c7f151 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -450,6 +450,99 @@ def script(socket: ServerConnection) -> None: acc.get_final_response() +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize( + "invalid,field", + [ + pytest.param({"type": "message", "content": "bad"}, "content", id="string-content"), + pytest.param({"type": "message", "content": 42}, "content", id="numeric-content"), + pytest.param({"type": "message", "content": {}}, "content", id="object-content"), + pytest.param({"type": "message", "content": ["bad"]}, "type", id="invalid-part"), + pytest.param({"type": "message", "content": [{"type": 42}]}, "type", id="invalid-part-type"), + pytest.param({"type": "function_call", "name": {"bad": True}}, "name", id="function-name"), + pytest.param({"type": "function_call", "call_id": 42}, "call_id", id="function-call-id"), + pytest.param({"type": "mcp_call", "name": 42}, "name", id="mcp-name"), + pytest.param({"type": "mcp_approval_request", "call_id": {"bad": True}}, "call_id", id="mcp-call-id"), + pytest.param({"type": "custom_tool_call", "name": 42}, "name", id="custom-name"), + ], +) +async def test_invalid_projected_item_fields_do_not_retire_or_replace_current( + mode: str, invalid: dict[str, object], field: str +) -> None: + first = { + "type": "response.output_text.delta", + "item_id": "original", + "output_index": 0, + "content_index": 0, + "delta": "kept", + } + replacement = {"id": "replacement", **invalid} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(first)) + for phase in ("added", "done"): + socket.send(json.dumps({"type": "response.output_item." + phase, "output_index": 0, "item": replacement})) + socket.send(json.dumps({**first, "delta": " more"})) + # A lifecycle carrying an invalid later item cannot commit the earlier + # valid item or a final, and must retain its validation error. + socket.send(json.dumps(response_event("completed", output=[{"type": "message", "id": "valid"}, replacement]))) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + before = acc.snapshot() + for _ in range(2): + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match=field): + acc.add_event(received) + assert acc.snapshot() == before + assert received == original + acc.add_event(await driver.call(lane, "recv")) + assert before.output_text == "kept" + assert acc.snapshot().output_text == "kept more" + before = acc.snapshot() + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + with pytest.raises(ValueError, match=field) as failure: + acc.add_event(received) + assert acc.snapshot() == before + assert received == original + with pytest.raises(ValueError) as final: + acc.get_final_response() + assert final.value is failure.value + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("nullable", [{}, {"content": None, "name": None, "call_id": None}]) +async def test_nullable_message_content_and_projected_metadata(mode: str, nullable: dict[str, object]) -> None: + message: dict[str, object] = {"type": "message", "id": "message", **nullable} + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps({"type": "response.output_item.done", "output_index": 0, "item": message})) + socket.send(json.dumps(response_event("completed", output=[message]))) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + for _ in range(2): + received = await driver.call(lane, "recv") + untouched = model_copy(received, deep=True) + acc.add_event(received) + projection = acc.snapshot() + assert projection.output_text == "" + assert projection.output[0].name is None and projection.output[0].call_id is None + assert received == untouched + assert acc.get_final_response().to_dict() == received.response.to_dict() + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize( "invalid", From 3a3a1e21be1197868fc65b92eacba3b11799217e Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:16:47 +0000 Subject: [PATCH 12/15] test: narrow the nullable content terminal before asserting --- tests/lib/responses/test_websocket_accumulator.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 6559c7f151..de808d3c4e 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -540,7 +540,8 @@ def script(socket: ServerConnection) -> None: assert projection.output_text == "" assert projection.output[0].name is None and projection.output[0].call_id is None assert received == untouched - assert acc.get_final_response().to_dict() == received.response.to_dict() + if received.type == "response.completed": + assert acc.get_final_response().to_dict() == received.response.to_dict() @pytest.mark.parametrize("mode", ["sync", "async"]) From 8b38e7cc4ff848d53268c927be90435add9bc800 Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:26:56 +0000 Subject: [PATCH 13/15] fix: validate lifecycle identity and preserve current websocket items --- .../lib/responses_websocket/_accumulator.py | 9 +- .../responses/test_websocket_accumulator.py | 151 +++++++++++++++++- 2 files changed, 153 insertions(+), 7 deletions(-) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index c789313abb..05d68266aa 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -122,7 +122,7 @@ def add_event(self, event: ResponsesServerEvent) -> None: Original event fields remain accessible to the caller. """ kind = _field(event, "type") - if kind not in { + if not isinstance(kind, str) or kind not in { "response.created", "response.in_progress", "response.completed", @@ -161,6 +161,11 @@ def add_event(self, event: ResponsesServerEvent) -> None: self._error = error raise error response_id = _field(response, "id") + if response_id is not None and not isinstance(response_id, str): + error = ValueError("WebSocket response id must be a string or null") + if terminal: + self._error = error + raise error if self._response_id is not None and response_id is not None and self._response_id != response_id: raise ValueError("Event belongs to another response") output = _field(response, "output") @@ -287,6 +292,8 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None: key = hex(index) previous = output.get(key) if previous is not None: + if item.item_id and item.item_id in previous.retired_ids: + return item.retired_ids = previous.retired_ids if previous.item_id and item.item_id and previous.item_id != item.item_id: item.retired_ids.add(previous.item_id) diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index de808d3c4e..98c1314a4d 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -4,11 +4,12 @@ import json from typing import cast +import httpx2 import pytest from websockets.sync.server import ServerConnection -from openai import omit -from openai._compat import model_copy +from openai import OpenAI, AsyncOpenAI, omit +from openai._compat import PYDANTIC_V1, model_copy from openai.types.responses import ResponseStreamEvent from openai.lib.responses_websocket import ResponsesWebSocketError, ResponsesWebSocketAccumulator from openai.lib.streaming.responses import ResponseStreamState @@ -660,21 +661,31 @@ def script(socket: ServerConnection) -> None: @pytest.mark.parametrize("mode", ["sync", "async"]) -async def test_retired_items_do_not_revert_full_or_incremental_replacements(mode: str) -> None: +@pytest.mark.parametrize("terminal_output", ["null", "authoritative"]) +async def test_retired_items_do_not_revert_full_or_incremental_replacements(mode: str, terminal_output: str) -> None: delta = {"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "item_id": "a", "delta": "A"} replacement = { "type": "response.output_item.done", "output_index": 0, "item": {"type": "message", "id": "b", "role": "assistant", "content": [{"type": "output_text", "text": "B"}]}, } + stale_a = { + "type": "message", + "id": "a", + "role": "assistant", + "content": [{"type": "output_text", "text": "Final A"}], + } frames = [ delta, replacement, {**delta, "delta": "stale A"}, + {"type": "response.output_item.added", "output_index": 0, "item": stale_a}, + {"type": "response.output_item.done", "output_index": 0, "item": stale_a}, + {**delta, "item_id": "b", "delta": " continued"}, {**delta, "item_id": "c", "delta": "C"}, {**delta, "item_id": "b", "delta": "stale B"}, {**delta, "delta": "stale A"}, - response_event("completed", output=None), + response_event("completed", output=None if terminal_output == "null" else [stale_a]), ] def script(socket: ServerConnection) -> None: @@ -687,7 +698,7 @@ def script(socket: ServerConnection) -> None: lane = driver.session.default await driver.call(lane, "send", {"type": "response.create", "input": "test"}) acc = ResponsesWebSocketAccumulator() - for expected in ("A", "B", "B", "C", "C", "C"): + for expected in ("A", "B", "B", "B", "B", "B continued", "C", "C", "C"): received = await driver.call(lane, "recv") original = model_copy(received, deep=True) acc.add_event(received) @@ -695,10 +706,138 @@ def script(socket: ServerConnection) -> None: assert received == original terminal = await driver.call(lane, "recv") acc.add_event(terminal) - assert acc.snapshot().output_text == "C" + assert acc.snapshot().output_text == ("C" if terminal_output == "null" else "Final A") assert acc.get_final_response().to_dict() == terminal.response.to_dict() +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("phase", ["created", "in_progress", "completed", "failed", "incomplete"]) +@pytest.mark.parametrize("bad_id", [42, ["not", "an", "id"], {"value": "bad"}]) +async def test_invalid_response_id_cannot_commit_lifecycle_output(mode: str, phase: str, bad_id: object) -> None: + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "saved", + } + ) + ) + socket.send(json.dumps(response_event(phase, id=bad_id, output=[]))) + if phase in {"created", "in_progress"}: + socket.send(json.dumps(response_event("completed", id="valid", 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + prior = acc.snapshot() + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + if PYDANTIC_V1 and bad_id == 42: + # V1 normalizes the wire number before it reaches this helper. + # Keep accepting the resulting string, as on any other event. + assert received.response.id == "42" + acc.add_event(received) + assert acc.snapshot().response_id == "42" + assert received == original + if phase in {"completed", "failed", "incomplete"}: + assert acc.get_final_response().to_dict() == received.response.to_dict() + return + with pytest.raises(ValueError, match="id"): + acc.add_event(received) + assert acc.snapshot() == prior + assert received == original + if phase in {"created", "in_progress"}: + acc.add_event(await driver.call(lane, "recv")) + assert acc.get_final_response().id == "valid" + assert acc.snapshot().output_text == "saved" + else: + with pytest.raises(ValueError, match="id"): + acc.get_final_response() + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("id_field", [{}, {"id": None}]) +async def test_nullable_response_id_keeps_exact_final(mode: str, id_field: dict[str, object]) -> None: + final = response_event("completed", output=None) + del final["response"]["id"] + final["response"].update(id_field) + + def script(socket: ServerConnection) -> None: + socket.recv(timeout=5) + socket.send(json.dumps(response_event("created", id="prior"))) + socket.send(json.dumps(final)) + + 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": "test"}) + acc = ResponsesWebSocketAccumulator() + acc.add_event(await driver.call(lane, "recv")) + received = await driver.call(lane, "recv") + original = model_copy(received, deep=True) + acc.add_event(received) + assert acc.snapshot().response_id == "prior" + assert acc.get_final_response().to_dict() == received.response.to_dict() + assert received == original + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("unknown_type", [["response.future"], {"type": "response.future"}]) +async def test_raw_connection_unknown_unhashable_events_are_ignored(mode: str, unknown_type: object) -> None: + def script(socket: ServerConnection) -> None: + socket.send( + json.dumps( + { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "saved", + } + ) + ) + socket.send(json.dumps({"type": unknown_type, "stream_id": "unregistered"})) + socket.send(json.dumps(response_event("completed", output=None))) + + with script_server(script) as url: + # Test the public raw connection: session.recv() has its own handling + # before it passes any decoded event to this opt-in helper. + if mode == "sync": + with ( + OpenAI( + api_key="fake-accumulator-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client, + client.responses.connect() as connection, + ): + events = [connection.recv() for _ in range(3)] + else: + async with ( + AsyncOpenAI( + api_key="fake-accumulator-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client, + async_client.responses.connect() as async_connection, + ): + events = [await async_connection.recv() for _ in range(3)] + acc = ResponsesWebSocketAccumulator() + acc.add_event(events[0]) + prior = acc.snapshot() + unknown = model_copy(events[1], deep=True) + acc.add_event(events[1]) + assert acc.snapshot() == prior + assert events[1] == unknown + acc.add_event(events[2]) + assert acc.snapshot().output_text == "saved" + assert acc.get_final_response().status == "completed" + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("phase", ["created", "in_progress", "completed", "failed", "incomplete"]) @pytest.mark.parametrize("bad_output", [{}, 42, "bad"]) From d18e1e17056ae8de02802a9dd876bc3c7fd9ec8a Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:35:08 +0000 Subject: [PATCH 14/15] fix: reject invalid stream IDs before accumulating websocket events --- .../lib/responses_websocket/_accumulator.py | 2 + .../responses/test_websocket_accumulator.py | 69 +++++++++++++++++++ 2 files changed, 71 insertions(+) diff --git a/src/openai/lib/responses_websocket/_accumulator.py b/src/openai/lib/responses_websocket/_accumulator.py index 05d68266aa..c9cecd0963 100644 --- a/src/openai/lib/responses_websocket/_accumulator.py +++ b/src/openai/lib/responses_websocket/_accumulator.py @@ -144,6 +144,8 @@ def add_event(self, event: ResponsesServerEvent) -> None: }: return stream_id = _field(event, "stream_id") + 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: raise ValueError("Event belongs to another WebSocket lane") if self._terminal_type is not None or self._error is not None: diff --git a/tests/lib/responses/test_websocket_accumulator.py b/tests/lib/responses/test_websocket_accumulator.py index 98c1314a4d..b12a798c74 100644 --- a/tests/lib/responses/test_websocket_accumulator.py +++ b/tests/lib/responses/test_websocket_accumulator.py @@ -789,6 +789,75 @@ def script(socket: ServerConnection) -> None: assert received == original +@pytest.mark.parametrize("mode", ["sync", "async"]) +@pytest.mark.parametrize("bad_id", [42, ["not", "an", "id"], {"value": "bad"}]) +@pytest.mark.parametrize("first_kind", ["delta", "completed"]) +@pytest.mark.parametrize("next_lane", [None, "next"]) +async def test_raw_connection_invalid_stream_id_cannot_bind( + mode: str, bad_id: object, first_kind: str, next_lane: str | None +) -> None: + delta = { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "a", + "delta": "saved", + } + first = {**(delta if first_kind == "delta" else response_event("completed")), "stream_id": bad_id} + + def script(socket: ServerConnection) -> None: + for frame in ( + first, + response_event("created", next_lane), + {**delta, "stream_id": next_lane}, + response_event("completed", next_lane, output=None), + ): + socket.send(json.dumps(frame)) + + with script_server(script) as url: + if mode == "sync": + with ( + OpenAI( + api_key="fake-accumulator-key", base_url=url, http_client=httpx2.Client(trust_env=False) + ) as client, + client.responses.connect() as connection, + ): + events = [connection.recv() for _ in range(4)] + else: + async with ( + AsyncOpenAI( + api_key="fake-accumulator-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) + ) as async_client, + async_client.responses.connect() as async_connection, + ): + events = [await async_connection.recv() for _ in range(4)] + acc = ResponsesWebSocketAccumulator() + prior = acc.snapshot() + malformed = events[0] + original = model_copy(malformed, deep=True) + if isinstance(malformed.stream_id, str): + # Pydantic may normalize a number before this helper sees it. + assert bad_id == 42 and malformed.stream_id == "42" + acc.add_event(malformed) + assert acc.snapshot().stream_id == "42" + acc.reset() + else: + with pytest.raises(ValueError, match="stream_id"): + acc.add_event(malformed) + assert acc.snapshot() == prior + with pytest.raises(RuntimeError, match="No terminal"): + acc.get_final_response() + assert malformed == original + for event in events[1:]: + original = model_copy(event, deep=True) + acc.add_event(event) + assert event == original + assert prior.stream_id is None and not prior.output + assert acc.snapshot().stream_id == next_lane + assert acc.snapshot().output_text == "saved" + assert acc.get_final_response().id == f"resp_{next_lane or 'default'}" + + @pytest.mark.parametrize("mode", ["sync", "async"]) @pytest.mark.parametrize("unknown_type", [["response.future"], {"type": "response.future"}]) async def test_raw_connection_unknown_unhashable_events_are_ignored(mode: str, unknown_type: object) -> None: From db90e317271737e4425a0449d7fa0afc5a0fd5cc Mon Sep 17 00:00:00 2001 From: markstuart-oai Date: Sun, 27 Sep 2026 21:46:44 +0000 Subject: [PATCH 15/15] docs: guard malformed event types in websocket accumulation example --- src/openai/lib/responses_websocket/README.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/openai/lib/responses_websocket/README.md b/src/openai/lib/responses_websocket/README.md index 32d87db646..41c3dc7f71 100644 --- a/src/openai/lib/responses_websocket/README.md +++ b/src/openai/lib/responses_websocket/README.md @@ -84,11 +84,12 @@ while True: if getattr(event, "stream_id", None) != expected_stream_id: continue accumulator.add_event(event) + event_type = event.type if event.type == "response.output_text.delta": pass # Update the application's UI with event.delta (provisional). elif event.type == "response.output_item.done": pass # Replace the provisional item at event.output_index with event.item. - elif event.type in {"response.completed", "response.failed", "response.incomplete"}: + elif isinstance(event_type, str) and event_type in {"response.completed", "response.failed", "response.incomplete"}: snapshot = accumulator.snapshot() # One full projection of all items. response = accumulator.get_final_response() break