Files
wecom_it_smart_desk/backend/tests/test_voice_asr.py
T

473 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# =============================================================================
# 百度 ASR cuid 修复 — 防御性检查逻辑验证测试
# =============================================================================
# 验证内容:
# 1. get_baidu_token() — api_key/secret_key 空值检查
# 2. transcribe_audio() — app_id 空值检查 + 配置诊断日志
# 3. 正常流程 — cuid 参数正确映射为 app_id
# 4. token 获取异常传播 — HTTPException 正确向上传播
#
# 关联修复:百度 ASR "url param cuid error" — cuid 为空因为
# BAIDU_ASR_APP_ID 环境变量未传入容器
# =============================================================================
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import HTTPException
from app.api.voice_asr import (
BAIDU_ASR_API_URL,
BAIDU_ASR_DEV_PID,
get_baidu_token,
transcribe_audio,
)
from app.config import settings
# =============================================================================
# 辅助类与函数
# =============================================================================
class _MockBaiduHttpResponse:
"""模拟 httpx.Response — 提供同步 .json() 方法。"""
def __init__(self, json_data: dict):
self._json_data = json_data
def json(self) -> dict:
return self._json_data
class _MockAsyncHTTPClient:
"""模拟 httpx.AsyncClient — 正确实现 async context manager 协议。
用真实类替代 AsyncMock,避免 AsyncMock 子属性不可 await 的问题。
"""
# 类级配置,测试前设置
_post_return_value = None
_post_side_effect = None
# 实例级记录
def __init__(self, *args, **kwargs):
self.post_calls = []
async def __aenter__(self):
return self
async def __aexit__(self, *args, **kwargs):
return False
async def post(self, *args, **kwargs):
self.post_calls.append({"args": args, "kwargs": kwargs})
if _MockAsyncHTTPClient._post_side_effect:
raise _MockAsyncHTTPClient._post_side_effect
return _MockAsyncHTTPClient._post_return_value
@classmethod
def reset(cls):
cls._post_return_value = None
cls._post_side_effect = None
def _make_mock_audio(data: bytes = b"fake_pcm_audio_data", filename: str = "test.pcm"):
"""创建模拟的 UploadFile 对象。
Args:
data: 模拟的 PCM 音频数据
filename: 模拟的文件名
Returns:
MagicMock: 模拟的 UploadFile,支持 await audio.read()
"""
mock_audio = MagicMock()
mock_audio.read = AsyncMock(return_value=data)
mock_audio.filename = filename
return mock_audio
# =============================================================================
# 测试组 1get_baidu_token() 防御性检查
# =============================================================================
class TestGetBaiduTokenDefensiveChecks:
"""验证 get_baidu_token() 在凭证缺失时抛出 HTTPException。"""
@pytest.mark.asyncio
async def test_raises_when_api_key_empty(self, monkeypatch):
"""BAIDU_ASR_API_KEY 为空时,应抛出 HTTPException(500)。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_api_key", "")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "valid_secret_key")
# Act & Assert
with pytest.raises(HTTPException) as exc_info:
await get_baidu_token()
assert exc_info.value.status_code == 500
assert "BAIDU_ASR_API_KEY" in exc_info.value.detail
@pytest.mark.asyncio
async def test_raises_when_secret_key_empty(self, monkeypatch):
"""BAIDU_ASR_SECRET_KEY 为空时,应抛出 HTTPException(500)。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_api_key", "valid_api_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "")
# Act & Assert
with pytest.raises(HTTPException) as exc_info:
await get_baidu_token()
assert exc_info.value.status_code == 500
assert "SECRET_KEY" in exc_info.value.detail
@pytest.mark.asyncio
async def test_raises_when_both_empty(self, monkeypatch):
"""API_KEY 和 SECRET_KEY 均为空时,应抛出 HTTPException(500)。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_api_key", "")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "")
# Act & Assert
with pytest.raises(HTTPException) as exc_info:
await get_baidu_token()
assert exc_info.value.status_code == 500
assert "BAIDU_ASR_API_KEY" in exc_info.value.detail
@pytest.mark.asyncio
async def test_error_message_contains_diagnostic_info(self, monkeypatch):
"""错误消息应包含可读的诊断信息,帮助运维定位问题。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_api_key", "")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "")
# Act & Assert
with pytest.raises(HTTPException) as exc_info:
await get_baidu_token()
detail = exc_info.value.detail
assert "百度ASR配置缺失" in detail
assert "BAIDU_ASR_API_KEY" in detail or "SECRET_KEY" in detail
# =============================================================================
# 测试组 2transcribe_audio() app_id 空值检查
# =============================================================================
class TestTranscribeAudioAppIdCheck:
"""验证 transcribe_audio() 在 app_id 为空时返回明确错误。"""
@pytest.mark.asyncio
async def test_returns_error_when_app_id_empty(self, monkeypatch):
"""BAIDU_ASR_APP_ID 为空时,应返回 error_response(code=3001)。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "")
mock_audio = _make_mock_audio()
# Act
result = await transcribe_audio(mock_audio)
# Assert
assert result["code"] == 3001
assert "BAIDU_ASR_APP_ID" in result["message"]
assert result["data"] is None
@pytest.mark.asyncio
async def test_returns_exact_error_message(self, monkeypatch):
"""验证错误消息的精确文本。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "")
mock_audio = _make_mock_audio()
# Act
result = await transcribe_audio(mock_audio)
# Assert
assert result["message"] == "百度ASR配置缺失:BAIDU_ASR_APP_ID未设置"
@pytest.mark.asyncio
async def test_app_id_check_does_not_call_get_baidu_token(self, monkeypatch):
"""app_id 为空时,不应调用 get_baidu_token()(快速失败)。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "")
mock_audio = _make_mock_audio()
# Act
with patch(
"app.api.voice_asr.get_baidu_token",
new_callable=AsyncMock,
) as mock_token:
result = await transcribe_audio(mock_audio)
# Assert
mock_token.assert_not_called()
assert result["code"] == 3001
@pytest.mark.asyncio
async def test_empty_audio_raises_400_before_app_id_check(self, monkeypatch):
"""空音频数据应在 app_id 检查之前抛出 HTTPException(400)。"""
# Arrange — 即使 app_id 有效,空音频也应先报错
monkeypatch.setattr(settings, "baidu_asr_app_id", "valid_app_id")
mock_audio = _make_mock_audio(data=b"")
# Act & Assert
with pytest.raises(HTTPException) as exc_info:
await transcribe_audio(mock_audio)
assert exc_info.value.status_code == 400
assert "未收到音频数据" in exc_info.value.detail
# =============================================================================
# 测试组 3:token 获取异常传播
# =============================================================================
class TestTokenExceptionPropagation:
"""验证 get_baidu_token() 的 HTTPException 在 transcribe_audio 中正确传播。"""
@pytest.mark.asyncio
async def test_http_exception_from_get_token_propagates(self, monkeypatch):
"""app_id 有效但 api_key 为空时,get_baidu_token() 抛出的
HTTPException 应被 transcribe_audio 重新抛出。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "valid_app_id")
monkeypatch.setattr(settings, "baidu_asr_api_key", "")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "")
mock_audio = _make_mock_audio()
# Act & Assert — get_baidu_token 抛 HTTPException(500)
# transcribe_audio 中 except HTTPException: raise 应将其重新抛出
with pytest.raises(HTTPException) as exc_info:
await transcribe_audio(mock_audio)
assert exc_info.value.status_code == 500
assert "BAIDU_ASR_API_KEY" in exc_info.value.detail
@pytest.mark.asyncio
async def test_non_http_exception_returns_error_response(self, monkeypatch):
"""get_baidu_token() 抛出非 HTTPException 异常时,
transcribe_audio 应捕获并返回 error_response。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "valid_app_id")
mock_audio = _make_mock_audio()
# Act
with patch(
"app.api.voice_asr.get_baidu_token",
new_callable=AsyncMock,
side_effect=ConnectionError("Redis 连接失败"),
):
result = await transcribe_audio(mock_audio)
# Assert
assert result["code"] == 3001
assert "token" in result["message"].lower() or "获取" in result["message"]
# =============================================================================
# 测试组 4:正常流程 — cuid 参数映射验证(核心修复点)
# =============================================================================
class TestTranscribeAudioNormalFlow:
"""验证正常流程中 cuid 参数正确映射为 app_id(修复的核心)。"""
def setup_method(self):
"""每个测试前重置 mock HTTP 客户端状态。"""
_MockAsyncHTTPClient.reset()
@pytest.mark.asyncio
async def test_cuid_set_to_app_id_in_api_call(self, monkeypatch):
"""正常流程中,百度 ASR API 请求的 cuid 参数应等于 settings.baidu_asr_app_id。
这是本次修复的核心:之前 cuid 为空导致百度返回 "url param cuid error"
"""
# Arrange
test_app_id = "12345678"
monkeypatch.setattr(settings, "baidu_asr_app_id", test_app_id)
monkeypatch.setattr(settings, "baidu_asr_api_key", "test_api_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "test_secret_key")
mock_audio = _make_mock_audio()
# 配置 mock HTTP 客户端
_MockAsyncHTTPClient._post_return_value = _MockBaiduHttpResponse({
"err_no": 0,
"result": ["你好世界"],
})
# Act
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="fake_access_token")):
with patch("app.api.voice_asr.httpx.AsyncClient", _MockAsyncHTTPClient):
result = await transcribe_audio(mock_audio)
# Assert — 响应正确
assert result["code"] == 0
assert result["data"]["text"] == "你好世界"
assert result["message"] == "success"
# Assert — cuid 参数正确映射为 app_id(核心验证点)
# _MockAsyncHTTPClient 被实例化后,post_calls 记录在实例上
# 通过检查 httpx.AsyncClient 的调用来验证参数
import app.api.voice_asr as voice_asr_module
# patch 已退出,恢复真实 httpx.AsyncClient,但实例化的 _MockAsyncHTTPClient
# 实例的 post_calls 已记录。由于类被替换,每次 AsyncClient(timeout=30)
# 创建一个新实例。我们无法直接拿到实例,但可以通过类变量检查。
# 改用另一种方式:在 patch 期间捕获实例。
# 重新设计:用 side_effect 捕获实例。
_MockAsyncHTTPClient.reset()
_MockAsyncHTTPClient._post_return_value = _MockBaiduHttpResponse({
"err_no": 0,
"result": ["你好世界"],
})
captured_instances = []
class _CapturingClient(_MockAsyncHTTPClient):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
captured_instances.append(self)
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="fake_access_token")):
with patch("app.api.voice_asr.httpx.AsyncClient", _CapturingClient):
result = await transcribe_audio(mock_audio)
assert result["code"] == 0
assert len(captured_instances) == 1
client_instance = captured_instances[0]
# 验证 post 调用参数
assert len(client_instance.post_calls) == 1
call = client_instance.post_calls[0]
called_url = call["args"][0] if call["args"] else call["kwargs"].get("url")
assert called_url == BAIDU_ASR_API_URL
called_params = call["kwargs"].get("params", {})
assert called_params["cuid"] == test_app_id
assert called_params["token"] == "fake_access_token"
assert called_params["dev_pid"] == BAIDU_ASR_DEV_PID
@pytest.mark.asyncio
async def test_successful_recognition_returns_text(self, monkeypatch):
"""正常识别成功时,应返回识别文字。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "test_app_id")
monkeypatch.setattr(settings, "baidu_asr_api_key", "test_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "test_secret")
mock_audio = _make_mock_audio()
_MockAsyncHTTPClient.reset()
_MockAsyncHTTPClient._post_return_value = _MockBaiduHttpResponse({
"err_no": 0,
"result": ["测试识别结果文字"],
})
# Act
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="token123")):
with patch("app.api.voice_asr.httpx.AsyncClient", _MockAsyncHTTPClient):
result = await transcribe_audio(mock_audio)
# Assert
assert result["code"] == 0
assert result["data"]["text"] == "测试识别结果文字"
@pytest.mark.asyncio
async def test_baidu_api_error_returns_error_response(self, monkeypatch):
"""百度 ASR 返回错误时(err_no != 0),应返回 error_response。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "test_app_id")
monkeypatch.setattr(settings, "baidu_asr_api_key", "test_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "test_secret")
mock_audio = _make_mock_audio()
_MockAsyncHTTPClient.reset()
_MockAsyncHTTPClient._post_return_value = _MockBaiduHttpResponse({
"err_no": 3301,
"err_msg": "audio quality too low",
"result": [],
})
# Act
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="token123")):
with patch("app.api.voice_asr.httpx.AsyncClient", _MockAsyncHTTPClient):
result = await transcribe_audio(mock_audio)
# Assert
assert result["code"] == 3001
assert "audio quality too low" in result["message"]
@pytest.mark.asyncio
async def test_timeout_returns_error_response(self, monkeypatch):
"""百度 ASR API 超时时,应返回超时错误响应。"""
# Arrange
import httpx as _httpx
monkeypatch.setattr(settings, "baidu_asr_app_id", "test_app_id")
monkeypatch.setattr(settings, "baidu_asr_api_key", "test_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "test_secret")
mock_audio = _make_mock_audio()
_MockAsyncHTTPClient.reset()
_MockAsyncHTTPClient._post_side_effect = _httpx.TimeoutException(
"Connection timed out"
)
# Act
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="token123")):
with patch("app.api.voice_asr.httpx.AsyncClient", _MockAsyncHTTPClient):
result = await transcribe_audio(mock_audio)
# Assert
assert result["code"] == 3001
assert "超时" in result["message"]
# =============================================================================
# 测试组 5:响应格式契约验证
# =============================================================================
class TestResponseFormatContract:
"""验证错误响应和成功响应符合统一信封格式。"""
def setup_method(self):
_MockAsyncHTTPClient.reset()
@pytest.mark.asyncio
async def test_error_response_has_three_keys(self, monkeypatch):
"""错误响应应包含 code、data、message 三个字段。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "")
mock_audio = _make_mock_audio()
# Act
result = await transcribe_audio(mock_audio)
# Assert
assert set(result.keys()) == {"code", "data", "message"}
assert result["code"] == 3001
assert result["data"] is None
assert isinstance(result["message"], str)
@pytest.mark.asyncio
async def test_success_response_has_three_keys(self, monkeypatch):
"""成功响应应包含 code、data、message 三个字段。"""
# Arrange
monkeypatch.setattr(settings, "baidu_asr_app_id", "test_app_id")
monkeypatch.setattr(settings, "baidu_asr_api_key", "test_key")
monkeypatch.setattr(settings, "baidu_asr_secret_key", "test_secret")
mock_audio = _make_mock_audio()
_MockAsyncHTTPClient._post_return_value = _MockBaiduHttpResponse({
"err_no": 0,
"result": ["成功识别"],
})
# Act
with patch("app.api.voice_asr.get_baidu_token", AsyncMock(return_value="token123")):
with patch("app.api.voice_asr.httpx.AsyncClient", _MockAsyncHTTPClient):
result = await transcribe_audio(mock_audio)
# Assert
assert set(result.keys()) == {"code", "data", "message"}
assert result["code"] == 0
assert "text" in result["data"]
assert result["message"] == "success"
if __name__ == "__main__":
pytest.main([__file__, "-v", "--tb=short"])