Files
kefu/wechat_rpa/test_voice_transcription.py
T
2026-09-21 10:34:06 +08:00

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)