ASR-demo/backend/realtime_websocket/funasr_server.py

531 lines
22 KiB
Python

"""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")))
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
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)))
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]})
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
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 = self.total_audio_ms
if final_text:
await self.emit_sentence(
final_text, final=True, start_time_ms=start_ms, end_time_ms=end_ms
)
# CAM++ runs off the native WS reader so ASR keeps draining.
await self.speaker_jobs.put(
SpeakerJob(
sentence_id=self.sentence_id,
text=final_text,
audio=audio,
start_time_ms=start_ms,
end_time_ms=end_ms,
)
)
self.sentence_id += 1
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:
"""Resolve final utterances in order to keep CAM++ cluster IDs stable."""
while True:
job = await self.speaker_jobs.get()
try:
if job is None:
return
speaker: dict[str, Any] = {"speaker_id": -1}
if 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:
speaker = resolved
except Exception:
LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id)
# Re-emit the final sentence with the CAM++ label for the unchanged UI.
await self.emit_sentence(
job.text,
final=True,
speaker=speaker,
sentence_id=job.sentence_id,
start_time_ms=job.start_time_ms,
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")
),
"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()