76 lines
2.6 KiB
Python
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()
|