Skip to content

Commit 453c012

Browse files
bojielicursoragent
andcommitted
Refactor ChatEngine for immediate event dispatch and stale event filtering
- Removed buffering for text_part and work_log_part events, ensuring they are dispatched immediately. - Introduced a mechanism to filter out stale events based on a timestamp cutoff. - Updated the _listen method to support skipping state precheck for more efficient event handling. - Added unit tests to verify immediate dispatch behavior and stale event filtering. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent 2ddff19 commit 453c012

2 files changed

Lines changed: 238 additions & 93 deletions

File tree

src/pine_assistant/chat.py

Lines changed: 44 additions & 93 deletions
Original file line numberDiff line numberDiff line change
@@ -1,28 +1,23 @@
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

107
import asyncio
11-
import time
8+
from datetime import datetime, timedelta, timezone
129
from typing import Any, AsyncGenerator, Callable, Coroutine, Optional
1310

1411
from pine_assistant.models.events import C2SEvent, S2CEvent
1512
from pine_assistant.transport.socketio import SocketIOManager
1613

1714
TERMINAL_STATES = {"task_finished", "task_cancelled", "task_stale"}
1815
DEFAULT_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
2418
SUBSTANTIVE_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-
6442
class 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

Comments
 (0)