Files
xiaoxia-saas/tests/unit/test_ai_service.py
T
xiaoxia 0a424bbc46
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 6s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m14s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m36s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 34s
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m3s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m4s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
Preview Deploy / Deploy Preview Environment (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m5s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m27s
style: 用black+isort重新格式化(与CI工具链对齐)
2026-07-23 17:32:34 +08:00

294 lines
11 KiB
Python
Executable File

"""AI 服务层单元测试.
测试覆盖:
- DoubaoAIClient 可用性检测
- 智能标题生成(降级模式)
- 标题解析(多种返回格式)
- 风格校验
- 参数边界
"""
from __future__ import annotations
import json
import sys
import unittest
from unittest.mock import MagicMock, patch
sys.path.insert(0, "apps/api")
from app.services.ai_service import ( # noqa: E402
TITLE_STYLES,
DoubaoAIClient,
_generate_titles_fallback,
_parse_titles_from_response,
generate_smart_titles,
)
class TestDoubaoAIClient(unittest.TestCase):
"""豆包客户端基础测试."""
def test_client_availability_without_key(self):
"""未配置 API Key 时不可用."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
client = DoubaoAIClient()
self.assertFalse(client.is_available)
def test_client_availability_with_key(self):
"""配置了 API Key 时可用."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="sk-test-123",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
client = DoubaoAIClient()
self.assertTrue(client.is_available)
def test_chat_completion_not_available_returns_none(self):
"""不可用时调用返回 None."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
client = DoubaoAIClient()
result = client._chat_completion([{"role": "user", "content": "hi"}])
self.assertIsNone(result)
class TestTitleParsing(unittest.TestCase):
"""标题解析测试 — 覆盖多种返回格式."""
def test_parse_json_array(self):
"""解析 JSON 数组格式."""
content = json.dumps(["标题一", "标题二", "标题三"])
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 3)
self.assertEqual(result[0], "标题一")
def test_parse_json_with_titles_key(self):
"""解析带 titles 字段的 JSON 对象."""
content = json.dumps({"titles": ["标题A", "标题B"]})
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 2)
def test_parse_markdown_code_block_json(self):
"""解析 markdown 代码块包裹的 JSON."""
content = '```json\n["标题1", "标题2"]\n```'
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 2)
def test_parse_numbered_list(self):
"""解析编号列表."""
content = "1. 第一个标题\n2. 第二个标题\n3. 第三个标题"
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 3)
self.assertIn("第一个标题", result)
def test_parse_dash_list(self):
"""解析破折号列表."""
content = "- 标题甲\n- 标题乙\n- 标题丙"
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 3)
def test_parse_chinese_numbered(self):
"""解析中文数字编号."""
content = "1、标题一\n2、标题二"
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 2)
def test_parse_empty_content(self):
"""空内容返回空列表."""
result = _parse_titles_from_response("")
self.assertEqual(result, [])
def test_parse_filters_long_lines(self):
"""过滤过长的行."""
long_title = "这是一个非常长的标题" * 15 # 超过100字
content = f"1. 正常标题\n2. {long_title}\n3. 另一个标题"
result = _parse_titles_from_response(content)
self.assertEqual(len(result), 2)
self.assertNotIn(long_title, result)
def test_parse_invalid_json_falls_back_to_lines(self):
"""无效 JSON 回退到按行解析."""
content = '["标题1", "标题2", 无效'
result = _parse_titles_from_response(content)
# 至少能解析出一些内容
self.assertTrue(len(result) >= 0)
class TestFallbackGeneration(unittest.TestCase):
"""降级生成测试."""
def test_fallback_returns_requested_count(self):
"""返回请求的数量."""
result = _generate_titles_fallback("测试内容", "viral", 5)
self.assertEqual(len(result), 5)
def test_fallback_max_10(self):
"""最多返回10个."""
result = _generate_titles_fallback("测试内容", "viral", 20)
self.assertEqual(len(result), 10)
def test_fallback_different_styles(self):
"""不同风格都能生成."""
for style in ["viral", "emotional", "informative"]:
result = _generate_titles_fallback("测试", style, 3)
self.assertEqual(len(result), 3)
for title in result:
self.assertTrue(len(title) > 0)
def test_fallback_contains_keyword(self):
"""标题包含关键词."""
result = _generate_titles_fallback("旅行攻略", "viral", 5)
has_keyword = any("旅行" in t for t in result)
self.assertTrue(has_keyword)
class TestGenerateSmartTitles(unittest.TestCase):
"""智能标题生成集成测试."""
def test_generate_without_api_key_fallback(self):
"""无 API Key 时走降级路径."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="",
DOUBAO_MODEL="test",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
result = generate_smart_titles("测试视频内容", "viral", 5)
self.assertEqual(result["source"], "fallback")
self.assertEqual(result["style"], "viral")
self.assertEqual(len(result["titles"]), 5)
def test_generate_invalid_style_defaults_to_viral(self):
"""无效风格默认 viral."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="",
DOUBAO_MODEL="test",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
result = generate_smart_titles("测试", "invalid_style", 5)
self.assertEqual(result["style"], "viral")
def test_generate_count_bounds(self):
"""数量边界处理."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="",
DOUBAO_MODEL="test",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=2,
)
# 小于最小值
result = generate_smart_titles("测试", "viral", 1)
self.assertEqual(len(result["titles"]), 3)
# 大于最大值
result = generate_smart_titles("测试", "viral", 100)
self.assertEqual(len(result["titles"]), 10)
def test_generate_with_api_success(self):
"""API 调用成功路径."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="sk-test-123",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=0,
)
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"choices": [
{"message": {"content": json.dumps(["AI标题1", "AI标题2", "AI标题3", "AI标题4", "AI标题5"])}}
]
}
mock_response.raise_for_status = MagicMock()
with patch("httpx.post", return_value=mock_response):
result = generate_smart_titles("测试视频", "viral", 5)
self.assertEqual(result["source"], "doubao")
self.assertEqual(len(result["titles"]), 5)
self.assertIn("AI标题1", result["titles"])
def test_generate_with_api_failure_fallback(self):
"""API 调用失败时降级."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="sk-test-123",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=1,
DOUBAO_MAX_RETRIES=0,
)
with patch("httpx.post", side_effect=Exception("API Error")):
result = generate_smart_titles("测试视频", "viral", 5)
self.assertEqual(result["source"], "fallback")
self.assertEqual(len(result["titles"]), 5)
def test_generate_api_returns_unparseable_fallback(self):
"""API 返回无法解析时降级."""
with patch("app.services.ai_service.get_settings") as mock_settings:
mock_settings.return_value = MagicMock(
DOUBAO_API_KEY="sk-test-123",
DOUBAO_MODEL="test-model",
DOUBAO_BASE_URL="https://test.com",
DOUBAO_TIMEOUT=30,
DOUBAO_MAX_RETRIES=0,
)
mock_response = MagicMock()
mock_response.status_code = 200
# 返回无法解析的内容(只有一个标题且格式异常)
mock_response.json.return_value = {"choices": [{"message": {"content": "一段文字说明,不是标题列表"}}]}
mock_response.raise_for_status = MagicMock()
with patch("httpx.post", return_value=mock_response):
result = generate_smart_titles("测试视频", "viral", 5)
# 只有1个有效标题,不足2个触发降级
self.assertEqual(result["source"], "fallback")
class TestTitleStyles(unittest.TestCase):
"""标题风格定义测试."""
def test_all_styles_have_required_fields(self):
"""所有风格都有必要字段."""
for key, info in TITLE_STYLES.items():
self.assertIn("name", info)
self.assertIn("description", info)
self.assertIn("examples", info)
self.assertTrue(len(info["examples"]) >= 2)
def test_three_styles_defined(self):
"""定义了三种风格."""
self.assertEqual(len(TITLE_STYLES), 3)
self.assertIn("viral", TITLE_STYLES)
self.assertIn("emotional", TITLE_STYLES)
self.assertIn("informative", TITLE_STYLES)
if __name__ == "__main__":
unittest.main()