531 lines
22 KiB
Python
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[1]
|
|
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()
|