"""
Full Pipeline Worker: VAD -> ASR -> LLM -> TTS
VAD: Silero VAD (ONNX) with dual-threshold + sliding window
VP: AEC/ANS/AGC handled by LiveKit WebRTC
ASR: Qwen3-ASR (local model, batch on speech-end)
LLM: streaming via OpenAI-compatible API (vLLM), tag filtering
TTS: Mimo API (remote), xiaozhi-style two-tier sentence segmentation
"""
from __future__ import annotations
import asyncio
import base64
import json
import logging
import os
import re
import sys
import time
from collections import deque
from dataclasses import dataclass, field
import numpy as np
try:
import onnxruntime
_HAS_ONNX = True
except ImportError:
_HAS_ONNX = False
import aiohttp
import jwt
from livekit import rtc
_root = os.path.join(os.path.dirname(os.path.abspath(__file__)), "whisper_asr")
sys.path.insert(0, _root)
from audio_processor import AudioBuffer
from qwen_engine import QwenASREngine
logger = logging.getLogger("worker")
# ── Env ──
LIVEKIT_URL = os.environ.get("LIVEKIT_URL", "ws://localhost:7880")
ASR_MODEL = os.environ.get("ASR_MODEL_PATH", "Qwen/Qwen3-ASR-0.6B")
VLLM_URL = os.environ.get("VLLM_URL", "http://127.0.0.1:8000/v1")
LLM_MODEL = os.environ.get("LLM_MODEL", "qwen3.6-35b-awq")
MIMO_KEY = os.environ.get("MIMO_KEY", "")
MIMO_API_BASE = os.environ.get("MIMO_API_BASE", "https://token-plan-cn.xiaomimimo.com/v1")
# Provider switches: "vllm" | "mimo" for LLM, "qwen" | "mimo" for ASR
LLM_PROVIDER = os.environ.get("LLM_PROVIDER", "vllm")
ASR_PROVIDER = os.environ.get("ASR_PROVIDER", "qwen")
# ── LLM ──
LLM_MAX_TOKENS = 128
LLM_TEMPERATURE = 0.7
LLM_TIMEOUT = 25
SYSTEM_PROMPT = (
"你是友好的中文语音助手,名字叫狄诺尼试验员。"
"回答简洁自然,2-3句即可,用口语化中文。"
"不要使用 Markdown、代码块、表格或特殊符号。"
"不要输出括号注释、不要使用英文缩写。"
)
# ── TTS segmentation (xiaozhi-style two-tier) ──
FIRST_SENTENCE_PUNCT = ",、,。!?;::\n"
SENTENCE_END_PUNCT = "。!?!?\n"
MIN_TTS_CHARS = 2
# ═══════════════════════════════════════════════════════════════
# JWT
# ═══════════════════════════════════════════════════════════════
def _token(room, ident):
n = int(time.time())
return jwt.encode({
"iss": "devkey", "sub": ident, "name": ident,
"nbf": n - 60, "exp": n + 6 * 3600,
"video": {"roomJoin": True, "room": room,
"canPublish": True, "canSubscribe": True, "canPublishData": True},
}, "secretsecretsecretsecretsecret12", algorithm="HS256")
# ═══════════════════════════════════════════════════════════════
# LLM (streaming, filtering)
# ═══════════════════════════════════════════════════════════════
async def _llm_stream(prompt: str, hist: list[dict]):
"""Stream LLM tokens, dispatching based on LLM_PROVIDER env."""
if LLM_PROVIDER == "mimo":
async for x in _llm_stream_mimo(prompt, hist):
yield x
return
async for x in _llm_stream_vllm(prompt, hist):
yield x
async def _llm_stream_vllm(prompt: str, hist: list[dict]):
"""vLLM streaming via OpenAI-compatible SSE."""
msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
msgs.extend(hist)
msgs.append({"role": "user", "content": prompt})
async with aiohttp.ClientSession() as s:
async with s.post(
f"{VLLM_URL}/chat/completions",
json={"model": LLM_MODEL, "messages": msgs,
"max_tokens": LLM_MAX_TOKENS, "temperature": LLM_TEMPERATURE,
"stream": True,
"chat_template_kwargs": {"enable_thinking": False}},
timeout=aiohttp.ClientTimeout(total=LLM_TIMEOUT),
) as r:
async for line in r.content:
line = line.decode().strip()
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
text = delta.get("content", "")
if text:
if "" in text:
text = text.split("")[-1]
if "" in text:
text = text.split("")[0]
if text.strip():
yield text, False
except Exception:
continue
yield "", True
async def _llm_stream_mimo(prompt: str, hist: list[dict]):
"""Mimo v2.5 LLM streaming via OpenAI-compatible SSE."""
if not MIMO_KEY:
yield "", True
return
msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
msgs.extend(hist)
msgs.append({"role": "user", "content": prompt})
h = {"api-key": MIMO_KEY, "Content-Type": "application/json"}
async with aiohttp.ClientSession() as s:
async with s.post(
f"{MIMO_API_BASE}/chat/completions",
json={"model": "mimo-v2.5", "messages": msgs,
"max_tokens": LLM_MAX_TOKENS, "temperature": LLM_TEMPERATURE,
"stream": True},
headers=h, timeout=aiohttp.ClientTimeout(total=LLM_TIMEOUT),
) as r:
async for line in r.content:
line = line.decode().strip()
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
text = delta.get("content", "")
if text and text.strip():
yield text, False
except Exception:
continue
yield "", True
# ═══════════════════════════════════════════════════════════════
# ASR (dispatch: qwen-local | mimo-api)
# ═══════════════════════════════════════════════════════════════
async def _asr_transcribe(audio_data: np.ndarray, asr_engine=None) -> str:
"""Dispatch ASR based on ASR_PROVIDER env. Return final text."""
if ASR_PROVIDER == "mimo":
text = ""
async for chunk, is_done in _asr_mimo_stream(audio_data):
if chunk:
text += chunk
return text.strip()
loop = asyncio.get_running_loop()
result = await loop.run_in_executor(None, asr_engine.transcribe_array, audio_data)
return result.get("text", "").strip()
async def _asr_mimo_stream(audio_data: np.ndarray):
"""Mimo v2.5 streaming ASR. Yields (text_snapshot, is_final)."""
if not MIMO_KEY:
yield "", True
return
import io, wave
buf = io.BytesIO()
with wave.open(buf, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(16000)
wf.writeframes((audio_data * 32767).astype(np.int16).tobytes())
audio_b64 = base64.b64encode(buf.getvalue()).decode()
h = {"api-key": MIMO_KEY, "Content-Type": "application/json"}
async with aiohttp.ClientSession() as s:
async with s.post(
f"{MIMO_API_BASE}/chat/completions",
json={
"model": "mimo-v2.5-asr",
"messages": [{"role": "user", "content": [
{"type": "input_audio", "input_audio": {
"data": f"data:audio/wav;base64,{audio_b64}"}}
]}],
"stream": True,
"asr_options": {"language": "auto"},
},
headers=h, timeout=aiohttp.ClientTimeout(total=20),
) as r:
async for line in r.content:
line = line.decode().strip()
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
yield "", True
return
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
text = delta.get("content", "")
if text:
yield text, False
except Exception:
continue
yield "", True
# ═══════════════════════════════════════════════════════════════
# TTS
# ═══════════════════════════════════════════════════════════════
async def _tts_stream(text: str, voice: str = "Chloe"):
"""Mimo v2.5 streaming TTS. Yields (pcm_chunk: np.ndarray, sr: int, is_final: bool)."""
if not MIMO_KEY:
yield np.array([], dtype=np.float32), 24000, True
return
if voice.startswith("data:"):
model = "mimo-v2.5-tts-voiceclone"
else:
model = "mimo-v2.5-tts"
h = {"api-key": MIMO_KEY, "Content-Type": "application/json"}
body: dict = {
"model": model,
"messages": [
{"role": "user", "content": ""},
{"role": "assistant", "content": text},
],
"audio": {"format": "wav", "voice": voice},
"stream": True,
}
async with aiohttp.ClientSession() as s:
try:
async with s.post(
f"{MIMO_API_BASE}/chat/completions",
json=body, headers=h,
timeout=aiohttp.ClientTimeout(total=30),
) as r:
async for line in r.content:
line = line.decode().strip()
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
audio = delta.get("audio")
if audio and isinstance(audio, dict) and "data" in audio:
pcm = np.frombuffer(
base64.b64decode(audio["data"]), dtype=np.int16
).astype(np.float32) / 32768.0
if len(pcm) > 0:
yield pcm, 24000, False
except Exception:
continue
except Exception:
logger.exception("TTS stream failed")
yield np.array([], dtype=np.float32), 24000, True
def _clean_tts_text(text: str) -> str:
text = re.sub(r"\*{1,3}(.*?)\*{1,3}", r"\1", text)
text = re.sub(r"`{1,3}.*?`{1,3}", "", text)
text = re.sub(r"\[([^\]]+)]\([^)]+\)", r"\1", text)
text = re.sub(r"#{1,6}\s*", "", text)
text = re.sub(r"[>\-]\s", "", text)
text = re.sub(r"[\U0001F300-\U0001F9FF]", "", text)
return text.strip()
# ═══════════════════════════════════════════════════════════════
# TTS Segmenter (xiaozhi-style: two-tier, processed_chars)
# ═══════════════════════════════════════════════════════════════
@dataclass
class _TTSSegmenter:
buffer: str = ""
processed: int = 0
is_first: bool = True
def feed(self, text: str) -> list[str]:
self.buffer += text
segments = []
puncts = FIRST_SENTENCE_PUNCT if self.is_first else SENTENCE_END_PUNCT
unprocessed = self.buffer[self.processed:]
last_pos = -1
for ch in puncts:
pos = unprocessed.rfind(ch)
if pos > last_pos:
last_pos = pos
if last_pos >= 0:
seg_raw = unprocessed[:last_pos + 1]
seg_clean = _clean_tts_text(seg_raw)
if len(seg_clean) >= MIN_TTS_CHARS:
segments.append(seg_clean)
self.processed += len(seg_raw)
if self.is_first:
self.is_first = False
return segments
def flush(self) -> str:
remaining = self.buffer[self.processed:].strip()
if remaining:
remaining = _clean_tts_text(remaining)
self.processed = len(self.buffer)
return remaining
# ── VAD ──
VAD_MODEL_PATH = os.environ.get("VAD_MODEL_PATH") or os.path.join(
os.path.dirname(os.path.abspath(__file__)), "silero_vad.onnx")
VAD_THRESHOLD_HIGH = float(os.environ.get("VAD_THRESHOLD_HIGH", "0.5"))
VAD_THRESHOLD_LOW = float(os.environ.get("VAD_THRESHOLD_LOW", "0.2"))
VAD_WINDOW_SIZE = int(os.environ.get("VAD_WINDOW_SIZE", "5"))
VAD_VOICE_RATIO = float(os.environ.get("VAD_VOICE_RATIO", "0.5"))
MIN_SPEECH_S = float(os.environ.get("MIN_SPEECH_S", "0.3"))
MIN_SILENCE_S = float(os.environ.get("MIN_SILENCE_S", "0.4"))
MAX_SPEECH_S = float(os.environ.get("MAX_SPEECH_S", "8.0"))
# Global ONNX session (shared across participants)
_vad_session = None
def _get_vad_session():
global _vad_session
if _vad_session is not None:
return _vad_session
if not _HAS_ONNX or not os.path.exists(VAD_MODEL_PATH):
return None
opts = onnxruntime.SessionOptions()
opts.inter_op_num_threads = 1
opts.intra_op_num_threads = 1
_vad_session = onnxruntime.InferenceSession(
VAD_MODEL_PATH, providers=["CPUExecutionProvider"], sess_options=opts,
)
logger.info("Silero VAD loaded (%s)", VAD_MODEL_PATH)
return _vad_session
# ═══════════════════════════════════════════════════════════════
# Silero VAD (ONNX, per-participant state)
# ═══════════════════════════════════════════════════════════════
class _SileroVAD:
"""Xiaozhi-style Silero VAD with dual-threshold + sliding window."""
def __init__(self):
self._sess = _get_vad_session()
self._state = np.zeros((2, 1, 128), dtype=np.float32)
self._context = np.zeros((1, 64), dtype=np.float32)
self._window: deque = deque(maxlen=max(VAD_WINDOW_SIZE, 1))
self.in_speech = False
self.speech_start_sample = 0
self.silence_counter = 0
self.total_samples = 0
self._sr = np.array(16000, dtype=np.int64)
def process(self, frame: np.ndarray) -> bool:
"""Process 160-sample (10ms) frame. Call every frame. Returns in_speech."""
n = len(frame)
if self._sess is not None:
# ── Silero ONNX inference ──
# Silero processes 512-sample windows. Accumulate frames.
if not hasattr(self, '_buf'):
self._buf = np.array([], dtype=np.float32)
self._buf = np.concatenate([self._buf, frame]).astype(np.float32)
is_voice = False
while len(self._buf) >= 512:
chunk = self._buf[:512]
self._buf = self._buf[512:]
audio_in = chunk.reshape(1, -1)
inp = np.concatenate([self._context, audio_in], axis=1).astype(np.float32)
out, self._state = self._sess.run(
None, {"input": inp, "state": self._state, "sr": self._sr},
)
self._context = inp[:, -64:]
prob = out.item()
is_voice = prob >= VAD_THRESHOLD_HIGH or (
prob > VAD_THRESHOLD_LOW and (
self._window[-1] if self._window else False
)
)
self._window.append(is_voice)
else:
# ── Fallback: energy-based ──
energy = float(np.sqrt(np.mean(frame ** 2)))
if energy > 0.015:
is_voice = True
elif energy < 0.005:
is_voice = False
else:
is_voice = self._window[-1] if self._window else False
self._window.append(is_voice)
have_voice = sum(self._window) / max(len(self._window), 1) >= VAD_VOICE_RATIO
if have_voice and not self.in_speech:
logger.info("VAD START")
self.in_speech = True
self.speech_start_sample = self.total_samples
self.silence_counter = 0
elif not have_voice and self.in_speech:
self.silence_counter += n
elif have_voice:
self.silence_counter = 0
self.total_samples += n
return self.in_speech
def should_end(self) -> bool:
return self.in_speech and self.silence_counter >= int(MIN_SILENCE_S * 16000)
def min_speech_met(self) -> bool:
return (self.total_samples - self.speech_start_sample) / 16000 >= MIN_SPEECH_S
def reset(self):
self._window.clear()
self.in_speech = False
self.speech_start_sample = 0
self.silence_counter = 0
self._state = np.zeros((2, 1, 128), dtype=np.float32)
self._context = np.zeros((1, 64), dtype=np.float32)
if hasattr(self, '_buf'):
self._buf = np.array([], dtype=np.float32)
# ═══════════════════════════════════════════════════════════════
# Worker
# ═══════════════════════════════════════════════════════════════
class Worker:
def __init__(self, room, identity="asr-bot"):
self.rn, self.id = room, identity
self.room = rtc.Room()
self.tasks: dict = {}
self.hist: list[dict] = []
async def run(self):
logger.info("Loading ASR model...")
self.asr = QwenASREngine(model_id=ASR_MODEL, language=None)
logger.info("ASR ready")
self.room.on("track_subscribed", self._on_track)
self.room.on("participant_disconnected", self._on_off)
self.room.on("participant_connected", lambda p: None)
self.room.on("track_published", lambda pub, p: None)
await self.room.connect(LIVEKIT_URL, _token(self.rn, self.id))
logger.info("Worker ready (VAD+ASR+LLM+TTS)")
for p in self.room.remote_participants.values():
for pub in p.track_publications.values():
if pub.track and pub.kind == rtc.TrackKind.KIND_AUDIO:
self._on_track(pub.track, pub, p)
try:
await asyncio.Future()
finally:
await self.room.disconnect()
def _on_track(self, t, _, p):
s = p.sid
if s in self.tasks:
return
logger.info("TRACK %s", p.identity)
self.tasks[s] = asyncio.create_task(self._run(s, t))
def _on_off(self, p):
x = self.tasks.pop(p.sid, None)
if x:
x.cancel()
async def _run(self, sid, track):
logger.info("_run %s", sid)
buf = AudioBuffer(max_duration=MAX_SPEECH_S + 4, sample_rate=16000)
sr = 16000
vad = _SileroVAD()
busy = False
# ── Single persistent TTS output track per session ──
# One track is reused for all turns so audio never cross-streams
# between turns/segments (old code published a new track each call).
tts_src = rtc.AudioSource(24000, 1)
tts_track = rtc.LocalAudioTrack.create_audio_track("tts", tts_src)
await self.room.local_participant.publish_track(tts_track)
play_q: "asyncio.Queue[tuple[int, np.ndarray]]" = asyncio.Queue()
gen = 0 # turn generation; bumped to flush in-flight audio
cur_gen = {"v": -1} # generation currently being played by _player
async def _player():
"""Drain TTS audio sequentially on the single track."""
while True:
g, audio = await play_q.get()
cur_gen["v"] = g
for i in range(0, len(audio), 24000 // 50):
if cur_gen["v"] != g:
break # a newer turn started; drop this stale audio
c = audio[i: i + 24000 // 50]
await tts_src.capture_frame(rtc.AudioFrame(
data=c.tobytes(), sample_rate=24000,
num_channels=1, samples_per_channel=len(c),
))
await asyncio.sleep(0.018)
cur_gen["v"] = -1
player_task = asyncio.create_task(_player())
async def _acc():
nonlocal busy, gen
stream = rtc.AudioStream(track, sample_rate=sr, num_channels=1)
fc = 0
async for ev in stream:
fc += 1
if busy:
continue
arr = np.frombuffer(ev.frame.data, dtype=np.int16).astype(np.float32) / 32768.0
buf.append(arr)
was_speech = vad.in_speech
vad.process(arr)
if fc % 100 == 0:
logger.info("[%s] af#%d e=%.4f vad=%s blen=%d",
sid, fc, float(np.sqrt(np.mean(arr ** 2))),
vad.in_speech, len(buf.buffer))
if was_speech and vad.should_end() and vad.min_speech_met():
logger.info("[%s] VAD END blen=%d", sid, len(buf.buffer))
gen += 1
cur_gen["v"] = -1 # stop current playback immediately
_drain_queue(play_q)
busy = True
try:
await self._transcribe_and_respond(buf, play_q, gen)
except Exception:
logger.exception("[%s] transcribe failed", sid)
busy = False
vad.reset()
vad.total_samples = len(buf.buffer)
if vad.in_speech and vad.total_samples >= int(MAX_SPEECH_S * sr):
logger.info("[%s] VAD MAX blen=%d", sid, len(buf.buffer))
gen += 1
cur_gen["v"] = -1
_drain_queue(play_q)
busy = True
try:
await self._transcribe_and_respond(buf, play_q, gen)
except Exception:
logger.exception("[%s] max failed", sid)
busy = False
vad.reset()
vad.total_samples = len(buf.buffer)
logger.info("[%s] stream end fc=%d", sid, fc)
try:
await _acc()
except asyncio.CancelledError:
pass
except Exception:
logger.exception("[%s] _acc failed", sid)
finally:
player_task.cancel()
try:
await self.room.local_participant.unpublish_track(tts_track)
except Exception:
pass
async def _transcribe_and_respond(self, buf: AudioBuffer, play_q, gen: int):
all_audio = buf.get_all()
buf.clear()
if len(all_audio) <= 1600:
return
txt = await _asr_transcribe(all_audio, self.asr)
if not txt:
return
logger.info("ASR: %s", txt[:80])
await self._send({"type": "utterance", "text": txt, "seq": 0})
# ── Streaming LLM + concurrent TTS generation (sequential playback) ──
# TTS is generated concurrently per segment, but the player consumes the
# queue in order on a single track, so segments never overlap.
reply_full = ""
seg = _TTSSegmenter()
try:
async for delta, is_final in _llm_stream(txt, self.hist):
if delta:
reply_full += delta
await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})
for s in seg.feed(delta):
logger.info("TTS seg(%d): %s", len(s), s[:40])
asyncio.create_task(self._tts_enqueue(s, play_q, gen))
if is_final:
remaining = seg.flush()
if remaining:
asyncio.create_task(self._tts_enqueue(remaining, play_q, gen))
break
except Exception:
logger.exception("LLM stream failed")
reply_full = "抱歉,我暂时无法回答。"
if reply_full:
await self._send({"type": "reply", "text": reply_full})
self.hist.extend([
{"role": "user", "content": txt},
{"role": "assistant", "content": reply_full},
])
self.hist[:] = self.hist[-20:]
async def _tts_enqueue(self, text: str, play_q, gen: int):
"""Stream TTS audio for one segment, then enqueue it for playback."""
try:
chunks: list[np.ndarray] = []
async for pcm, sr, is_final in _tts_stream(text):
if len(pcm) > 0:
chunks.append(pcm)
if is_final:
break
if chunks:
audio = np.concatenate(chunks)
if len(audio) > 0:
await play_q.put((gen, audio))
except asyncio.CancelledError:
pass
except Exception:
logger.exception("TTS failed")
pass
async def main():
import argparse
p = argparse.ArgumentParser()
p.add_argument("--room", required=True)
a = p.parse_args()
logging.basicConfig(level=logging.INFO)
await Worker(a.room).run()
if __name__ == "__main__":
asyncio.run(main())
def _drain_queue(q: "asyncio.Queue") -> None:
"""Discard any queued-but-not-yet-played TTS audio."""
while not q.empty():
try:
q.get_nowait()
except Exception:
break