1313# limitations under the License.
1414
1515from abc import ABC , abstractmethod
16- from typing import Any , AsyncIterator , cast
16+ from typing import Any , AsyncIterator
1717import time
1818
1919from .types import ContentItem , FinishReason , UniConfig , UniEvent , UniMessage , UsageMetadata
@@ -123,15 +123,13 @@ def concat_uni_events_to_uni_message(self, events: list[UniEvent]) -> UniMessage
123123 finish_reason = event .get ("finish_reason" ) # finish_reason is taken from the last event
124124 created_at = event .get ("created_at" ) # created_at is taken from the last event
125125
126- result : UniMessage = {
126+ return {
127127 "role" : "assistant" ,
128128 "content_items" : content_items ,
129129 "usage_metadata" : usage_metadata ,
130130 "finish_reason" : finish_reason ,
131+ "created_at" : created_at ,
131132 }
132- if created_at is not None :
133- result ["created_at" ] = created_at
134- return result
135133
136134 @abstractmethod
137135 async def _streaming_response_internal (
@@ -173,14 +171,29 @@ async def streaming_response(
173171 Yields:
174172 Universal events from the streaming response
175173 """
174+ # Stamp any messages that don't yet have a created_at timestamp
175+ for msg in messages :
176+ if "created_at" not in msg :
177+ msg ["created_at" ] = int (time .time () * 1000 )
178+
176179 last_event : UniEvent | None = None
180+ events = []
177181 async for event in self ._streaming_response_internal (messages , config ):
178182 event ["created_at" ] = int (time .time () * 1000 )
179183 last_event = event
184+ events .append (event )
180185 yield event
181186
182187 self ._validate_last_event (last_event )
183188
189+ # Save history to file if trace_id is specified
190+ if config .get ("trace_id" ) and events :
191+ from .integration .tracer import Tracer
192+
193+ assistant_message = self .concat_uni_events_to_uni_message (events )
194+ tracer = Tracer ()
195+ tracer .save_history (self ._model , messages + [assistant_message ], config ["trace_id" ], config )
196+
184197 async def streaming_response_stateful (
185198 self ,
186199 message : UniMessage ,
@@ -200,10 +213,6 @@ async def streaming_response_stateful(
200213 Yields:
201214 Universal events from the streaming response
202215 """
203- # Stamp input message with current time if not provided
204- if "created_at" not in message :
205- message = cast (UniMessage , {** message , "created_at" : int (time .time () * 1000 )})
206-
207216 # Build a temporary messages list for inference without mutating history yet
208217 temp_messages = self ._history + [message ]
209218
@@ -214,18 +223,12 @@ async def streaming_response_stateful(
214223 yield event
215224
216225 # Only update history after successful inference
226+ # temp_messages[-1] is the user message, now stamped with created_at by streaming_response
217227 if events :
218228 assistant_message = self .concat_uni_events_to_uni_message (events )
219- self ._history .append (message )
229+ self ._history .append (temp_messages [ - 1 ] )
220230 self ._history .append (assistant_message )
221231
222- # Save history to file if trace_id is specified
223- if config .get ("trace_id" ):
224- from .integration .tracer import Tracer
225-
226- tracer = Tracer ()
227- tracer .save_history (self ._model , self ._history , config ["trace_id" ], config )
228-
229232 @staticmethod
230233 def _validate_last_event (last_event : UniEvent | None ) -> None :
231234 """Validate that the last event has usage_metadata and finish_reason.
0 commit comments