ASR-demo/backend/realtime_websocket/funasr_server.py

741 lines
31 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

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

"""将未修改的腾讯演示协议转换为 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()