Files
wecom_it_smart_desk/backend/tests/test_voice_asr.py
T

473 lines
18 KiB
Python
Raw Normal View History

# =============================================================================
# 百度 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"])