Use FunASR native realtime WebSocket flow
parent
fbaf56fd9f
commit
13278d3aea
|
|
@ -1,30 +1,32 @@
|
|||
# Local model assets; the backend never downloads models at startup.
|
||||
# Local model assets; model loading never downloads weights at startup.
|
||||
MODEL_DIR=models
|
||||
# Set this if CAM++ is not at a model_manifest.json path.
|
||||
# Set this if CAM++ is not under the path declared in model_manifest.json.
|
||||
# CAM_MODEL_PATH=iic/speech_campplus_sv_zh-cn_16k-common
|
||||
|
||||
# FunASR streaming ASR model. The backend launcher requires it to exist locally.
|
||||
# Local FunASR streaming ASR and VAD models.
|
||||
FUNASR_ASR_MODEL=paraformer-zh-streaming
|
||||
FUNASR_VAD_MODEL=fsmn-vad
|
||||
|
||||
# ASR and VAD may use different devices to limit peak GPU memory.
|
||||
FUNASR_DEVICE=cuda:0
|
||||
FUNASR_VAD_DEVICE=cpu
|
||||
|
||||
# [left context, current chunk, right lookahead]; 10 * 60ms = 600ms.
|
||||
# Native FunASR WSS chunk settings. The middle chunk is sent as 10 x 60 ms.
|
||||
FUNASR_CHUNK_SIZE=0,10,5
|
||||
FUNASR_CHUNK_INTERVAL=10
|
||||
FUNASR_ENCODER_LOOK_BACK=4
|
||||
FUNASR_DECODER_LOOK_BACK=1
|
||||
FUNASR_VAD_CHUNK_MS=200
|
||||
FUNASR_MAX_SEGMENT_SEC=30
|
||||
FUNASR_FINALIZE_TIMEOUT_SECONDS=300
|
||||
FUNASR_NATIVE_WS_HOST=127.0.0.1
|
||||
FUNASR_NATIVE_WS_PORT=10095
|
||||
|
||||
# CAM++ is required. The backend launcher starts this service itself.
|
||||
# CAM++ is required and started by scripts/run_backend.py.
|
||||
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010
|
||||
AUXILIARY_DEVICE=cpu
|
||||
AUXILIARY_PRELOAD_KINDS=speaker_verification
|
||||
|
||||
# Standalone frontend and browser-facing backend origin.
|
||||
# Frontend and backend run independently. The frontend proxies /ws and /api/stop.
|
||||
FRONTEND_HOST=127.0.0.1
|
||||
FRONTEND_PORT=8080
|
||||
FRONTEND_ORIGIN=http://127.0.0.1:8080
|
||||
BACKEND_INTERNAL_URL=http://127.0.0.1:8082
|
||||
BACKEND_PUBLIC_URL=http://127.0.0.1:8082
|
||||
WEB_HOST=0.0.0.0
|
||||
WEB_PORT=8082
|
||||
|
|
|
|||
|
|
@ -1,59 +1,43 @@
|
|||
# FunASR realtime browser demo
|
||||
|
||||
The frontend and backend start separately. The backend launcher loads three
|
||||
required local assets: streaming ASR, FSMN VAD, and CAM++ speaker verification.
|
||||
It starts the CAM++ model service and the WebSocket service and stops both
|
||||
together. No model is downloaded by the launcher. Use the shared model manifest
|
||||
and downloader to prepare the three required snapshots:
|
||||
The frontend and backend start separately. The backend starts these supervised processes:
|
||||
|
||||
~~~powershell
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
- FunASR native online WebSocket server, using the local streaming ASR and FSMN VAD models. The browser adapter sends FunASR's `mode=online`, chunk/look-back settings, fixed 60 ms PCM frames, and `is_speaking=false` end-of-input flush.
|
||||
- CAM++ auxiliary service, which assigns stable speaker labels to finalized utterances.
|
||||
- A small browser protocol adapter that translates the unchanged Tencent demo message format to FunASR's native WS format. It does not run a second ASR/VAD segmentation pipeline.
|
||||
|
||||
The frontend serves the exact files from the local `tencent-demo/static` directory and proxies the page's same-origin `/ws` and `/api/stop` requests to the backend.
|
||||
|
||||
## Model directories
|
||||
|
||||
Put the assets under models/ or set MODEL_DIR in .env. The default
|
||||
FunASR names resolve to these local directories:
|
||||
Put assets under `models/` or set `MODEL_DIR` in `.env`. The default names resolve to these local directories:
|
||||
|
||||
- models/iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online
|
||||
- models/damo/speech_fsmn_vad_zh-cn-16k-common-pytorch
|
||||
- models/iic/speech_campplus_sv_zh-cn_16k-common (or the damo/ variant)
|
||||
- `models/iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online`
|
||||
- `models/damo/speech_fsmn_vad_zh-cn-16k-common-pytorch`
|
||||
- `models/iic/speech_campplus_sv_zh-cn_16k-common` (or the configured `damo/` variant)
|
||||
|
||||
If your directories have different names, set FUNASR_ASR_MODEL and
|
||||
FUNASR_VAD_MODEL to their paths. The CAM++ directories follow
|
||||
model_manifest.json; set CAM_MODEL_PATH for another location. Backend startup reports every checked path when a
|
||||
required asset is missing.
|
||||
If directories have different names, set `FUNASR_ASR_MODEL` and `FUNASR_VAD_MODEL` to their local paths. CAM++ follows `model_manifest.json`; set `CAM_MODEL_PATH` for another location. Startup checks all three assets before exposing the browser bridge.
|
||||
|
||||
## Start
|
||||
## Install and start
|
||||
|
||||
First install a torch/torchaudio build suitable for the host CPU or CUDA,
|
||||
then install the project dependencies:
|
||||
Install a torch/torchaudio build suitable for the host, then install the project dependencies and models:
|
||||
|
||||
~~~powershell
|
||||
cd D:\github-project\ASR\Asr-demo
|
||||
python -m pip install -r requirements-funasr.txt
|
||||
python -m pip install -r requirements-auxiliary.txt
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.funasr.example .env }
|
||||
~~~
|
||||
|
||||
In one terminal start the backend:
|
||||
Start the backend and frontend in separate terminals:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_backend.py
|
||||
~~~
|
||||
|
||||
In another terminal start the frontend:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_frontend.py
|
||||
~~~
|
||||
|
||||
Open http://127.0.0.1:8080/. The backend WebSocket listens on port 8082 and
|
||||
the CAM++ service on port 8010. BACKEND_PUBLIC_URL must be reachable from
|
||||
the browser. Set FRONTEND_ORIGIN to the exact frontend origin if it differs
|
||||
from the default.
|
||||
Defaults are frontend port 8080, browser backend port 8082, CAM++ HTTP port 8010, and native FunASR WS port 10095 bound to loopback. Change `FRONTEND_PORT`, `WEB_PORT`, `AUXILIARY_SERVICE_URL`, `FUNASR_NATIVE_WS_HOST`, and `FUNASR_NATIVE_WS_PORT` together when needed. `BACKEND_INTERNAL_URL` is the backend origin reachable from the frontend process; it defaults to `http://127.0.0.1:${WEB_PORT}`.
|
||||
|
||||
The browser sends start, PCM16 frames, and stop/eof over WebSocket. Speaker
|
||||
labels are required. Short or unusable speech may still receive an unknown
|
||||
speaker label, but missing CAM++ prevents backend startup.
|
||||
`FUNASR_DEVICE` and `FUNASR_VAD_DEVICE` control ASR and VAD placement independently. Set `AUXILIARY_DEVICE=cpu` when CAM++ should not share the ASR GPU. `FUNASR_CHUNK_SIZE` and `FUNASR_CHUNK_INTERVAL` control FunASR's native chunk protocol; the default `[0,10,5]` and interval 10 send the current chunk in 600 ms groups.
|
||||
|
||||
Real-time file input currently accepts raw PCM16 or 16 kHz mono PCM WAV, matching the backend's available decoder. Speaker labels are computed for every finalized utterance; very short or silent segments can still be marked as unknown by CAM++.
|
||||
|
|
|
|||
13
README.md
13
README.md
|
|
@ -1,20 +1,16 @@
|
|||
# FunASR realtime ASR demo
|
||||
|
||||
This branch runs streaming ASR, VAD, and CAM++ from local model assets.
|
||||
The browser UI starts as a separate service.
|
||||
This branch runs FunASR's native online WebSocket flow with local streaming ASR and FSMN VAD models. CAM++ speaker embeddings are handled by the separate auxiliary service. The browser page, JavaScript, and CSS are copied byte-for-byte from the local `tencent-demo/static` directory.
|
||||
|
||||
## Start
|
||||
|
||||
Install a torch/torchaudio build for the host, then install
|
||||
requirements-funasr.txt and requirements-auxiliary.txt. Copy
|
||||
.env.funasr.example to .env if no .env exists, then download the three
|
||||
required FunASR models with the shared project manifest:
|
||||
Install a torch/torchaudio build for the host, then install `requirements-funasr.txt` and `requirements-auxiliary.txt`. Copy `.env.funasr.example` to `.env` if no `.env` exists, then download the three required FunASR assets:
|
||||
|
||||
~~~powershell
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
|
||||
Start the backend (CAM++ model service plus WebSocket) in one terminal:
|
||||
Start the backend (CAM++, native FunASR WSS, and the browser protocol bridge) in one terminal:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_backend.py
|
||||
|
|
@ -26,5 +22,4 @@ Start the frontend in another terminal:
|
|||
python scripts\run_frontend.py
|
||||
~~~
|
||||
|
||||
Open http://127.0.0.1:8080/. Model directory layout and address settings
|
||||
are described in FUNASR_README.md.
|
||||
Open the URL configured by `FRONTEND_HOST` and `FRONTEND_PORT`. The frontend keeps the Tencent demo's same-origin `/ws` and `/api/stop` calls and proxies them to `BACKEND_INTERNAL_URL`. The public backend bridge defaults to port 8082; the native FunASR WS socket defaults to loopback port 10095. Model paths and device settings are described in [FUNASR_README.md](FUNASR_README.md).
|
||||
|
|
|
|||
|
|
@ -0,0 +1,888 @@
|
|||
"""FunASR realtime WebSocket server adapted from the local FunASR checkout.
|
||||
|
||||
Source: ``runtime/python/websocket/funasr_wss_server.py``. The browser adapter
|
||||
uses FunASR's online WS protocol. Offline ASR, punctuation, and in-process
|
||||
speaker verification are optional; CAM++ runs in the separate auxiliary service.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import websockets
|
||||
import time
|
||||
import numpy as np
|
||||
import argparse
|
||||
import ssl
|
||||
import os
|
||||
import wave
|
||||
import functools
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from scipy.spatial.distance import cosine
|
||||
|
||||
import torch # 保留不影响
|
||||
|
||||
|
||||
def to_python(obj):
|
||||
"""递归地把 numpy / torch 等类型转成纯 Python,可 JSON 序列化。"""
|
||||
try:
|
||||
import numpy as np # noqa
|
||||
import torch # noqa
|
||||
except Exception:
|
||||
np = None
|
||||
torch = None
|
||||
|
||||
if np is not None and isinstance(obj, np.generic):
|
||||
return obj.item()
|
||||
if np is not None and isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
if torch is not None and isinstance(obj, torch.Tensor):
|
||||
return obj.cpu().tolist()
|
||||
|
||||
if isinstance(obj, dict):
|
||||
return {k: to_python(v) for k, v in obj.items()}
|
||||
if isinstance(obj, (list, tuple)):
|
||||
return [to_python(v) for v in obj]
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0", required=False, help="host ip")
|
||||
parser.add_argument("--port", type=int, default=10095, required=False, help="grpc server port")
|
||||
|
||||
parser.add_argument(
|
||||
"--asr_model",
|
||||
type=str,
|
||||
default="",
|
||||
help="Optional offline ASR model; empty means online-only mode.",
|
||||
)
|
||||
parser.add_argument("--asr_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--asr_model_online",
|
||||
type=str,
|
||||
default="iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--asr_model_online_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--vad_model",
|
||||
type=str,
|
||||
default="iic/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--vad_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--punc_model",
|
||||
type=str,
|
||||
default="",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--punc_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument("--ngpu", type=int, default=1, help="0 for cpu, 1 for gpu")
|
||||
parser.add_argument("--device", type=str, default="cuda", help="cuda, cpu")
|
||||
parser.add_argument("--vad_device", type=str, default=None, help="Optional VAD device override")
|
||||
parser.add_argument("--ncpu", type=int, default=4, help="cpu cores")
|
||||
parser.add_argument(
|
||||
"--enable_speaker_verification",
|
||||
action="store_true",
|
||||
help="Load native CAM++; disabled when CAM++ is hosted by the auxiliary service.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--certfile",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="certfile for ssl",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keyfile",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="keyfile for ssl",
|
||||
)
|
||||
|
||||
# ====== 保存 2pass 离线阶段送入 ASR 的音频片段(排查 VAD 切分)======
|
||||
parser.add_argument(
|
||||
"--save_offline_segments",
|
||||
action="store_true",
|
||||
help="Save each offline (2pass) audio segment sent to offline ASR as wav for debugging VAD split.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_offline_segments_dir",
|
||||
type=str,
|
||||
default="./offline_segments",
|
||||
help="Directory to save offline wav segments when --save_offline_segments is enabled.",
|
||||
)
|
||||
|
||||
# ====== 并发控制:核心新增 ======
|
||||
parser.add_argument(
|
||||
"--worker_threads",
|
||||
type=int,
|
||||
default=max(4, (os.cpu_count() or 4)),
|
||||
help="ThreadPoolExecutor max_workers. Used to offload blocking inference so event loop won't be blocked.",
|
||||
)
|
||||
parser.add_argument("--concurrent_vad", type=int, default=4, help="Max concurrent VAD generate() calls.")
|
||||
parser.add_argument("--concurrent_asr_online", type=int, default=4, help="Max concurrent streaming ASR generate() calls.")
|
||||
parser.add_argument("--concurrent_asr_offline", type=int, default=2, help="Max concurrent offline ASR generate() calls.")
|
||||
parser.add_argument("--concurrent_punc", type=int, default=1, help="Max concurrent punctuation generate() calls.")
|
||||
parser.add_argument("--concurrent_sv", type=int, default=1, help="Max concurrent speaker verification generate() calls.")
|
||||
parser.add_argument(
|
||||
"--speaker_db_reload_sec",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Reload speaker_db.json at most once every N seconds (avoid frequent disk IO).",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
websocket_users = set()
|
||||
SPEAKER_DB_PATH = os.path.join(os.path.dirname(__file__), "speaker_db.json")
|
||||
|
||||
|
||||
def _ensure_dir(p: str):
|
||||
try:
|
||||
os.makedirs(p, exist_ok=True)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _pcm_duration_ms(pcm_bytes: bytes, fs: int, ch: int = 1, sampwidth: int = 2) -> int:
|
||||
"""根据 fs/ch/sampwidth 计算 PCM 时长,避免写死 16k -> 32 bytes/ms。"""
|
||||
if not pcm_bytes:
|
||||
return 0
|
||||
bytes_per_ms = (fs * ch * sampwidth) / 1000.0
|
||||
if bytes_per_ms <= 0:
|
||||
return 0
|
||||
return int(len(pcm_bytes) / bytes_per_ms)
|
||||
|
||||
|
||||
def _safe_int(v, default):
|
||||
try:
|
||||
return int(v)
|
||||
except Exception:
|
||||
return default
|
||||
|
||||
|
||||
# ========= speaker db:加缓存,避免每段都读盘 =========
|
||||
_SPEAKER_DB_CACHE = {}
|
||||
_SPEAKER_DB_CACHE_TS = 0.0
|
||||
|
||||
|
||||
def _load_speaker_db_sync():
|
||||
if not os.path.exists(SPEAKER_DB_PATH):
|
||||
return {}
|
||||
try:
|
||||
with open(SPEAKER_DB_PATH, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def get_speaker_db_cached(now_ts: float, reload_sec: int):
|
||||
global _SPEAKER_DB_CACHE, _SPEAKER_DB_CACHE_TS
|
||||
if (now_ts - _SPEAKER_DB_CACHE_TS) >= max(1, int(reload_sec)):
|
||||
_SPEAKER_DB_CACHE = _load_speaker_db_sync()
|
||||
_SPEAKER_DB_CACHE_TS = now_ts
|
||||
return _SPEAKER_DB_CACHE or {}
|
||||
|
||||
|
||||
def _save_wav_sync(out_path: str, audio_bytes: bytes, fs: int, ch: int, sampwidth: int):
|
||||
with wave.open(out_path, "wb") as wf:
|
||||
wf.setnchannels(ch)
|
||||
wf.setsampwidth(sampwidth)
|
||||
wf.setframerate(fs)
|
||||
wf.writeframes(audio_bytes)
|
||||
|
||||
|
||||
def save_offline_wav_segment_sync(websocket, audio_bytes: bytes, reason: str = "offline"):
|
||||
"""
|
||||
保存离线阶段送入 ASR 的音频片段,方便人工试听排查 VAD 切分是否正确。
|
||||
约定:audio_bytes 为 单声道 PCM16 little-endian(默认 16k)。
|
||||
(注意:这是同步函数,外层会放线程池执行)
|
||||
"""
|
||||
if not getattr(websocket, "save_offline_segments", False):
|
||||
return
|
||||
if "2pass" not in (getattr(websocket, "mode", "") or ""):
|
||||
return
|
||||
if not audio_bytes:
|
||||
return
|
||||
|
||||
fs = int(getattr(websocket, "audio_fs", 16000) or 16000)
|
||||
ch = 1
|
||||
sampwidth = 2 # int16
|
||||
|
||||
# int16 对齐
|
||||
if len(audio_bytes) % 2 == 1:
|
||||
audio_bytes = audio_bytes[:-1]
|
||||
if not audio_bytes:
|
||||
return
|
||||
|
||||
seg_idx = int(getattr(websocket, "offline_seg_idx", 0))
|
||||
websocket.offline_seg_idx = seg_idx + 1
|
||||
|
||||
duration_ms = _pcm_duration_ms(audio_bytes, fs=fs, ch=ch, sampwidth=sampwidth)
|
||||
|
||||
base_dir = getattr(websocket, "offline_save_dir", args.save_offline_segments_dir)
|
||||
_ensure_dir(base_dir)
|
||||
|
||||
wav_name = (getattr(websocket, "wav_name", "microphone") or "microphone").replace("/", "_")
|
||||
ts = int(time.time() * 1000)
|
||||
fname = f"{wav_name}_{ts}_seg{seg_idx:04d}_{reason}_{duration_ms}ms.wav"
|
||||
out_path = os.path.join(base_dir, fname)
|
||||
|
||||
try:
|
||||
_save_wav_sync(out_path, audio_bytes, fs=fs, ch=ch, sampwidth=sampwidth)
|
||||
print(f"[SAVE_OFFLINE_SEG] {out_path} ({duration_ms} ms, {len(audio_bytes)} bytes)")
|
||||
except Exception as e:
|
||||
print(f"[SAVE_OFFLINE_SEG] failed: {e}")
|
||||
|
||||
|
||||
print("model loading")
|
||||
from funasr import AutoModel # noqa
|
||||
|
||||
# ====== 离线 ASR ======
|
||||
# Online deployments leave the offline model unloaded to conserve memory.
|
||||
model_asr = (
|
||||
AutoModel(
|
||||
model=args.asr_model,
|
||||
model_revision=args.asr_model_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
if args.asr_model
|
||||
else None
|
||||
)
|
||||
|
||||
# streaming asr
|
||||
model_asr_streaming = AutoModel(
|
||||
model=args.asr_model_online,
|
||||
model_revision=args.asr_model_online_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
|
||||
# vad
|
||||
model_vad = AutoModel(
|
||||
model=args.vad_model,
|
||||
model_revision=args.vad_model_revision,
|
||||
ngpu=args.ngpu if (args.vad_device or args.device).startswith("cuda") else 0,
|
||||
ncpu=args.ncpu,
|
||||
device=args.vad_device or args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
|
||||
# punc
|
||||
if args.punc_model != "":
|
||||
model_punc = AutoModel(
|
||||
model=args.punc_model,
|
||||
model_revision=args.punc_model_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
else:
|
||||
model_punc = None
|
||||
|
||||
# CAM++ is loaded by the auxiliary service, avoiding a second GPU copy here.
|
||||
model_sv = (
|
||||
AutoModel(
|
||||
model="iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
ngpu=args.ngpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
if args.enable_speaker_verification
|
||||
else None
|
||||
)
|
||||
|
||||
print("model loaded! (now supports multi-client with non-blocking inference)")
|
||||
|
||||
|
||||
# ====== 线程池 + 并发阈值(核心)======
|
||||
EXECUTOR = ThreadPoolExecutor(max_workers=int(args.worker_threads))
|
||||
|
||||
SEM_VAD = asyncio.Semaphore(max(1, int(args.concurrent_vad)))
|
||||
SEM_ASR_ONLINE = asyncio.Semaphore(max(1, int(args.concurrent_asr_online)))
|
||||
SEM_ASR_OFFLINE = asyncio.Semaphore(max(1, int(args.concurrent_asr_offline)))
|
||||
SEM_PUNC = asyncio.Semaphore(max(1, int(args.concurrent_punc)))
|
||||
SEM_SV = asyncio.Semaphore(max(1, int(args.concurrent_sv)))
|
||||
SEM_WAV = asyncio.Semaphore(max(1, 4)) # 保存 wav 一般不需要太大
|
||||
|
||||
|
||||
async def run_blocking(fn, *a, sem: asyncio.Semaphore | None = None, **kw):
|
||||
"""
|
||||
把阻塞函数丢线程池执行,避免卡 event loop。
|
||||
sem 用于限流(避免 GPU / 模型被打爆)。
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
call = functools.partial(fn, *a, **kw)
|
||||
if sem is None:
|
||||
return await loop.run_in_executor(EXECUTOR, call)
|
||||
async with sem:
|
||||
return await loop.run_in_executor(EXECUTOR, call)
|
||||
|
||||
|
||||
def _generate_sync(model, audio_or_text, status_dict):
|
||||
# 注意:status_dict 里包含 cache,会被 generate 更新
|
||||
return model.generate(input=audio_or_text, **status_dict)
|
||||
|
||||
|
||||
async def ws_reset(websocket):
|
||||
print("ws reset now, total num is ", len(websocket_users))
|
||||
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = True
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
websocket.status_dict_vad["is_final"] = True
|
||||
websocket.status_dict_punc["cache"] = {}
|
||||
|
||||
await websocket.close()
|
||||
|
||||
|
||||
async def clear_websocket():
|
||||
for websocket in list(websocket_users):
|
||||
await ws_reset(websocket)
|
||||
websocket_users.clear()
|
||||
|
||||
|
||||
async def ws_serve(websocket, path=None):
|
||||
# websockets 新版本不会传 path,这里做兼容
|
||||
if path is None:
|
||||
path = getattr(websocket, "path", None)
|
||||
frames = []
|
||||
frames_asr = []
|
||||
frames_asr_online = []
|
||||
pending_offline_audio = []
|
||||
global websocket_users
|
||||
websocket_users.add(websocket)
|
||||
|
||||
websocket.status_dict_asr = {} # hotword 等
|
||||
websocket.status_dict_asr_online = {"cache": {}, "is_final": False}
|
||||
websocket.status_dict_vad = {"cache": {}, "is_final": False}
|
||||
websocket.status_dict_punc = {"cache": {}}
|
||||
|
||||
websocket.chunk_interval = 10
|
||||
websocket.vad_pre_idx = 0
|
||||
speech_start = False
|
||||
speech_end_i = -1
|
||||
online_needs_finalization = False
|
||||
session_errors = []
|
||||
|
||||
websocket.wav_name = "microphone"
|
||||
websocket.mode = "2pass"
|
||||
websocket.is_speaking = True # ✅ 默认初始化,避免 AttributeError
|
||||
|
||||
# 保存离线片段
|
||||
websocket.audio_fs = 16000
|
||||
websocket.offline_seg_idx = 0
|
||||
websocket.save_offline_segments = bool(args.save_offline_segments)
|
||||
websocket.offline_save_dir = args.save_offline_segments_dir
|
||||
if websocket.save_offline_segments:
|
||||
_ensure_dir(websocket.offline_save_dir)
|
||||
print(f"[SAVE_OFFLINE_SEG] enabled, dir={websocket.offline_save_dir}")
|
||||
|
||||
print("new user connected", flush=True)
|
||||
|
||||
def record_error(message):
|
||||
if message not in session_errors:
|
||||
session_errors.append(message)
|
||||
|
||||
async def finalize_online_segment():
|
||||
nonlocal frames_asr_online, online_needs_finalization
|
||||
|
||||
if websocket.mode not in ("2pass", "online") or not online_needs_finalization:
|
||||
return
|
||||
|
||||
websocket.status_dict_asr_online["is_final"] = True
|
||||
try:
|
||||
await async_asr_online(websocket, b"".join(frames_asr_online))
|
||||
except Exception as e:
|
||||
print("error in final asr streaming:", e)
|
||||
record_error(f"online inference failed: {e}")
|
||||
|
||||
frames_asr_online = []
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = False
|
||||
online_needs_finalization = False
|
||||
|
||||
async def finish_input(send_end_ack):
|
||||
nonlocal frames, frames_asr, frames_asr_online, pending_offline_audio
|
||||
nonlocal speech_start, speech_end_i, online_needs_finalization
|
||||
|
||||
await finalize_online_segment()
|
||||
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
audio_in = b"".join(frames_asr)
|
||||
if not audio_in:
|
||||
audio_in = b"".join(pending_offline_audio)
|
||||
|
||||
if audio_in:
|
||||
if websocket.save_offline_segments and audio_in:
|
||||
try:
|
||||
await run_blocking(
|
||||
save_offline_wav_segment_sync,
|
||||
websocket,
|
||||
audio_in,
|
||||
"not_speaking",
|
||||
sem=SEM_WAV,
|
||||
)
|
||||
except Exception as e:
|
||||
print("[SAVE_OFFLINE_SEG] async failed:", e)
|
||||
|
||||
try:
|
||||
await async_asr(websocket, audio_in)
|
||||
pending_offline_audio = []
|
||||
except Exception as e:
|
||||
print("error in final asr offline:", e)
|
||||
record_error(f"offline inference failed: {e}")
|
||||
|
||||
errors = list(session_errors)
|
||||
|
||||
frames = []
|
||||
frames_asr = []
|
||||
frames_asr_online = []
|
||||
pending_offline_audio = []
|
||||
speech_start = False
|
||||
speech_end_i = -1
|
||||
online_needs_finalization = False
|
||||
websocket.vad_pre_idx = 0
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
|
||||
if send_end_ack:
|
||||
acknowledgement = {
|
||||
"mode": websocket.mode,
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": not errors,
|
||||
"is_end": True,
|
||||
}
|
||||
if errors:
|
||||
acknowledgement["error"] = "; ".join(errors)
|
||||
await websocket.send(
|
||||
json.dumps(acknowledgement, ensure_ascii=False)
|
||||
)
|
||||
session_errors.clear()
|
||||
elif errors:
|
||||
raise RuntimeError("; ".join(errors))
|
||||
|
||||
try:
|
||||
async for message in websocket:
|
||||
# ========== 1) 先处理“文本配置消息” ==========
|
||||
if isinstance(message, str):
|
||||
try:
|
||||
messagejson = json.loads(message)
|
||||
except Exception as e:
|
||||
print("bad json message:", e, message[:200])
|
||||
continue
|
||||
|
||||
# Avoid per-message logging during long-running audio sessions.
|
||||
|
||||
end_of_input = False
|
||||
if "is_speaking" in messagejson:
|
||||
websocket.is_speaking = bool(messagejson["is_speaking"])
|
||||
websocket.status_dict_asr_online["is_final"] = (not websocket.is_speaking)
|
||||
end_of_input = not websocket.is_speaking
|
||||
|
||||
if "chunk_interval" in messagejson:
|
||||
websocket.chunk_interval = _safe_int(
|
||||
messagejson["chunk_interval"], websocket.chunk_interval
|
||||
)
|
||||
|
||||
if "wav_name" in messagejson:
|
||||
websocket.wav_name = messagejson.get("wav_name") or websocket.wav_name
|
||||
|
||||
if "chunk_size" in messagejson:
|
||||
chunk_size = messagejson["chunk_size"]
|
||||
if isinstance(chunk_size, str):
|
||||
chunk_size = [x.strip() for x in chunk_size.split(",") if x.strip()]
|
||||
websocket.status_dict_asr_online["chunk_size"] = [int(x) for x in chunk_size]
|
||||
|
||||
if "encoder_chunk_look_back" in messagejson:
|
||||
websocket.status_dict_asr_online["encoder_chunk_look_back"] = messagejson[
|
||||
"encoder_chunk_look_back"
|
||||
]
|
||||
|
||||
if "decoder_chunk_look_back" in messagejson:
|
||||
websocket.status_dict_asr_online["decoder_chunk_look_back"] = messagejson[
|
||||
"decoder_chunk_look_back"
|
||||
]
|
||||
|
||||
if "hotwords" in messagejson:
|
||||
hotword_data = messagejson["hotwords"]
|
||||
websocket.status_dict_asr["hotword"] = hotword_data
|
||||
websocket.status_dict_asr_online["hotword"] = hotword_data
|
||||
print(f"热词已更新: {hotword_data}")
|
||||
|
||||
if "mode" in messagejson:
|
||||
requested_mode = messagejson["mode"]
|
||||
if requested_mode and requested_mode not in ("online", "offline", "2pass"):
|
||||
websocket.mode = requested_mode
|
||||
record_error(f"unsupported mode: {requested_mode!r}")
|
||||
else:
|
||||
websocket.mode = requested_mode or websocket.mode
|
||||
|
||||
if "audio_fs" in messagejson:
|
||||
websocket.audio_fs = _safe_int(messagejson["audio_fs"], 16000)
|
||||
|
||||
if end_of_input:
|
||||
await finish_input(send_end_ack=bool(messagejson.get("is_end")))
|
||||
|
||||
continue
|
||||
|
||||
# ========== 2) 处理“二进制音频消息” ==========
|
||||
if websocket.mode not in ("online", "offline", "2pass"):
|
||||
continue
|
||||
|
||||
if "chunk_size" not in websocket.status_dict_asr_online:
|
||||
print("[WARN] chunk_size not set yet, skip audio frame (send config first).")
|
||||
record_error("audio frame discarded: chunk_size is not configured")
|
||||
continue
|
||||
|
||||
try:
|
||||
websocket.status_dict_vad["chunk_size"] = int(
|
||||
websocket.status_dict_asr_online["chunk_size"][1] * 60 / websocket.chunk_interval
|
||||
)
|
||||
except Exception as e:
|
||||
print("[WARN] set vad chunk_size failed:", e)
|
||||
record_error(f"audio frame discarded: invalid VAD chunk_size: {e}")
|
||||
continue
|
||||
|
||||
pcm = message
|
||||
frames.append(pcm)
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
pending_offline_audio.append(pcm)
|
||||
|
||||
duration_ms = _pcm_duration_ms(pcm, fs=websocket.audio_fs, ch=1, sampwidth=2)
|
||||
websocket.vad_pre_idx += duration_ms
|
||||
|
||||
# online asr
|
||||
frames_asr_online.append(pcm)
|
||||
if websocket.mode in ("2pass", "online"):
|
||||
online_needs_finalization = True
|
||||
websocket.status_dict_asr_online["is_final"] = (speech_end_i != -1)
|
||||
|
||||
if (len(frames_asr_online) % websocket.chunk_interval == 0) or websocket.status_dict_asr_online["is_final"]:
|
||||
if websocket.mode in ("2pass", "online"):
|
||||
audio_in = b"".join(frames_asr_online)
|
||||
try:
|
||||
await async_asr_online(websocket, audio_in)
|
||||
except Exception as e:
|
||||
print(f"error in asr streaming, {websocket.status_dict_asr_online}")
|
||||
record_error(f"online inference failed: {e}")
|
||||
frames_asr_online = []
|
||||
|
||||
if speech_start:
|
||||
frames_asr.append(pcm)
|
||||
|
||||
# vad online
|
||||
try:
|
||||
speech_start_i, speech_end_i = await async_vad(websocket, pcm)
|
||||
except Exception as e:
|
||||
print("error in vad:", e)
|
||||
record_error(f"vad inference failed: {e}")
|
||||
speech_start_i, speech_end_i = -1, -1
|
||||
|
||||
if speech_start_i != -1:
|
||||
speech_start = True
|
||||
if duration_ms > 0:
|
||||
beg_bias = (websocket.vad_pre_idx - speech_start_i) // duration_ms
|
||||
else:
|
||||
beg_bias = 0
|
||||
frames_pre = frames[-beg_bias:] if beg_bias > 0 else []
|
||||
frames_asr = []
|
||||
frames_asr.extend(frames_pre)
|
||||
|
||||
# ========== 3) 2pass:离线阶段触发点 ==========
|
||||
if (speech_end_i != -1) or (not websocket.is_speaking):
|
||||
await finalize_online_segment()
|
||||
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
audio_in = b"".join(frames_asr)
|
||||
if not audio_in and speech_end_i != -1:
|
||||
audio_in = b"".join(pending_offline_audio)
|
||||
reason = "vad_end" if speech_end_i != -1 else "not_speaking"
|
||||
|
||||
# 保存 wav:放线程池,避免磁盘 IO 卡 loop
|
||||
if websocket.save_offline_segments and audio_in:
|
||||
try:
|
||||
await run_blocking(
|
||||
save_offline_wav_segment_sync,
|
||||
websocket,
|
||||
audio_in,
|
||||
reason,
|
||||
sem=SEM_WAV,
|
||||
)
|
||||
except Exception as e:
|
||||
print("[SAVE_OFFLINE_SEG] async failed:", e)
|
||||
|
||||
if audio_in:
|
||||
try:
|
||||
await async_asr(websocket, audio_in)
|
||||
pending_offline_audio = []
|
||||
except Exception as e:
|
||||
print("error in asr offline:", e)
|
||||
record_error(f"offline inference failed: {e}")
|
||||
|
||||
frames_asr = []
|
||||
speech_start = False
|
||||
frames_asr_online = []
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = False
|
||||
online_needs_finalization = False
|
||||
speech_end_i = -1
|
||||
|
||||
if not websocket.is_speaking:
|
||||
websocket.vad_pre_idx = 0
|
||||
frames = []
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
else:
|
||||
frames = frames[-20:]
|
||||
|
||||
except websockets.ConnectionClosed:
|
||||
print("ConnectionClosed...", websocket_users, flush=True)
|
||||
await ws_reset(websocket)
|
||||
if websocket in websocket_users:
|
||||
websocket_users.remove(websocket)
|
||||
except websockets.InvalidState:
|
||||
print("InvalidState...")
|
||||
try:
|
||||
await ws_reset(websocket)
|
||||
except Exception:
|
||||
pass
|
||||
websocket_users.discard(websocket)
|
||||
except Exception as e:
|
||||
print("Exception:", e)
|
||||
try:
|
||||
await ws_reset(websocket)
|
||||
except Exception:
|
||||
pass
|
||||
if websocket in websocket_users:
|
||||
websocket_users.remove(websocket)
|
||||
|
||||
|
||||
# ===================== 推理:全部改为“线程池 + 限流” =====================
|
||||
|
||||
async def async_vad(websocket, audio_in: bytes):
|
||||
# model_vad.generate 是阻塞的,必须 offload
|
||||
out = await run_blocking(_generate_sync, model_vad, audio_in, websocket.status_dict_vad, sem=SEM_VAD)
|
||||
segments_result = out[0].get("value", [])
|
||||
|
||||
speech_start = -1
|
||||
speech_end = -1
|
||||
|
||||
if len(segments_result) == 0 or len(segments_result) > 1:
|
||||
return speech_start, speech_end
|
||||
if segments_result[0][0] != -1:
|
||||
speech_start = segments_result[0][0]
|
||||
if segments_result[0][1] != -1:
|
||||
speech_end = segments_result[0][1]
|
||||
return speech_start, speech_end
|
||||
|
||||
|
||||
def _sv_and_match_sync(audio_in: bytes, reload_sec: int):
|
||||
"""
|
||||
同步执行:SV embedding + speaker_db 匹配
|
||||
返回 (spk_name, best_score)
|
||||
"""
|
||||
spk_name = "unknown"
|
||||
best_score = 0.0
|
||||
|
||||
sv_out = model_sv.generate(input=audio_in, embedding=True)[0]
|
||||
embedding = sv_out["spk_embedding"][0].cpu().numpy()
|
||||
|
||||
now_ts = time.time()
|
||||
local_speaker_db = get_speaker_db_cached(now_ts, reload_sec=reload_sec)
|
||||
if local_speaker_db:
|
||||
for name, ref_embedding in local_speaker_db.items():
|
||||
if ref_embedding is None:
|
||||
continue
|
||||
arr = np.array(ref_embedding, dtype=np.float32)
|
||||
similarity = 1.0 - cosine(embedding, arr)
|
||||
print("sv similarity with {}: {}".format(name, similarity))
|
||||
if similarity > best_score and similarity > 0.2:
|
||||
best_score = similarity
|
||||
spk_name = name
|
||||
|
||||
return spk_name, float(best_score)
|
||||
|
||||
|
||||
async def async_asr(websocket, audio_in: bytes):
|
||||
mode = "2pass-offline" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
if model_asr is None:
|
||||
raise RuntimeError("offline ASR is disabled; use FunASR online mode")
|
||||
|
||||
if len(audio_in) <= 0:
|
||||
message = {
|
||||
"mode": mode,
|
||||
"text": "",
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
return
|
||||
|
||||
# 1) ASR(阻塞,线程池执行)
|
||||
rec_result_list = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr,
|
||||
audio_in,
|
||||
websocket.status_dict_asr,
|
||||
sem=SEM_ASR_OFFLINE,
|
||||
)
|
||||
rec_result = rec_result_list[0]
|
||||
|
||||
print("offline_asr, raw:", rec_result)
|
||||
print("offline_asr, keys:", rec_result.keys())
|
||||
|
||||
text = rec_result.get("text", "")
|
||||
timestamp = rec_result.get("timestamp", None)
|
||||
sentence_info = rec_result.get("sentence_info", None)
|
||||
|
||||
# 2) 声纹识别(阻塞,线程池执行)
|
||||
spk_name = "unknown"
|
||||
best_score = 0.0
|
||||
try:
|
||||
spk_name, best_score = await run_blocking(
|
||||
_sv_and_match_sync,
|
||||
audio_in,
|
||||
int(args.speaker_db_reload_sec),
|
||||
sem=SEM_SV,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"声纹识别失败: {e}")
|
||||
|
||||
# 3) 标点(阻塞,线程池执行)
|
||||
punc_array = None
|
||||
if model_punc is not None and len(text) > 0:
|
||||
try:
|
||||
# punc 只对文本处理
|
||||
punc_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_punc,
|
||||
text,
|
||||
websocket.status_dict_punc,
|
||||
sem=SEM_PUNC,
|
||||
)
|
||||
punc_result = punc_out[0]
|
||||
print("offline, after punc", punc_result)
|
||||
|
||||
if "text" in punc_result and punc_result["text"]:
|
||||
text = punc_result["text"]
|
||||
if "punc_array" in punc_result:
|
||||
punc_array = punc_result["punc_array"]
|
||||
except Exception as e:
|
||||
print("punc failed:", e)
|
||||
|
||||
# 4) 构造最终 message
|
||||
if len(text) > 0:
|
||||
print("======offline final text:", text)
|
||||
message = {
|
||||
"mode": mode,
|
||||
"spk_name": spk_name,
|
||||
"spk_score": float(best_score),
|
||||
"text": text,
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
if timestamp is not None:
|
||||
message["timestamp"] = to_python(timestamp)
|
||||
if sentence_info is not None:
|
||||
message["sentence_info"] = to_python(sentence_info)
|
||||
if punc_array is not None:
|
||||
message["punc_array"] = to_python(punc_array)
|
||||
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
else:
|
||||
message = {
|
||||
"mode": mode,
|
||||
"spk_name": spk_name,
|
||||
"spk_score": float(best_score),
|
||||
"text": "",
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
|
||||
|
||||
async def async_asr_online(websocket, audio_in: bytes):
|
||||
if len(audio_in) <= 0 and not websocket.status_dict_asr_online.get("is_final", False):
|
||||
return
|
||||
|
||||
# streaming generate 也是阻塞:线程池执行
|
||||
rec_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr_streaming,
|
||||
audio_in,
|
||||
websocket.status_dict_asr_online,
|
||||
sem=SEM_ASR_ONLINE,
|
||||
)
|
||||
rec_result = rec_out[0]
|
||||
|
||||
# 2pass:online 只要 partial,不发 final(final 交给 offline)
|
||||
if websocket.mode == "2pass" and websocket.status_dict_asr_online.get("is_final", False):
|
||||
return
|
||||
|
||||
if rec_result.get("text"):
|
||||
mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
message = {
|
||||
"mode": mode,
|
||||
"text": rec_result["text"],
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": bool(
|
||||
websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking)
|
||||
),
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
|
||||
|
||||
# ===================== 启动服务 =====================
|
||||
|
||||
async def main():
|
||||
if len(args.certfile) > 0:
|
||||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
ssl_context.load_cert_chain(args.certfile, keyfile=args.keyfile)
|
||||
server = await websockets.serve(
|
||||
ws_serve,
|
||||
args.host,
|
||||
args.port,
|
||||
subprotocols=["binary"],
|
||||
ping_interval=None,
|
||||
ssl=ssl_context,
|
||||
)
|
||||
else:
|
||||
server = await websockets.serve(
|
||||
ws_serve,
|
||||
args.host,
|
||||
args.port,
|
||||
subprotocols=["binary"],
|
||||
ping_interval=None,
|
||||
)
|
||||
|
||||
print(f"WS server started at ws(s)://{args.host}:{args.port}")
|
||||
await server.wait_closed()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
asyncio.run(main())
|
||||
finally:
|
||||
try:
|
||||
EXECUTOR.shutdown(wait=False, cancel_futures=True)
|
||||
except Exception:
|
||||
pass
|
||||
|
|
@ -1,9 +1,4 @@
|
|||
"""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.
|
||||
"""
|
||||
"""Translate the unchanged Tencent demo protocol to FunASR's native online WS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -12,9 +7,7 @@ import asyncio
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import webbrowser
|
||||
from dataclasses import dataclass, replace
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
|
@ -22,620 +15,515 @@ 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
|
||||
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")
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
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__)
|
||||
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")))
|
||||
|
||||
|
||||
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),
|
||||
}
|
||||
AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
||||
SESSION_REGISTRY_KEY = web.AppKey("sessions", dict)
|
||||
|
||||
|
||||
class IncrementalWavDecoder:
|
||||
"""Strip a streamed RIFF header before forwarding PCM16 to FunASR."""
|
||||
"""Read a streamed PCM WAV header and yield its 16 kHz mono PCM payload."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.buffer = bytearray()
|
||||
self.payload_started = False
|
||||
self.riff_read = False
|
||||
self.header_done = 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:
|
||||
def feed(self, data: bytes) -> bytes:
|
||||
if self.header_done:
|
||||
if self.data_remaining is None:
|
||||
return chunk
|
||||
payload = chunk[: self.data_remaining]
|
||||
return data
|
||||
payload = data[: self.data_remaining]
|
||||
self.data_remaining -= len(payload)
|
||||
return payload
|
||||
|
||||
self.buffer.extend(chunk)
|
||||
if not self.riff_read:
|
||||
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("文件不是有效的 RIFF/WAV 音频")
|
||||
raise ValueError("WAV file must use a RIFF/WAVE container")
|
||||
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:
|
||||
raise ValueError("WAV header chunk is unexpectedly large")
|
||||
full_size = 8 + size + (size % 2)
|
||||
if len(self.buffer) < full_size:
|
||||
return b""
|
||||
body = self.buffer[8 : 8 + size]
|
||||
body = bytes(self.buffer[8 : 8 + size])
|
||||
if kind == b"fmt ":
|
||||
if size < 16:
|
||||
raise ValueError("WAV fmt 区块不完整")
|
||||
fields = (
|
||||
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 fields != (1, 1, 16000, 16):
|
||||
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
|
||||
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 必须为 16kHz、单声道、PCM16")
|
||||
self.payload_started = True
|
||||
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[:chunk_size]
|
||||
del self.buffer[:full_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 文件不完整,未收到全部音频数据")
|
||||
if not self.header_done or self.data_remaining not in (None, 0):
|
||||
raise ValueError("WAV ended before its complete PCM data chunk arrived")
|
||||
|
||||
|
||||
class RealtimeSession:
|
||||
"""Translate FunASR events into the existing sentence/display protocol."""
|
||||
@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,
|
||||
ws: web.WebSocketResponse,
|
||||
model_service: FunASRModelService,
|
||||
auxiliary_service: AuxiliaryModelService | None,
|
||||
browser_ws: web.WebSocketResponse,
|
||||
auxiliary: AuxiliaryModelService,
|
||||
start: dict[str, Any],
|
||||
voice_id: str,
|
||||
) -> None:
|
||||
self.ws = ws
|
||||
self.model_service = model_service
|
||||
self.auxiliary_service = auxiliary_service
|
||||
self.browser_ws = browser_ws
|
||||
self.auxiliary = auxiliary
|
||||
self.start = start
|
||||
self.engine = model_service.create_session()
|
||||
self.voice_id = voice_id
|
||||
self.session_id = uuid4().hex
|
||||
self.stop_event = asyncio.Event()
|
||||
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.speaker_jobs: asyncio.Queue[SpeakerJob | None] = asyncio.Queue()
|
||||
self.wav_decoder = (
|
||||
IncrementalWavDecoder()
|
||||
if self.source == "file" and Path(self.file_name).suffix.lower() == ".wav"
|
||||
if str(start.get("source") or "mic") == "file"
|
||||
and Path(str(start.get("file_name") or "")).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)
|
||||
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:
|
||||
"""Send ordered JSON while the browser connection is still alive."""
|
||||
"""Serialize browser writes because ASR and CAM++ finish independently."""
|
||||
async with self.send_lock:
|
||||
if not self.ws.closed:
|
||||
await self.ws.send_json(payload)
|
||||
if not self.browser_ws.closed:
|
||||
await self.browser_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(),
|
||||
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"),
|
||||
}
|
||||
if sentence is not None:
|
||||
await self.emit(
|
||||
{
|
||||
"type": "sentences",
|
||||
"sentences": [sentence],
|
||||
"metrics": self.metrics.snapshot(),
|
||||
}
|
||||
)
|
||||
await self.emit(state)
|
||||
await self.emit({"type": "sentences", "sentences": [sentence]})
|
||||
|
||||
async def warn_speaker(self, message: str) -> None:
|
||||
"""Expose speaker-service failures without interrupting ASR."""
|
||||
if self.speaker_warning_sent:
|
||||
async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None:
|
||||
if not frame:
|
||||
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
|
||||
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 _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()
|
||||
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)
|
||||
|
||||
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 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 _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++ 声纹并匹配说话人")
|
||||
async def read_native(self, native_ws: Any) -> None:
|
||||
"""Consume native FunASR events and keep its per-utterance partial cache."""
|
||||
try:
|
||||
speaker = await self.auxiliary_service.resolve_speaker(
|
||||
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,
|
||||
)
|
||||
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"},
|
||||
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(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(
|
||||
async def config_handler(_: web.Request) -> web.Response:
|
||||
"""Expose a small readiness response for the launcher and diagnostics."""
|
||||
return 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
|
||||
"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"
|
||||
),
|
||||
}
|
||||
)
|
||||
# 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 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:
|
||||
"""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
|
||||
"""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 ws.receive()
|
||||
first = await browser_ws.receive()
|
||||
if first.type != WSMsgType.TEXT:
|
||||
await ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||||
return ws
|
||||
try:
|
||||
await browser_ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||||
return browser_ws
|
||||
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
|
||||
|
||||
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 ws.send_json(
|
||||
await browser_ws.send_json(
|
||||
{
|
||||
"type": "error",
|
||||
"message": "实时流式测试的文件模式只支持 PCM 或 WAV,请改用麦克风、PCM 或 WAV",
|
||||
"message": "FunASR realtime mode accepts PCM or 16 kHz mono PCM WAV files.",
|
||||
}
|
||||
)
|
||||
return ws
|
||||
return browser_ws
|
||||
if source == "file" and suffix == ".pcm":
|
||||
start = dict(start)
|
||||
start["file_name"] = str(start.get("file_name") or "audio.pcm")
|
||||
|
||||
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}"
|
||||
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"})
|
||||
|
||||
await session.emit(
|
||||
# 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(
|
||||
{
|
||||
"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
|
||||
"mode": "online",
|
||||
"chunk_size": list(CHUNK_SIZE),
|
||||
"chunk_interval": CHUNK_INTERVAL,
|
||||
"encoder_chunk_look_back": int(
|
||||
os.getenv("FUNASR_ENCODER_LOOK_BACK", "4")
|
||||
),
|
||||
"speaker_service_health": speaker_health,
|
||||
"speaker_gap_enabled": False,
|
||||
"sentence_strategy": start.get("sentence_strategy", 0),
|
||||
"display_state_supported": True,
|
||||
}
|
||||
"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,
|
||||
)
|
||||
if speaker_health_error:
|
||||
await session.warn_speaker(speaker_health_error)
|
||||
)
|
||||
native_reader = asyncio.create_task(session.read_native(native_ws))
|
||||
speaker_worker = asyncio.create_task(session.resolve_speakers())
|
||||
|
||||
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]
|
||||
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, *workers],
|
||||
[receive_task, stop_task, native_reader],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if receive_task in done:
|
||||
return await receive_task
|
||||
if stop_task in done:
|
||||
receive_task.cancel()
|
||||
await asyncio.gather(receive_task, return_exceptions=True)
|
||||
for worker in workers:
|
||||
if worker in done:
|
||||
await worker
|
||||
raise RuntimeError("FunASR 实时处理任务意外结束")
|
||||
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
|
||||
|
||||
while not ws.closed:
|
||||
message = await receive_or_raise()
|
||||
message = await receive_task
|
||||
if message.type == WSMsgType.BINARY:
|
||||
await session.audio_queue.put(bytes(message.data))
|
||||
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 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)
|
||||
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 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()
|
||||
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": "end",
|
||||
"metrics": session.metrics.snapshot(),
|
||||
"sentences": session.assembler.raw_snapshot(),
|
||||
"display_blocks": session.assembler.display_blocks(session.merge_adjacent),
|
||||
"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("FunASR WebSocket session failed")
|
||||
if not ws.closed:
|
||||
await ws.send_json({"type": "error", "message": str(exc)})
|
||||
LOGGER.exception("Tencent-compatible WebSocket session failed")
|
||||
if not browser_ws.closed:
|
||||
await browser_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 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:
|
||||
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
|
||||
await session.auxiliary.reset_speaker_session(session.session_id)
|
||||
if not browser_ws.closed:
|
||||
await browser_ws.close()
|
||||
return browser_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)
|
||||
|
||||
async def create_app() -> web.Application:
|
||||
"""Create a light protocol bridge; model inference belongs to native FunASR."""
|
||||
app = web.Application()
|
||||
app[MODEL_SERVICE_KEY] = FunASRModelService(config)
|
||||
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
|
||||
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[MODEL_SERVICE_KEY].start()
|
||||
await application[AUXILIARY_SERVICE_KEY].start()
|
||||
await application[AUXILIARY_KEY].start()
|
||||
try:
|
||||
health = await application[AUXILIARY_SERVICE_KEY].health()
|
||||
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_SERVICE_KEY].close()
|
||||
await application[MODEL_SERVICE_KEY].close()
|
||||
await application[AUXILIARY_KEY].close()
|
||||
raise
|
||||
yield
|
||||
await application[AUXILIARY_SERVICE_KEY].close()
|
||||
await application[MODEL_SERVICE_KEY].close()
|
||||
await application[AUXILIARY_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("/api/stop", stop_handler)
|
||||
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,
|
||||
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__":
|
||||
|
|
|
|||
|
|
@ -4,3 +4,4 @@ funasr==1.4.16
|
|||
modelscope[framework]==1.34.0
|
||||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
websockets>=12,<14
|
||||
|
|
|
|||
|
|
@ -1,8 +1,5 @@
|
|||
// ===== 页面元素 =====
|
||||
// ===== DOM Elements =====
|
||||
const elEngineModel = document.getElementById('engineModel');
|
||||
const elModelServiceUrl = document.getElementById('modelServiceUrl');
|
||||
const elSpeakerStatus = document.getElementById('speakerStatus');
|
||||
const elDisplayMerge = document.getElementById('displayMerge');
|
||||
const elSpeakerDiarization = document.getElementById('speakerDiarization');
|
||||
const elDiarizationLabel = document.getElementById('diarizationLabel');
|
||||
const elSentenceStrategy = document.getElementById('sentenceStrategy');
|
||||
|
|
@ -22,13 +19,13 @@ const elMicStatus = document.getElementById('micStatus');
|
|||
const elMicTimer = document.getElementById('micTimer');
|
||||
const elMicElapsed = document.getElementById('micElapsed');
|
||||
|
||||
// 输入模式标签页
|
||||
// Input mode tabs
|
||||
const elTabMic = document.getElementById('tabMic');
|
||||
const elTabFile = document.getElementById('tabFile');
|
||||
const elPanelMic = document.getElementById('panelMic');
|
||||
const elPanelFile = document.getElementById('panelFile');
|
||||
|
||||
// 文件选择区域
|
||||
// File upload
|
||||
const elAudioFile = document.getElementById('audioFile');
|
||||
const elFileInfo = document.getElementById('fileInfo');
|
||||
const elAudioMeta = document.getElementById('audioMeta');
|
||||
|
|
@ -39,12 +36,12 @@ const elSpeedControl = document.getElementById('speedControl');
|
|||
const elSpeedSlider = document.getElementById('speedSlider');
|
||||
const elSpeedValue = document.getElementById('speedValue');
|
||||
|
||||
// ===== 说话人分离开关 =====
|
||||
// ===== Speaker Diarization Toggle =====
|
||||
elSpeakerDiarization.addEventListener('change', () => {
|
||||
elDiarizationLabel.textContent = elSpeakerDiarization.checked ? '开启' : '关闭';
|
||||
});
|
||||
|
||||
// ===== 日志区域 =====
|
||||
// ===== Log Area =====
|
||||
elBtnClearLog.addEventListener('click', () => { elLogArea.innerHTML = ''; });
|
||||
|
||||
function appendLog(msg) {
|
||||
|
|
@ -55,20 +52,12 @@ function appendLog(msg) {
|
|||
const typeClass = 'log-type-' + (msg.type || 'unknown');
|
||||
const entry = document.createElement('div');
|
||||
entry.className = 'log-entry';
|
||||
// 原始文本不作为 HTML 解释,转写中的标签也应原样显示。
|
||||
const stamp = document.createElement('span');
|
||||
stamp.className = 'log-time';
|
||||
stamp.textContent = ts;
|
||||
const content = document.createElement('span');
|
||||
content.className = typeClass;
|
||||
content.textContent = JSON.stringify(msg);
|
||||
entry.append(stamp, content);
|
||||
entry.innerHTML = `<span class="log-time">${ts}</span><span class="${typeClass}">${JSON.stringify(msg)}</span>`;
|
||||
elLogArea.appendChild(entry);
|
||||
while (elLogArea.childNodes.length > 300) elLogArea.firstChild.remove();
|
||||
elLogArea.scrollTop = elLogArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== 会话状态 =====
|
||||
// ===== State =====
|
||||
let ws = null;
|
||||
let sending = false;
|
||||
let stoppingByUser = false;
|
||||
|
|
@ -82,7 +71,7 @@ let micWorklet = null;
|
|||
let micTimerInterval = null;
|
||||
let micStartTime = 0;
|
||||
|
||||
// 输入模式(麦克风 / 文件)
|
||||
// Input mode (mic / file)
|
||||
let inputMode = 'mic';
|
||||
let selectedFile = null;
|
||||
|
||||
|
|
@ -91,28 +80,25 @@ const EXT_FORMAT_MAP = {
|
|||
'pcm': 1, 'wav': 12, 'mp3': 8, 'm4a': 14,
|
||||
'aac': 16, 'opus': 10, 'ogg': 10, 'silk': 6, 'speex': 4
|
||||
};
|
||||
// PCM/WAV 的默认发送倍速;实时验证默认按 1 倍速输入。
|
||||
// 不同格式的默认发送倍速:PCM/WAV 实时速度 1x,压缩格式解压快可提速
|
||||
const DEFAULT_SPEED = {
|
||||
'pcm': 1.0, 'wav': 1.0,
|
||||
'mp3': 2.0, 'm4a': 2.0, 'aac': 2.0,
|
||||
'opus': 3.0, 'ogg': 3.0, 'silk': 3.0, 'speex': 3.0
|
||||
};
|
||||
const MAX_SPEED = 3.0;
|
||||
// 实时 WebSocket 需要服务端逐帧读取音频;压缩格式必须等文件完整后才能解码,
|
||||
// 因此本次流式验证只允许 PCM/WAV,避免把整段上传伪装成实时识别。
|
||||
const STREAMABLE_AUDIO_EXTENSIONS = new Set(['pcm', 'wav']);
|
||||
const UNSUPPORTED_STREAMING_EXTENSIONS = new Set(['m4a']);
|
||||
|
||||
const SPEAKER_COLORS = ['#4a7dff', '#52c41a', '#faad14', '#ff4d4f', '#9254de', '#13c2c2'];
|
||||
|
||||
let sentenceMap = {};
|
||||
let speakerOrderMap = {};
|
||||
let speakerOrderCounter = 0;
|
||||
let displayStateSupported = false;
|
||||
let displayRevision = -1;
|
||||
let lastConfirmedBubbleEl = null;
|
||||
let pendingSpanMap = {};
|
||||
|
||||
// ===== 输入模式标签页 =====
|
||||
// ===== Input Mode Tabs =====
|
||||
function switchMode(mode) {
|
||||
if (ws) return;
|
||||
inputMode = mode;
|
||||
elTabMic.classList.toggle('active', mode === 'mic');
|
||||
elTabFile.classList.toggle('active', mode === 'file');
|
||||
|
|
@ -127,20 +113,20 @@ function switchMode(mode) {
|
|||
elTabMic.addEventListener('click', () => switchMode('mic'));
|
||||
elTabFile.addEventListener('click', () => switchMode('file'));
|
||||
|
||||
// ===== 文件选择 =====
|
||||
// ===== File Selection =====
|
||||
elAudioFile.addEventListener('change', (e) => {
|
||||
const file = e.target.files[0];
|
||||
if (!file) return;
|
||||
const ext = getFileExt(file.name);
|
||||
if (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
selectedFile = null;
|
||||
e.target.value = '';
|
||||
elFileInfo.textContent = '实时测试只支持 PCM 或 WAV,请先转换音频格式';
|
||||
elFileInfo.textContent = 'M4A 暂不支持直接上传,请先转成 WAV 或 MP3';
|
||||
elFileInfo.classList.remove('has-file');
|
||||
elAudioMeta.style.display = 'none';
|
||||
elSpeedControl.style.display = 'none';
|
||||
elBtnStart.disabled = true;
|
||||
showToast('压缩音频不能按当前实时 WebSocket 逐帧识别,请转成 PCM 或 WAV', true);
|
||||
showToast('M4A 容器格式无法按当前实时切片方式直接识别,请转成 WAV 或 MP3', true);
|
||||
return;
|
||||
}
|
||||
selectedFile = file;
|
||||
|
|
@ -151,12 +137,12 @@ elAudioFile.addEventListener('change', (e) => {
|
|||
parseAudioMeta(file);
|
||||
});
|
||||
|
||||
// 发送速度滑块
|
||||
// Speed slider
|
||||
elSpeedSlider.addEventListener('input', () => {
|
||||
elSpeedValue.textContent = parseFloat(elSpeedSlider.value).toFixed(1) + 'x';
|
||||
});
|
||||
|
||||
// ===== 音频元数据解析 =====
|
||||
// ===== Audio Meta Parsing =====
|
||||
function getFileExt(filename) {
|
||||
const parts = filename.split('.');
|
||||
return parts.length > 1 ? parts[parts.length - 1].toLowerCase() : '';
|
||||
|
|
@ -221,7 +207,7 @@ async function parseAudioMeta(file) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== 复制和提示 =====
|
||||
// ===== Copy & Toast =====
|
||||
function showToast(message, isError) {
|
||||
const toast = document.createElement('div');
|
||||
toast.className = 'copy-toast' + (isError ? ' copy-toast-error' : '');
|
||||
|
|
@ -246,7 +232,7 @@ function handleCopyClick(btn, textEl) {
|
|||
|
||||
elBtnCopyVoiceId.addEventListener('click', () => handleCopyClick(elBtnCopyVoiceId, elVoiceIdDisplay));
|
||||
|
||||
// ===== WAV 导出 =====
|
||||
// ===== WAV Export =====
|
||||
function buildWavBlob(pcmChunks) {
|
||||
let totalLen = 0;
|
||||
for (const c of pcmChunks) totalLen += c.byteLength;
|
||||
|
|
@ -292,7 +278,7 @@ elBtnExportWav.addEventListener('click', () => {
|
|||
showToast('WAV 已导出');
|
||||
});
|
||||
|
||||
// ===== 辅助函数 =====
|
||||
// ===== Helpers =====
|
||||
function formatTime(ms) {
|
||||
const totalSec = Math.floor(ms / 1000);
|
||||
const min = String(Math.floor(totalSec / 60)).padStart(2, '0');
|
||||
|
|
@ -309,10 +295,10 @@ function setStatus(state, text) {
|
|||
elStatusText.textContent = text;
|
||||
}
|
||||
|
||||
// ===== 渲染字幕(关闭说话人分离) =====
|
||||
// ===== Render: Subtitle (no diarization) =====
|
||||
// 每个 sentence_id 对应一个独立气泡:
|
||||
// - 中间态(sentence_type===0):实时更新该气泡文本,末尾加 "..." 表示未定。
|
||||
// - 稳态(sentence_type===1):去掉 "..." 并定格,下一个 sentence_id 自动创建新气泡。
|
||||
// - 中间态(sentence_type===0):实时更新该气泡文本,末尾加 "..." 表示未定
|
||||
// - 稳态(sentence_type===1):去掉 "..." 并定格,下一个 sentence_id 自动创建新气泡
|
||||
function renderSubtitle(sentence) {
|
||||
const id = 'subtitle-' + sentence.sentence_id;
|
||||
const isInterim = sentence.sentence_type === 0;
|
||||
|
|
@ -347,85 +333,146 @@ function renderSubtitle(sentence) {
|
|||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== 渲染说话人气泡 =====
|
||||
// 未确认片段独立展示,不能临时塞进上一位说话人的气泡。
|
||||
// ===== Render: Speaker Bubble =====
|
||||
function renderBubble(sentence) {
|
||||
const id = 'sent-' + sentence.sentence_id;
|
||||
const speakerId = Number(sentence.speaker_id);
|
||||
const trusted = Number.isInteger(speakerId) && speakerId >= 0
|
||||
&& ['fresh', 'confirmed'].includes(sentence.speaker_evidence);
|
||||
const speakerId = sentence.speaker_id;
|
||||
const isUnknown = speakerId < 0;
|
||||
const isInterim = sentence.sentence_type === 0;
|
||||
|
||||
if (isUnknown) {
|
||||
const pendingText = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
let pending = pendingSpanMap[id];
|
||||
if (pending) {
|
||||
pending.spanEl.textContent = pendingText;
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
return;
|
||||
}
|
||||
if (lastConfirmedBubbleEl) {
|
||||
const body = lastConfirmedBubbleEl.querySelector('.bubble-body');
|
||||
const span = document.createElement('span');
|
||||
span.className = 'pending-text';
|
||||
span.dataset.sentenceId = id;
|
||||
span.textContent = pendingText;
|
||||
body.appendChild(span);
|
||||
pendingSpanMap[id] = { hostEl: lastConfirmedBubbleEl, spanEl: span };
|
||||
} else {
|
||||
renderFallbackPendingBubble(id, pendingText, isInterim);
|
||||
}
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
return;
|
||||
}
|
||||
|
||||
if (pendingSpanMap[id]) {
|
||||
pendingSpanMap[id].spanEl.remove();
|
||||
delete pendingSpanMap[id];
|
||||
}
|
||||
|
||||
let entry = sentenceMap[id];
|
||||
let insertBefore = null;
|
||||
if (entry && entry.speakerId !== speakerId) {
|
||||
insertBefore = entry.el.nextSibling;
|
||||
entry.el.remove();
|
||||
entry = null;
|
||||
delete sentenceMap[id];
|
||||
}
|
||||
|
||||
if (!entry) {
|
||||
if (!(speakerId in speakerOrderMap)) {
|
||||
speakerOrderMap[speakerId] = speakerOrderCounter++;
|
||||
}
|
||||
const orderIdx = speakerOrderMap[speakerId];
|
||||
const side = orderIdx % 2 === 0 ? 'left' : 'right';
|
||||
const colorIdx = orderIdx % SPEAKER_COLORS.length;
|
||||
|
||||
const el = document.createElement('div');
|
||||
el.id = id;
|
||||
el.className = `bubble-row speaker-${side} speaker-${colorIdx}`;
|
||||
const wrapper = document.createElement('div');
|
||||
wrapper.className = 'bubble-wrapper';
|
||||
const header = document.createElement('div');
|
||||
header.className = 'bubble-header';
|
||||
for (const name of ['speaker-badge', 'speaker-name', 'bubble-time']) {
|
||||
const span = document.createElement('span');
|
||||
span.className = name;
|
||||
header.appendChild(span);
|
||||
}
|
||||
const badge = document.createElement('span');
|
||||
badge.className = `speaker-badge speaker-color-${colorIdx}`;
|
||||
const nameSpan = document.createElement('span');
|
||||
nameSpan.className = 'speaker-name';
|
||||
nameSpan.textContent = `说话人 ${speakerId}`;
|
||||
const timeSpan = document.createElement('span');
|
||||
timeSpan.className = 'bubble-time';
|
||||
header.appendChild(badge);
|
||||
header.appendChild(nameSpan);
|
||||
header.appendChild(timeSpan);
|
||||
const body = document.createElement('div');
|
||||
body.className = 'bubble-body';
|
||||
wrapper.append(header, body);
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
el.appendChild(wrapper);
|
||||
elResultArea.appendChild(el);
|
||||
entry = { el };
|
||||
|
||||
if (insertBefore) elResultArea.insertBefore(el, insertBefore);
|
||||
else elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: speakerId };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
if (trusted && !(speakerId in speakerOrderMap)) speakerOrderMap[speakerId] = speakerOrderCounter++;
|
||||
const order = trusted ? speakerOrderMap[speakerId] : 0;
|
||||
const color = order % SPEAKER_COLORS.length;
|
||||
|
||||
const el = entry.el;
|
||||
el.className = trusted ? `bubble-row speaker-${order % 2 ? 'right' : 'left'} speaker-${color}`
|
||||
: 'bubble-row speaker-left speaker-unknown';
|
||||
el.querySelector('.speaker-badge').className = 'speaker-badge speaker-color-' + (trusted ? color : 'unknown');
|
||||
// 未获得当前片段的可靠声纹证据时,标题保持简短;详细原因放到悬停提示,
|
||||
// 这样不会把“有效语音不足……”等内部诊断信息挤进说话人名称区域。
|
||||
const speakerName = el.querySelector('.speaker-name');
|
||||
speakerName.textContent = trusted
|
||||
? (sentence.speaker_name || `说话人 ${speakerId + 1}`)
|
||||
: '未知说话人';
|
||||
speakerName.title = trusted ? '' : (sentence.speaker_reason || '未匹配到说话人');
|
||||
el.querySelector('.bubble-time').textContent = formatTimeRange(sentence.start_time, sentence.end_time);
|
||||
const timeSpan = el.querySelector('.bubble-time');
|
||||
const body = el.querySelector('.bubble-body');
|
||||
body.textContent = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
timeSpan.textContent = formatTimeRange(sentence.start_time, sentence.end_time);
|
||||
body.querySelectorAll('.pending-text').forEach(s => s.remove());
|
||||
Array.from(body.childNodes).filter(n => n.nodeType === Node.TEXT_NODE).forEach(n => n.remove());
|
||||
const textNode = document.createTextNode(sentence.sentence + (isInterim ? ' ...' : ''));
|
||||
body.insertBefore(textNode, body.firstChild);
|
||||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
entry.speakerId = speakerId;
|
||||
lastConfirmedBubbleEl = el;
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// 按完整快照重建相邻块;序号防止两个后台 worker 的旧快照覆盖新状态。
|
||||
function renderDisplayState(msg, useSpeaker) {
|
||||
if (msg.revision != null && msg.revision <= displayRevision) return;
|
||||
if (msg.revision != null) displayRevision = msg.revision;
|
||||
elResultArea.replaceChildren();
|
||||
sentenceMap = {};
|
||||
const raw = msg.raw_segments || msg.sentences || [];
|
||||
if (useSpeaker) {
|
||||
for (const block of msg.display_blocks || []) renderBubble({ ...block, sentence_id: block.block_id });
|
||||
const confirmed = raw.filter(s => s.speaker_status === 'confirmed').length;
|
||||
const failed = raw.filter(s => ['service_error', 'service_unavailable', 'no_embedding', 'evidence_rejected'].includes(s.speaker_status)).length;
|
||||
elSpeakerStatus.textContent = `说话人:已确认 ${confirmed} / ${raw.length} 段` + (failed ? `,${failed} 段未识别成功(原因见气泡及日志)` : '');
|
||||
} else {
|
||||
raw.forEach(renderSubtitle);
|
||||
elSpeakerStatus.textContent = '说话人分离已关闭';
|
||||
function renderFallbackPendingBubble(id, text, isInterim) {
|
||||
let entry = sentenceMap[id];
|
||||
if (!entry) {
|
||||
const el = document.createElement('div');
|
||||
el.id = id;
|
||||
el.className = 'bubble-row speaker-left speaker-unknown';
|
||||
const wrapper = document.createElement('div');
|
||||
wrapper.className = 'bubble-wrapper';
|
||||
const header = document.createElement('div');
|
||||
header.className = 'bubble-header';
|
||||
const badge = document.createElement('span');
|
||||
badge.className = 'speaker-badge speaker-color-unknown';
|
||||
const nameSpan = document.createElement('span');
|
||||
nameSpan.className = 'speaker-name';
|
||||
nameSpan.textContent = '说话人不确定';
|
||||
const timeSpan = document.createElement('span');
|
||||
timeSpan.className = 'bubble-time';
|
||||
header.appendChild(badge);
|
||||
header.appendChild(nameSpan);
|
||||
header.appendChild(timeSpan);
|
||||
const body = document.createElement('div');
|
||||
body.className = 'bubble-body';
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
el.appendChild(wrapper);
|
||||
elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: -1 };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
const body = entry.el.querySelector('.bubble-body');
|
||||
body.textContent = text;
|
||||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
}
|
||||
|
||||
// ===== Start Recognition =====
|
||||
elBtnStart.addEventListener('click', () => {
|
||||
if (inputMode === 'file' && !selectedFile) return;
|
||||
startRecognition();
|
||||
});
|
||||
|
||||
async function startRecognition() {
|
||||
if (ws) return;
|
||||
if (inputMode === 'file' && selectedFile) {
|
||||
const ext = getFileExt(selectedFile.name);
|
||||
if (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
showToast('实时流式测试只支持 PCM 或 WAV,请转换后再试', true);
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
showToast('当前 demo 不支持直接流式上传 M4A,请转成 WAV 或 MP3', true);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
|
@ -435,9 +482,8 @@ async function startRecognition() {
|
|||
sentenceMap = {};
|
||||
speakerOrderMap = {};
|
||||
speakerOrderCounter = 0;
|
||||
displayStateSupported = false;
|
||||
displayRevision = -1;
|
||||
elSpeakerStatus.textContent = '正在检查说话人服务…';
|
||||
lastConfirmedBubbleEl = null;
|
||||
pendingSpanMap = {};
|
||||
audioChunks = [];
|
||||
elBtnExportWav.disabled = true;
|
||||
elResultPlaceholder?.remove();
|
||||
|
|
@ -450,10 +496,9 @@ async function startRecognition() {
|
|||
sending = true;
|
||||
|
||||
const currentSession = ++sessionId;
|
||||
const useSpeaker = true; // Speaker labels are mandatory for this deployment.
|
||||
let receivedTerminal = false;
|
||||
const useSpeaker = elSpeakerDiarization.checked;
|
||||
|
||||
// 构造 WebSocket 首条 start 消息。
|
||||
// 构造 start 消息
|
||||
let voiceFormat = 0, fileName = '', speedFactor = 0;
|
||||
if (inputMode === 'file') {
|
||||
const ext = getFileExt(selectedFile.name);
|
||||
|
|
@ -464,9 +509,7 @@ async function startRecognition() {
|
|||
|
||||
const startPayload = {
|
||||
type: 'start',
|
||||
model: elEngineModel.value,
|
||||
model_service_url: elModelServiceUrl.value,
|
||||
display_merge: elDisplayMerge.checked,
|
||||
engine_model_type: elEngineModel.value,
|
||||
speaker_diarization: useSpeaker ? 1 : 0,
|
||||
sentence_strategy: parseInt(elSentenceStrategy.value),
|
||||
source: inputMode,
|
||||
|
|
@ -475,23 +518,23 @@ async function startRecognition() {
|
|||
speed_factor: speedFactor
|
||||
};
|
||||
|
||||
const backendUrl = (window.ASR_BACKEND_URL || location.origin).replace(/\/$/, '');
|
||||
const websocketUrl = new URL(backendUrl + '/ws');
|
||||
websocketUrl.protocol = backendUrl.startsWith('https:') ? 'wss:' : 'ws:';
|
||||
ws = new WebSocket(websocketUrl);
|
||||
const protocol = location.protocol === 'https:' ? 'wss:' : 'ws:';
|
||||
ws = new WebSocket(`${protocol}//${location.host}/ws`);
|
||||
ws.binaryType = 'arraybuffer';
|
||||
|
||||
ws.onopen = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
ws.send(JSON.stringify(startPayload));
|
||||
// 在连接尚未建立时点击停止,也要在 start 后补发停止信号。
|
||||
if (!sending) ws.send(JSON.stringify({ type: 'stop' }));
|
||||
if (inputMode === 'file') {
|
||||
sendAudioFile(selectedFile);
|
||||
} else {
|
||||
startMicCapture();
|
||||
}
|
||||
};
|
||||
|
||||
ws.onmessage = (event) => {
|
||||
if (currentSession !== sessionId) return;
|
||||
const msg = JSON.parse(event.data);
|
||||
if (msg.type === 'end' || msg.type === 'error') receivedTerminal = true;
|
||||
if (msg.type !== 'sentences') {
|
||||
console.log('[ws] type=' + msg.type, msg);
|
||||
}
|
||||
|
|
@ -510,7 +553,8 @@ async function startRecognition() {
|
|||
ws.onclose = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
stopMicCapture();
|
||||
if (!receivedTerminal) setStatus('error', '连接中断,最终识别结果可能尚未完成');
|
||||
if (sending) setStatus('error', '连接意外断开');
|
||||
else if (stoppingByUser) setStatus('done', '已停止');
|
||||
if (audioChunks.length > 0) elBtnExportWav.disabled = false;
|
||||
ws = null;
|
||||
resetControls();
|
||||
|
|
@ -527,28 +571,10 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
break;
|
||||
|
||||
case 'start':
|
||||
displayStateSupported = Boolean(msg.display_state_supported);
|
||||
currentVoiceId = msg.session_id;
|
||||
elVoiceIdDisplay.textContent = currentVoiceId || '—';
|
||||
elSpeakerStatus.textContent = useSpeaker
|
||||
? `说话人服务:${msg.speaker_service_url || '未配置'};片段结束后提取声纹`
|
||||
: '说话人分离已关闭';
|
||||
if (!sending) break;
|
||||
setStatus('running', '识别中...');
|
||||
if (inputMode === 'file') sendAudioFile(selectedFile).catch(handleInputError);
|
||||
else startMicCapture().catch(handleInputError);
|
||||
break;
|
||||
|
||||
case 'display_state':
|
||||
renderDisplayState(msg, useSpeaker);
|
||||
break;
|
||||
|
||||
case 'draining':
|
||||
setStatus('running', msg.message || '等待最终识别结果…');
|
||||
break;
|
||||
|
||||
case 'sentences':
|
||||
if (displayStateSupported) break;
|
||||
if (msg.sentences) {
|
||||
msg.sentences.forEach(s => {
|
||||
if (useSpeaker) renderBubble(s);
|
||||
|
|
@ -557,14 +583,7 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
break;
|
||||
|
||||
case 'speaker_warning':
|
||||
// ASR 仍可继续输出,但必须让测试人员立即知道说话人链路没有生效。
|
||||
elSpeakerStatus.textContent = '说话人服务异常:' + msg.message;
|
||||
showToast('说话人服务异常,详见状态和片段原因', true);
|
||||
break;
|
||||
|
||||
case 'end':
|
||||
if (msg.display_blocks) renderDisplayState(msg, useSpeaker);
|
||||
setStatus('done', '识别完成');
|
||||
sending = false;
|
||||
if (audioChunks.length > 0) elBtnExportWav.disabled = false;
|
||||
|
|
@ -582,7 +601,7 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== 停止识别 =====
|
||||
// ===== Stop =====
|
||||
elBtnStop.addEventListener('click', () => stopRecognition());
|
||||
|
||||
function stopRecognition() {
|
||||
|
|
@ -592,12 +611,26 @@ function stopRecognition() {
|
|||
setStatus('running', '停止中...');
|
||||
elBtnStop.disabled = true;
|
||||
|
||||
if (currentVoiceId) {
|
||||
fetch(`/api/stop?voice_id=${encodeURIComponent(currentVoiceId)}`)
|
||||
.then(r => r.json())
|
||||
.catch(err => console.error('[stop] error:', err));
|
||||
}
|
||||
|
||||
if (ws && ws.readyState === WebSocket.OPEN) {
|
||||
try { ws.send(JSON.stringify({ type: 'stop' })); } catch (e) {}
|
||||
}
|
||||
|
||||
// 等待服务端排空 ASR/声纹队列后发送 end,不能用五秒计时器截断更新。
|
||||
|
||||
const stopSession = sessionId;
|
||||
setTimeout(() => {
|
||||
if (sessionId !== stopSession) return;
|
||||
if (ws) {
|
||||
try { ws.close(); } catch (e) {}
|
||||
ws = null;
|
||||
setStatus('done', '已停止(超时)');
|
||||
resetControls();
|
||||
}
|
||||
}, 5000);
|
||||
}
|
||||
|
||||
function resetControls() {
|
||||
|
|
@ -612,53 +645,36 @@ function resetControls() {
|
|||
elBtnStop.disabled = true;
|
||||
}
|
||||
|
||||
// ===== 发送音频文件 =====
|
||||
// 按 16KB 切片发送,并按照音频实际时长等待,确保文件模式也是真实的
|
||||
// 实时输入,而不是瞬间上传完整文件后再由服务端批量切片。
|
||||
// ===== Send Audio File =====
|
||||
// 按 16KB 切片发送,后端会缓冲成 6400 字节块并按 speed_factor 限流
|
||||
const UPLOAD_CHUNK_SIZE = 16000;
|
||||
async function sendAudioFile(file) {
|
||||
const ownerSession = sessionId;
|
||||
const buffer = await file.arrayBuffer();
|
||||
if (ownerSession !== sessionId || !sending) return;
|
||||
const totalBytes = buffer.byteLength;
|
||||
let offset = 0;
|
||||
const ext = getFileExt(file.name);
|
||||
const isPcm = (ext === 'pcm');
|
||||
let bytesPerSecond = 16000 * 2;
|
||||
if (ext === 'wav' && totalBytes >= 44) {
|
||||
const header = new DataView(buffer, 0, 44);
|
||||
const byteRate = header.getUint32(28, true);
|
||||
if (byteRate > 0) bytesPerSecond = byteRate;
|
||||
}
|
||||
const speedFactor = Math.max(parseFloat(elSpeedSlider.value) || 1.0, 0.1);
|
||||
while (ownerSession === sessionId && offset < totalBytes && sending && ws && ws.readyState === WebSocket.OPEN) {
|
||||
while (offset < totalBytes && sending && ws && ws.readyState === WebSocket.OPEN) {
|
||||
const end = Math.min(offset + UPLOAD_CHUNK_SIZE, totalBytes);
|
||||
const chunk = buffer.slice(offset, end);
|
||||
// 仅 PCM 数据可直接拼成 WAV 导出;当前实时模式不会接收压缩格式。
|
||||
// 仅 PCM 数据可直接拼成 WAV 导出;压缩格式跳过
|
||||
if (isPcm) audioChunks.push(chunk.slice(0));
|
||||
ws.send(chunk);
|
||||
offset = end;
|
||||
const chunkDurationMs = (chunk.byteLength / bytesPerSecond) * 1000 / speedFactor;
|
||||
await new Promise(r => setTimeout(r, Math.max(0, Math.round(chunkDurationMs))));
|
||||
await new Promise(r => setTimeout(r, 0));
|
||||
}
|
||||
if (ownerSession === sessionId && ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
sending = false;
|
||||
setStatus('running', '音频已发送,等待最终结果…');
|
||||
if (ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
ws.send(JSON.stringify({ type: 'eof' }));
|
||||
}
|
||||
}
|
||||
|
||||
// ===== 麦克风采集 =====
|
||||
// ===== Microphone Capture =====
|
||||
async function startMicCapture() {
|
||||
const ownerSession = sessionId;
|
||||
let stream;
|
||||
try {
|
||||
stream = await navigator.mediaDevices.getUserMedia({
|
||||
micStream = await navigator.mediaDevices.getUserMedia({
|
||||
audio: { sampleRate: 16000, channelCount: 1, echoCancellation: true, noiseSuppression: true }
|
||||
});
|
||||
} catch (err) {
|
||||
if (ownerSession !== sessionId) return;
|
||||
handleInputError(err);
|
||||
console.error('getUserMedia error:', err);
|
||||
setStatus('error', '无法获取麦克风权限');
|
||||
elMicStatus.textContent = '无法获取麦克风: ' + err.message;
|
||||
|
|
@ -666,11 +682,6 @@ async function startMicCapture() {
|
|||
return;
|
||||
}
|
||||
|
||||
if (ownerSession !== sessionId || !sending) {
|
||||
stream.getTracks().forEach(track => track.stop());
|
||||
return;
|
||||
}
|
||||
micStream = stream;
|
||||
micAudioContext = new (window.AudioContext || window.webkitAudioContext)({ sampleRate: 16000 });
|
||||
const source = micAudioContext.createMediaStreamSource(micStream);
|
||||
const processor = micAudioContext.createScriptProcessor(4096, 1, 1);
|
||||
|
|
@ -725,23 +736,3 @@ function stopMicCapture() {
|
|||
elMicStatus.textContent = '点击下方按钮开始录音';
|
||||
elMicElapsed.textContent = '00:00';
|
||||
}
|
||||
|
||||
// 展示实际部署端点及模型,避免沿用旧 SDK 的无效引擎配置。
|
||||
// The static frontend reads API and WS addresses from runtime-config.js.
|
||||
const backendUrl = (window.ASR_BACKEND_URL || location.origin).replace(/\/$/, '');
|
||||
fetch(backendUrl + '/api/config').then(response => response.json()).then(config => {
|
||||
if (!ws) {
|
||||
elEngineModel.value = config.model;
|
||||
elModelServiceUrl.value = config.model_service_url;
|
||||
elSpeakerStatus.textContent = '说话人辅助服务:' + (config.speaker_service_url || '未配置');
|
||||
}
|
||||
}).catch(error => { elSpeakerStatus.textContent = '读取服务配置失败:' + error.message; });
|
||||
|
||||
// 输入端失败必须释放空会话,避免用户再次开始时留下旧连接。
|
||||
function handleInputError(error) {
|
||||
sending = false;
|
||||
setStatus('error', error.message);
|
||||
stopMicCapture();
|
||||
if (ws) { ws.close(); ws = null; }
|
||||
resetControls();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,28 +20,22 @@
|
|||
<div class="form-stack">
|
||||
<div class="form-group">
|
||||
<label for="engineModel">引擎模型</label>
|
||||
<input type="text" id="engineModel" value="paraformer-zh-streaming">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="modelServiceUrl">FunASR 本地引擎</label>
|
||||
<input type="text" id="modelServiceUrl" value="local://funasr" readonly>
|
||||
<input type="text" id="engineModel" value="16k_zh_en_speaker" readonly>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="sentenceStrategy">分句策略</label>
|
||||
<select id="sentenceStrategy">
|
||||
<option value="0" selected>短停顿(800ms)</option>
|
||||
<option value="1">长停顿(1400ms)</option>
|
||||
<option value="0" selected>语义单句</option>
|
||||
<option value="1">段落</option>
|
||||
</select>
|
||||
</div>
|
||||
<!-- 展示合并与内部音频切段分开配置,便于对照原始片段。 -->
|
||||
<label><input type="checkbox" id="displayMerge" checked> 合并相邻且已确认的同一说话人</label>
|
||||
<div class="form-row">
|
||||
<div class="form-group">
|
||||
<label>话者分离</label>
|
||||
<label class="toggle">
|
||||
<input type="checkbox" id="speakerDiarization" checked disabled>
|
||||
<input type="checkbox" id="speakerDiarization" checked>
|
||||
<span class="toggle-slider"></span>
|
||||
<span class="toggle-label" id="diarizationLabel">必需</span>
|
||||
<span class="toggle-label" id="diarizationLabel">开启</span>
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
|
|
@ -60,9 +54,9 @@
|
|||
<div class="input-mode-panel" id="panelFile" style="display:none">
|
||||
<div class="file-select">
|
||||
<label class="btn btn-outline" for="audioFile">选择音频文件</label>
|
||||
<input type="file" id="audioFile" accept=".pcm,.wav" hidden>
|
||||
<input type="file" id="audioFile" accept=".pcm,.wav,.mp3,.aac,.opus,.ogg,.silk,.speex" hidden>
|
||||
</div>
|
||||
<span class="file-info" id="fileInfo">支持 16kHz、单声道、PCM16 的 PCM/WAV</span>
|
||||
<span class="file-info" id="fileInfo">支持 pcm/wav/mp3/aac/opus/silk/speex,m4a 请先转 wav/mp3</span>
|
||||
<div class="audio-meta" id="audioMeta" style="display:none">
|
||||
<div class="audio-meta-row"><span class="meta-k">格式</span><span class="meta-v" id="metaFormat">—</span></div>
|
||||
<div class="audio-meta-row"><span class="meta-k">采样率</span><span class="meta-v" id="metaSampleRate">—</span></div>
|
||||
|
|
@ -85,7 +79,6 @@
|
|||
<main class="panel-right">
|
||||
<section class="card result-card">
|
||||
<h2>识别结果</h2>
|
||||
<p id="speakerStatus" role="status">正在读取服务配置…</p>
|
||||
<div class="result-meta" id="resultMeta" style="display:none">
|
||||
<div class="meta-item">
|
||||
<span class="meta-label">VoiceID:</span>
|
||||
|
|
@ -113,8 +106,6 @@
|
|||
</main>
|
||||
</div>
|
||||
|
||||
<!-- 版本号变更用于刷新浏览器缓存,确保加载“未知说话人”标签逻辑。 -->
|
||||
<script src="runtime-config.js"></script>
|
||||
<script src="app.js?v=215"></script>
|
||||
<script src="app.js?v=210"></script>
|
||||
</body>
|
||||
</html>
|
||||
|
|
|
|||
|
|
@ -801,5 +801,3 @@ input[type="text"][readonly]:focus {
|
|||
.log-type-end { color: #f9e2af; }
|
||||
.log-type-error { color: #f38ba8; }
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,93 +1,102 @@
|
|||
"""WebSocket contract tests for the FunASR browser adapter."""
|
||||
"""Contract test for the unchanged Tencent UI to FunASR native WS bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
import json
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import AioHTTPTestCase
|
||||
|
||||
from realtime_websocket.funasr_engine import FunASRSegment
|
||||
from realtime_websocket.funasr_server import (
|
||||
AUXILIARY_SERVICE_KEY,
|
||||
MODEL_SERVICE_KEY,
|
||||
config_handler,
|
||||
AUXILIARY_KEY,
|
||||
SESSION_REGISTRY_KEY,
|
||||
websocket_handler,
|
||||
)
|
||||
|
||||
|
||||
class FakeFunASRSession:
|
||||
class FakeNativeWebSocket:
|
||||
"""Stand in for FunASR's native WSS process without loading model weights."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.sent = False
|
||||
self.incoming: asyncio.Queue[str] = asyncio.Queue()
|
||||
self.audio_bytes = 0
|
||||
self.config: dict[str, object] = {}
|
||||
|
||||
async def feed(self, audio: bytes):
|
||||
if self.sent:
|
||||
return []
|
||||
self.sent = True
|
||||
return [
|
||||
FunASRSegment(
|
||||
text="实时片段",
|
||||
start_time_ms=0,
|
||||
end_time_ms=len(audio) / 32,
|
||||
audio=b"",
|
||||
voiced_ms=len(audio) / 32,
|
||||
is_final=False,
|
||||
sentence_id=0,
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return None
|
||||
|
||||
async def send(self, payload: str | bytes) -> None:
|
||||
if isinstance(payload, bytes):
|
||||
self.audio_bytes += len(payload)
|
||||
return
|
||||
control = json.loads(payload)
|
||||
if "mode" in control:
|
||||
self.config = control
|
||||
return
|
||||
if control.get("is_end"):
|
||||
await self.incoming.put(
|
||||
json.dumps({"mode": "online", "text": "hello", "is_final": False})
|
||||
)
|
||||
]
|
||||
|
||||
async def finish(self):
|
||||
return [
|
||||
FunASRSegment(
|
||||
text="最终片段",
|
||||
start_time_ms=0,
|
||||
end_time_ms=1000,
|
||||
audio=b"\x01\x00" * 8000,
|
||||
voiced_ms=1000,
|
||||
is_final=True,
|
||||
sentence_id=0,
|
||||
reason="eof",
|
||||
await self.incoming.put(
|
||||
json.dumps({"mode": "online", "text": " world", "is_final": True})
|
||||
)
|
||||
await self.incoming.put(
|
||||
json.dumps({"is_end": True, "is_final": True})
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class FakeFunASRService:
|
||||
config = SimpleNamespace(model="fake-funasr")
|
||||
|
||||
def create_session(self):
|
||||
return FakeFunASRSession()
|
||||
async def recv(self) -> str:
|
||||
return await self.incoming.get()
|
||||
|
||||
|
||||
class FakeAuxiliaryService:
|
||||
config = SimpleNamespace(base_url="http://fake-speaker")
|
||||
async def resolve_speaker(self, _audio, _session_id, _start, _end):
|
||||
return {
|
||||
"speaker_id": 0,
|
||||
"speaker_name": "speaker 1",
|
||||
"speaker_confidence": 0.9,
|
||||
"speaker_status": "confirmed",
|
||||
}
|
||||
|
||||
async def health(self):
|
||||
return {"ready": True, "speaker_embedding_ready": True}
|
||||
|
||||
async def reset_speaker_session(self, session_id):
|
||||
async def reset_speaker_session(self, _session_id):
|
||||
return None
|
||||
|
||||
|
||||
class FunASRWebSocketTests(AioHTTPTestCase):
|
||||
class FunASRBridgeTests(AioHTTPTestCase):
|
||||
def get_app(self):
|
||||
app = web.Application()
|
||||
app[MODEL_SERVICE_KEY] = FakeFunASRService()
|
||||
app[AUXILIARY_SERVICE_KEY] = FakeAuxiliaryService()
|
||||
app.router.add_get("/api/config", config_handler)
|
||||
app[AUXILIARY_KEY] = FakeAuxiliaryService()
|
||||
app[SESSION_REGISTRY_KEY] = {}
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
return app
|
||||
|
||||
async def test_frontend_contract_uses_funasr_streaming_mode(self):
|
||||
async def test_tencent_ui_messages_use_native_funasr_and_keep_speaker_label(self):
|
||||
native = FakeNativeWebSocket()
|
||||
with patch(
|
||||
"realtime_websocket.funasr_server.websocket_connect",
|
||||
return_value=native,
|
||||
):
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "speaker_diarization": 0})
|
||||
start = await ws.receive_json()
|
||||
self.assertEqual(start["type"], "start")
|
||||
self.assertEqual(start["engine"], "funasr")
|
||||
self.assertEqual(start["partial_mode"], "funasr_streaming_cache")
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
# The server keeps speaker labeling enabled even if this flag is false.
|
||||
"speaker_diarization": 0,
|
||||
}
|
||||
)
|
||||
first = await ws.receive_json()
|
||||
second = await ws.receive_json()
|
||||
self.assertEqual(first["type"], "voice_id")
|
||||
self.assertEqual(second["type"], "start")
|
||||
|
||||
await ws.send_bytes(b"\x01\x00" * 16000)
|
||||
pcm = b"\x01\x00" * 16000
|
||||
await ws.send_bytes(pcm)
|
||||
await ws.send_json({"type": "eof"})
|
||||
|
||||
messages = []
|
||||
|
|
@ -98,13 +107,19 @@ class FunASRWebSocketTests(AioHTTPTestCase):
|
|||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
self.assertEqual(native.config["mode"], "online")
|
||||
self.assertEqual(native.config["audio_fs"], 16000)
|
||||
self.assertEqual(native.audio_bytes, len(pcm))
|
||||
sentence_events = [
|
||||
item for item in messages
|
||||
if item["type"] == "sentences" and item["sentences"]
|
||||
sentence
|
||||
for message in messages
|
||||
if message["type"] == "sentences"
|
||||
for sentence in message["sentences"]
|
||||
]
|
||||
self.assertTrue(any(item["sentences"][0]["sentence_type"] == 0 for item in sentence_events))
|
||||
self.assertEqual(sentence_events[-1]["sentences"][0]["sentence"], "最终片段")
|
||||
self.assertEqual(messages[-1]["type"], "end")
|
||||
self.assertTrue(any(sentence["sentence_type"] == 0 for sentence in sentence_events))
|
||||
final_events = [sentence for sentence in sentence_events if sentence["sentence_type"] == 1]
|
||||
self.assertEqual(final_events[-1]["sentence"], "hello world")
|
||||
self.assertEqual(final_events[-1]["speaker_id"], 0)
|
||||
await ws.close()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -6,3 +6,4 @@ modelscope[framework]==1.34.0
|
|||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
numpy>=1.24
|
||||
websockets>=12,<14
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start the local CAM++ model service and FunASR WebSocket backend together."""
|
||||
"""Start local CAM++, FunASR native realtime WSS, and the browser protocol bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
|
@ -30,7 +31,7 @@ load_dotenv(PROJECT_ROOT / ".env")
|
|||
|
||||
|
||||
def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
||||
"""Resolve ASR/VAD IDs through the shared model manifest."""
|
||||
"""Resolve a local ASR or VAD model through the project manifest."""
|
||||
name = requested.strip()
|
||||
configured_path = Path(name)
|
||||
direct_candidates = (
|
||||
|
|
@ -56,7 +57,7 @@ def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
|||
|
||||
|
||||
def local_cam_model(models_dir: Path) -> Path:
|
||||
"""Require one complete CAM++ speaker verification asset from the manifest."""
|
||||
"""Require a complete CAM++ speaker verification asset from the manifest."""
|
||||
manifest = load_manifest()
|
||||
override = os.getenv("CAM_MODEL_PATH", "").strip()
|
||||
if override:
|
||||
|
|
@ -86,7 +87,7 @@ def local_cam_model(models_dir: Path) -> Path:
|
|||
|
||||
|
||||
def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int = 300) -> None:
|
||||
"""Wait until a child has loaded its models, failing when it exits early."""
|
||||
"""Wait for an HTTP child to become ready, reporting early process exit."""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
|
|
@ -103,8 +104,29 @@ def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int
|
|||
raise TimeoutError(f"Service did not become ready within {seconds}s: {url}")
|
||||
|
||||
|
||||
def wait_for_tcp(host: str, port: int, process: subprocess.Popen, seconds: int = 300) -> None:
|
||||
"""Wait until FunASR has loaded its models and opened the internal WS socket."""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"FunASR native WS exited before readiness (exit={code})")
|
||||
try:
|
||||
with socket.create_connection((host, port), timeout=0.5):
|
||||
pass
|
||||
# Catch an address-in-use failure instead of accepting another process's port.
|
||||
time.sleep(0.5)
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"FunASR native WS exited before readiness (exit={code})")
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(0.5)
|
||||
raise TimeoutError(f"FunASR native WS did not open {host}:{port} within {seconds}s")
|
||||
|
||||
|
||||
def stop_child(process: subprocess.Popen | None) -> None:
|
||||
"""Stop a supervised model or WebSocket process on shutdown."""
|
||||
"""Stop a supervised model or WebSocket process on launcher shutdown."""
|
||||
if process is None or process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
|
|
@ -116,7 +138,7 @@ def stop_child(process: subprocess.Popen | None) -> None:
|
|||
|
||||
|
||||
def main() -> None:
|
||||
"""Require ASR, VAD, and CAM++ before exposing the backend WebSocket."""
|
||||
"""Require ASR, VAD, and CAM++ before exposing the public WS bridge."""
|
||||
models_dir = Path(os.getenv("MODEL_DIR", "models"))
|
||||
if not models_dir.is_absolute():
|
||||
models_dir = PROJECT_ROOT / models_dir
|
||||
|
|
@ -126,35 +148,102 @@ def main() -> None:
|
|||
)
|
||||
vad = local_model(os.getenv("FUNASR_VAD_MODEL", "fsmn-vad"), models_dir, "vad")
|
||||
cam = local_cam_model(models_dir)
|
||||
print(f"Local models: ASR={asr}; VAD={vad}; CAM++={cam}", flush=True)
|
||||
|
||||
native_host = os.getenv("FUNASR_NATIVE_WS_HOST", "127.0.0.1")
|
||||
native_port = int(os.getenv("FUNASR_NATIVE_WS_PORT", "10095"))
|
||||
native_url = f"ws://{native_host}:{native_port}"
|
||||
aux_url = os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010").rstrip("/")
|
||||
device = os.getenv("FUNASR_DEVICE", "cuda:0")
|
||||
vad_device = os.getenv("FUNASR_VAD_DEVICE", "cpu")
|
||||
ngpu = "1" if device.startswith("cuda") else "0"
|
||||
ncpu = os.getenv("FUNASR_NCPU", str(os.cpu_count() or 4))
|
||||
web_port = int(os.getenv("WEB_PORT", "8082"))
|
||||
|
||||
print(f"Local models: ASR={asr}; VAD={vad}; CAM++={cam}", flush=True)
|
||||
print(
|
||||
f"Runtime: ASR device={device}; VAD device={vad_device}; native WS={native_url}",
|
||||
flush=True,
|
||||
)
|
||||
env = os.environ.copy()
|
||||
env.update({
|
||||
preload_kinds = {
|
||||
item.strip() for item in env.get("AUXILIARY_PRELOAD_KINDS", "speaker_verification").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
# Native FunASR owns realtime VAD; don't load a duplicate VAD in the CAM++ process.
|
||||
preload_kinds.discard("vad")
|
||||
preload_kinds.add("speaker_verification")
|
||||
env.update(
|
||||
{
|
||||
"AUXILIARY_PRELOAD_KINDS": ",".join(sorted(preload_kinds)),
|
||||
"MODEL_DIR": str(models_dir),
|
||||
"FUNASR_ASR_MODEL": str(asr),
|
||||
"FUNASR_VAD_MODEL": str(vad),
|
||||
"FUNASR_DEVICE": device,
|
||||
"FUNASR_VAD_DEVICE": vad_device,
|
||||
"FUNASR_NATIVE_WS_URL": native_url,
|
||||
"CAM_MODEL_PATH": str(cam),
|
||||
"AUXILIARY_SERVICE_URL": "http://127.0.0.1:8010",
|
||||
})
|
||||
"AUXILIARY_SERVICE_URL": aux_url,
|
||||
}
|
||||
)
|
||||
|
||||
auxiliary = None
|
||||
native = None
|
||||
websocket = None
|
||||
try:
|
||||
auxiliary = subprocess.Popen(
|
||||
[sys.executable, "-m", "scripts.auxiliary_server"],
|
||||
cwd=PROJECT_ROOT, env=env,
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
wait_for_health("http://127.0.0.1:8010/health", auxiliary, "speaker_embedding_ready")
|
||||
wait_for_health(f"{aux_url}/health", auxiliary, "speaker_embedding_ready")
|
||||
|
||||
native_args = [
|
||||
sys.executable,
|
||||
str(PROJECT_ROOT / "realtime_websocket" / "funasr_native_wss.py"),
|
||||
"--host",
|
||||
native_host,
|
||||
"--port",
|
||||
str(native_port),
|
||||
"--asr_model",
|
||||
"",
|
||||
"--asr_model_online",
|
||||
str(asr),
|
||||
"--vad_model",
|
||||
str(vad),
|
||||
"--punc_model",
|
||||
"",
|
||||
"--device",
|
||||
device,
|
||||
"--vad_device",
|
||||
vad_device,
|
||||
"--ngpu",
|
||||
ngpu,
|
||||
"--ncpu",
|
||||
ncpu,
|
||||
"--certfile",
|
||||
"",
|
||||
"--keyfile",
|
||||
"",
|
||||
]
|
||||
native = subprocess.Popen(native_args, cwd=PROJECT_ROOT, env=env)
|
||||
probe_host = "127.0.0.1" if native_host in {"0.0.0.0", "::"} else native_host
|
||||
wait_for_tcp(probe_host, native_port, native)
|
||||
|
||||
websocket = subprocess.Popen(
|
||||
[sys.executable, "-m", "scripts.run_funasr_demo", "--no-browser"],
|
||||
cwd=PROJECT_ROOT, env=env,
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
web_port = int(env.get("WEB_PORT", "8082"))
|
||||
wait_for_health(
|
||||
f"http://127.0.0.1:{web_port}/api/config", websocket, "engine"
|
||||
)
|
||||
print(f"Backend ready: ws://127.0.0.1:{web_port}/ws", flush=True)
|
||||
print(f"FunASR browser backend ready: ws://127.0.0.1:{web_port}/ws", flush=True)
|
||||
while True:
|
||||
for label, process in (("CAM++", auxiliary), ("WebSocket", websocket)):
|
||||
for label, process in (
|
||||
("CAM++", auxiliary),
|
||||
("FunASR native WS", native),
|
||||
("browser WS adapter", websocket),
|
||||
):
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"{label} service exited (exit={code})")
|
||||
|
|
@ -163,6 +252,7 @@ def main() -> None:
|
|||
pass
|
||||
finally:
|
||||
stop_child(websocket)
|
||||
stop_child(native)
|
||||
stop_child(auxiliary)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,56 +1,131 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Serve the browser UI independently of the FunASR backend."""
|
||||
"""Serve the unchanged Tencent demo UI and proxy its API to the backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import os
|
||||
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlsplit
|
||||
from typing import Any
|
||||
|
||||
from aiohttp import ClientSession, ClientTimeout, WSMsgType, web
|
||||
from dotenv import load_dotenv
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
STATIC_ROOT = PROJECT_ROOT / "realtime_websocket" / "static"
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
FRONTEND_HOST = os.getenv("FRONTEND_HOST", "127.0.0.1")
|
||||
FRONTEND_PORT = int(os.getenv("FRONTEND_PORT", "8080"))
|
||||
BACKEND_BASE_URL = os.getenv(
|
||||
"BACKEND_INTERNAL_URL",
|
||||
f"http://127.0.0.1:{os.getenv('WEB_PORT', '8082')}",
|
||||
).rstrip("/")
|
||||
HTTP_SESSION = web.AppKey("http_session", ClientSession)
|
||||
|
||||
|
||||
class FrontendHandler(SimpleHTTPRequestHandler):
|
||||
"""Serve static assets and a browser-visible backend URL at runtime."""
|
||||
def backend_url(request: web.Request) -> str:
|
||||
"""Keep the original path and query while routing through the backend port."""
|
||||
return f"{BACKEND_BASE_URL}{request.rel_url}"
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, directory=str(STATIC_ROOT), **kwargs)
|
||||
|
||||
def do_GET(self) -> None:
|
||||
if self.path.split("?", 1)[0] != "/runtime-config.js":
|
||||
return super().do_GET()
|
||||
backend_url = os.getenv("BACKEND_PUBLIC_URL", "http://127.0.0.1:8082").rstrip("/")
|
||||
parsed = urlsplit(backend_url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc or parsed.path:
|
||||
self.send_error(500, "BACKEND_PUBLIC_URL must be an HTTP origin")
|
||||
async def index_handler(_: web.Request) -> web.FileResponse:
|
||||
return web.FileResponse(
|
||||
STATIC_ROOT / "index.html", headers={"Cache-Control": "no-store"}
|
||||
)
|
||||
|
||||
|
||||
async def api_stop_proxy(request: web.Request) -> web.Response:
|
||||
"""Forward the Tencent page's existing stop request to the WS backend."""
|
||||
async with request.app[HTTP_SESSION].get(
|
||||
backend_url(request), timeout=ClientTimeout(total=5)
|
||||
) as response:
|
||||
body = await response.read()
|
||||
return web.Response(
|
||||
status=response.status,
|
||||
body=body,
|
||||
headers={"Content-Type": response.headers.get("Content-Type", "application/json")},
|
||||
)
|
||||
|
||||
|
||||
async def websocket_proxy(request: web.Request) -> web.WebSocketResponse:
|
||||
"""Relay text and binary frames without changing the Tencent browser protocol."""
|
||||
browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30)
|
||||
await browser_ws.prepare(request)
|
||||
try:
|
||||
backend_ws = await request.app[HTTP_SESSION].ws_connect(
|
||||
backend_url(request),
|
||||
max_msg_size=64 * 1024 * 1024,
|
||||
heartbeat=30,
|
||||
autoping=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
await browser_ws.send_json({"type": "error", "message": f"backend unavailable: {exc}"})
|
||||
await browser_ws.close(code=1011, message=b"backend unavailable")
|
||||
return browser_ws
|
||||
|
||||
async def relay(source: Any, destination: Any) -> None:
|
||||
async for message in source:
|
||||
if message.type == WSMsgType.TEXT:
|
||||
await destination.send_str(message.data)
|
||||
elif message.type == WSMsgType.BINARY:
|
||||
await destination.send_bytes(message.data)
|
||||
elif message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}:
|
||||
if not destination.closed:
|
||||
code = message.data if isinstance(message.data, int) else 1000
|
||||
reason = message.extra or ""
|
||||
await destination.close(code=code, message=str(reason).encode("utf-8"))
|
||||
return
|
||||
# JSON escaping also yields a valid JavaScript string literal.
|
||||
body = ("window.ASR_BACKEND_URL = " + json.dumps(backend_url) + ";\n").encode("utf-8")
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "application/javascript; charset=utf-8")
|
||||
self.send_header("Cache-Control", "no-store")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(relay(browser_ws, backend_ws)),
|
||||
asyncio.create_task(relay(backend_ws, browser_ws)),
|
||||
]
|
||||
try:
|
||||
done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
await asyncio.gather(*done, *pending, return_exceptions=True)
|
||||
finally:
|
||||
if not backend_ws.closed:
|
||||
await backend_ws.close()
|
||||
if not browser_ws.closed:
|
||||
await browser_ws.close()
|
||||
return browser_ws
|
||||
|
||||
|
||||
async def create_app() -> web.Application:
|
||||
if not STATIC_ROOT.is_dir():
|
||||
raise FileNotFoundError(f"Tencent demo static directory is missing: {STATIC_ROOT}")
|
||||
parsed = urlsplit(BACKEND_BASE_URL)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc or parsed.path:
|
||||
raise ValueError("BACKEND_INTERNAL_URL must contain only an HTTP origin")
|
||||
|
||||
app = web.Application()
|
||||
|
||||
async def lifecycle(application: web.Application):
|
||||
# An unbounded total timeout allows long recordings and slow model loads.
|
||||
application[HTTP_SESSION] = ClientSession(
|
||||
timeout=ClientTimeout(total=None, connect=10, sock_connect=10, sock_read=None)
|
||||
)
|
||||
yield
|
||||
await application[HTTP_SESSION].close()
|
||||
|
||||
app.cleanup_ctx.append(lifecycle)
|
||||
app.router.add_get("/", index_handler)
|
||||
app.router.add_get("/ws", websocket_proxy)
|
||||
app.router.add_get("/api/stop", api_stop_proxy)
|
||||
app.router.add_static("/", STATIC_ROOT, show_index=False)
|
||||
return app
|
||||
|
||||
|
||||
def main() -> None:
|
||||
host = os.getenv("FRONTEND_HOST", "127.0.0.1")
|
||||
port = int(os.getenv("FRONTEND_PORT", "8080"))
|
||||
server = ThreadingHTTPServer((host, port), FrontendHandler)
|
||||
print(f"Frontend: http://{host}:{port}/", flush=True)
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
finally:
|
||||
server.server_close()
|
||||
print(
|
||||
f"Tencent demo frontend: http://{FRONTEND_HOST}:{FRONTEND_PORT}/ "
|
||||
f"(backend proxy: {BACKEND_BASE_URL})",
|
||||
flush=True,
|
||||
)
|
||||
web.run_app(create_app(), host=FRONTEND_HOST, port=FRONTEND_PORT)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Reference in New Issue