Skip to content

Commit a380cf2

Browse files
feat(responses): preserve detailed WebSocket accumulator snapshots (openai#3981)
The opt-in Responses WebSocket accumulator currently retains selected text/tool fields but loses response metadata, streamed logprobs, citations and non-text parts. This adds `detailed_snapshot()`, an independent mutable view of observed fields and sparse indexed output/content/annotation rows. Existing immutable, hashable `snapshot()`, raw events, and the exact server terminal response keep their contracts. Done events replace supplied logprobs and whole part/item/response events replace the relevant projections. Missing and null are preserved, and unknown parts or tool fields remain provisional data instead of fabricated validated final models. Annotations only enrich a matching item on the bound lane; they cannot bind a turn or retire an unrelated item. The helper still never reads, sends, runs tools or changes sockets. Handwritten helper, docs and tests only. Validation: - Before change: 12 new real-wire tests failed. Review regressions reproduced: 8 failures (foreign or omitted annotation, null/missing content). - Current accumulator/session suite: 438 passed on each of Pydantic 1 and 2, sync and async local WebSockets. - SDK mypy: 1,865 files passed; Ruff and formatting pass. - Base's verified custom-code check: unchanged 7,826 / 10,000; no schema or generator change.
1 parent 80e9686 commit a380cf2

3 files changed

Lines changed: 525 additions & 9 deletions

File tree

‎src/openai/lib/responses_websocket/README.md‎

Lines changed: 40 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -108,8 +108,47 @@ a changed nonempty item ID starts fresh at its index. A supplied response
108108
output list overrides the projected items, including an explicit empty list;
109109
omitted or null output retains only the helper's earlier projections.
110110

111+
`detailed_snapshot()` returns a separate mutable view when you need the fields
112+
beyond selected text and tool inputs. It contains `stream_id`, `response_id`,
113+
`terminal_type`, `response` and `output`. `response` is the last observed
114+
lifecycle response metadata (excluding `output`), or `None` if none arrived.
115+
`output` is a list of `{"output_index": index, "item": metadata, "content": rows}`.
116+
Each content row is `{"content_index": index, "part": observed_fields}`. The
117+
part's `annotations`, when present as a list, use
118+
`{"annotation_index": index, "annotation": observed_fields}` rows. Indices may
119+
be sparse; list position is not the API index.
120+
For a message that has no projected content, the row's `content` is omitted,
121+
null or empty according to what was actually received.
122+
123+
For example, after collecting an item or terminal as above:
124+
125+
```python
126+
details = accumulator.detailed_snapshot()
127+
for output in details["output"]:
128+
for content in output.get("content") or []:
129+
part = content["part"]
130+
# Fields exist only if received: a WS delta may have no part/item type.
131+
text = part.get("text")
132+
citations = part.get("annotations")
133+
token_scores = part.get("logprobs")
134+
```
135+
136+
Logprobs accumulate with text deltas and a supplied `output_text.done.logprobs`
137+
replaces them, including empty or null values. Content/item/lifecycle replacements
138+
also replace their corresponding annotations and other metadata; text-only done
139+
events do not erase citations. Refusal, tool/MCP and unknown item/part fields are
140+
retained as observed, with unset and null distinct. Unknown standalone events still
141+
pass through unchanged and are not accumulated. Annotation events enrich matching
142+
known items on the accumulator's lane; annotations received before any matching
143+
item/text remain available in the original event and do not start or replace a turn.
144+
A partial field or unknown type
145+
is provisional, not a fabricated validated response or a successful tool result.
146+
You may mutate this returned view without changing the accumulator, events, or
147+
earlier snapshots. The original `snapshot()` remains immutable and hashable.
148+
111149
`snapshot()` materializes the entire current projection and joins retained
112-
fragments. It is proportional to the accumulated output, so requesting it after
150+
fragments. `detailed_snapshot()` also materializes and copies its whole projection.
151+
Both are proportional to the accumulated output, so requesting either after
113152
every delta or completed item repeatedly rebuilds growing prefixes. Use the
114153
original event for progress, including the final item itself on
115154
`response.output_item.done`. Read the full snapshot at a terminal or on explicit

‎src/openai/lib/responses_websocket/_accumulator.py‎

Lines changed: 138 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
from __future__ import annotations
22

3-
from typing import cast
3+
from copy import deepcopy
4+
from typing import Any, cast
45
from dataclasses import field, dataclass
56

67
from ._session import ResponsesWebSocketError, _field
7-
from ..._compat import model_copy
8+
from ..._compat import PYDANTIC_V1, model_copy
9+
from ..._models import BaseModel
810
from ...types.responses import Response
911
from ...types.responses.responses_server_event import ResponsesServerEvent
1012

@@ -47,6 +49,9 @@ class _Output:
4749
arguments: list[str] = field(default_factory=list[str])
4850
input: list[str] = field(default_factory=list[str])
4951
text: dict[str, list[str]] = field(default_factory=dict[str, list[str]])
52+
data: dict[str, object] = field(default_factory=dict[str, object])
53+
content: dict[str, dict[str, object]] = field(default_factory=dict[str, dict[str, object]])
54+
annotations: dict[str, dict[str, object]] = field(default_factory=dict[str, dict[str, object]])
5055

5156

5257
class ResponsesWebSocketAccumulator:
@@ -67,6 +72,7 @@ def __init__(self) -> None:
6772
# String keys get Python's randomized hash. Hex preserves arbitrary-size
6873
# non-negative indices without decimal string conversion limits.
6974
self._output: dict[str, _Output] = {}
75+
self._response: dict[str, object] | None = None
7076
self._final: Response | None = None
7177
self._error: Exception | None = None
7278

@@ -75,6 +81,7 @@ def reset(self) -> None:
7581
self._stream_id = self._response_id = self._terminal_type = None
7682
self._bound = False
7783
self._output.clear()
84+
self._response = None
7885
self._final = self._error = None
7986

8087
def snapshot(self) -> ResponsesWebSocketSnapshot:
@@ -105,6 +112,61 @@ def snapshot(self) -> ResponsesWebSocketSnapshot:
105112
),
106113
)
107114

115+
def detailed_snapshot(self) -> dict[str, Any]:
116+
"""Return an independent mutable projection of observed response, item and part data.
117+
118+
Output and content are lists of indexed rows, even when wire indices are
119+
sparse. Part annotations use the same indexed-row form. Unset and null
120+
fields stay distinct; missing scaffolding never invents a Response or
121+
an item/part type. Cost is proportional to the full accumulated data.
122+
"""
123+
output: list[dict[str, object]] = []
124+
for index, item in sorted(self._output.items(), key=lambda pair: int(pair[0], 16)):
125+
data = deepcopy(item.data)
126+
if item.item_id:
127+
data["id"] = item.item_id
128+
if item.arguments:
129+
data["arguments"] = "".join(item.arguments)
130+
if item.input:
131+
data["input"] = "".join(item.input)
132+
content: list[dict[str, object]] = []
133+
for pos in sorted(
134+
item.content.keys() | item.text.keys() | item.annotations.keys(), key=lambda k: int(k, 16)
135+
):
136+
part = deepcopy(item.content.get(pos, {}))
137+
if pos in item.text:
138+
part["text"] = "".join(item.text[pos])
139+
original = part.get("annotations")
140+
if isinstance(original, list) or pos in item.annotations:
141+
annotations = (
142+
{hex(i): value for i, value in enumerate(cast("list[object]", original))}
143+
if isinstance(original, list)
144+
else {}
145+
)
146+
annotations.update(deepcopy(item.annotations.get(pos, {})))
147+
part["annotations"] = [
148+
{"annotation_index": int(i, 16), "annotation": annotation}
149+
for i, annotation in sorted(annotations.items(), key=lambda pair: int(pair[0], 16))
150+
]
151+
content.append({"content_index": int(pos, 16), "part": part})
152+
row: dict[str, object] = {"output_index": int(index, 16), "item": data}
153+
if item.type == "message":
154+
# Message content is extracted into indexed rows. Retain the
155+
# presence marker until a part/delta supplies projected content.
156+
if "content" in data or content:
157+
original_content = data.pop("content", None)
158+
row["content"] = content if content or isinstance(original_content, list) else None
159+
else:
160+
row["content"] = content
161+
output.append(row)
162+
return {
163+
"stream_id": self._stream_id,
164+
"response_id": self._response_id,
165+
"terminal_type": self._terminal_type,
166+
"response": deepcopy(self._response),
167+
"output": output,
168+
}
169+
108170
def get_final_response(self) -> Response:
109171
"""Return an independent copy of the received terminal response, never a partial success."""
110172
error = self._error
@@ -134,6 +196,7 @@ def add_event(self, event: ResponsesServerEvent) -> None:
134196
"response.content_part.done",
135197
"response.output_text.delta",
136198
"response.output_text.done",
199+
"response.output_text.annotation.added",
137200
"response.function_call_arguments.delta",
138201
"response.function_call_arguments.done",
139202
"response.mcp_call_arguments.delta",
@@ -144,6 +207,31 @@ def add_event(self, event: ResponsesServerEvent) -> None:
144207
}:
145208
return
146209
stream_id = _field(event, "stream_id")
210+
# Annotations were historically ignored. They may enrich the matching
211+
# known item, but never bind a lane/turn, replace or retire an item, or
212+
# change the errors observed by existing accumulator callers.
213+
if kind == "response.output_text.annotation.added":
214+
output_pos = _field(event, "output_index")
215+
content_pos = _field(event, "content_index")
216+
annotation_pos = _field(event, "annotation_index")
217+
if (
218+
self._bound
219+
and self._stream_id == stream_id
220+
and self._terminal_type is None
221+
and self._error is None
222+
and all(
223+
isinstance(value, int) and not isinstance(value, bool) and value >= 0
224+
for value in (output_pos, content_pos, annotation_pos)
225+
)
226+
):
227+
existing = self._output.get(hex(output_pos))
228+
if existing is not None and existing.item_id == _field(event, "item_id"):
229+
value_data = _data(event, include={"annotation"})
230+
if "annotation" in value_data:
231+
existing.annotations.setdefault(hex(content_pos), {})[hex(annotation_pos)] = value_data[
232+
"annotation"
233+
]
234+
return
147235
if stream_id is not None and not isinstance(stream_id, str):
148236
raise ValueError("WebSocket stream_id must be a string or null")
149237
if self._bound and self._stream_id != stream_id:
@@ -186,6 +274,7 @@ def add_event(self, event: ResponsesServerEvent) -> None:
186274
self._error = error
187275
raise
188276
self._output = replacement
277+
self._response = _data(response, exclude={"output"})
189278
self._response_id = response_id or self._response_id
190279
self._bound, self._stream_id = True, stream_id
191280
if terminal:
@@ -248,16 +337,28 @@ def add_event(self, event: ResponsesServerEvent) -> None:
248337
"response.output_text.delta",
249338
"response.output_text.done",
250339
}:
340+
position = hex(pos)
251341
if kind == "response.output_text.delta":
252-
item.text.setdefault(hex(pos), []).append(value)
342+
item.text.setdefault(position, []).append(value)
343+
data = item.content.setdefault(position, {})
344+
prob_data = _data(event, include={"logprobs"})
345+
previous = data.get("logprobs")
346+
incoming = prob_data.get("logprobs")
347+
if isinstance(previous, list) and isinstance(incoming, list):
348+
cast("list[object]", previous).extend(cast("list[object]", incoming))
349+
else:
350+
data.update(prob_data)
253351
elif kind == "response.output_text.done":
254-
item.text[hex(pos)] = [value]
352+
item.text[position] = [value]
353+
item.content.setdefault(position, {}).update(_data(event, include={"logprobs"}))
255354
else:
256355
part = _field(event, "part")
356+
item.content[position] = _data(part)
357+
item.annotations.pop(position, None)
257358
if _field(part, "type") == "output_text":
258-
item.text[hex(pos)] = [value]
359+
item.text[position] = [value]
259360
else:
260-
item.text.pop(hex(pos), None)
361+
item.text.pop(position, None)
261362
elif kind in {"response.function_call_arguments.delta", "response.mcp_call_arguments.delta"}:
262363
item.arguments.append(value)
263364
elif kind in {"response.function_call_arguments.done", "response.mcp_call_arguments.done"}:
@@ -276,21 +377,32 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None:
276377
type=_text_field(source, "type"),
277378
name=_optional_text_field(source, "name"),
278379
call_id=_optional_text_field(source, "call_id"),
380+
data=_data(source, exclude={"content"} if _field(source, "type") == "message" else None),
279381
)
280382
if item.type == "message":
281383
content = _field(source, "content")
282384
if content is not None:
283385
if not isinstance(content, list):
284386
raise ValueError("WebSocket message content must be a list or null")
387+
# Marker only; keep actual parts indexed once, never copied twice.
388+
item.data["content"] = []
285389
for pos, part in enumerate(cast("list[object]", content)):
390+
if part is not None:
391+
item.content[hex(pos)] = _data(part)
286392
if part is not None and _text_field(part, "type") == "output_text":
287393
text = _optional_text_field(part, "text")
288394
if text is not None:
289395
item.text[hex(pos)] = [text]
396+
else:
397+
item.data.update(_data(source, include={"content"}))
290398
elif item.type in {"function_call", "mcp_call", "mcp_approval_request"}:
291-
item.arguments = [_optional_text_field(source, "arguments") or ""]
399+
arguments = _optional_text_field(source, "arguments")
400+
if arguments is not None:
401+
item.arguments = [arguments]
292402
elif item.type == "custom_tool_call":
293-
item.input = [_optional_text_field(source, "input") or ""]
403+
input_text = _optional_text_field(source, "input")
404+
if input_text is not None:
405+
item.input = [input_text]
294406
key = hex(index)
295407
previous = output.get(key)
296408
if previous is not None:
@@ -302,6 +414,24 @@ def _add_item(output: dict[str, _Output], index: int, source: object) -> None:
302414
output[key] = item
303415

304416

417+
def _data(source: object, *, include: set[str] | None = None, exclude: set[str] | None = None) -> dict[str, object]:
418+
if isinstance(source, BaseModel):
419+
return deepcopy(
420+
source.model_dump(
421+
mode="python", by_alias=True, exclude_unset=True, include=include, exclude=exclude, warnings=PYDANTIC_V1
422+
)
423+
)
424+
if isinstance(source, dict):
425+
return deepcopy(
426+
{
427+
key: value
428+
for key, value in cast("dict[str, object]", source).items()
429+
if (include is None or key in include) and (exclude is None or key not in exclude)
430+
}
431+
)
432+
return {}
433+
434+
305435
def _text_field(value: object, name: str) -> str:
306436
text = _field(value, name)
307437
if not isinstance(text, str):

0 commit comments

Comments
 (0)