157 lines
5.1 KiB
Python
157 lines
5.1 KiB
Python
#!/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)
|