|
|
@@ -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)
|