From 2626077d81ced35270ebd7949403f95c7b807391 Mon Sep 17 00:00:00 2001 From: Chris Brien Date: Tue, 11 Aug 2026 12:17:08 -0400 Subject: [PATCH] fix: serialize PiVoice RPC turns Correlate new_session acknowledgements by request ID and surface reset failures instead of accepting stale responses. Co-Authored-By: OpenAI Codex --- custom-providers/pi_voice/pi_client.py | 5 ++ custom-providers/pi_voice/pi_voice.py | 11 +++ .../pi_voice/tests/test_pi_client.py | 77 +++++++++++++++++++ .../pi_voice/tests/test_pi_voice.py | 50 ++++++++++++ 4 files changed, 143 insertions(+) diff --git a/custom-providers/pi_voice/pi_client.py b/custom-providers/pi_voice/pi_client.py index 7431bf3..94fa8b5 100644 --- a/custom-providers/pi_voice/pi_client.py +++ b/custom-providers/pi_voice/pi_client.py @@ -204,7 +204,12 @@ def new_session(self) -> None: if ( frame.get("type") == "response" and frame.get("command") == "new_session" + and frame.get("id") == req_id ): + if not frame.get("success", False): + raise PiClientError( + f"pi rejected new_session: {frame.get('error', 'unknown')}" + ) return raise PiClientError("new_session timed out waiting for response") diff --git a/custom-providers/pi_voice/pi_voice.py b/custom-providers/pi_voice/pi_voice.py index 39036be..1695339 100644 --- a/custom-providers/pi_voice/pi_voice.py +++ b/custom-providers/pi_voice/pi_voice.py @@ -35,6 +35,7 @@ import json import os +import threading import unicodedata from pathlib import Path from typing import Iterator @@ -214,6 +215,11 @@ def __init__(self, config: dict, *, client: PiClient | None = None): # `client` is injected by tests; production passes None to get # the env-configured default. self._client: PiClient = client if client is not None else make_default_pi_client() + # A connection can submit multiple chat jobs to its thread pool. Pi RPC + # is one ordered stream, so keep the complete new_session -> prompt -> + # agent_end transaction exclusive; per-write locking cannot prevent one + # caller from consuming another caller's response frames. + self._turn_lock = threading.Lock() self._first_turn = True msg = f"PiVoiceLLM ready (container={self._container} kid_mode={self._kid_mode})" try: @@ -224,6 +230,11 @@ def __init__(self, config: dict, *, client: PiClient | None = None): # xiaozhi-server's voice loop calls this as a sync generator. # Each yielded string becomes a TTS chunk. def response(self, session_id, dialogue, **kwargs) -> Iterator[str]: + with self._turn_lock: + yield from self._response_serialized(session_id, dialogue, **kwargs) + + def _response_serialized(self, session_id, dialogue, **kwargs) -> Iterator[str]: + """Run one complete Pi RPC transaction while ``_turn_lock`` is held.""" self._kid_mode = _read_kid_mode() user_text = _last_user_text(dialogue) if not user_text: diff --git a/custom-providers/pi_voice/tests/test_pi_client.py b/custom-providers/pi_voice/tests/test_pi_client.py index 50e2b5e..9e563c4 100644 --- a/custom-providers/pi_voice/tests/test_pi_client.py +++ b/custom-providers/pi_voice/tests/test_pi_client.py @@ -156,6 +156,7 @@ def ns_responder(): for cmd in fake.stdin_lines: if cmd.get("type") == "new_session": fake.emit({ + "id": cmd["id"], "type": "response", "command": "new_session", "success": True, }) @@ -172,6 +173,82 @@ def ns_responder(): finally: client.close() + def test_new_session_ignores_stale_response_until_matching_ack(self): + fake = FakePopen() + client = make_client(fake) + stale_emitted = threading.Event() + release_matching = threading.Event() + errors: list[BaseException] = [] + + def responder(): + while True: + time.sleep(0.01) + for cmd in fake.stdin_lines: + if cmd.get("type") == "new_session": + fake.emit({ + "id": "nsess-stale", + "type": "response", + "command": "new_session", + "success": False, + "error": "stale failure", + }) + stale_emitted.set() + release_matching.wait(timeout=2) + fake.emit({ + "id": cmd["id"], + "type": "response", + "command": "new_session", + "success": True, + }) + return + + def reset_session(): + try: + client.new_session() + except BaseException as exc: # captured for the main test thread + errors.append(exc) + + threading.Thread(target=responder, daemon=True).start() + reset = threading.Thread(target=reset_session) + reset.start() + try: + self.assertTrue(stale_emitted.wait(timeout=1)) + time.sleep(0.05) + self.assertTrue(reset.is_alive(), "stale response must not complete reset") + release_matching.set() + reset.join(timeout=2) + self.assertFalse(reset.is_alive()) + self.assertEqual(errors, []) + finally: + release_matching.set() + reset.join(timeout=2) + client.close() + + def test_new_session_raises_on_matching_failure(self): + fake = FakePopen() + client = make_client(fake) + + def responder(): + while True: + time.sleep(0.01) + for cmd in fake.stdin_lines: + if cmd.get("type") == "new_session": + fake.emit({ + "id": cmd["id"], + "type": "response", + "command": "new_session", + "success": False, + "error": "reset refused", + }) + return + + threading.Thread(target=responder, daemon=True).start() + try: + with self.assertRaisesRegex(PiClientError, "reset refused"): + client.new_session() + finally: + client.close() + class TestThinkingFilter(unittest.TestCase): def test_thinking_deltas_are_dropped(self): diff --git a/custom-providers/pi_voice/tests/test_pi_voice.py b/custom-providers/pi_voice/tests/test_pi_voice.py index 7725fa7..fd5abdc 100644 --- a/custom-providers/pi_voice/tests/test_pi_voice.py +++ b/custom-providers/pi_voice/tests/test_pi_voice.py @@ -10,6 +10,8 @@ import os import sys import tempfile +import threading +import time import unittest from pathlib import Path from typing import Iterator @@ -193,6 +195,54 @@ def test_first_turn_skips_new_session(self): list(provider.response("s", [{"role": "user", "content": "b"}])) self.assertEqual(client.new_session_calls, 1, "new_session on second turn") + def test_concurrent_responses_are_serialized_through_agent_end(self): + class OverlapDetectingClient(FakeClient): + def __init__(self): + super().__init__() + self.active = 0 + self.max_active = 0 + self.first_started = threading.Event() + self.release_first = threading.Event() + + def iter_turn_text(self, prompt: str) -> Iterator[str]: + self.prompts.append(prompt) + self.active += 1 + self.max_active = max(self.max_active, self.active) + try: + if len(self.prompts) == 1: + self.first_started.set() + self.release_first.wait(timeout=2) + yield "😊 ok" + finally: + self.active -= 1 + + os.environ["DOTTY_KID_MODE"] = "false" + client = OverlapDetectingClient() + provider = LLMProvider({}, client=client) # type: ignore[arg-type] + outputs: list[list[str]] = [] + + def run(text: str) -> None: + outputs.append(list(provider.response( + "s", [{"role": "user", "content": text}], + ))) + + first = threading.Thread(target=run, args=("first",)) + second = threading.Thread(target=run, args=("second",)) + first.start() + self.assertTrue(client.first_started.wait(timeout=1)) + second.start() + time.sleep(0.05) + self.assertEqual(len(client.prompts), 1, "second turn must wait") + client.release_first.set() + first.join(timeout=2) + second.join(timeout=2) + + self.assertFalse(first.is_alive()) + self.assertFalse(second.is_alive()) + self.assertEqual(client.max_active, 1) + self.assertEqual(len(outputs), 2) + self.assertEqual(client.new_session_calls, 1) + class TestErrorFallback(unittest.TestCase): def test_client_error_yields_fallback(self):