741 lines
31 KiB
Python
741 lines
31 KiB
Python
"""将未修改的腾讯演示协议转换为 FunASR 原生在线 WebSocket 协议。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
import logging
|
||
import os
|
||
import re
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from uuid import uuid4
|
||
|
||
from aiohttp import WSMsgType, web
|
||
from dotenv import load_dotenv
|
||
|
||
try:
|
||
from websockets.asyncio.client import connect as websocket_connect
|
||
except ImportError: # websockets 13 之前的版本在包根目录提供相同客户端。
|
||
from websockets import connect as websocket_connect
|
||
|
||
try:
|
||
from .auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||
except ImportError:
|
||
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||
|
||
|
||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||
load_dotenv(PROJECT_ROOT / ".env")
|
||
LOGGER = logging.getLogger(__name__)
|
||
WEB_HOST = os.getenv("WEB_HOST", "0.0.0.0")
|
||
WEB_PORT = int(os.getenv("WEB_PORT", "8082"))
|
||
NATIVE_WS_URL = os.getenv("FUNASR_NATIVE_WS_URL", "ws://127.0.0.1:10095")
|
||
CHUNK_SIZE = tuple(
|
||
int(part.strip()) for part in os.getenv("FUNASR_CHUNK_SIZE", "0,10,5").split(",")
|
||
)
|
||
if len(CHUNK_SIZE) != 3:
|
||
CHUNK_SIZE = (0, 10, 5)
|
||
CHUNK_INTERVAL = max(1, int(os.getenv("FUNASR_CHUNK_INTERVAL", "10")))
|
||
SAMPLE_RATE = 16000
|
||
PCM_BYTES_PER_MS = SAMPLE_RATE * 2 / 1000
|
||
FRAME_BYTES = max(2, round(60 * CHUNK_SIZE[1] / CHUNK_INTERVAL * PCM_BYTES_PER_MS))
|
||
MAX_SPEAKER_AUDIO_BYTES = 60 * SAMPLE_RATE * 2
|
||
MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2)
|
||
FINALIZE_TIMEOUT_SECONDS = max(30, int(os.getenv("FUNASR_FINALIZE_TIMEOUT_SECONDS", "300")))
|
||
TURN_SENTENCE_ID_STRIDE = 100
|
||
|
||
AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
||
SESSION_REGISTRY_KEY = web.AppKey("sessions", dict)
|
||
|
||
|
||
def normalize_hotwords(value: Any) -> str:
|
||
"""去重并规范化浏览器输入,供 Contextual Paraformer 按空格解析词条。"""
|
||
if isinstance(value, list):
|
||
source = ",".join(str(item) for item in value if item is not None)
|
||
elif isinstance(value, str):
|
||
source = value
|
||
else:
|
||
return ""
|
||
|
||
terms = re.split(r"[,\uff0c\u3001;\uff1b\r\n]+", source)
|
||
# 保持用户输入顺序,并去掉重复项,避免同一热词被重复传给解码端。
|
||
unique_terms = list(dict.fromkeys(term.strip() for term in terms if term.strip()))
|
||
return " ".join(unique_terms)
|
||
|
||
|
||
class IncrementalWavDecoder:
|
||
"""读取流式 PCM WAV 文件头,并逐段返回 16 kHz 单声道 PCM 负载。"""
|
||
|
||
def __init__(self) -> None:
|
||
self.buffer = bytearray()
|
||
self.header_done = False
|
||
self.format_valid = False
|
||
self.data_remaining: int | None = None
|
||
|
||
def feed(self, data: bytes) -> bytes:
|
||
if self.header_done:
|
||
if self.data_remaining is None:
|
||
return data
|
||
payload = data[: self.data_remaining]
|
||
self.data_remaining -= len(payload)
|
||
return payload
|
||
|
||
self.buffer.extend(data)
|
||
if len(self.buffer) < 12:
|
||
return b""
|
||
if self.buffer[:4] != b"RIFF" or self.buffer[8:12] != b"WAVE":
|
||
raise ValueError("WAV file must use a RIFF/WAVE container")
|
||
del self.buffer[:12]
|
||
|
||
while len(self.buffer) >= 8:
|
||
kind = bytes(self.buffer[:4])
|
||
size = int.from_bytes(self.buffer[4:8], "little")
|
||
if size > 1024 * 1024:
|
||
raise ValueError("WAV header chunk is unexpectedly large")
|
||
full_size = 8 + size + (size % 2)
|
||
if len(self.buffer) < full_size:
|
||
return b""
|
||
body = bytes(self.buffer[8 : 8 + size])
|
||
if kind == b"fmt ":
|
||
if size < 16:
|
||
raise ValueError("WAV fmt chunk is incomplete")
|
||
fmt = (
|
||
int.from_bytes(body[0:2], "little"),
|
||
int.from_bytes(body[2:4], "little"),
|
||
int.from_bytes(body[4:8], "little"),
|
||
int.from_bytes(body[14:16], "little"),
|
||
)
|
||
if fmt != (1, 1, SAMPLE_RATE, 16):
|
||
raise ValueError("WAV must be PCM16, mono, 16 kHz")
|
||
self.format_valid = True
|
||
if kind == b"data":
|
||
if not self.format_valid or size % 2:
|
||
raise ValueError("WAV must be PCM16, mono, 16 kHz")
|
||
self.header_done = True
|
||
self.data_remaining = size
|
||
del self.buffer[:8]
|
||
payload = bytes(self.buffer[:size])
|
||
del self.buffer[: min(size, len(self.buffer))]
|
||
self.data_remaining -= len(payload)
|
||
return payload
|
||
del self.buffer[:full_size]
|
||
return b""
|
||
|
||
def finish(self) -> None:
|
||
if not self.header_done or self.data_remaining not in (None, 0):
|
||
raise ValueError("WAV ended before its complete PCM data chunk arrived")
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class SpeakerJob:
|
||
sentence_id: int
|
||
text: str
|
||
audio: bytes
|
||
start_time_ms: float
|
||
end_time_ms: float
|
||
|
||
|
||
def split_text_by_speaker_segments(
|
||
text: str,
|
||
segments: list[dict[str, Any]],
|
||
turn_start_ms: float,
|
||
turn_end_ms: float,
|
||
) -> list[dict[str, Any]]:
|
||
"""按 CAM++ 时间段近似拆分 ASR 文本,并保留每段稳定说话人标签。"""
|
||
usable: list[dict[str, Any]] = []
|
||
for segment in sorted(
|
||
segments,
|
||
key=lambda item: float(item.get("start_time", turn_start_ms)),
|
||
):
|
||
try:
|
||
start = max(turn_start_ms, float(segment.get("start_time", turn_start_ms)))
|
||
end = min(turn_end_ms, float(segment.get("end_time", turn_end_ms)))
|
||
speaker_id = int(segment.get("speaker_id", -1))
|
||
except (TypeError, ValueError):
|
||
continue
|
||
if end <= start or speaker_id < 0:
|
||
continue
|
||
speaker = {
|
||
"speaker_id": speaker_id,
|
||
"speaker_name": str(segment.get("speaker_name") or f"说话人 {speaker_id + 1}"),
|
||
"speaker_confidence": float(segment.get("speaker_confidence") or 0),
|
||
"speaker_status": str(segment.get("speaker_status") or "confirmed"),
|
||
}
|
||
if usable and usable[-1]["speaker"]["speaker_id"] == speaker_id:
|
||
usable[-1]["end_time_ms"] = end
|
||
else:
|
||
usable.append(
|
||
{"start_time_ms": start, "end_time_ms": end, "speaker": speaker}
|
||
)
|
||
|
||
if not usable:
|
||
return []
|
||
if len(usable) == 1:
|
||
usable[0]["text"] = text
|
||
return usable
|
||
|
||
# 双重保护:服务端有最短段长滤波,桥接层也拒绝意外的短片段响应。
|
||
min_split_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_MIN_SEGMENT_MS", "3000")))
|
||
while len(usable) > 1:
|
||
short_index = next(
|
||
(
|
||
index
|
||
for index, item in enumerate(usable)
|
||
if item["end_time_ms"] - item["start_time_ms"] < min_split_ms
|
||
),
|
||
None,
|
||
)
|
||
if short_index is None:
|
||
break
|
||
if short_index == 0:
|
||
usable[1]["start_time_ms"] = usable[0]["start_time_ms"]
|
||
del usable[0]
|
||
elif short_index == len(usable) - 1:
|
||
usable[-2]["end_time_ms"] = usable[-1]["end_time_ms"]
|
||
del usable[-1]
|
||
else:
|
||
previous = usable[short_index - 1]
|
||
following = usable[short_index + 1]
|
||
if previous["end_time_ms"] - previous["start_time_ms"] >= (
|
||
following["end_time_ms"] - following["start_time_ms"]
|
||
):
|
||
previous["end_time_ms"] = usable[short_index]["end_time_ms"]
|
||
del usable[short_index]
|
||
else:
|
||
following["start_time_ms"] = usable[short_index]["start_time_ms"]
|
||
del usable[short_index]
|
||
|
||
if len(usable) == 1 or len(text) < len(usable):
|
||
usable = [max(usable, key=lambda item: item["end_time_ms"] - item["start_time_ms"])]
|
||
usable[0]["text"] = text
|
||
return usable
|
||
|
||
total_duration = sum(item["end_time_ms"] - item["start_time_ms"] for item in usable)
|
||
char_start = 0
|
||
elapsed = 0.0
|
||
for index, item in enumerate(usable):
|
||
elapsed += item["end_time_ms"] - item["start_time_ms"]
|
||
if index == len(usable) - 1:
|
||
char_end = len(text)
|
||
else:
|
||
char_end = round(len(text) * elapsed / total_duration)
|
||
char_end = max(char_start + 1, min(char_end, len(text) - (len(usable) - index - 1)))
|
||
item["text"] = text[char_start:char_end].strip()
|
||
char_start = char_end
|
||
return [item for item in usable if item.get("text")]
|
||
|
||
|
||
class BrowserSession:
|
||
"""管理浏览器与原生 WebSocket 连接,并转换双方的消息格式。"""
|
||
|
||
def __init__(
|
||
self,
|
||
browser_ws: web.WebSocketResponse,
|
||
auxiliary: AuxiliaryModelService,
|
||
start: dict[str, Any],
|
||
voice_id: str,
|
||
) -> None:
|
||
self.browser_ws = browser_ws
|
||
self.auxiliary = auxiliary
|
||
self.start = start
|
||
self.voice_id = voice_id
|
||
self.session_id = uuid4().hex
|
||
self.stop_event = asyncio.Event()
|
||
self.send_lock = asyncio.Lock()
|
||
self.speaker_jobs: asyncio.Queue[SpeakerJob | None] = asyncio.Queue()
|
||
self.wav_decoder = (
|
||
IncrementalWavDecoder()
|
||
if str(start.get("source") or "mic") == "file"
|
||
and Path(str(start.get("file_name") or "")).suffix.lower() == ".wav"
|
||
else None
|
||
)
|
||
self.speed_factor = max(0.5, min(3.0, float(start.get("speed_factor") or 1.0)))
|
||
try:
|
||
sentence_strategy = int(start.get("sentence_strategy", 0))
|
||
except (TypeError, ValueError):
|
||
sentence_strategy = 0
|
||
self.sentence_strategy = sentence_strategy if sentence_strategy in (0, 1) else 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:
|
||
"""ASR 和 CAM++ 的完成时机不同,因此需要串行写入浏览器连接。"""
|
||
async with self.send_lock:
|
||
if not self.browser_ws.closed:
|
||
await self.browser_ws.send_json(payload)
|
||
|
||
async def emit_sentence(
|
||
self,
|
||
text: str,
|
||
final: bool,
|
||
speaker: dict[str, Any] | None = None,
|
||
sentence_id: int | None = None,
|
||
start_time_ms: float | None = None,
|
||
end_time_ms: float | None = None,
|
||
) -> None:
|
||
speaker = speaker or {}
|
||
sentence = {
|
||
"sentence_id": self.sentence_id if sentence_id is None else sentence_id,
|
||
"sentence": text,
|
||
"sentence_type": 1 if final else 0,
|
||
"start_time": round(self.turn_start_ms if start_time_ms is None else start_time_ms),
|
||
"end_time": round(self.total_audio_ms if end_time_ms is None else end_time_ms),
|
||
# 未修改的腾讯界面使用 speaker_id 选择对应的说话人气泡。
|
||
"speaker_id": int(speaker.get("speaker_id", -1)),
|
||
"speaker_name": str(speaker.get("speaker_name") or ""),
|
||
"speaker_confidence": float(speaker.get("speaker_confidence") or 0),
|
||
"speaker_status": str(speaker.get("speaker_status") or "pending"),
|
||
}
|
||
await self.emit({"type": "sentences", "sentences": [sentence]})
|
||
|
||
def align_turn_audio_to_vad(self, message: dict[str, Any]) -> None:
|
||
"""用原生 VAD 的绝对时间裁掉当前 turn 的前后静音样本。"""
|
||
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
|
||
try:
|
||
speech_start_ms = message.get("speech_start_ms")
|
||
if speech_start_ms is not None:
|
||
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_start_ms)))
|
||
trim_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
|
||
trim_bytes -= trim_bytes % 2
|
||
del self.turn_audio[:trim_bytes]
|
||
self.turn_start_ms += trim_bytes / PCM_BYTES_PER_MS
|
||
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
|
||
|
||
speech_end_ms = message.get("speech_end_ms")
|
||
if speech_end_ms is not None:
|
||
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_end_ms)))
|
||
keep_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
|
||
keep_bytes -= keep_bytes % 2
|
||
del self.turn_audio[keep_bytes:]
|
||
except (TypeError, ValueError):
|
||
LOGGER.warning("Ignoring invalid native FunASR VAD boundary: %s", message)
|
||
|
||
async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None:
|
||
if not frame:
|
||
return
|
||
if len(frame) % 2:
|
||
raise ValueError("PCM16 audio ended on an incomplete sample")
|
||
self.total_audio_ms += len(frame) / PCM_BYTES_PER_MS
|
||
self.turn_audio.extend(frame)
|
||
if len(self.turn_audio) > MAX_SPEAKER_AUDIO_BYTES:
|
||
# 即使 VAD 始终没有报告端点,也要限制每轮音频的内存占用。
|
||
trim = len(self.turn_audio) - MAX_SPEAKER_AUDIO_BYTES
|
||
del self.turn_audio[:trim]
|
||
self.turn_start_ms += trim / PCM_BYTES_PER_MS
|
||
await native_ws.send(frame)
|
||
if pace_file:
|
||
await asyncio.sleep(len(frame) / (SAMPLE_RATE * 2) / self.speed_factor)
|
||
|
||
async def accept_audio(self, native_ws: Any, data: bytes) -> None:
|
||
pcm = self.wav_decoder.feed(data) if self.wav_decoder else data
|
||
if not pcm:
|
||
return
|
||
self.pending_pcm.extend(pcm)
|
||
is_file = str(self.start.get("source") or "mic") == "file"
|
||
while len(self.pending_pcm) >= FRAME_BYTES:
|
||
frame = bytes(self.pending_pcm[:FRAME_BYTES])
|
||
del self.pending_pcm[:FRAME_BYTES]
|
||
await self.send_pcm_frame(native_ws, frame, is_file)
|
||
|
||
async def finish_audio(self, native_ws: Any) -> None:
|
||
if self.wav_decoder:
|
||
self.wav_decoder.finish()
|
||
if self.pending_pcm:
|
||
frame = bytes(self.pending_pcm)
|
||
self.pending_pcm.clear()
|
||
await self.send_pcm_frame(
|
||
native_ws, frame, str(self.start.get("source") or "mic") == "file"
|
||
)
|
||
|
||
async def read_native(self, native_ws: Any) -> None:
|
||
"""接收 FunASR 原生事件,并维护每段话语的中间结果缓存。"""
|
||
try:
|
||
while True:
|
||
raw = await native_ws.recv()
|
||
message = json.loads(raw)
|
||
if message.get("is_end"):
|
||
self.native_ack = message
|
||
return
|
||
if message.get("event") == "vad":
|
||
self.align_turn_audio_to_vad(message)
|
||
continue
|
||
text = str(message.get("text") or "")
|
||
native_mode = str(message.get("mode") or "")
|
||
# 2-pass 在线文本只是预览,离线 Contextual ASR 的文本才是热词定稿。
|
||
offline_final = native_mode in {"2pass-offline", "offline"}
|
||
if text:
|
||
if offline_final:
|
||
self.turn_text = text
|
||
else:
|
||
# 在线模式逐块发送新增文本;2-pass 在线结果只用于临时预览。
|
||
self.turn_text += text
|
||
if message.get("is_final"):
|
||
final_text = (
|
||
text.strip() if offline_final else self.turn_text.strip()
|
||
)
|
||
audio = bytes(self.turn_audio)
|
||
start_ms = self.turn_start_ms
|
||
end_ms = start_ms + len(audio) / PCM_BYTES_PER_MS
|
||
turn_sentence_id = self.sentence_id
|
||
if final_text:
|
||
# 先结束前端中间结果气泡;标点和 CAM++ 在独立工作线程中完成,
|
||
# 不阻塞原生 WebSocket 继续读取后续音频帧。
|
||
await self.emit_sentence(
|
||
final_text,
|
||
final=True,
|
||
sentence_id=turn_sentence_id,
|
||
start_time_ms=start_ms,
|
||
end_time_ms=end_ms,
|
||
)
|
||
await self.speaker_jobs.put(
|
||
SpeakerJob(
|
||
sentence_id=turn_sentence_id,
|
||
text=final_text,
|
||
audio=audio,
|
||
start_time_ms=start_ms,
|
||
end_time_ms=end_ms,
|
||
)
|
||
)
|
||
# 为同一个 VAD 轮次内可能拆出的多个气泡预留独立 ID。
|
||
self.sentence_id += TURN_SENTENCE_ID_STRIDE
|
||
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:
|
||
"""按序定稿标点并用 CAM++ 滑窗恢复 turn 内说话人切换。"""
|
||
while True:
|
||
job = await self.speaker_jobs.get()
|
||
try:
|
||
if job is None:
|
||
return
|
||
final_text = job.text
|
||
try:
|
||
punctuation = await self.auxiliary.punctuate(final_text)
|
||
if not punctuation.get("available", False):
|
||
LOGGER.warning(
|
||
"FunASR punctuation is unavailable: %s",
|
||
punctuation.get("error", "model is not configured"),
|
||
)
|
||
else:
|
||
punctuated = str(punctuation.get("text") or "").strip()
|
||
if punctuated:
|
||
final_text = punctuated
|
||
except Exception:
|
||
LOGGER.exception("FunASR punctuation request failed; keeping raw text")
|
||
|
||
subsegments: list[dict[str, Any]] = []
|
||
if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
|
||
# 必需的 CAM++ 按 FunASR 的 1.5 秒窗口和 0.75 秒步长,
|
||
# 在同一个 VAD 轮次内识别多个说话人片段。
|
||
track = getattr(self.auxiliary, "track_speakers", None)
|
||
if callable(track):
|
||
try:
|
||
tracked = await track(
|
||
job.audio,
|
||
self.session_id,
|
||
job.start_time_ms,
|
||
job.end_time_ms,
|
||
)
|
||
subsegments = split_text_by_speaker_segments(
|
||
final_text,
|
||
tracked,
|
||
job.start_time_ms,
|
||
job.end_time_ms,
|
||
)
|
||
except Exception:
|
||
LOGGER.exception(
|
||
"CAM++ sliding-window diarization failed; falling back to whole-turn speaker"
|
||
)
|
||
|
||
if not subsegments and len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
|
||
try:
|
||
resolved = await self.auxiliary.resolve_speaker(
|
||
job.audio,
|
||
self.session_id,
|
||
job.start_time_ms,
|
||
job.end_time_ms,
|
||
)
|
||
if resolved and int(resolved.get("speaker_id", -1)) >= 0:
|
||
subsegments = [
|
||
{
|
||
"text": final_text,
|
||
"speaker": resolved,
|
||
"start_time_ms": job.start_time_ms,
|
||
"end_time_ms": job.end_time_ms,
|
||
}
|
||
]
|
||
except Exception:
|
||
LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id)
|
||
|
||
if not subsegments:
|
||
subsegments = [
|
||
{
|
||
"text": final_text,
|
||
"speaker": {"speaker_id": -1},
|
||
"start_time_ms": job.start_time_ms,
|
||
"end_time_ms": job.end_time_ms,
|
||
}
|
||
]
|
||
|
||
for index, segment in enumerate(subsegments):
|
||
await self.emit_sentence(
|
||
str(segment["text"]),
|
||
final=True,
|
||
speaker=segment.get("speaker"),
|
||
sentence_id=job.sentence_id + index,
|
||
start_time_ms=float(segment.get("start_time_ms", job.start_time_ms)),
|
||
end_time_ms=float(segment.get("end_time_ms", job.end_time_ms)),
|
||
)
|
||
finally:
|
||
self.speaker_jobs.task_done()
|
||
|
||
|
||
async def config_handler(_: web.Request) -> web.Response:
|
||
"""提供简洁的就绪状态响应,供启动器和诊断使用。"""
|
||
return web.json_response(
|
||
{
|
||
"engine": "funasr-native-online-ws",
|
||
"model": os.getenv("FUNASR_ASR_MODEL", ""),
|
||
"native_ws_url": NATIVE_WS_URL,
|
||
"speaker_service_url": os.getenv(
|
||
"AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"
|
||
),
|
||
}
|
||
)
|
||
|
||
|
||
async def stop_handler(request: web.Request) -> web.Response:
|
||
voice_id = request.query.get("voice_id", "").strip()
|
||
if not voice_id:
|
||
return web.json_response({"ok": False, "error": "missing voice_id"}, status=400)
|
||
stop_event = request.app[SESSION_REGISTRY_KEY].get(voice_id)
|
||
if stop_event is None:
|
||
return web.json_response({"ok": False, "error": "session not found"})
|
||
stop_event.set()
|
||
return web.json_response({"ok": True})
|
||
|
||
|
||
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
||
"""将腾讯前端消息桥接到 FunASR 原生实时协议。"""
|
||
browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30)
|
||
await browser_ws.prepare(request)
|
||
session: BrowserSession | None = None
|
||
native_reader: asyncio.Task[None] | None = None
|
||
speaker_worker: asyncio.Task[None] | None = None
|
||
voice_id = ""
|
||
registered = False
|
||
|
||
try:
|
||
first = await browser_ws.receive()
|
||
if first.type != WSMsgType.TEXT:
|
||
await browser_ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||
return browser_ws
|
||
start = json.loads(first.data)
|
||
if not isinstance(start, dict) or start.get("type") != "start":
|
||
await browser_ws.send_json({"type": "error", "message": "first message must have type=start"})
|
||
return browser_ws
|
||
source = str(start.get("source") or "mic")
|
||
suffix = Path(str(start.get("file_name") or "")).suffix.lower()
|
||
if source == "file" and suffix not in {".pcm", ".wav"}:
|
||
await browser_ws.send_json(
|
||
{
|
||
"type": "error",
|
||
"message": "FunASR realtime mode accepts PCM or 16 kHz mono PCM WAV files.",
|
||
}
|
||
)
|
||
return browser_ws
|
||
if source == "file" and suffix == ".pcm":
|
||
start = dict(start)
|
||
start["file_name"] = str(start.get("file_name") or "audio.pcm")
|
||
|
||
voice_id = uuid4().hex
|
||
auxiliary: AuxiliaryModelService = request.app[AUXILIARY_KEY]
|
||
session = BrowserSession(browser_ws, auxiliary, start, voice_id)
|
||
request.app[SESSION_REGISTRY_KEY][voice_id] = session.stop_event
|
||
registered = True
|
||
await session.emit({"type": "voice_id", "voice_id": voice_id})
|
||
await session.emit({"type": "start"})
|
||
|
||
# FunASR 的 Contextual Paraformer 在 2-pass 离线定稿阶段应用会话热词。
|
||
async with websocket_connect(
|
||
NATIVE_WS_URL,
|
||
subprotocols=["binary"],
|
||
ping_interval=None,
|
||
close_timeout=3,
|
||
max_size=None,
|
||
) as native_ws:
|
||
native_start = {
|
||
"mode": "online",
|
||
"chunk_size": list(CHUNK_SIZE),
|
||
"chunk_interval": CHUNK_INTERVAL,
|
||
"encoder_chunk_look_back": int(
|
||
os.getenv("FUNASR_ENCODER_LOOK_BACK", "4")
|
||
),
|
||
"decoder_chunk_look_back": int(
|
||
os.getenv("FUNASR_DECODER_LOOK_BACK", "1")
|
||
),
|
||
"sentence_strategy": session.sentence_strategy,
|
||
"audio_fs": SAMPLE_RATE,
|
||
"wav_name": voice_id,
|
||
"is_speaking": True,
|
||
}
|
||
# Contextual 模型仅在 2-pass 最终解码阶段使用热词;普通在线预览仍走流式模型。
|
||
hotwords = normalize_hotwords(start.get("hotwords"))
|
||
if hotwords:
|
||
native_start["mode"] = "2pass"
|
||
native_start["hotwords"] = hotwords
|
||
await native_ws.send(json.dumps(native_start, ensure_ascii=False))
|
||
native_reader = asyncio.create_task(session.read_native(native_ws))
|
||
speaker_worker = asyncio.create_task(session.resolve_speakers())
|
||
|
||
while not browser_ws.closed:
|
||
receive_task = asyncio.create_task(browser_ws.receive())
|
||
stop_task = asyncio.create_task(session.stop_event.wait())
|
||
done, _ = await asyncio.wait(
|
||
[receive_task, stop_task, native_reader],
|
||
return_when=asyncio.FIRST_COMPLETED,
|
||
)
|
||
if stop_task in done:
|
||
receive_task.cancel()
|
||
await asyncio.gather(receive_task, return_exceptions=True)
|
||
break
|
||
stop_task.cancel()
|
||
await asyncio.gather(stop_task, return_exceptions=True)
|
||
if native_reader in done:
|
||
receive_task.cancel()
|
||
await asyncio.gather(receive_task, return_exceptions=True)
|
||
if session.native_ack is None and session.native_error is None:
|
||
session.native_error = "FunASR native WebSocket ended before EOF acknowledgement"
|
||
break
|
||
|
||
message = await receive_task
|
||
if message.type == WSMsgType.BINARY:
|
||
await session.accept_audio(native_ws, bytes(message.data))
|
||
continue
|
||
if message.type == WSMsgType.TEXT:
|
||
try:
|
||
control = json.loads(message.data)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
if isinstance(control, dict) and control.get("type") in {"eof", "stop"}:
|
||
break
|
||
if isinstance(control, dict) and control.get("type") == "abort":
|
||
session.native_error = "session aborted by browser"
|
||
break
|
||
if message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}:
|
||
break
|
||
|
||
if session.native_error is None and not browser_ws.closed:
|
||
await session.finish_audio(native_ws)
|
||
# FunASR 会刷新在线缓存,并在输出最终结果后才返回确认。
|
||
await native_ws.send(
|
||
json.dumps({"is_speaking": False, "is_end": True}, ensure_ascii=False)
|
||
)
|
||
try:
|
||
await asyncio.wait_for(native_reader, timeout=FINALIZE_TIMEOUT_SECONDS)
|
||
except asyncio.TimeoutError:
|
||
session.native_error = (
|
||
f"FunASR did not acknowledge end-of-input within "
|
||
f"{FINALIZE_TIMEOUT_SECONDS}s"
|
||
)
|
||
|
||
if speaker_worker is not None:
|
||
await session.speaker_jobs.join()
|
||
await session.speaker_jobs.put(None)
|
||
await speaker_worker
|
||
speaker_worker = None
|
||
if session.native_error:
|
||
await session.emit({"type": "error", "message": session.native_error})
|
||
elif session.native_ack and not session.native_ack.get("is_final", False):
|
||
await session.emit(
|
||
{
|
||
"type": "error",
|
||
"message": str(session.native_ack.get("error") or "FunASR did not finalize the stream"),
|
||
}
|
||
)
|
||
elif session.native_ack:
|
||
await session.emit({"type": "end"})
|
||
except asyncio.CancelledError:
|
||
raise
|
||
except Exception as exc:
|
||
LOGGER.exception("Tencent-compatible WebSocket session failed")
|
||
if not browser_ws.closed:
|
||
await browser_ws.send_json({"type": "error", "message": str(exc)})
|
||
finally:
|
||
if native_reader is not None and not native_reader.done():
|
||
native_reader.cancel()
|
||
await asyncio.gather(native_reader, return_exceptions=True)
|
||
if speaker_worker is not None and not speaker_worker.done():
|
||
speaker_worker.cancel()
|
||
await asyncio.gather(speaker_worker, return_exceptions=True)
|
||
if registered:
|
||
request.app[SESSION_REGISTRY_KEY].pop(voice_id, None)
|
||
if session is not None:
|
||
await session.auxiliary.reset_speaker_session(session.session_id)
|
||
if not browser_ws.closed:
|
||
await browser_ws.close()
|
||
return browser_ws
|
||
|
||
|
||
async def create_app() -> web.Application:
|
||
"""创建轻量协议桥接层;模型推理由 FunASR 原生服务负责。"""
|
||
app = web.Application()
|
||
app[AUXILIARY_KEY] = AuxiliaryModelService(
|
||
AuxiliaryServiceConfig(
|
||
base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010")
|
||
)
|
||
)
|
||
app[SESSION_REGISTRY_KEY] = {}
|
||
|
||
async def lifecycle(application: web.Application):
|
||
await application[AUXILIARY_KEY].start()
|
||
try:
|
||
health = await application[AUXILIARY_KEY].health()
|
||
if not health.get("speaker_embedding_ready"):
|
||
raise RuntimeError("CAM++ speaker service is not ready")
|
||
except Exception:
|
||
await application[AUXILIARY_KEY].close()
|
||
raise
|
||
yield
|
||
await application[AUXILIARY_KEY].close()
|
||
|
||
app.cleanup_ctx.append(lifecycle)
|
||
app.router.add_get("/api/config", config_handler)
|
||
app.router.add_get("/api/stop", stop_handler)
|
||
app.router.add_get("/ws", websocket_handler)
|
||
return app
|
||
|
||
|
||
def main() -> None:
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument("--no-browser", action="store_true")
|
||
args = parser.parse_args()
|
||
logging.basicConfig(level=logging.INFO)
|
||
print(
|
||
f"FunASR browser bridge: http://{os.getenv('WEB_DISPLAY_HOST', '127.0.0.1')}:{WEB_PORT}/api/config",
|
||
flush=True,
|
||
)
|
||
web.run_app(create_app(), host=WEB_HOST, port=WEB_PORT)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|