128 lines
4.1 KiB
Python
128 lines
4.1 KiB
Python
"""验证未修改的腾讯界面与 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()
|