ASR-demo/realtime_websocket/funasr_server.py

643 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""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()