"""Translate the unchanged Tencent demo protocol to FunASR's native online WS.""" from __future__ import annotations import argparse import asyncio import json import logging import os from dataclasses import dataclass from pathlib import Path from typing import Any from uuid import uuid4 from aiohttp import WSMsgType, web from dotenv import load_dotenv try: from websockets.asyncio.client import connect as websocket_connect except ImportError: # websockets before 13 exposes the same client at package root. from websockets import connect as websocket_connect try: from .auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig except ImportError: from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig PROJECT_ROOT = Path(__file__).resolve().parents[2] load_dotenv(PROJECT_ROOT / ".env") LOGGER = logging.getLogger(__name__) WEB_HOST = os.getenv("WEB_HOST", "0.0.0.0") WEB_PORT = int(os.getenv("WEB_PORT", "8082")) NATIVE_WS_URL = os.getenv("FUNASR_NATIVE_WS_URL", "ws://127.0.0.1:10095") CHUNK_SIZE = tuple( int(part.strip()) for part in os.getenv("FUNASR_CHUNK_SIZE", "0,10,5").split(",") ) if len(CHUNK_SIZE) != 3: CHUNK_SIZE = (0, 10, 5) CHUNK_INTERVAL = max(1, int(os.getenv("FUNASR_CHUNK_INTERVAL", "10"))) SAMPLE_RATE = 16000 PCM_BYTES_PER_MS = SAMPLE_RATE * 2 / 1000 FRAME_BYTES = max(2, round(60 * CHUNK_SIZE[1] / CHUNK_INTERVAL * PCM_BYTES_PER_MS)) MAX_SPEAKER_AUDIO_BYTES = 60 * SAMPLE_RATE * 2 MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2) FINALIZE_TIMEOUT_SECONDS = max(30, int(os.getenv("FUNASR_FINALIZE_TIMEOUT_SECONDS", "300"))) TURN_SENTENCE_ID_STRIDE = 100 AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService) SESSION_REGISTRY_KEY = web.AppKey("sessions", dict) class IncrementalWavDecoder: """Read a streamed PCM WAV header and yield its 16 kHz mono PCM payload.""" def __init__(self) -> None: self.buffer = bytearray() self.header_done = False self.format_valid = False self.data_remaining: int | None = None def feed(self, data: bytes) -> bytes: if self.header_done: if self.data_remaining is None: return data payload = data[: self.data_remaining] self.data_remaining -= len(payload) return payload self.buffer.extend(data) if len(self.buffer) < 12: return b"" if self.buffer[:4] != b"RIFF" or self.buffer[8:12] != b"WAVE": raise ValueError("WAV file must use a RIFF/WAVE container") del self.buffer[:12] while len(self.buffer) >= 8: kind = bytes(self.buffer[:4]) size = int.from_bytes(self.buffer[4:8], "little") if size > 1024 * 1024: raise ValueError("WAV header chunk is unexpectedly large") full_size = 8 + size + (size % 2) if len(self.buffer) < full_size: return b"" body = bytes(self.buffer[8 : 8 + size]) if kind == b"fmt ": if size < 16: raise ValueError("WAV fmt chunk is incomplete") fmt = ( int.from_bytes(body[0:2], "little"), int.from_bytes(body[2:4], "little"), int.from_bytes(body[4:8], "little"), int.from_bytes(body[14:16], "little"), ) if fmt != (1, 1, SAMPLE_RATE, 16): raise ValueError("WAV must be PCM16, mono, 16 kHz") self.format_valid = True if kind == b"data": if not self.format_valid or size % 2: raise ValueError("WAV must be PCM16, mono, 16 kHz") self.header_done = True self.data_remaining = size del self.buffer[:8] payload = bytes(self.buffer[:size]) del self.buffer[: min(size, len(self.buffer))] self.data_remaining -= len(payload) return payload del self.buffer[:full_size] return b"" def finish(self) -> None: if not self.header_done or self.data_remaining not in (None, 0): raise ValueError("WAV ended before its complete PCM data chunk arrived") @dataclass(frozen=True) class SpeakerJob: sentence_id: int text: str audio: bytes start_time_ms: float end_time_ms: float def split_text_by_speaker_segments( text: str, segments: list[dict[str, Any]], turn_start_ms: float, turn_end_ms: float, ) -> list[dict[str, Any]]: """按 CAM++ 时间段近似拆分 ASR 文本,并保留每段稳定说话人标签。""" usable: list[dict[str, Any]] = [] for segment in sorted( segments, key=lambda item: float(item.get("start_time", turn_start_ms)), ): try: start = max(turn_start_ms, float(segment.get("start_time", turn_start_ms))) end = min(turn_end_ms, float(segment.get("end_time", turn_end_ms))) speaker_id = int(segment.get("speaker_id", -1)) except (TypeError, ValueError): continue if end <= start or speaker_id < 0: continue speaker = { "speaker_id": speaker_id, "speaker_name": str(segment.get("speaker_name") or f"说话人 {speaker_id + 1}"), "speaker_confidence": float(segment.get("speaker_confidence") or 0), "speaker_status": str(segment.get("speaker_status") or "confirmed"), } if usable and usable[-1]["speaker"]["speaker_id"] == speaker_id: usable[-1]["end_time_ms"] = end else: usable.append( {"start_time_ms": start, "end_time_ms": end, "speaker": speaker} ) if not usable: return [] if len(usable) == 1: usable[0]["text"] = text return usable # 双重保护:服务端有最短段长滤波,桥接层也拒绝意外的短片段响应。 min_split_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_MIN_SEGMENT_MS", "3000"))) while len(usable) > 1: short_index = next( ( index for index, item in enumerate(usable) if item["end_time_ms"] - item["start_time_ms"] < min_split_ms ), None, ) if short_index is None: break if short_index == 0: usable[1]["start_time_ms"] = usable[0]["start_time_ms"] del usable[0] elif short_index == len(usable) - 1: usable[-2]["end_time_ms"] = usable[-1]["end_time_ms"] del usable[-1] else: previous = usable[short_index - 1] following = usable[short_index + 1] if previous["end_time_ms"] - previous["start_time_ms"] >= ( following["end_time_ms"] - following["start_time_ms"] ): previous["end_time_ms"] = usable[short_index]["end_time_ms"] del usable[short_index] else: following["start_time_ms"] = usable[short_index]["start_time_ms"] del usable[short_index] if len(usable) == 1 or len(text) < len(usable): usable = [max(usable, key=lambda item: item["end_time_ms"] - item["start_time_ms"])] usable[0]["text"] = text return usable total_duration = sum(item["end_time_ms"] - item["start_time_ms"] for item in usable) char_start = 0 elapsed = 0.0 for index, item in enumerate(usable): elapsed += item["end_time_ms"] - item["start_time_ms"] if index == len(usable) - 1: char_end = len(text) else: char_end = round(len(text) * elapsed / total_duration) char_end = max(char_start + 1, min(char_end, len(text) - (len(usable) - index - 1))) item["text"] = text[char_start:char_end].strip() char_start = char_end return [item for item in usable if item.get("text")] class BrowserSession: """Own one browser/native WS pair and translate their message contracts.""" def __init__( self, browser_ws: web.WebSocketResponse, auxiliary: AuxiliaryModelService, start: dict[str, Any], voice_id: str, ) -> None: self.browser_ws = browser_ws self.auxiliary = auxiliary self.start = start self.voice_id = voice_id self.session_id = uuid4().hex self.stop_event = asyncio.Event() self.send_lock = asyncio.Lock() self.speaker_jobs: asyncio.Queue[SpeakerJob | None] = asyncio.Queue() self.wav_decoder = ( IncrementalWavDecoder() if str(start.get("source") or "mic") == "file" and Path(str(start.get("file_name") or "")).suffix.lower() == ".wav" else None ) self.speed_factor = max(0.5, min(3.0, float(start.get("speed_factor") or 1.0))) try: sentence_strategy = int(start.get("sentence_strategy", 0)) except (TypeError, ValueError): sentence_strategy = 0 self.sentence_strategy = sentence_strategy if sentence_strategy in (0, 1) else 0 self.pending_pcm = bytearray() self.turn_audio = bytearray() self.total_audio_ms = 0.0 self.turn_start_ms = 0.0 self.turn_text = "" self.sentence_id = 0 self.native_ack: dict[str, Any] | None = None self.native_error: str | None = None async def emit(self, payload: dict[str, Any]) -> None: """Serialize browser writes because ASR and CAM++ finish independently.""" async with self.send_lock: if not self.browser_ws.closed: await self.browser_ws.send_json(payload) async def emit_sentence( self, text: str, final: bool, speaker: dict[str, Any] | None = None, sentence_id: int | None = None, start_time_ms: float | None = None, end_time_ms: float | None = None, ) -> None: speaker = speaker or {} sentence = { "sentence_id": self.sentence_id if sentence_id is None else sentence_id, "sentence": text, "sentence_type": 1 if final else 0, "start_time": round(self.turn_start_ms if start_time_ms is None else start_time_ms), "end_time": round(self.total_audio_ms if end_time_ms is None else end_time_ms), # The unchanged Tencent UI uses speaker_id to choose its speaker bubble. "speaker_id": int(speaker.get("speaker_id", -1)), "speaker_name": str(speaker.get("speaker_name") or ""), "speaker_confidence": float(speaker.get("speaker_confidence") or 0), "speaker_status": str(speaker.get("speaker_status") or "pending"), } await self.emit({"type": "sentences", "sentences": [sentence]}) def align_turn_audio_to_vad(self, message: dict[str, Any]) -> None: """用原生 VAD 的绝对时间裁掉当前 turn 的前后静音样本。""" buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS try: speech_start_ms = message.get("speech_start_ms") if speech_start_ms is not None: target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_start_ms))) trim_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS)) trim_bytes -= trim_bytes % 2 del self.turn_audio[:trim_bytes] self.turn_start_ms += trim_bytes / PCM_BYTES_PER_MS buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS speech_end_ms = message.get("speech_end_ms") if speech_end_ms is not None: target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_end_ms))) keep_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS)) keep_bytes -= keep_bytes % 2 del self.turn_audio[keep_bytes:] except (TypeError, ValueError): LOGGER.warning("Ignoring invalid native FunASR VAD boundary: %s", message) async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None: if not frame: return if len(frame) % 2: raise ValueError("PCM16 audio ended on an incomplete sample") self.total_audio_ms += len(frame) / PCM_BYTES_PER_MS self.turn_audio.extend(frame) if len(self.turn_audio) > MAX_SPEAKER_AUDIO_BYTES: # Bound per-turn RAM even if VAD never reports an endpoint. trim = len(self.turn_audio) - MAX_SPEAKER_AUDIO_BYTES del self.turn_audio[:trim] self.turn_start_ms += trim / PCM_BYTES_PER_MS await native_ws.send(frame) if pace_file: await asyncio.sleep(len(frame) / (SAMPLE_RATE * 2) / self.speed_factor) async def accept_audio(self, native_ws: Any, data: bytes) -> None: pcm = self.wav_decoder.feed(data) if self.wav_decoder else data if not pcm: return self.pending_pcm.extend(pcm) is_file = str(self.start.get("source") or "mic") == "file" while len(self.pending_pcm) >= FRAME_BYTES: frame = bytes(self.pending_pcm[:FRAME_BYTES]) del self.pending_pcm[:FRAME_BYTES] await self.send_pcm_frame(native_ws, frame, is_file) async def finish_audio(self, native_ws: Any) -> None: if self.wav_decoder: self.wav_decoder.finish() if self.pending_pcm: frame = bytes(self.pending_pcm) self.pending_pcm.clear() await self.send_pcm_frame( native_ws, frame, str(self.start.get("source") or "mic") == "file" ) async def read_native(self, native_ws: Any) -> None: """Consume native FunASR events and keep its per-utterance partial cache.""" try: while True: raw = await native_ws.recv() message = json.loads(raw) if message.get("is_end"): self.native_ack = message return if message.get("event") == "vad": self.align_turn_audio_to_vad(message) continue text = str(message.get("text") or "") if text: # FunASR online sends the newly decoded text for each chunk. self.turn_text += text if message.get("is_final"): final_text = self.turn_text.strip() audio = bytes(self.turn_audio) start_ms = self.turn_start_ms end_ms = start_ms + len(audio) / PCM_BYTES_PER_MS turn_sentence_id = self.sentence_id if final_text: # 先结束前端 interim 气泡;标点和 CAM++ 在独立 worker 中完成, # 不阻塞 native WS 继续读取后续音频帧。 await self.emit_sentence( final_text, final=True, sentence_id=turn_sentence_id, start_time_ms=start_ms, end_time_ms=end_ms, ) await self.speaker_jobs.put( SpeakerJob( sentence_id=turn_sentence_id, text=final_text, audio=audio, start_time_ms=start_ms, end_time_ms=end_ms, ) ) # 为同一个 VAD turn 内可能拆出的多个气泡预留独立 ID。 self.sentence_id += TURN_SENTENCE_ID_STRIDE self.turn_text = "" self.turn_audio.clear() self.turn_start_ms = self.total_audio_ms elif self.turn_text: await self.emit_sentence(self.turn_text, final=False) except asyncio.CancelledError: raise except Exception as exc: self.native_error = str(exc) LOGGER.exception("native FunASR WebSocket closed unexpectedly") await self.emit({"type": "error", "message": f"FunASR realtime WS: {exc}"}) async def resolve_speakers(self) -> None: """按序定稿标点并用 CAM++ 滑窗恢复 turn 内说话人切换。""" while True: job = await self.speaker_jobs.get() try: if job is None: return final_text = job.text try: punctuation = await self.auxiliary.punctuate(final_text) if not punctuation.get("available", False): LOGGER.warning( "FunASR punctuation is unavailable: %s", punctuation.get("error", "model is not configured"), ) else: punctuated = str(punctuation.get("text") or "").strip() if punctuated: final_text = punctuated except Exception: LOGGER.exception("FunASR punctuation request failed; keeping raw text") subsegments: list[dict[str, Any]] = [] if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES: # required CAM++ 按 FunASR 1.5s/0.75s 滑窗识别同一 VAD turn 内的 # 多人切换;若窗长不足或接口暂不可用,再退回整段声纹验证。 track = getattr(self.auxiliary, "track_speakers", None) if callable(track): try: tracked = await track( job.audio, self.session_id, job.start_time_ms, job.end_time_ms, ) subsegments = split_text_by_speaker_segments( final_text, tracked, job.start_time_ms, job.end_time_ms, ) except Exception: LOGGER.exception( "CAM++ sliding-window diarization failed; falling back to whole-turn speaker" ) if not subsegments and len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES: try: resolved = await self.auxiliary.resolve_speaker( job.audio, self.session_id, job.start_time_ms, job.end_time_ms, ) if resolved and int(resolved.get("speaker_id", -1)) >= 0: subsegments = [ { "text": final_text, "speaker": resolved, "start_time_ms": job.start_time_ms, "end_time_ms": job.end_time_ms, } ] except Exception: LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id) if not subsegments: subsegments = [ { "text": final_text, "speaker": {"speaker_id": -1}, "start_time_ms": job.start_time_ms, "end_time_ms": job.end_time_ms, } ] for index, segment in enumerate(subsegments): await self.emit_sentence( str(segment["text"]), final=True, speaker=segment.get("speaker"), sentence_id=job.sentence_id + index, start_time_ms=float(segment.get("start_time_ms", job.start_time_ms)), end_time_ms=float(segment.get("end_time_ms", job.end_time_ms)), ) finally: self.speaker_jobs.task_done() async def config_handler(_: web.Request) -> web.Response: """Expose a small readiness response for the launcher and diagnostics.""" return web.json_response( { "engine": "funasr-native-online-ws", "model": os.getenv("FUNASR_ASR_MODEL", ""), "native_ws_url": NATIVE_WS_URL, "speaker_service_url": os.getenv( "AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010" ), } ) async def stop_handler(request: web.Request) -> web.Response: voice_id = request.query.get("voice_id", "").strip() if not voice_id: return web.json_response({"ok": False, "error": "missing voice_id"}, status=400) stop_event = request.app[SESSION_REGISTRY_KEY].get(voice_id) if stop_event is None: return web.json_response({"ok": False, "error": "session not found"}) stop_event.set() return web.json_response({"ok": True}) async def websocket_handler(request: web.Request) -> web.WebSocketResponse: """Bridge Tencent's browser messages to FunASR's native realtime protocol.""" browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30) await browser_ws.prepare(request) session: BrowserSession | None = None native_reader: asyncio.Task[None] | None = None speaker_worker: asyncio.Task[None] | None = None voice_id = "" registered = False try: first = await browser_ws.receive() if first.type != WSMsgType.TEXT: await browser_ws.send_json({"type": "error", "message": "first message must be JSON start"}) return browser_ws start = json.loads(first.data) if not isinstance(start, dict) or start.get("type") != "start": await browser_ws.send_json({"type": "error", "message": "first message must have type=start"}) return browser_ws source = str(start.get("source") or "mic") suffix = Path(str(start.get("file_name") or "")).suffix.lower() if source == "file" and suffix not in {".pcm", ".wav"}: await browser_ws.send_json( { "type": "error", "message": "FunASR realtime mode accepts PCM or 16 kHz mono PCM WAV files.", } ) return browser_ws if source == "file" and suffix == ".pcm": start = dict(start) start["file_name"] = str(start.get("file_name") or "audio.pcm") voice_id = uuid4().hex auxiliary: AuxiliaryModelService = request.app[AUXILIARY_KEY] session = BrowserSession(browser_ws, auxiliary, start, voice_id) request.app[SESSION_REGISTRY_KEY][voice_id] = session.stop_event registered = True await session.emit({"type": "voice_id", "voice_id": voice_id}) await session.emit({"type": "start"}) # This is FunASR's native WSS message contract; PCM frames follow at 60 ms. async with websocket_connect( NATIVE_WS_URL, subprotocols=["binary"], ping_interval=None, close_timeout=3, max_size=None, ) as native_ws: await native_ws.send( json.dumps( { "mode": "online", "chunk_size": list(CHUNK_SIZE), "chunk_interval": CHUNK_INTERVAL, "encoder_chunk_look_back": int( os.getenv("FUNASR_ENCODER_LOOK_BACK", "4") ), "decoder_chunk_look_back": int( os.getenv("FUNASR_DECODER_LOOK_BACK", "1") ), "sentence_strategy": session.sentence_strategy, "audio_fs": SAMPLE_RATE, "wav_name": voice_id, "is_speaking": True, }, ensure_ascii=False, ) ) native_reader = asyncio.create_task(session.read_native(native_ws)) speaker_worker = asyncio.create_task(session.resolve_speakers()) while not browser_ws.closed: receive_task = asyncio.create_task(browser_ws.receive()) stop_task = asyncio.create_task(session.stop_event.wait()) done, _ = await asyncio.wait( [receive_task, stop_task, native_reader], return_when=asyncio.FIRST_COMPLETED, ) if stop_task in done: receive_task.cancel() await asyncio.gather(receive_task, return_exceptions=True) break stop_task.cancel() await asyncio.gather(stop_task, return_exceptions=True) if native_reader in done: receive_task.cancel() await asyncio.gather(receive_task, return_exceptions=True) if session.native_ack is None and session.native_error is None: session.native_error = "FunASR native WebSocket ended before EOF acknowledgement" break message = await receive_task if message.type == WSMsgType.BINARY: await session.accept_audio(native_ws, bytes(message.data)) continue if message.type == WSMsgType.TEXT: try: control = json.loads(message.data) except json.JSONDecodeError: continue if isinstance(control, dict) and control.get("type") in {"eof", "stop"}: break if isinstance(control, dict) and control.get("type") == "abort": session.native_error = "session aborted by browser" break if message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}: break if session.native_error is None and not browser_ws.closed: await session.finish_audio(native_ws) # FunASR flushes its online cache and acknowledges only after final output. await native_ws.send( json.dumps({"is_speaking": False, "is_end": True}, ensure_ascii=False) ) try: await asyncio.wait_for(native_reader, timeout=FINALIZE_TIMEOUT_SECONDS) except asyncio.TimeoutError: session.native_error = ( f"FunASR did not acknowledge end-of-input within " f"{FINALIZE_TIMEOUT_SECONDS}s" ) if speaker_worker is not None: await session.speaker_jobs.join() await session.speaker_jobs.put(None) await speaker_worker speaker_worker = None if session.native_error: await session.emit({"type": "error", "message": session.native_error}) elif session.native_ack and not session.native_ack.get("is_final", False): await session.emit( { "type": "error", "message": str(session.native_ack.get("error") or "FunASR did not finalize the stream"), } ) elif session.native_ack: await session.emit({"type": "end"}) except asyncio.CancelledError: raise except Exception as exc: LOGGER.exception("Tencent-compatible WebSocket session failed") if not browser_ws.closed: await browser_ws.send_json({"type": "error", "message": str(exc)}) finally: if native_reader is not None and not native_reader.done(): native_reader.cancel() await asyncio.gather(native_reader, return_exceptions=True) if speaker_worker is not None and not speaker_worker.done(): speaker_worker.cancel() await asyncio.gather(speaker_worker, return_exceptions=True) if registered: request.app[SESSION_REGISTRY_KEY].pop(voice_id, None) if session is not None: await session.auxiliary.reset_speaker_session(session.session_id) if not browser_ws.closed: await browser_ws.close() return browser_ws async def create_app() -> web.Application: """Create a light protocol bridge; model inference belongs to native FunASR.""" app = web.Application() app[AUXILIARY_KEY] = AuxiliaryModelService( AuxiliaryServiceConfig( base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010") ) ) app[SESSION_REGISTRY_KEY] = {} async def lifecycle(application: web.Application): await application[AUXILIARY_KEY].start() try: health = await application[AUXILIARY_KEY].health() if not health.get("speaker_embedding_ready"): raise RuntimeError("CAM++ speaker service is not ready") except Exception: await application[AUXILIARY_KEY].close() raise yield await application[AUXILIARY_KEY].close() app.cleanup_ctx.append(lifecycle) app.router.add_get("/api/config", config_handler) app.router.add_get("/api/stop", stop_handler) app.router.add_get("/ws", websocket_handler) return app def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--no-browser", action="store_true") args = parser.parse_args() logging.basicConfig(level=logging.INFO) print( f"FunASR browser bridge: http://{os.getenv('WEB_DISPLAY_HOST', '127.0.0.1')}:{WEB_PORT}/api/config", flush=True, ) web.run_app(create_app(), host=WEB_HOST, port=WEB_PORT) if __name__ == "__main__": main()