ASR-demo/realtime_websocket/tests/test_funasr_server.py

113 lines
3.2 KiB
Python

"""WebSocket contract tests for the FunASR browser adapter."""
from __future__ import annotations
import asyncio
from types import SimpleNamespace
import unittest
from aiohttp import web
from aiohttp.test_utils import AioHTTPTestCase
from realtime_websocket.funasr_engine import FunASRSegment
from realtime_websocket.funasr_server import (
AUXILIARY_SERVICE_KEY,
MODEL_SERVICE_KEY,
config_handler,
websocket_handler,
)
class FakeFunASRSession:
def __init__(self) -> None:
self.sent = False
async def feed(self, audio: bytes):
if self.sent:
return []
self.sent = True
return [
FunASRSegment(
text="实时片段",
start_time_ms=0,
end_time_ms=len(audio) / 32,
audio=b"",
voiced_ms=len(audio) / 32,
is_final=False,
sentence_id=0,
)
]
async def finish(self):
return [
FunASRSegment(
text="最终片段",
start_time_ms=0,
end_time_ms=1000,
audio=b"\x01\x00" * 8000,
voiced_ms=1000,
is_final=True,
sentence_id=0,
reason="eof",
)
]
class FakeFunASRService:
config = SimpleNamespace(model="fake-funasr")
def create_session(self):
return FakeFunASRSession()
class FakeAuxiliaryService:
config = SimpleNamespace(base_url="http://fake-speaker")
async def health(self):
return {"ready": True, "speaker_embedding_ready": True}
async def reset_speaker_session(self, session_id):
return None
class FunASRWebSocketTests(AioHTTPTestCase):
def get_app(self):
app = web.Application()
app[MODEL_SERVICE_KEY] = FakeFunASRService()
app[AUXILIARY_SERVICE_KEY] = FakeAuxiliaryService()
app.router.add_get("/api/config", config_handler)
app.router.add_get("/ws", websocket_handler)
return app
async def test_frontend_contract_uses_funasr_streaming_mode(self):
ws = await self.client.ws_connect("/ws")
await ws.send_json({"type": "start", "speaker_diarization": 0})
start = await ws.receive_json()
self.assertEqual(start["type"], "start")
self.assertEqual(start["engine"], "funasr")
self.assertEqual(start["partial_mode"], "funasr_streaming_cache")
await ws.send_bytes(b"\x01\x00" * 16000)
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
sentence_events = [
item for item in messages
if item["type"] == "sentences" and item["sentences"]
]
self.assertTrue(any(item["sentences"][0]["sentence_type"] == 0 for item in sentence_events))
self.assertEqual(sentence_events[-1]["sentences"][0]["sentence"], "最终片段")
self.assertEqual(messages[-1]["type"], "end")
await ws.close()
if __name__ == "__main__":
unittest.main()