11"""
2- Chat engine — send messages and yield complete events via async generator.
2+ Chat engine — send messages and yield events via async generator.
33
4- Stream buffering:
5- - Tier 1: session:text_part buffered until final: true
6- - Tier 2: Non-streaming events yield immediately
7- - Tier 3: session:work_log_part debounced (3s silence)
4+ All events are dispatched immediately as they arrive from the server.
85"""
96
107import asyncio
11- import time
8+ from datetime import datetime , timedelta , timezone
129from typing import Any , AsyncGenerator , Callable , Coroutine , Optional
1310
1411from pine_assistant .models .events import C2SEvent , S2CEvent
1512from pine_assistant .transport .socketio import SocketIOManager
1613
1714TERMINAL_STATES = {"task_finished" , "task_cancelled" , "task_stale" }
1815DEFAULT_IDLE_TIMEOUT_S = 120.0
16+ DEFAULT_RESPONSE_IDLE_TIMEOUT_S = 2.0
1917
20- # Events that are buffered/debounced — NOT dispatched immediately
21- BUFFERED_EVENTS = {S2CEvent .SESSION_TEXT_PART , S2CEvent .SESSION_WORK_LOG_PART }
22-
23- # Substantive response events — track for waiting_input termination gating
2418SUBSTANTIVE_EVENTS = {
25- S2CEvent .SESSION_TEXT , S2CEvent .SESSION_FORM_TO_USER ,
19+ S2CEvent .SESSION_TEXT , S2CEvent .SESSION_TEXT_PART ,
20+ S2CEvent .SESSION_FORM_TO_USER ,
2621 S2CEvent .SESSION_ASK_FOR_LOCATION , S2CEvent .SESSION_TASK_READY ,
2722 S2CEvent .SESSION_TASK_FINISHED , S2CEvent .SESSION_INTERACTIVE_AUTH_CONFIRMATION ,
2823 S2CEvent .SESSION_THREE_WAY_CALL , S2CEvent .SESSION_REWARD ,
@@ -44,33 +39,18 @@ def __repr__(self) -> str:
4439 return f"ChatEvent(type={ self .type !r} , session_id={ self .session_id !r} )"
4540
4641
47- class TextPartBuffer :
48- """Buffer text_part chunks per message_id, flush on final: true."""
49-
50- def __init__ (self ) -> None :
51- self ._parts : dict [str , list [str ]] = {}
52-
53- def collect (self , message_id : str , content : str , final : bool ) -> Optional [str ]:
54- if message_id not in self ._parts :
55- self ._parts [message_id ] = []
56- if content :
57- self ._parts [message_id ].append (content )
58- if final :
59- merged = "" .join (self ._parts .pop (message_id , []))
60- return merged
61- return None
62-
63-
6442class ChatEngine :
6543 def __init__ (
6644 self ,
6745 sio : SocketIOManager ,
6846 check_session_state : Optional [Callable [[str ], Coroutine [Any , Any , dict [str , Any ]]]] = None ,
6947 idle_timeout_s : float = DEFAULT_IDLE_TIMEOUT_S ,
48+ response_idle_timeout_s : float = DEFAULT_RESPONSE_IDLE_TIMEOUT_S ,
7049 ):
7150 self ._sio = sio
7251 self ._check_session_state = check_session_state
7352 self ._idle_timeout_s = idle_timeout_s
53+ self ._response_idle_timeout_s = response_idle_timeout_s
7454
7555 async def join_session (self , session_id : str ) -> dict [str , Any ]:
7656 """Join a session room — spec 5.1.1.
@@ -117,14 +97,32 @@ async def chat(
11797 """Send a message and yield events with stream buffering.
11898 Production handler reads payload.data as {content, attachments, ...}.
11999 """
100+ cutoff = datetime .now (timezone .utc ) - timedelta (seconds = 5 )
120101 self ._sio .emit (
121102 C2SEvent .SESSION_MESSAGE ,
122103 self ._build_message_data (content , attachments , referenced_sessions , action ),
123104 session_id ,
124105 )
125- async for event in self ._listen (session_id ):
106+ async for event in self ._listen (session_id , _skip_state_precheck = True ):
107+ if self ._is_stale_event (event , cutoff ):
108+ continue
126109 yield event
127110
111+ @staticmethod
112+ def _is_stale_event (event : "ChatEvent" , cutoff : datetime ) -> bool :
113+ """Return True if the event's metadata timestamp predates the cutoff."""
114+ meta = event .metadata
115+ if not isinstance (meta , dict ):
116+ return False
117+ ts_str = meta .get ("timestamp" )
118+ if not ts_str :
119+ return False
120+ try :
121+ event_ts = datetime .fromisoformat (ts_str .replace ("Z" , "+00:00" ))
122+ return event_ts < cutoff
123+ except (ValueError , TypeError ):
124+ return False
125+
128126 def send_message (
129127 self ,
130128 session_id : str ,
@@ -141,10 +139,11 @@ def send_message(
141139 session_id ,
142140 )
143141
144- async def _listen (self , session_id : str ) -> AsyncGenerator [ChatEvent , None ]:
145- """Listen for events with stream buffering."""
146- # Check session state before entering loop — don't hang on completed sessions
147- if self ._check_session_state :
142+ async def _listen (
143+ self , session_id : str , * , _skip_state_precheck : bool = False ,
144+ ) -> AsyncGenerator [ChatEvent , None ]:
145+ """Listen for events — all events dispatched immediately."""
146+ if not _skip_state_precheck and self ._check_session_state :
148147 try :
149148 session = await self ._check_session_state (session_id )
150149 if session .get ("state" ) in TERMINAL_STATES :
@@ -153,79 +152,31 @@ async def _listen(self, session_id: str) -> AsyncGenerator[ChatEvent, None]:
153152 except Exception :
154153 pass # best effort
155154
156- text_buffer = TextPartBuffer ()
157155 queue : asyncio .Queue [Optional [ChatEvent ]] = asyncio .Queue ()
158156 done = False
159157 received_agent_response = False
160158
161- # Work log debounce state
162- wl_timers : dict [str , asyncio .TimerHandle ] = {}
163- wl_buffers : dict [str , dict [str , Any ]] = {}
164-
165- def flush_wl (step_id : str ) -> None :
166- buf = wl_buffers .pop (step_id , None )
167- wl_timers .pop (step_id , None )
168- if buf :
169- queue .put_nowait (ChatEvent (
170- type = S2CEvent .SESSION_WORK_LOG_PART ,
171- session_id = session_id ,
172- data = {"step_id" : step_id , "text" : buf .get ("text" , "" ), "status" : buf .get ("status" )},
173- ))
174-
175159 def handler (event : str , raw : dict [str , Any ]) -> None :
176- nonlocal done
160+ nonlocal done , received_agent_response
177161 payload = raw .get ("payload" , {})
178162 p_session_id = payload .get ("session_id" )
179163 if p_session_id and p_session_id != session_id :
180- return # Not our session — other handlers will process it
181-
182- message_id = payload .get ("message_id" )
183- data = payload .get ("data" )
184- metadata = raw .get ("metadata" )
185-
186- # Tier 1: text_part
187- if event == S2CEvent .SESSION_TEXT_PART :
188- if isinstance (data , dict ):
189- content = data .get ("content" , "" )
190- final = data .get ("final" , False )
191- merged = text_buffer .collect (message_id or "unknown" , content , final )
192- if merged is not None :
193- queue .put_nowait (ChatEvent (
194- type = S2CEvent .SESSION_TEXT , session_id = session_id ,
195- message_id = message_id , data = {"content" : merged }, metadata = metadata ,
196- ))
197- return
198-
199- # Tier 3: work_log_part debounce
200- if event == S2CEvent .SESSION_WORK_LOG_PART :
201- if isinstance (data , dict ):
202- step_id = data .get ("step_id" , "unknown" )
203- existing = wl_buffers .get (step_id , {"text" : "" })
204- existing ["text" ] = existing .get ("text" , "" ) + (data .get ("text_delta" , "" ) or "" )
205- if data .get ("status" ):
206- existing ["status" ] = data ["status" ]
207- wl_buffers [step_id ] = existing
208- old_timer = wl_timers .pop (step_id , None )
209- if old_timer :
210- old_timer .cancel ()
211- loop = asyncio .get_running_loop ()
212- wl_timers [step_id ] = loop .call_later (3.0 , flush_wl , step_id )
213164 return
214165
215- # All other events: dispatch immediately (pass-through for agent)
216- nonlocal received_agent_response
217166 queue .put_nowait (ChatEvent (
218167 type = event , session_id = session_id ,
219- message_id = message_id , data = data , metadata = metadata ,
168+ message_id = payload .get ("message_id" ),
169+ data = payload .get ("data" ),
170+ metadata = raw .get ("metadata" ),
220171 ))
221172 if event in SUBSTANTIVE_EVENTS :
222173 received_agent_response = True
223- if event == S2CEvent .SESSION_INPUT_STATE and isinstance (data , dict ):
224- if data .get ("content" ) == "waiting_input" and received_agent_response :
174+ if event == S2CEvent .SESSION_INPUT_STATE and isinstance (payload . get ( " data" ) , dict ):
175+ if payload [ " data" ] .get ("content" ) == "waiting_input" and received_agent_response :
225176 done = True
226177 queue .put_nowait (None )
227- if event == S2CEvent .SESSION_STATE and isinstance (data , dict ):
228- state = data .get ("content" , "" )
178+ if event == S2CEvent .SESSION_STATE and isinstance (payload . get ( " data" ) , dict ):
179+ state = payload [ " data" ] .get ("content" , "" )
229180 if state in TERMINAL_STATES :
230181 done = True
231182 queue .put_nowait (None )
@@ -234,10 +185,12 @@ def handler(event: str, raw: dict[str, Any]) -> None:
234185
235186 try :
236187 while not done :
188+ timeout = self ._response_idle_timeout_s if received_agent_response else self ._idle_timeout_s
237189 try :
238- evt = await asyncio .wait_for (queue .get (), timeout = self . _idle_timeout_s )
190+ evt = await asyncio .wait_for (queue .get (), timeout = timeout )
239191 except asyncio .TimeoutError :
240- # Idle timeout — check session state via REST
192+ if received_agent_response :
193+ break
241194 if self ._check_session_state :
242195 try :
243196 session = await self ._check_session_state (session_id )
@@ -255,8 +208,6 @@ def handler(event: str, raw: dict[str, Any]) -> None:
255208 if evt is not None :
256209 yield evt
257210 finally :
258- for t in wl_timers .values ():
259- t .cancel ()
260211 remove_handler ()
261212
262213 def send_form_response (self , session_id : str , message_id : str , form_data : dict [str , Any ]) -> None :
0 commit comments