Przeglądaj źródła

feat(tts): concurrent TTS playback + user interruption (xiaozhi-style abort)

- Each TTS segment plays immediately (background asyncio tasks)
- No more waiting for all LLM output before playing audio
- User speaking cancels any running TTS playback (interruption)
- Replace cur_tts tracking with play_tasks list
wenhongquan 2 miesięcy temu
rodzic
commit
6afa008037
1 zmienionych plików z 30 dodań i 23 usunięć
  1. 30 23
      asr_agent/conversation_worker.py

+ 30 - 23
asr_agent/conversation_worker.py

@@ -364,10 +364,10 @@ class Worker:
         sr = 16000
         vad = _SileroVAD()
         busy = False
-        cur_tts = None
+        play_tasks: list[asyncio.Task] = []
 
         async def _acc():
-            nonlocal busy, cur_tts
+            nonlocal busy, play_tasks
             stream = rtc.AudioStream(track, sample_rate=sr, num_channels=1)
             fc = 0
             async for ev in stream:
@@ -387,11 +387,14 @@ class Worker:
 
                 if was_speech and vad.should_end() and vad.min_speech_met():
                     logger.info("[%s] VAD END blen=%d", sid, len(buf.buffer))
-                    if cur_tts:
-                        cur_tts = None
+                    # Cancel any ongoing TTS (interruption)
+                    for pt in play_tasks:
+                        if not pt.done():
+                            pt.cancel()
+                    play_tasks.clear()
                     busy = True
                     try:
-                        await self._transcribe_and_respond(buf)
+                        await self._transcribe_and_respond(buf, play_tasks)
                     except Exception:
                         logger.exception("[%s] transcribe failed", sid)
                     busy = False
@@ -400,9 +403,13 @@ class Worker:
 
                 if vad.in_speech and vad.total_samples >= int(MAX_SPEECH_S * sr):
                     logger.info("[%s] VAD MAX blen=%d", sid, len(buf.buffer))
+                    for pt in play_tasks:
+                        if not pt.done():
+                            pt.cancel()
+                    play_tasks.clear()
                     busy = True
                     try:
-                        await self._transcribe_and_respond(buf)
+                        await self._transcribe_and_respond(buf, play_tasks)
                     except Exception:
                         logger.exception("[%s] max failed", sid)
                     busy = False
@@ -418,7 +425,9 @@ class Worker:
         except Exception:
             logger.exception("[%s] _acc failed", sid)
 
-    async def _transcribe_and_respond(self, buf: AudioBuffer):
+    async def _transcribe_and_respond(self, buf: AudioBuffer, play_tasks: list[asyncio.Task] | None = None):
+        if play_tasks is None:
+            play_tasks = []
         all_audio = buf.get_all()
         buf.clear()
         if len(all_audio) <= 1600:
@@ -433,10 +442,9 @@ class Worker:
         logger.info("ASR: %s", txt[:80])
         await self._send({"type": "utterance", "text": txt, "seq": 0})
 
-        # ── Streaming LLM + xiaozhi-style TTS ──
+        # ── Streaming LLM + concurrent TTS playback ──
         reply_full = ""
         seg = _TTSSegmenter()
-        tts_task = None
 
         try:
             async for delta, is_final in _llm_stream(txt, self.hist):
@@ -444,30 +452,18 @@ class Worker:
                     reply_full += delta
                     await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})
                     for s in seg.feed(delta):
-                        if tts_task:
-                            await tts_task
                         logger.info("TTS seg(%d): %s", len(s), s[:40])
-                        tts_task = asyncio.create_task(_tts(s))
+                        play_tasks.append(asyncio.create_task(self._tts_and_play(s)))
 
                 if is_final:
                     remaining = seg.flush()
                     if remaining:
-                        if tts_task:
-                            await tts_task
-                        tts_task = asyncio.create_task(_tts(remaining))
+                        play_tasks.append(asyncio.create_task(self._tts_and_play(remaining)))
                     break
         except Exception:
             logger.exception("LLM stream failed")
             reply_full = "抱歉,我暂时无法回答。"
 
-        if tts_task:
-            try:
-                wav = await asyncio.wait_for(tts_task, timeout=15)
-                if wav is not None and len(wav) > 0:
-                    await self._play(wav, 24000)
-            except (asyncio.TimeoutError, Exception):
-                pass
-
         if reply_full:
             await self._send({"type": "reply", "text": reply_full})
             self.hist.extend([
@@ -476,6 +472,17 @@ class Worker:
             ])
             self.hist[:] = self.hist[-20:]
 
+    async def _tts_and_play(self, text: str):
+        """Fetch TTS audio and play it immediately (background task)."""
+        try:
+            wav = await asyncio.wait_for(_tts(text), timeout=15)
+            if wav is not None and len(wav) > 0:
+                await self._play(wav, 24000)
+        except (asyncio.TimeoutError, asyncio.CancelledError):
+            pass
+        except Exception:
+            logger.exception("TTS+play failed")
+
     async def _play(self, audio: np.ndarray, sr: int):
         src = rtc.AudioSource(sr, 1)
         tk = rtc.LocalAudioTrack.create_audio_track("tts", src)