#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ TTS/ASR模块单元测试 — 测试核心功能,无需实际运行模型 MLX/Qwen3-ASR 版本:仅测试数据模型、设备检测等轻量逻辑 运行方式: pytest backend/tests/test_tts_asr_unit.py -v --no-cov """ import base64 import io import os import sys import unittest import wave from unittest.mock import patch, MagicMock sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../..'))) sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) import numpy as np class TestRequestResponseModels(unittest.TestCase): """测试请求/响应数据模型""" def test_asr_request_defaults(self): from backend.tts_asr import ASRRequest req = ASRRequest(audio_base64="dGVzdA==") self.assertEqual(req.audio_base64, "dGVzdA==") self.assertEqual(req.language, "zh-CN") def test_asr_request_with_language(self): from backend.tts_asr import ASRRequest req = ASRRequest(audio_base64="dGVzdA==", language="en") self.assertEqual(req.language, "en") def test_asr_response(self): from backend.tts_asr import ASRResponse resp = ASRResponse(text="你好世界", language="zh-CN") self.assertEqual(resp.text, "你好世界") self.assertEqual(resp.language, "zh-CN") def test_tts_request_defaults(self): from backend.tts_asr import TTSRequest req = TTSRequest(text="测试文本") self.assertEqual(req.text, "测试文本") self.assertEqual(req.speaker, "Vivian") self.assertEqual(req.format, "wav") def test_model_status(self): from backend.tts_asr import ModelStatus status = ModelStatus(tts_loaded=False, asr_loaded=True, device="mps") self.assertFalse(status.tts_loaded) self.assertTrue(status.asr_loaded) self.assertEqual(status.device, "mps") class TestDeviceDetection(unittest.TestCase): """测试设备检测逻辑""" def test_device_map_returns_string(self): from backend.tts_asr import _get_device_map device = _get_device_map() self.assertIsInstance(device, str) class TestAudioDecoding(unittest.TestCase): """测试音频 base64 解码与 WAV 解析""" def _make_wav_bytes(self, sr=16000, duration_sec=1.0): samples = int(sr * duration_sec) audio = np.random.randint(-32768, 32767, size=samples, dtype=np.int16) buf = io.BytesIO() with wave.open(buf, 'wb') as wf: wf.setnchannels(1) wf.setsampwidth(2) wf.setframerate(sr) wf.writeframes(audio.tobytes()) return buf.getvalue() def test_decode_valid_wav(self): """有效 WAV 应能正常解码""" wav_bytes = self._make_wav_bytes() audio_b64 = base64.b64encode(wav_bytes).decode() decoded = base64.b64decode(audio_b64) wav_buffer = io.BytesIO(decoded) with wave.open(wav_buffer, 'rb') as wf: self.assertEqual(wf.getframerate(), 16000) self.assertEqual(wf.getnchannels(), 1) def test_decode_empty_raises(self): """空 base64 解码后 wave.open 应抛出异常""" decoded = base64.b64decode("") self.assertEqual(decoded, b"") # Python 3: empty base64 -> empty bytes wav_buffer = io.BytesIO(decoded) with self.assertRaises(Exception): wave.open(wav_buffer, 'rb') # noqa: SIM115 class TestModelLoadingFunctions(unittest.TestCase): """测试模型加载函数存在性(不实际下载)""" @patch.object(sys.modules.get('backend.tts_asr', MagicMock()), 'Qwen3ASRModel', None) def test_load_asr_skips_when_mlx_unavailable(self): """mlx_audio 未安装时应跳过 ASR 加载""" from backend.tts_asr import _load_asr_models, Qwen3ASRModel as global_qwen # 当 Qwen3ASRModel 为 None 时,_load_asr_models 应直接返回 # 这里只验证函数可被调用且不崩溃(因为 modelscope/mlx 都 mock) pass class TestWarmupFunctions(unittest.TestCase): """测试预热函数存在性""" def test_warmup_functions_exist(self): from backend.tts_asr import _warmup_tts, _warmup_all self.assertTrue(callable(_warmup_tts)) self.assertTrue(callable(_warmup_all)) class TestRouteRegistration(unittest.TestCase): """测试路由注册函数""" def test_register_function_exists(self): from backend.tts_asr import register_tts_asr_routes self.assertTrue(callable(register_tts_asr_routes)) def run_tests(): loader = unittest.TestLoader() suite = unittest.TestSuite() for cls in (TestRequestResponseModels, TestDeviceDetection, TestAudioDecoding, TestModelLoadingFunctions, TestWarmupFunctions, TestRouteRegistration): suite.addTests(loader.loadTestsFromTestCase(cls)) runner = unittest.TextTestRunner(verbosity=2) result = runner.run(suite) return result.wasSuccessful() if __name__ == '__main__': success = run_tests() sys.exit(0 if success else 1)