113 lines
3.2 KiB
Python
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()
|