11from __future__ import annotations
22
3- from typing import cast
3+ from copy import deepcopy
4+ from typing import Any , cast
45from dataclasses import field , dataclass
56
67from ._session import ResponsesWebSocketError , _field
7- from ..._compat import model_copy
8+ from ..._compat import PYDANTIC_V1 , model_copy
9+ from ..._models import BaseModel
810from ...types .responses import Response
911from ...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
5257class 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+
305435def _text_field (value : object , name : str ) -> str :
306436 text = _field (value , name )
307437 if not isinstance (text , str ):
0 commit comments