Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions custom-providers/pi_voice/pi_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
11 changes: 11 additions & 0 deletions custom-providers/pi_voice/pi_voice.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@

import json
import os
import threading
import unicodedata
from pathlib import Path
from typing import Iterator
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
77 changes: 77 additions & 0 deletions custom-providers/pi_voice/tests/test_pi_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
})
Expand All @@ -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):
Expand Down
50 changes: 50 additions & 0 deletions custom-providers/pi_voice/tests/test_pi_voice.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
import os
import sys
import tempfile
import threading
import time
import unittest
from pathlib import Path
from typing import Iterator
Expand Down Expand Up @@ -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):
Expand Down