Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
16 commits
Select commit Hold shift + click to select a range
08f9e73
feat(responses): add opt-in websocket accumulation
markstuart-oai Sep 27, 2026
12ac7ed
fix(responses): reject malformed accumulation data before updating sn…
markstuart-oai Sep 27, 2026
a9a072e
fix(responses): retain valid projections and avoid per-delta snapshots
markstuart-oai Sep 27, 2026
a67024a
fix(responses): preserve sparse websocket output and document default…
markstuart-oai Sep 27, 2026
3c86dea
fix(responses): validate WebSocket item IDs and project MCP arguments
markstuart-oai Sep 27, 2026
a4446ef
fix(responses): reject invalid indices and atomically replace lifecyc…
markstuart-oai Sep 27, 2026
a86c897
fix(responses): reject negative indexes in opt-in WebSocket accumulation
markstuart-oai Sep 27, 2026
80973d7
fix(responses): keep sparse WebSocket index collection scalable
markstuart-oai Sep 27, 2026
fd3319b
fix(responses): retain valid projections across item and part correct…
markstuart-oai Sep 27, 2026
27bd5b4
fix(responses): validate WebSocket parts before replacing state and r…
markstuart-oai Sep 27, 2026
9075dab
fix: preserve response projection on malformed full items
markstuart-oai Sep 27, 2026
3a3a1e2
test: narrow the nullable content terminal before asserting
markstuart-oai Sep 27, 2026
8b38e7c
fix: validate lifecycle identity and preserve current websocket items
markstuart-oai Sep 27, 2026
d18e1e1
fix: reject invalid stream IDs before accumulating websocket events
markstuart-oai Sep 27, 2026
db90e31
docs: guard malformed event types in websocket accumulation example
markstuart-oai Sep 27, 2026
38f42e9
Merge branch 'main' into codex/sdk-1021-python-responses-accumulator
markstuart-oai Sep 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions src/openai/lib/responses_websocket/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,62 @@ 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()
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.
# Original event fields, including unknown variants/fields, remain available.
# 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)
event_type = event.type

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Guard missing event types in the example

Fresh evidence beyond the resolved unhashable-type case: when a default-lane or direct-connection consumer receives a malformed top-level JSON value or object without type, the loose parser can return a value with no .type attribute. add_event() safely ignores it, but this access then raises AttributeError and terminates the documented receive loop; retrieve the type with getattr(event, "type", None) and use that guarded value for all subsequent comparisons.

Useful? React with 👍 / 👎.

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 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
accumulator.reset() # Does not close the lane or connection.
```

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
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 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. 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.

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
Expand Down
5 changes: 5 additions & 0 deletions src/openai/lib/responses_websocket/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,3 +7,8 @@
ResponsesWebSocketBufferError as ResponsesWebSocketBufferError,
AsyncResponsesWebSocketSession as AsyncResponsesWebSocketSession,
)
from ._accumulator import (
ResponsesWebSocketOutput as ResponsesWebSocketOutput,
ResponsesWebSocketSnapshot as ResponsesWebSocketSnapshot,
ResponsesWebSocketAccumulator as ResponsesWebSocketAccumulator,
)
315 changes: 315 additions & 0 deletions src/openai/lib/responses_websocket/_accumulator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,315 @@
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
retired_ids: set[str] = field(default_factory=set[str])
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[str, list[str]] = field(default_factory=dict[str, list[str]])


class ResponsesWebSocketAccumulator:
"""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,
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
# 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

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:
"""Materialize the entire current immutable projection.

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,
response_id=self._response_id,
terminal_type=self._terminal_type,
output=tuple(
ResponsesWebSocketOutput(
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(
(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(), key=lambda pair: int(pair[0], 16))
),
)

def get_final_response(self) -> Response:
"""Return an independent copy of the received terminal response, never a partial success."""
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)

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 not isinstance(kind, str) or 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.mcp_call_arguments.delta",
"response.mcp_call_arguments.done",
"response.custom_tool_call_input.delta",
"response.custom_tool_call_input.done",
Comment thread
markstuart-oai marked this conversation as resolved.
"error",
}:
return
stream_id = _field(event, "stream_id")
Comment thread
markstuart-oai marked this conversation as resolved.
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")
Comment thread
markstuart-oai marked this conversation as resolved.
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":
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")
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")
Comment thread
markstuart-oai marked this conversation as resolved.
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")
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):
Comment thread
markstuart-oai marked this conversation as resolved.
replacement: dict[str, _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

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Preserve empty response IDs when binding state

When a compatible provider supplies an empty-string response ID, it passes the preceding string validation but this truthiness fallback stores None instead. Besides exposing the wrong snapshot().response_id, that leaves the accumulator unbound to the response identity, so a later lifecycle event with a different nonempty ID bypasses the mismatch check and can replace or finalize the existing projection. Assign the new ID whenever it is not None, rather than only when it is truthy.

Useful? React with 👍 / 👎.

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) 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"))
Comment thread
markstuart-oai marked this conversation as resolved.
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.
value = ""
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 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")
elif kind in {"response.content_part.added", "response.content_part.done"}:
part = _field(event, "part")
if _text_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) 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
key = hex(index)
item = self._output.get(key)
if item is not None and item_id in item.retired_ids:
return
if item is None:
item = _Output(item_id=item_id)
Comment thread
markstuart-oai marked this conversation as resolved.
self._output[key] = item
elif item_id and item.item_id and item_id != 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
if kind in {
"response.content_part.added",
"response.content_part.done",
"response.output_text.delta",
"response.output_text.done",
}:
if kind == "response.output_text.delta":
item.text.setdefault(hex(pos), []).append(value)
elif kind == "response.output_text.done":
item.text[hex(pos)] = [value]
else:
part = _field(event, "part")
if _field(part, "type") == "output_text":
Comment thread
markstuart-oai marked this conversation as resolved.
item.text[hex(pos)] = [value]
else:
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"}:
item.arguments = [value]
elif kind == "response.custom_tool_call_input.delta":
item.input.append(value)
elif kind == "response.custom_tool_call_input.done":
item.input = [value]

@staticmethod
def _add_item(output: dict[str, _Output], index: int, source: object) -> None:
if source is None:
return
item = _Output(
Comment thread
markstuart-oai marked this conversation as resolved.
item_id=_optional_text_field(source, "id"),
type=_text_field(source, "type"),
name=_optional_text_field(source, "name"),
call_id=_optional_text_field(source, "call_id"),
)
if item.type == "message":
content = _field(source, "content")
if content is not None:
if not isinstance(content, list):
raise ValueError("WebSocket message content must be a list or null")
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"}:
item.arguments = [_optional_text_field(source, "arguments") or ""]
elif item.type == "custom_tool_call":
item.input = [_optional_text_field(source, "input") or ""]
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)
output[key] = 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


def _optional_text_field(value: object, name: str) -> str | None:
if _field(value, name) is None:
return None
return _text_field(value, name)
Loading
Loading