ASR-demo/backend/realtime_websocket/tests/test_funasr_server.py

128 lines
4.1 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 asyncio
import json
import unittest
from unittest.mock import patch
from aiohttp import web
from aiohttp.test_utils import AioHTTPTestCase
from backend.realtime_websocket.funasr_server import (
AUXILIARY_KEY,
SESSION_REGISTRY_KEY,
websocket_handler,
)
class FakeNativeWebSocket:
"""用模拟进程代替 FunASR 原生 WSS 服务,不加载模型权重。"""
def __init__(self) -> None:
self.incoming: asyncio.Queue[str] = asyncio.Queue()
self.audio_bytes = 0
self.config: dict[str, object] = {}
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})
)
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})
)
async def recv(self) -> str:
return await self.incoming.get()
class FakeAuxiliaryService:
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 reset_speaker_session(self, _session_id):
return None
class FunASRBridgeTests(AioHTTPTestCase):
def get_app(self):
app = web.Application()
app[AUXILIARY_KEY] = FakeAuxiliaryService()
app[SESSION_REGISTRY_KEY] = {}
app.router.add_get("/ws", websocket_handler)
return app
async def test_tencent_ui_messages_use_native_funasr_and_keep_speaker_label(self):
native = FakeNativeWebSocket()
with patch(
"backend.realtime_websocket.funasr_server.websocket_connect",
return_value=native,
):
ws = await self.client.ws_connect("/ws")
await ws.send_json(
{
"type": "start",
"source": "mic",
# 即使该标志为 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")
pcm = b"\x01\x00" * 16000
await ws.send_bytes(pcm)
await ws.send_json({"type": "eof"})
messages = []
async with asyncio.timeout(5):
while True:
message = await ws.receive_json()
messages.append(message)
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 = [
sentence
for message in messages
if message["type"] == "sentences"
for sentence in message["sentences"]
]
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()
if __name__ == "__main__":
unittest.main()