128 lines
4.1 KiB
Python
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 backend.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(
|
|
"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",
|
|
# 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()
|