293 lines
14 KiB
Python
293 lines
14 KiB
Python
"""Synthetic-only offline ASR regressions; no account data, downloads or sends."""
|
|
import io
|
|
import json
|
|
import math
|
|
import os
|
|
from pathlib import Path
|
|
import socket
|
|
import struct
|
|
import tempfile
|
|
import threading
|
|
import time
|
|
import unittest
|
|
from unittest import mock
|
|
import wave
|
|
|
|
import voice_transcription as voice
|
|
|
|
|
|
def wav_bytes(seconds=0.3, frequency=440, rate=16000, silent=False):
|
|
pcm = b"".join(struct.pack("<h", 0 if silent else int(6000 * math.sin(i * frequency * 2 * math.pi / rate)))
|
|
for i in range(int(seconds * rate)))
|
|
data = io.BytesIO()
|
|
with wave.open(data, "wb") as stream:
|
|
stream.setnchannels(1)
|
|
stream.setsampwidth(2)
|
|
stream.setframerate(rate)
|
|
stream.writeframes(pcm)
|
|
return data.getvalue()
|
|
|
|
|
|
class VoiceTranscriptionTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temporary = tempfile.TemporaryDirectory(prefix="synthetic-voice-test-")
|
|
self.root = Path(self.temporary.name).resolve()
|
|
self.path = self.root / "synthetic.wav"
|
|
self.path.write_bytes(wav_bytes())
|
|
self.services = []
|
|
self.external = mock.patch.object(socket.socket, "connect", side_effect=AssertionError("Network forbidden"))
|
|
self.external.start()
|
|
|
|
def tearDown(self):
|
|
for service in self.services:
|
|
service.close()
|
|
if service._worker_thread:
|
|
service._worker_thread.join(timeout=2)
|
|
self.external.stop()
|
|
self.temporary.cleanup()
|
|
|
|
def service(self, **kwargs):
|
|
service = voice.VoiceTranscriptionService(allowed_roots=[self.root], **kwargs)
|
|
self.services.append(service)
|
|
return service
|
|
|
|
def wait(self, service, account="111", conv="M:222", sid="100", path=None):
|
|
deadline = time.monotonic() + 3
|
|
while time.monotonic() < deadline:
|
|
result = service.request(account, conv, sid, path or self.path)
|
|
if result["status"] != "pending":
|
|
return result
|
|
time.sleep(0.005)
|
|
self.fail("Worker did not finish")
|
|
|
|
def test_wav_transcribes_and_preserves_digest(self):
|
|
result = voice.transcribe_file(self.path, allowed_roots=[self.root], recognizer=lambda audio: "合成转写")
|
|
self.assertEqual(result["text"], "合成转写")
|
|
self.assertEqual(len(result["audio_sha256"]), 64)
|
|
self.assertEqual(result["duration_seconds"], 0.3)
|
|
self.assertEqual(result["source"], "local_asr")
|
|
|
|
def test_missing_allowed_root_rejected(self):
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
voice.transcribe_file(self.path, recognizer=mock.Mock())
|
|
self.assertEqual(error.exception.code, "path_not_allowed")
|
|
|
|
def test_outside_root_rejected(self):
|
|
service = voice.VoiceTranscriptionService(allowed_roots=[self.root / "allowed"])
|
|
self.services.append(service)
|
|
self.assertEqual(service.request("111", "M:222", "1", self.path)["error_code"], "path_not_allowed")
|
|
|
|
def test_missing_file_has_readable_error_without_path(self):
|
|
result = self.service().request("111", "M:222", "1", self.root / "private-name.wav")
|
|
self.assertEqual(result["error_code"], "audio_missing")
|
|
self.assertNotIn("private-name", json.dumps(result))
|
|
|
|
def test_group_wrong_account_self_invalid_server_denied(self):
|
|
service = self.service()
|
|
for account, conv, sid in (("111", "R:2", "1"), ("111", "S:222_333", "1"),
|
|
("111", "S:111_111", "1"), ("111", "M:111", "1"),
|
|
("111", "M:2", "0"), ("", "M:2", "1")):
|
|
with self.subTest(identity=(account, conv, sid)):
|
|
self.assertEqual(service.request(account, conv, sid, self.path)["error_code"], "invalid_identity")
|
|
|
|
def test_request_is_fast_while_worker_waits(self):
|
|
entered, release = threading.Event(), threading.Event()
|
|
def recognize(audio):
|
|
entered.set()
|
|
release.wait(2)
|
|
return "合成文字"
|
|
service = self.service(recognizer=recognize)
|
|
started = time.monotonic()
|
|
self.assertEqual(service.request("111", "S:111_222", "1", self.path)["status"], "pending")
|
|
self.assertLess(time.monotonic() - started, 0.1)
|
|
self.assertTrue(entered.wait(2))
|
|
self.assertEqual(service.request("111", "S:111_222", "1", self.path)["status"], "pending")
|
|
release.set()
|
|
self.assertEqual(self.wait(service, conv="S:111_222", sid="1")["status"], "ready")
|
|
|
|
def test_accounts_conversations_and_message_ids_do_not_share_results(self):
|
|
recognize = mock.Mock(side_effect=["第一段", "第二段", "第三段", "第四段"])
|
|
service = self.service(recognizer=recognize)
|
|
values = [("111", "M:222", "1"), ("333", "M:222", "1"),
|
|
("111", "M:444", "1"), ("111", "M:222", "2")]
|
|
texts = [self.wait(service, account=a, conv=c, sid=s)["text"] for a, c, s in values]
|
|
self.assertEqual(texts, ["第一段", "第二段", "第三段", "第四段"])
|
|
self.assertEqual(recognize.call_count, 4)
|
|
|
|
def test_same_message_same_audio_is_recognized_once(self):
|
|
recognize = mock.Mock(return_value="合成文字")
|
|
service = self.service(recognizer=recognize)
|
|
self.assertEqual(self.wait(service)["status"], "ready")
|
|
self.assertEqual(self.wait(service)["status"], "ready")
|
|
self.assertEqual(recognize.call_count, 1)
|
|
|
|
def test_same_message_changed_audio_is_not_cached(self):
|
|
recognize = mock.Mock(side_effect=["旧语音", "新语音"])
|
|
service = self.service(recognizer=recognize)
|
|
first = self.wait(service)
|
|
self.path.write_bytes(wav_bytes(frequency=880))
|
|
second = self.wait(service)
|
|
self.assertEqual(second["text"], "新语音")
|
|
self.assertNotEqual(first["audio_sha256"], second["audio_sha256"])
|
|
|
|
def test_change_during_recognition_drops_old_result(self):
|
|
entered, release = threading.Event(), threading.Event()
|
|
def recognize(audio):
|
|
entered.set()
|
|
release.wait(2)
|
|
return "旧语音内容"
|
|
service = self.service(recognizer=recognize)
|
|
service.request("111", "M:222", "100", self.path)
|
|
self.assertTrue(entered.wait(2))
|
|
self.path.write_bytes(wav_bytes(frequency=880))
|
|
release.set()
|
|
deadline = time.monotonic() + 2
|
|
while service._jobs.unfinished_tasks and time.monotonic() < deadline:
|
|
time.sleep(0.005)
|
|
result = service._states[("111", "M:222", "100")]["result"]
|
|
self.assertEqual(result["error_code"], "audio_changed")
|
|
self.assertEqual(result["text"], "")
|
|
|
|
def test_failures_back_off_and_do_not_leak_exception_details(self):
|
|
recognize = mock.Mock(side_effect=RuntimeError("private transcript/path"))
|
|
service = self.service(recognizer=recognize, retry_seconds=30)
|
|
result = self.wait(service)
|
|
self.assertEqual(result["error_code"], "recognition_failed")
|
|
self.assertGreater(result["retry_at"], time.time())
|
|
self.assertNotIn("private", json.dumps(result))
|
|
self.assertEqual(self.wait(service), result)
|
|
self.assertEqual(recognize.call_count, 1)
|
|
|
|
def test_empty_recognition_never_becomes_ready(self):
|
|
result = self.wait(self.service(recognizer=lambda audio: " "))
|
|
self.assertEqual(result["error_code"], "audio_empty")
|
|
self.assertEqual(result["text"], "")
|
|
|
|
def test_queue_is_bounded(self):
|
|
entered, release = threading.Event(), threading.Event()
|
|
def recognize(audio):
|
|
entered.set()
|
|
release.wait(2)
|
|
return "合成"
|
|
service = self.service(recognizer=recognize, max_pending=1)
|
|
service.request("111", "M:222", "1", self.path)
|
|
self.assertTrue(entered.wait(2))
|
|
self.assertEqual(service.request("111", "M:222", "2", self.path)["error_code"], "queue_full")
|
|
release.set()
|
|
|
|
def test_close_discards_text_and_stops_new_jobs(self):
|
|
service = self.service(recognizer=lambda audio: "合成转写")
|
|
self.assertEqual(self.wait(service)["status"], "ready")
|
|
service.close()
|
|
self.assertEqual(service.request("111", "M:222", "100", self.path)["error_code"], "stopped")
|
|
self.assertFalse(service._states)
|
|
self.assertFalse(service._results)
|
|
|
|
def test_bad_header_is_not_sent_to_model(self):
|
|
self.path.write_bytes(b"not audio")
|
|
recognize = mock.Mock()
|
|
result = self.wait(self.service(recognizer=recognize))
|
|
self.assertEqual(result["error_code"], "audio_format")
|
|
recognize.assert_not_called()
|
|
|
|
def test_silence_is_not_sent_to_model(self):
|
|
self.path.write_bytes(wav_bytes(silent=True))
|
|
recognize = mock.Mock()
|
|
result = self.wait(self.service(recognizer=recognize))
|
|
self.assertEqual(result["error_code"], "audio_empty")
|
|
recognize.assert_not_called()
|
|
|
|
def test_size_limit_prevents_worker(self):
|
|
with mock.patch.object(voice, "MAX_AUDIO_BYTES", 100):
|
|
result = self.service().request("111", "M:222", "100", self.path)
|
|
self.assertEqual(result["error_code"], "audio_too_large")
|
|
|
|
def test_duration_limit_rejects_instead_of_truncating(self):
|
|
with mock.patch.object(voice, "MAX_DURATION_SECONDS", 0.1):
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
voice.transcribe_file(self.path, allowed_roots=[self.root], recognizer=mock.Mock())
|
|
self.assertEqual(error.exception.code, "audio_too_long")
|
|
|
|
def test_silk_real_decode_of_synthetic_pcm(self):
|
|
import pysilk
|
|
encoded = io.BytesIO()
|
|
raw = wav_bytes()
|
|
with wave.open(io.BytesIO(raw), "rb") as stream:
|
|
pcm = stream.readframes(stream.getnframes())
|
|
pysilk.encode(io.BytesIO(pcm), encoded, 16000, 16000, tencent=True)
|
|
self.path.write_bytes(encoded.getvalue())
|
|
result = voice.transcribe_file(self.path, allowed_roots=[self.root], recognizer=lambda audio: "合成SILK")
|
|
self.assertEqual(result["text"], "合成SILK")
|
|
|
|
def test_missing_model_fails_without_network(self):
|
|
recognize = voice.OfflineRecognizer(self.root / "absent-model")
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
voice.transcribe_file(self.path, allowed_roots=[self.root], recognizer=recognize)
|
|
self.assertEqual(error.exception.code, "model_missing")
|
|
|
|
def test_invalid_model_manifest_fails_without_network(self):
|
|
directory = self.root / "model"
|
|
directory.mkdir()
|
|
(directory / "manifest.json").write_text('{"revision":"wrong"}', encoding="utf-8")
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
voice._verified_model_directory(directory)
|
|
self.assertEqual(error.exception.code, "model_invalid")
|
|
|
|
def test_close_during_recognition_never_repopulates_cache(self):
|
|
entered, release = threading.Event(), threading.Event()
|
|
def recognize(audio):
|
|
entered.set()
|
|
release.wait(2)
|
|
return "不应保留的合成文本"
|
|
service = self.service(recognizer=recognize)
|
|
service.request("111", "M:222", "100", self.path)
|
|
self.assertTrue(entered.wait(2))
|
|
service.close()
|
|
release.set()
|
|
service._worker_thread.join(timeout=2)
|
|
self.assertFalse(service._worker_thread.is_alive())
|
|
self.assertFalse(service._states)
|
|
self.assertFalse(service._results)
|
|
|
|
def test_failed_job_retries_only_after_backoff(self):
|
|
recognize = mock.Mock(side_effect=[RuntimeError("synthetic failure"), "恢复转写"])
|
|
service = self.service(recognizer=recognize, retry_seconds=30)
|
|
failure = self.wait(service)
|
|
self.assertEqual(failure["status"], "error")
|
|
with mock.patch.object(voice.time, "time", return_value=failure["retry_at"] + 1):
|
|
ready = self.wait(service)
|
|
self.assertEqual(ready["text"], "恢复转写")
|
|
self.assertEqual(recognize.call_count, 2)
|
|
|
|
def test_silk_pcm_output_limit_is_enforced(self):
|
|
output = voice._BoundedPCM()
|
|
with mock.patch.object(voice, "MAX_DURATION_SECONDS", 0.001):
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
output.write(b"x" * 100)
|
|
self.assertEqual(error.exception.code, "audio_too_long")
|
|
|
|
def test_low_confidence_segment_is_rejected(self):
|
|
from types import SimpleNamespace
|
|
recognizer = voice.OfflineRecognizer(self.root)
|
|
recognizer._model = mock.Mock()
|
|
recognizer._model.transcribe.return_value = ([SimpleNamespace(
|
|
text="可能猜测的文字", avg_logprob=-1.2, no_speech_prob=0.1, compression_ratio=1.0)], None)
|
|
with self.assertRaises(voice.VoiceTranscriptionError) as error:
|
|
voice.transcribe_file(self.path, allowed_roots=[self.root], recognizer=recognizer)
|
|
self.assertEqual(error.exception.code, "low_confidence")
|
|
|
|
def test_model_loading_is_forced_offline(self):
|
|
import numpy as np
|
|
constructor = mock.Mock()
|
|
constructor.return_value.transcribe.return_value = ([], None)
|
|
with mock.patch.object(voice, "_verified_model_directory", return_value=self.root), mock.patch.dict(
|
|
"sys.modules", {"faster_whisper": type("FakeWhisper", (), {"WhisperModel": constructor})}):
|
|
voice.OfflineRecognizer(self.root)(np.ones(1600, dtype=np.float32))
|
|
self.assertTrue(constructor.call_args.kwargs["local_files_only"])
|
|
self.assertFalse(constructor.call_args.kwargs["use_auth_token"])
|
|
self.assertEqual(constructor.call_args.kwargs["device"], "cpu")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main(verbosity=2)
|