ASR-demo/realtime_websocket/tests/test_funasr_server.py

128 lines
4.1 KiB
Python

"""Contract test for the unchanged Tencent UI to FunASR native WS bridge."""
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 realtime_websocket.funasr_server import (
AUXILIARY_KEY,
SESSION_REGISTRY_KEY,
websocket_handler,
)
class FakeNativeWebSocket:
"""Stand in for FunASR's native WSS process without loading model weights."""
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(
"realtime_websocket.funasr_server.websocket_connect",
return_value=native,
):
ws = await self.client.ws_connect("/ws")
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")
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()