69 lines
2.0 KiB
Python
69 lines
2.0 KiB
Python
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from assistant_bot.speech import SpeechRecognitionError, SpeechRecognizer
|
|
|
|
|
|
class FakeSegment:
|
|
def __init__(self, text: str) -> None:
|
|
self.text = text
|
|
|
|
|
|
class SpeechRecognizerTests(unittest.TestCase):
|
|
def test_transcribe_joins_non_empty_segments(self) -> None:
|
|
factory_calls = []
|
|
transcribe_calls = []
|
|
|
|
class FakeModel:
|
|
def transcribe(self, audio_path, **kwargs):
|
|
transcribe_calls.append((audio_path, kwargs))
|
|
return iter(
|
|
[FakeSegment(" Привет "), FakeSegment(""), FakeSegment("мир")]
|
|
), None
|
|
|
|
def model_factory(*args, **kwargs):
|
|
factory_calls.append((args, kwargs))
|
|
return FakeModel()
|
|
|
|
recognizer = SpeechRecognizer(
|
|
model_name="small",
|
|
device="cpu",
|
|
compute_type="int8",
|
|
language="ru",
|
|
model_factory=model_factory,
|
|
)
|
|
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
audio_path = Path(directory) / "voice.ogg"
|
|
result = recognizer.transcribe(audio_path)
|
|
|
|
self.assertEqual(result, "Привет мир")
|
|
self.assertEqual(
|
|
factory_calls[0],
|
|
(("small",), {"device": "cpu", "compute_type": "int8"}),
|
|
)
|
|
self.assertEqual(
|
|
transcribe_calls[0][1],
|
|
{"language": "ru", "beam_size": 5, "vad_filter": True},
|
|
)
|
|
|
|
def test_transcribe_wraps_model_errors(self) -> None:
|
|
def failing_factory(*_args, **_kwargs):
|
|
raise RuntimeError("model is unavailable")
|
|
|
|
recognizer = SpeechRecognizer(
|
|
model_name="small",
|
|
device="cpu",
|
|
compute_type="int8",
|
|
language=None,
|
|
model_factory=failing_factory,
|
|
)
|
|
|
|
with self.assertRaises(SpeechRecognitionError):
|
|
recognizer.transcribe("voice.ogg")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|