ASR-demo/realtime_websocket/tests/test_funasr_engine.py

76 lines
2.6 KiB
Python

"""Unit tests for the migrated FunASR streaming lifecycle."""
from __future__ import annotations
import asyncio
import unittest
from unittest.mock import patch
try:
from realtime_websocket.funasr_engine import (
FunASRRealtimeSession,
FunASRServiceConfig,
_result_text,
_vad_events,
)
except ModuleNotFoundError:
from funasr_engine import (
FunASRRealtimeSession,
FunASRServiceConfig,
_result_text,
_vad_events,
)
class FakeService:
def __init__(self) -> None:
self.config = FunASRServiceConfig(
device="cpu",
vad_device="cpu",
vad_chunk_ms=1000,
chunk_size=(0, 2, 1),
)
self.vad_calls = 0
self.asr_calls: list[dict[str, object]] = []
async def generate_vad(self, audio, status, chunk_ms):
self.vad_calls += 1
return [(0, -1)] if self.vad_calls == 1 else [(-1, 2000)]
async def generate_asr(self, audio, status):
self.asr_calls.append(status)
return "第一块" if len(self.asr_calls) == 1 else "第二块"
class FunASREngineTests(unittest.TestCase):
def test_result_normalization(self) -> None:
self.assertEqual(_result_text([{"text": "你好"}]), "你好")
self.assertEqual(_result_text({"value": "片段"}), "片段")
self.assertEqual(
_vad_events([{"value": [[0, -1], [-1, 800]]}]),
[(0.0, -1.0), (-1.0, 800.0)],
)
def test_stream_has_independent_cache_and_final_flush(self) -> None:
async def exercise() -> None:
service = FakeService()
with patch.object(FunASRRealtimeSession, "_to_float32", staticmethod(lambda value: value)):
first = FunASRRealtimeSession(service)
second = FunASRRealtimeSession(service)
first_events = await first.feed(b"\\x01\\x00" * 16000)
second_events = await second.feed(b"\\x01\\x00" * 16000)
first_events += await first.feed(b"\\x01\\x00" * 16000)
first_events += await first.finish()
self.assertTrue(any(not event.is_final for event in first_events))
self.assertTrue(any(event.is_final for event in first_events))
self.assertEqual(first.segment_id, 1)
self.assertEqual(second.segment_id, 0)
self.assertTrue(any(status["is_final"] is True for status in service.asr_calls))
self.assertGreaterEqual(len(service.asr_calls), 2)
asyncio.run(exercise())
if __name__ == "__main__":
unittest.main()