Skip to content

Commit 2ea6a16

Browse files
committed
perf(chat): wait on turn completion notifications
Signed-off-by: duanjialing.777 <duanjialing.777@bytedance.com>
1 parent f72fea4 commit 2ea6a16

2 files changed

Lines changed: 44 additions & 5 deletions

File tree

‎loopx/chat_runtime.py‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1398,14 +1398,16 @@ def interrupt_turn(self, *, session_id: str, turn_id: str) -> dict[str, Any]:
13981398

13991399
def wait_for_turn(self, *, session_id: str, turn_id: str, timeout_sec: float = 920.0) -> dict[str, Any]:
14001400
deadline = time.monotonic() + timeout_sec
1401-
while time.monotonic() < deadline:
1402-
turn = self.store.load_turn(session_id, turn_id)
1403-
if turn is None:
1401+
while True:
1402+
if (turn := self.store.load_turn(session_id, turn_id)) is None:
14041403
raise KeyError("chat turn was not found")
14051404
if turn.get("status") in TERMINAL_TURN_STATES:
14061405
return turn
1407-
time.sleep(0.02)
1408-
raise TimeoutError("chat turn wait timed out")
1406+
if (remaining := deadline - time.monotonic()) <= 0:
1407+
raise TimeoutError("chat turn wait timed out")
1408+
with self.lock:
1409+
done_event = self.turn_done_events.get((session_id, turn_id))
1410+
(done_event.wait if done_event else time.sleep)(remaining if done_event else min(0.02, remaining))
14091411

14101412
def close_session(self, session_id: str) -> bool:
14111413
with self._session_adapter_lock(session_id):

‎tests/test_chat_turn_wait.py‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
from __future__ import annotations
2+
3+
import threading
4+
from unittest.mock import Mock
5+
6+
from loopx.chat_runtime import ChatRuntimeController
7+
8+
9+
def _runtime() -> ChatRuntimeController:
10+
runtime = ChatRuntimeController.__new__(ChatRuntimeController)
11+
runtime.store = Mock() # type: ignore[assignment]
12+
runtime.lock = threading.RLock()
13+
runtime.turn_done_events = {}
14+
return runtime
15+
16+
17+
def test_wait_for_turn_uses_managed_completion_event() -> None:
18+
runtime = _runtime()
19+
runtime.store.load_turn.side_effect = [{"status": "running"}, {"status": "completed"}]
20+
completion = Mock()
21+
runtime.turn_done_events[("session", "turn")] = completion # type: ignore[assignment]
22+
23+
turn = runtime.wait_for_turn(session_id="session", turn_id="turn", timeout_sec=0.1)
24+
25+
assert turn["status"] == "completed"
26+
completion.wait.assert_called_once()
27+
assert runtime.store.load_turn.call_count == 2
28+
29+
30+
def test_wait_for_turn_performs_final_fallback_read_at_deadline() -> None:
31+
runtime = _runtime()
32+
runtime.store.load_turn.side_effect = [{"status": "running"}, {"status": "completed"}]
33+
34+
turn = runtime.wait_for_turn(session_id="session", turn_id="turn", timeout_sec=0.001)
35+
36+
assert turn["status"] == "completed"
37+
assert runtime.store.load_turn.call_count == 2

0 commit comments

Comments
 (0)