"""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()