"""Browser-facing WebSocket adapter for the migrated FunASR engine. The frontend contract remains the existing start/binary/stop protocol. All speech activity detection and streaming ASR state now come from FunASR; this module only translates engine events into the existing sentence snapshots. """ from __future__ import annotations import argparse import asyncio import json import logging import os import time import webbrowser from dataclasses import dataclass, replace from pathlib import Path from typing import Any from uuid import uuid4 from aiohttp import WSMsgType, web from dotenv import load_dotenv try: from .auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig from .funasr_engine import FunASRModelService, FunASRSegment, FunASRServiceConfig from .speaker_assembler import SegmentAssembler except ImportError: from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig from funasr_engine import FunASRModelService, FunASRSegment, FunASRServiceConfig from speaker_assembler import SegmentAssembler PROJECT_ROOT = Path(__file__).resolve().parents[1] load_dotenv(PROJECT_ROOT / ".env") WEB_HOST = os.getenv("WEB_HOST", "0.0.0.0") WEB_PORT = int(os.getenv("WEB_PORT", "8082")) WEB_DISPLAY_HOST = os.getenv("WEB_DISPLAY_HOST", "127.0.0.1") LOCAL_ENGINE_URL = "local://funasr" PARTIAL_BYTES_PER_SECOND = 16000 * 2 MIN_SPEAKER_VOICE_MS = 800 LOGGER = logging.getLogger(__name__) class EndOfStream: """Queue marker that cannot be confused with an audio frame.""" EOF = EndOfStream() MODEL_SERVICE_KEY = web.AppKey("model_service", FunASRModelService) AUXILIARY_SERVICE_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService) @dataclass(frozen=True) class SpeakerJob: """Completed FunASR turn waiting for optional CAM++ speaker matching.""" sentence_id: int audio: bytes start_time_ms: float end_time_ms: float voiced_ms: float @dataclass class SessionMetrics: """Small runtime snapshot shown in the existing frontend.""" started_at: float audio_bytes: int = 0 input_chunks: int = 0 partial_count: int = 0 partial_revisions: int = 0 first_partial_ms: float | None = None final_ms: float | None = None def snapshot(self) -> dict[str, Any]: """Return JSON-safe metrics relative to session start.""" return { "audio_bytes": self.audio_bytes, "input_chunks": self.input_chunks, "partial_count": self.partial_count, "partial_revisions": self.partial_revisions, "first_partial_ms": self.first_partial_ms, "final_ms": self.final_ms, "elapsed_ms": round((time.perf_counter() - self.started_at) * 1000, 1), } class IncrementalWavDecoder: """Strip a streamed RIFF header before forwarding PCM16 to FunASR.""" def __init__(self) -> None: self.buffer = bytearray() self.payload_started = False self.riff_read = False self.format_valid = False self.data_remaining: int | None = None def feed(self, chunk: bytes) -> bytes: """Parse complete RIFF chunks without buffering the whole recording.""" if self.payload_started: if self.data_remaining is None: return chunk payload = chunk[: self.data_remaining] self.data_remaining -= len(payload) return payload self.buffer.extend(chunk) if not self.riff_read: if len(self.buffer) < 12: return b"" if self.buffer[:4] != b"RIFF" or self.buffer[8:12] != b"WAVE": raise ValueError("文件不是有效的 RIFF/WAV 音频") del self.buffer[:12] self.riff_read = True 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 元数据头过大,请转换为标准 PCM WAV") chunk_size = 8 + size + (size % 2) if len(self.buffer) < chunk_size: return b"" body = self.buffer[8 : 8 + size] if kind == b"fmt ": if size < 16: raise ValueError("WAV fmt 区块不完整") fields = ( 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 fields != (1, 1, 16000, 16): raise ValueError("WAV 必须为 16kHz、单声道、PCM16") self.format_valid = True if kind == b"data": if not self.format_valid or size % 2: raise ValueError("WAV 必须为 16kHz、单声道、PCM16") self.payload_started = 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[:chunk_size] return b"" def finish(self) -> None: """Reject a truncated or header-only WAV before final inference.""" if not self.payload_started: raise ValueError("WAV 文件不完整,未收到 data 音频区块") if self.data_remaining not in (None, 0): raise ValueError("WAV 文件不完整,未收到全部音频数据") class RealtimeSession: """Translate FunASR events into the existing sentence/display protocol.""" def __init__( self, ws: web.WebSocketResponse, model_service: FunASRModelService, auxiliary_service: AuxiliaryModelService | None, start: dict[str, Any], ) -> None: self.ws = ws self.model_service = model_service self.auxiliary_service = auxiliary_service self.start = start self.engine = model_service.create_session() self.session_id = uuid4().hex self.send_lock = asyncio.Lock() self.audio_queue: asyncio.Queue[bytes | EndOfStream] = asyncio.Queue(maxsize=128) self.speaker_queue: asyncio.Queue[SpeakerJob | EndOfStream] = asyncio.Queue(maxsize=64) self.assembler = SegmentAssembler() self.metrics = SessionMetrics(time.perf_counter()) self.source = str(start.get("source") or "mic") self.file_name = str(start.get("file_name") or "audio.pcm") self.wav_decoder = ( IncrementalWavDecoder() if self.source == "file" and Path(self.file_name).suffix.lower() == ".wav" else None ) self.merge_adjacent = self._parse_flag(start.get("display_merge"), True) # Speaker labels are required for every browser session. self.speaker_enabled = True self.speaker_warning_sent = False self.input_stopped = False @staticmethod def _parse_flag(value: Any, default: bool) -> bool: """Accept booleans, 0/1 and string flags from old frontend clients.""" if value is None: return default if isinstance(value, str): return value.strip().lower() not in {"", "0", "false", "no", "off"} return bool(value) async def emit(self, payload: dict[str, Any]) -> None: """Send ordered JSON while the browser connection is still alive.""" async with self.send_lock: if not self.ws.closed: await self.ws.send_json(payload) async def emit_state(self, sentence: dict[str, Any] | None = None) -> None: """Send both the legacy sentence event and the full display snapshot.""" state = { "type": "display_state", "raw_segments": self.assembler.raw_snapshot(), "display_blocks": self.assembler.display_blocks(self.merge_adjacent), "metrics": self.metrics.snapshot(), } if sentence is not None: await self.emit( { "type": "sentences", "sentences": [sentence], "metrics": self.metrics.snapshot(), } ) await self.emit(state) async def warn_speaker(self, message: str) -> None: """Expose speaker-service failures without interrupting ASR.""" if self.speaker_warning_sent: return await self.emit( { "type": "speaker_warning", "session_id": self.session_id, "speaker_service_url": getattr( getattr(self.auxiliary_service, "config", None), "base_url", None ), "message": message, } ) self.speaker_warning_sent = True async def _emit_engine_segment(self, segment: FunASRSegment) -> None: """Map one FunASR partial/final event to a stable sentence ID.""" if not segment.text: if segment.is_final and self.assembler.segments.pop(segment.sentence_id, None) is not None: await self.emit_state() return sentence_type = 1 if segment.is_final else 0 sentence = self.assembler.apply_sentence( { "sentence_id": segment.sentence_id, "sentence": segment.text, "sentence_type": sentence_type, "start_time": segment.start_time_ms, "end_time": segment.end_time_ms, "speaker_id": -1, "speaker_name": "", "speaker_evidence": "pending", "speaker_confidence": 0.0, "speaker_strategy": "funasr_pending", "commit_reason": segment.reason, "speaker_status": ( "queued" if segment.is_final else "waiting_final" ) if self.speaker_enabled else "disabled", "speaker_reason": ( "等待 CAM++ 声纹处理" if segment.is_final else "FunASR 流式结果,等待片段结束" ) if self.speaker_enabled else "说话人分离已关闭", } ) if sentence_type == 0: self.metrics.partial_count += 1 if sentence["revision_count"] > 0: self.metrics.partial_revisions += 1 if self.metrics.first_partial_ms is None: self.metrics.first_partial_ms = round( (time.perf_counter() - self.metrics.started_at) * 1000, 1 ) else: self.metrics.final_ms = round( (time.perf_counter() - self.metrics.started_at) * 1000, 1 ) await self.emit_state(sentence) if segment.is_final and self.speaker_enabled: await self.speaker_queue.put( SpeakerJob( sentence_id=segment.sentence_id, audio=segment.audio, start_time_ms=segment.start_time_ms, end_time_ms=segment.end_time_ms, voiced_ms=segment.voiced_ms, ) ) async def _resolve_speaker(self, job: SpeakerJob) -> None: """Keep the existing optional CAM++ display integration.""" if not self.speaker_enabled: return async def update_status(status: str, reason: str) -> None: updated = self.assembler.apply_speaker_update( { "sentence_id": job.sentence_id, "speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0, "speaker_status": status, "speaker_reason": reason, } ) await self.emit_state(updated) if job.voiced_ms < MIN_SPEAKER_VOICE_MS: await update_status( "insufficient_audio", f"有效语音不足 {MIN_SPEAKER_VOICE_MS}ms,不继承上一位说话人", ) return if self.auxiliary_service is None: await update_status("service_unavailable", "未配置说话人辅助模型服务") await self.warn_speaker("未配置辅助模型服务,无法执行说话人分离") return await update_status("processing", "正在提取 CAM++ 声纹并匹配说话人") try: speaker = await self.auxiliary_service.resolve_speaker( job.audio, self.session_id, job.start_time_ms, job.end_time_ms, ) except Exception as exc: LOGGER.exception("speaker resolve failed: session=%s sentence=%s", self.session_id, job.sentence_id) await update_status("service_error", str(exc)) await self.warn_speaker(str(exc)) return if not speaker: await update_status("no_embedding", "辅助服务未返回可用声纹结果") return update = dict(speaker) update["sentence_id"] = job.sentence_id update["speaker_name"] = str(update.get("speaker_name") or "") updated = self.assembler.apply_speaker_update(update) if updated is not None: await self.emit_state(updated) async def process_speakers(self) -> None: """Process completed turns in order so speaker clusters stay stable.""" while True: item = await self.speaker_queue.get() if isinstance(item, EndOfStream): return await self._resolve_speaker(item) async def _feed(self, chunk: bytes) -> None: """Forward raw PCM to FunASR and publish all returned events.""" if not chunk: return self.metrics.audio_bytes += len(chunk) for segment in await self.engine.feed(chunk): await self._emit_engine_segment(segment) async def process_audio(self) -> None: """Consume audio until EOF, with no local RMS/VLLM segmentation path.""" while True: item = await self.audio_queue.get() if isinstance(item, EndOfStream): break self.metrics.input_chunks += 1 chunk = self.wav_decoder.feed(item) if self.wav_decoder is not None else item await self._feed(chunk) if self.wav_decoder is not None: self.wav_decoder.finish() for segment in await self.engine.finish(): await self._emit_engine_segment(segment) await self.emit({"type": "metrics", "metrics": self.metrics.snapshot()}) async def close(self) -> None: """Drop per-session FunASR caches and temporary state.""" self.engine = None # type: ignore[assignment] async def index_handler(_: web.Request) -> web.FileResponse: """Serve the browser page without stale-cache surprises.""" return web.FileResponse( Path(__file__).parent / "static" / "index.html", headers={"Cache-Control": "no-store"}, ) async def config_handler(request: web.Request) -> web.Response: """Expose local FunASR settings using the old response field names.""" model_service: FunASRModelService = request.app[MODEL_SERVICE_KEY] auxiliary = request.app.get(AUXILIARY_SERVICE_KEY) response = web.json_response( { "model_service_url": LOCAL_ENGINE_URL, "model": model_service.config.model, "engine": "funasr", "speaker_service_url": getattr( getattr(auxiliary, "config", None), "base_url", None ), } ) # The standalone frontend reads this API from its own origin. response.headers["Access-Control-Allow-Origin"] = os.getenv( "FRONTEND_ORIGIN", "http://127.0.0.1:8080" ) return response async def websocket_handler(request: web.Request) -> web.WebSocketResponse: """Keep the old browser protocol while using FunASR internally.""" ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024) await ws.prepare(request) model_service: FunASRModelService = request.app[MODEL_SERVICE_KEY] auxiliary_service = request.app.get(AUXILIARY_SERVICE_KEY) processing: asyncio.Task[None] | None = None speaker_processing: asyncio.Task[None] | None = None session: RealtimeSession | None = None input_finished = False try: first = await ws.receive() if first.type != WSMsgType.TEXT: await ws.send_json({"type": "error", "message": "first message must be JSON start"}) return ws try: start = json.loads(first.data) except json.JSONDecodeError: await ws.send_json({"type": "error", "message": "invalid start JSON"}) return ws if not isinstance(start, dict) or start.get("type") != "start": await ws.send_json({"type": "error", "message": "first message must have type=start"}) return 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 ws.send_json( { "type": "error", "message": "实时流式测试的文件模式只支持 PCM 或 WAV,请改用麦克风、PCM 或 WAV", } ) return ws session = RealtimeSession(ws, model_service, auxiliary_service, start) speaker_health: dict[str, Any] | None = None speaker_health_error: str | None = None if session.speaker_enabled and auxiliary_service is not None: try: speaker_health = await asyncio.wait_for(auxiliary_service.health(), timeout=5) if speaker_health.get("speaker_embedding_ready", speaker_health.get("ready")) is False: speaker_health_error = "辅助模型服务未就绪,请检查 /health 返回的 models 状态" except Exception as exc: speaker_health_error = f"说话人辅助服务不可用:{exc}" await session.emit( { "type": "start", "model_service_url": LOCAL_ENGINE_URL, "model": model_service.config.model, "engine": "funasr", "session_id": session.session_id, "enable_native_partial_stream": True, "native_partial_supported": True, "partial_mode": "funasr_streaming_cache", "speaker_diarization_enabled": session.speaker_enabled, "speaker_service_url": getattr( getattr(auxiliary_service, "config", None), "base_url", None ), "speaker_service_health": speaker_health, "speaker_gap_enabled": False, "sentence_strategy": start.get("sentence_strategy", 0), "display_state_supported": True, } ) if speaker_health_error: await session.warn_speaker(speaker_health_error) processing = asyncio.create_task(session.process_audio()) if session.speaker_enabled: speaker_processing = asyncio.create_task(session.process_speakers()) async def receive_or_raise() -> Any: """Wake up immediately when the engine worker fails.""" receive_task = asyncio.create_task(ws.receive()) workers = [task for task in (processing, speaker_processing) if task is not None] done, _ = await asyncio.wait( [receive_task, *workers], return_when=asyncio.FIRST_COMPLETED, ) if receive_task in done: return await receive_task receive_task.cancel() await asyncio.gather(receive_task, return_exceptions=True) for worker in workers: if worker in done: await worker raise RuntimeError("FunASR 实时处理任务意外结束") while not ws.closed: message = await receive_or_raise() if message.type == WSMsgType.BINARY: await session.audio_queue.put(bytes(message.data)) continue if message.type == WSMsgType.TEXT: try: control = json.loads(message.data) except json.JSONDecodeError: continue if not isinstance(control, dict): continue if control.get("type") in {"eof", "stop"}: input_finished = True session.input_stopped = control.get("type") == "stop" await session.emit({"type": "draining", "message": "FunASR 正在完成最终识别"}) await session.audio_queue.put(EOF) break if control.get("type") == "abort": return ws if message.type in {WSMsgType.ERROR, WSMsgType.CLOSE, WSMsgType.CLOSED}: return ws if input_finished: await processing if speaker_processing is not None: await session.speaker_queue.put(EOF) await speaker_processing await session.emit_state() await session.emit( { "type": "end", "metrics": session.metrics.snapshot(), "sentences": session.assembler.raw_snapshot(), "display_blocks": session.assembler.display_blocks(session.merge_adjacent), } ) except asyncio.CancelledError: raise except Exception as exc: LOGGER.exception("FunASR WebSocket session failed") if not ws.closed: await ws.send_json({"type": "error", "message": str(exc)}) finally: for task in (processing, speaker_processing): if task is not None and not task.done(): task.cancel() await asyncio.gather( *(task for task in (processing, speaker_processing) if task is not None), return_exceptions=True, ) if session is not None: reset = getattr(session.auxiliary_service, "reset_speaker_session", None) await session.close() if reset is not None: try: await asyncio.wait_for(reset(session.session_id), timeout=5) except Exception: LOGGER.warning("speaker session cleanup failed: %s", session.session_id, exc_info=True) if not ws.closed: await ws.close() return ws async def start_app( model: str | None = None, device: str | None = None, ) -> web.Application: """Create the FunASR-backed frontend application.""" config = FunASRServiceConfig.from_env() if model: config = replace(config, model=model) if device: config = replace(config, device=device) app = web.Application() app[MODEL_SERVICE_KEY] = FunASRModelService(config) app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService( AuxiliaryServiceConfig( base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010") ) ) async def lifecycle(application: web.Application): await application[MODEL_SERVICE_KEY].start() await application[AUXILIARY_SERVICE_KEY].start() try: health = await application[AUXILIARY_SERVICE_KEY].health() if not health.get("speaker_embedding_ready"): raise RuntimeError("CAM++ speaker service is not ready") except Exception: await application[AUXILIARY_SERVICE_KEY].close() await application[MODEL_SERVICE_KEY].close() raise yield await application[AUXILIARY_SERVICE_KEY].close() await application[MODEL_SERVICE_KEY].close() app.cleanup_ctx.append(lifecycle) app.router.add_get("/", index_handler) app.router.add_get("/api/config", config_handler) app.router.add_static("/static/", Path(__file__).parent / "static") app.router.add_get("/ws", websocket_handler) app.router.add_static("/", Path(__file__).parent / "static", show_index=False) return app def main() -> None: """Start the FunASR-backed browser demo.""" parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", default=os.getenv("FUNASR_ASR_MODEL")) parser.add_argument("--device", default=os.getenv("FUNASR_DEVICE")) parser.add_argument("--no-browser", action="store_true") args = parser.parse_args() logging.basicConfig(level=logging.INFO) if not args.no_browser: webbrowser.open(f"http://{WEB_DISPLAY_HOST}:{WEB_PORT}/") print(f"FunASR demo: http://{WEB_DISPLAY_HOST}:{WEB_PORT}/", flush=True) print(f"FunASR model: {args.model or FunASRServiceConfig.from_env().model}", flush=True) web.run_app( start_app(model=args.model, device=args.device), host=WEB_HOST, port=WEB_PORT, ) if __name__ == "__main__": main()