Просмотр исходного кода

fix(livekit): add participant_connected and track_published handlers to fix track_subscribed not firing

In livekit-python 1.1.13, the track_subscribed event dispatch depends on
_remote_participants and track_publications dicts populated by the
participant_connected and track_published event handlers. Without
registering these handlers, track_subscribed never fires for participants
joining after the worker.

- asr_agent/worker.py: add no-op participant_connected + track_published handlers
- asr_agent/conversation_worker.py: same fix (VAD+LLM+TTS full pipeline)
wenhongquan 3 недель назад
Родитель
Сommit
a221c1ae02
2 измененных файлов с 434 добавлено и 0 удалено
  1. 214 0
      asr_agent/conversation_worker.py
  2. 220 0
      asr_agent/worker.py

+ 214 - 0
asr_agent/conversation_worker.py

@@ -0,0 +1,214 @@
+"""
+Full Pipeline Worker: VAD → ASR → LLM → TTS
+  VAD: energy-based voice activity detection (local)
+  VP:  AEC/ANS/AGC handled by LiveKit WebRTC
+  ASR: Qwen3-ASR (local model)
+  LLM: vLLM (local, or Mimo fallback)
+  TTS: Mimo API (remote)
+"""
+
+from __future__ import annotations
+
+import asyncio
+import base64
+import json
+import logging
+import os
+import sys
+import time
+
+import aiohttp
+import jwt
+import numpy as np
+from livekit import rtc
+
+_root = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "whisper-qt-client", "whisper_asr")
+if not os.path.isdir(_root):
+    _root = os.path.join(os.path.dirname(os.path.abspath(__file__)), "whisper_asr")
+sys.path.insert(0, _root)
+from audio_processor import AudioBuffer, VADProcessor
+from qwen_engine import QwenASREngine
+
+logger = logging.getLogger("worker")
+
+LIVEKIT_URL = os.environ.get("LIVEKIT_URL", "ws://localhost:7888")
+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")
+MIMO_KEY = os.environ.get("MIMO_KEY", "")
+
+# VAD params
+VAD_THRESHOLD = 0.015       # energy threshold
+MIN_SPEECH_S = 0.3           # minimum speech before considering utterance
+MIN_SILENCE_S = 1.2          # silence to end utterance
+
+
+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")
+
+
+async def _llm(prompt: str, hist: list[dict]) -> str:
+    msgs = [{"role": "system", "content": "你是友好的中文语音助手,回答简洁自然,2-3句即可。"}]
+    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": "default", "messages": msgs, "max_tokens": 128, "temperature": 0.7},
+            timeout=aiohttp.ClientTimeout(total=15)) as r:
+            return (await r.json())["choices"][0]["message"]["content"]
+
+
+async def _tts(text: str) -> np.ndarray | None:
+    key = os.environ.get("MIMO_KEY", "")
+    if not key: return None
+    H = {"api-key": key, "Content-Type": "application/json"}
+    async with aiohttp.ClientSession() as s:
+        async with s.post("https://token-plan-cn.xiaomimimo.com/v1/chat/completions",
+            json={"model": "mimo-v2.5-tts", "messages": [
+                {"role": "user", "content": "请用自然语速朗读。"},
+                {"role": "assistant", "content": text}], "audio": {"format": "wav", "voice": "Chloe"}},
+            headers=H, timeout=aiohttp.ClientTimeout(total=20)) as r:
+            try:
+                b64 = (await r.json())["choices"][0]["message"]["audio"]["data"]
+                return np.frombuffer(base64.b64decode(b64), dtype=np.int16)
+            except: return None
+
+
+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
+        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):
+        buf = AudioBuffer(max_duration=10, sample_rate=16000)
+        vad = VADProcessor(
+            threshold=VAD_THRESHOLD,
+            min_speech_duration=MIN_SPEECH_S,
+            min_silence_duration=MIN_SILENCE_S,
+            sample_rate=16000,
+        )
+        busy = False
+        cur_tts = None
+
+        async def _acc():
+            async for ev in rtc.AudioStream(track, sample_rate=16000, num_channels=1):
+                if not busy:
+                    data = ev.frame.data
+                    arr = np.frombuffer(data, dtype=np.int16).astype(np.float32) / 32768.0
+                    buf.append(arr)
+
+                    # Feed VAD with audio frames
+                    in_speech = vad.in_speech
+                    segments = vad.process(arr)
+                    # Speech just ended
+                    if not vad.in_speech and in_speech and len(buf.buffer) > 1600 * MIN_SPEECH_S:
+                        nonlocal cur_tts
+                        if cur_tts: cur_tts = None
+                        busy = True
+                        await self._transcribe_and_respond(buf)
+                        busy = False
+
+        async def _loop():
+            # Fallback: periodic transcription when VAD detects speech
+            last_len = 0
+            while True:
+                await asyncio.sleep(1.5)
+                if busy: continue
+                if len(buf.buffer) <= 1600: continue
+                if not vad.in_speech: continue
+                # Only transcribe if buffer grew (new speech happened)
+                if len(buf.buffer) == last_len: continue
+                last_len = len(buf.buffer)
+
+                a = buf.get_last(4.0)
+                if len(a) <= 1600: continue
+                loop = asyncio.get_running_loop()
+                result = await loop.run_in_executor(None, self.asr.transcribe_array, a)
+                txt = result.get("text", "").strip()
+                if txt:
+                    await self._send({"type": "partial", "text": txt, "seq": 0})
+
+        t = asyncio.create_task(_loop())
+        try: await _acc()
+        finally: t.cancel()
+
+    async def _transcribe_and_respond(self, buf: AudioBuffer):
+        all_audio = buf.get_all()
+        buf.clear()
+        if len(all_audio) <= 1600: return
+
+        loop = asyncio.get_running_loop()
+        result = await loop.run_in_executor(None, self.asr.transcribe_array, all_audio)
+        txt = result.get("text", "").strip()
+        if not txt: return
+
+        await self._send({"type": "utterance", "text": txt, "seq": 0})
+        await self._reply(txt)
+
+    async def _reply(self, text: str):
+        try: rep = await _llm(text, self.hist)
+        except: rep = "抱歉。"
+        await self._send({"type": "reply", "text": rep})
+        try:
+            a = await _tts(rep)
+            if a is not None and len(a) > 0:
+                await self._play(a, 24000)
+        except (asyncio.CancelledError, Exception):
+            pass
+        self.hist.extend([{"role": "user", "content": text}, {"role": "assistant", "content": rep}])
+        self.hist[:] = self.hist[-20:]
+
+    async def _play(self, audio: np.ndarray, sr: int):
+        src = rtc.AudioSource(sr, 1)
+        tk = rtc.LocalAudioTrack.create_audio_track("tts", src)
+        await self.room.local_participant.publish_track(tk)
+        for i in range(0, len(audio), sr // 50):
+            c = audio[i : i + sr // 50]
+            await src.capture_frame(rtc.AudioFrame(data=c.tobytes(), sample_rate=sr, num_channels=1, samples_per_channel=len(c)))
+            await asyncio.sleep(0.018)
+        return tk
+
+    async def _send(self, msg: dict):
+        try: await self.room.local_participant.publish_data(json.dumps(msg, ensure_ascii=False).encode(), reliable=True, topic="transcription")
+        except: 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())

+ 220 - 0
asr_agent/worker.py

@@ -0,0 +1,220 @@
+"""
+LiveKit ASR Worker — subscribes to participant audio, runs Qwen ASR,
+and sends transcriptions back as room data messages.
+"""
+
+from __future__ import annotations
+
+import asyncio
+import json
+import logging
+import os
+import sys
+import time
+import traceback
+
+import jwt
+import numpy as np
+from livekit import rtc
+
+# Import whisper_asr modules (same directory or adjacent).
+_asr_root = os.path.join(os.path.dirname(__file__), "whisper_asr")
+sys.path.insert(0, _asr_root)
+from audio_processor import AudioBuffer      # noqa: E402
+from qwen_engine import QwenASREngine        # noqa: E402
+from transcript_processor import TranscriptProcessor  # noqa: E402
+
+logger = logging.getLogger("asr-worker")
+
+LIVEKIT_URL = os.environ.get("LIVEKIT_URL", "ws://localhost:10003")
+API_KEY = "devkey"
+API_SECRET = "secretsecretsecretsecretsecret12"
+ASR_MODEL_PATH = os.environ.get("ASR_MODEL_PATH", "/data/models/Qwen3-ASR")
+TTS_MODEL_PATH = os.environ.get("TTS_MODEL_PATH", "/data/models/Qwen3-TTS")
+
+TRANSCRIBE_INTERVAL = 2.0
+BUFFER_DURATION = 5.0
+SILENCE_REPEATS = 3
+SILENCE_SECONDS = 3.0
+
+
+def create_token(room_name: str, identity: str) -> str:
+    now = int(time.time())
+    return jwt.encode(
+        {
+            "iss": API_KEY,
+            "sub": identity,
+            "name": identity,
+            "nbf": now - 60,
+            "exp": now + 6 * 3600,
+            "video": {
+                "roomJoin": True,
+                "room": room_name,
+                "canPublish": True,
+                "canSubscribe": True,
+                "canPublishData": True,
+            },
+        },
+        API_SECRET,
+        algorithm="HS256",
+    )
+
+
+class ASRWorker:
+    def __init__(self, room_name: str, identity: str = "asr-bot"):
+        self._room_name = room_name
+        self._identity = identity
+        self._room = rtc.Room()
+        self._engine: QwenASREngine | None = None
+        self._processors: dict[str, TranscriptProcessor] = {}
+        self._buffers: dict[str, AudioBuffer] = {}
+        self._tasks: dict[str, asyncio.Task] = {}
+
+    async def run(self):
+        logger.info("Loading Qwen ASR model...")
+        self._engine = QwenASREngine(
+            model_id=ASR_MODEL_PATH, language=None
+        )
+
+        self._room.on("track_subscribed", self._on_track_subscribed)
+        self._room.on(
+            "participant_disconnected", self._on_participant_disconnected
+        )
+        self._room.on("participant_connected", lambda p: None)
+        self._room.on("track_published", lambda pub, p: None)
+
+        token = create_token(self._room_name, self._identity)
+        logger.info(f"Connecting to room: {self._room_name}")
+        await self._room.connect(LIVEKIT_URL, token)
+        logger.info(f"ASR worker ready. Participants: {len(self._room.remote_participants)}")
+
+        # Subscribe to any tracks that already exist in the room.
+        for p in self._room.remote_participants.values():
+            logger.info(f"Existing participant: {p.identity}, tracks: {len(p.track_publications)}")
+            for pub in p.track_publications.values():
+                logger.info(f"  pub kind={pub.kind}, source={pub.source.name if pub.source else 'N/A'}, track={type(pub.track).__name__ if pub.track else 'None'}, subscribed={pub.subscribed}")
+                if pub.track is not None and pub.kind == rtc.TrackKind.KIND_AUDIO:
+                    self._on_track_subscribed(pub.track, pub, p)
+
+        logger.info(f"Subscribed tracks: {list(self._buffers.keys())}")
+
+        try:
+            await asyncio.Future()
+        finally:
+            await self._room.disconnect()
+
+    def _on_track_subscribed(
+        self,
+        track: rtc.RemoteAudioTrack,
+        publication: rtc.RemoteTrackPublication,
+        participant: rtc.RemoteParticipant,
+    ):
+        sid = participant.sid
+        if sid in self._buffers:
+            return
+
+        logger.info(f"Audio track subscribed: {participant.identity} ({sid})")
+        logger.info(f"  track kind={track.kind}, muted={track.muted}, stream_state={track.stream_state}")
+
+        self._buffers[sid] = AudioBuffer(
+            max_duration=BUFFER_DURATION * 2,
+            sample_rate=16000,
+        )
+        self._processors[sid] = TranscriptProcessor(
+            silence_repeats=SILENCE_REPEATS,
+            silence_seconds=SILENCE_SECONDS,
+        )
+        self._tasks[sid] = asyncio.create_task(self._process(sid, track))
+
+    def _on_participant_disconnected(self, participant: rtc.RemoteParticipant):
+        sid = participant.sid
+        t = self._tasks.pop(sid, None)
+        if t:
+            t.cancel()
+        self._buffers.pop(sid, None)
+        self._processors.pop(sid, None)
+
+    async def _process(self, sid: str, track: rtc.RemoteAudioTrack):
+        buf = self._buffers[sid]
+        proc = self._processors[sid]
+
+        async def _accumulate():
+            try:
+                stream = rtc.AudioStream(track, sample_rate=16000, num_channels=1)
+                count = 0
+                async for ev in stream:
+                    frame = ev.frame
+                    data = frame.data
+                    count += 1
+                    if count == 1:
+                        logger.info(f"[{sid}] first audio frame: {len(data)} bytes, sr={frame.sample_rate}, ch={frame.num_channels}")
+                    arr = (
+                        np.frombuffer(data, dtype=np.int16).astype(np.float32)
+                        / 32768.0
+                    )
+                    buf.append(arr)
+                logger.info(f"[{sid}] stream ended after {count} frames")
+            except Exception:
+                logger.exception(f"[{sid}] accumulate failed")
+
+        async def _transcribe():
+            tick = 0
+            while True:
+                await asyncio.sleep(TRANSCRIBE_INTERVAL)
+                tick += 1
+                audio = buf.get_last(BUFFER_DURATION)
+                if len(audio) <= 1600:
+                    continue
+
+                loop = asyncio.get_running_loop()
+                try:
+                    result = await loop.run_in_executor(
+                        None, self._engine.transcribe_array, audio
+                    )
+                except Exception:
+                    logger.exception("transcription failed")
+                    continue
+
+                raw = result.get("text", "").strip()
+                if tick % 5 == 0:
+                    logger.info(f"[{sid}] transcribe tick={tick} raw='{raw}' buffer_samples={len(audio)}")
+                if not raw:
+                    continue
+
+                for msg in proc.feed(raw):
+                    logger.info(f"[{sid}] msg: {msg['type']} '{msg['text'][:50]}'")
+                    await self._send(msg)
+
+        t = asyncio.create_task(_transcribe())
+        try:
+            await _accumulate()
+        finally:
+            t.cancel()
+
+    async def _send(self, message: dict):
+        try:
+            await self._room.local_participant.publish_data(
+                json.dumps(message, ensure_ascii=False).encode(),
+                reliable=True,
+                topic="transcription",
+            )
+        except Exception:
+            logger.exception("publish_data failed")
+
+
+async def main():
+    import argparse
+
+    parser = argparse.ArgumentParser()
+    parser.add_argument("--room", required=True)
+    parser.add_argument("--identity", default="asr-bot")
+    args = parser.parse_args()
+
+    logging.basicConfig(level=logging.INFO)
+
+    worker = ASRWorker(room_name=args.room, identity=args.identity)
+    await worker.run()
+
+
+if __name__ == "__main__":
+    asyncio.run(main())