"""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, _generate_titles_fallback, _parse_semantic_match_response, _parse_titles_from_response, _semantic_match_fallback, generate_smart_titles, semantic_match_assets, ) class TestAIClientAvailability(unittest.TestCase): """AI客户端可用性检测(通过mock get_doubao_client).""" def test_generate_fallback_when_client_unavailable(self): """客户端不可用时走降级.""" mock_client = MagicMock() mock_client.is_available = False mock_client.chat_completion = MagicMock(return_value=None) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = generate_smart_titles("测试内容", "viral", 5) self.assertEqual(result["source"], "fallback") self.assertEqual(len(result["titles"]), 5) # 不可用时不应调用 chat_completion mock_client.chat_completion.assert_not_called() def test_generate_calls_client_when_available(self): """客户端可用时调用API.""" mock_client = MagicMock() mock_client.is_available = True mock_client.chat_completion = MagicMock( return_value=json.dumps(["AI标题1", "AI标题2", "AI标题3", "AI标题4", "AI标题5"]) ) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = generate_smart_titles("测试", "viral", 5) self.assertEqual(result["source"], "doubao") mock_client.chat_completion.assert_called_once() 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 时走降级路径.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): 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.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = generate_smart_titles("测试", "invalid_style", 5) self.assertEqual(result["style"], "viral") def test_generate_count_bounds(self): """数量边界处理.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): # 小于最小值 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 调用成功路径.""" mock_client = MagicMock() mock_client.is_available = True mock_client.chat_completion = MagicMock( return_value=json.dumps(["AI标题1", "AI标题2", "AI标题3", "AI标题4", "AI标题5"]) ) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): 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 调用失败时降级.""" mock_client = MagicMock() mock_client.is_available = True mock_client.chat_completion = MagicMock(return_value=None) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): 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 返回无法解析时降级.""" mock_client = MagicMock() mock_client.is_available = True # 返回无法解析的内容(只有一个标题且格式异常) mock_client.chat_completion = MagicMock(return_value="一段文字说明,不是标题列表") with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): 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) # ── 语义匹配测试 ────────────────────────────────────────────────────────────── class TestSemanticMatchFallback(unittest.TestCase): """降级关键词匹配测试.""" def _make_assets(self): return [ {"id": "a1", "name": "海边日落风景", "tags": ["风景", "海边", "日落"], "description": "美丽的海边日落"}, {"id": "a2", "name": "城市夜景航拍", "tags": ["城市", "夜景", "航拍"], "description": "城市夜景航拍素材"}, {"id": "a3", "name": "美食制作过程", "tags": ["美食", "烹饪", "教程"], "description": "美食制作教程"}, ] def test_fallback_returns_sorted_scores(self): """返回按匹配度降序排列.""" assets = self._make_assets() result = _semantic_match_fallback("海边日落风景视频", assets) self.assertEqual(len(result), 3) # 第一个应该是海边日落 self.assertEqual(result[0]["id"], "a1") self.assertGreater(result[0]["match_score"], result[2]["match_score"]) def test_fallback_each_has_match_score(self): """每个素材都有 match_score.""" assets = self._make_assets() result = _semantic_match_fallback("测试", assets) for item in result: self.assertIn("match_score", item) self.assertGreaterEqual(item["match_score"], 0.0) self.assertLessEqual(item["match_score"], 1.0) self.assertIn("match_reason", item) def test_fallback_unrelated_desc_low_scores(self): """完全不相关的描述得分低.""" assets = self._make_assets() result = _semantic_match_fallback("篮球比赛运动", assets) # 所有素材得分都应该较低 for item in result: self.assertLess(item["match_score"], 0.8) def test_fallback_empty_keywords_default_score(self): """无有效关键词时给默认分.""" assets = self._make_assets() result = _semantic_match_fallback("a", assets) # 单字符无有效关键词 for item in result: self.assertEqual(item["match_score"], 0.5) self.assertEqual(item["match_reason"], "fallback_default") def test_fallback_name_match_higher(self): """名称命中得分更高.""" assets = [ {"id": "a1", "name": "美食探店vlog", "tags": [], "description": ""}, {"id": "a2", "name": "风景视频", "tags": ["美食"], "description": ""}, ] result = _semantic_match_fallback("美食", assets) # a1名称含美食,a2标签含美食,名称命中应有额外加分 self.assertEqual(result[0]["id"], "a1") self.assertGreater(result[0]["match_score"], result[1]["match_score"]) class TestSemanticMatchParsing(unittest.TestCase): """语义匹配返回解析测试.""" def test_parse_dict_format(self): """解析 {id: score} 格式.""" content = json.dumps({"asset1": 0.85, "asset2": 0.62, "asset3": 0.3}) result = _parse_semantic_match_response(content, ["asset1", "asset2", "asset3"]) self.assertIsNotNone(result) self.assertEqual(len(result), 3) self.assertAlmostEqual(result["asset1"], 0.85) def test_parse_matches_list_format(self): """解析 {matches: [...]} 格式.""" content = json.dumps( { "matches": [ {"asset_id": "a1", "score": 0.9}, {"asset_id": "a2", "score": 0.7}, ] } ) result = _parse_semantic_match_response(content, ["a1", "a2"]) self.assertIsNotNone(result) self.assertAlmostEqual(result["a1"], 0.9) self.assertAlmostEqual(result["a2"], 0.7) def test_parse_array_format(self): """解析数组格式.""" content = json.dumps( [ {"id": "x1", "score": 0.5}, {"id": "x2", "score": 0.88}, ] ) result = _parse_semantic_match_response(content, ["x1", "x2"]) self.assertIsNotNone(result) self.assertAlmostEqual(result["x1"], 0.5) def test_parse_score_clamped(self): """分数被限制在0-1.""" content = json.dumps({"a1": 1.5, "a2": -0.2}) result = _parse_semantic_match_response(content, ["a1", "a2"]) self.assertIsNotNone(result) self.assertAlmostEqual(result["a1"], 1.0) self.assertAlmostEqual(result["a2"], 0.0) def test_parse_markdown_code_block(self): """解析markdown代码块.""" content = '```json\n{"a1": 0.7}\n```' result = _parse_semantic_match_response(content, ["a1", "a2"]) # 只有1个素材评分,少于一半(需要至少1个,max(1, 2//2)=1) self.assertIsNotNone(result) self.assertAlmostEqual(result["a1"], 0.7) def test_parse_empty_returns_none(self): """空内容返回None.""" result = _parse_semantic_match_response("", ["a1"]) self.assertIsNone(result) def test_parse_invalid_json_returns_none(self): """无效JSON返回None.""" result = _parse_semantic_match_response("不是json", ["a1", "a2", "a3"]) self.assertIsNone(result) class TestSemanticMatchAssets(unittest.TestCase): """semantic_match_assets 集成测试.""" def _make_assets(self): return [ {"id": "a1", "name": "海边日落", "tags": ["风景"], "description": ""}, {"id": "a2", "name": "城市夜景", "tags": ["城市"], "description": ""}, {"id": "a3", "name": "美食制作", "tags": ["美食"], "description": ""}, ] def test_fallback_mode_without_api_key(self): """无API Key时走降级.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = semantic_match_assets("海边", self._make_assets()) self.assertEqual(result["source"], "fallback") self.assertEqual(result["total"], 3) self.assertEqual(len(result["matches"]), 3) def test_empty_assets(self): """空素材列表.""" result = semantic_match_assets("test", []) self.assertEqual(result["total"], 0) self.assertEqual(len(result["matches"]), 0) def test_top_k_limit(self): """top_k 限制返回数量.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = semantic_match_assets("测试", self._make_assets(), top_k=2) self.assertEqual(len(result["matches"]), 2) def test_with_doubao_success(self): """豆包调用成功路径.""" mock_client = MagicMock() mock_client.is_available = True mock_client.chat_completion = MagicMock(return_value=json.dumps({"a1": 0.9, "a2": 0.5, "a3": 0.2})) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = semantic_match_assets("风景视频", self._make_assets()) self.assertEqual(result["source"], "doubao") self.assertEqual(len(result["matches"]), 3) # 按分数降序,a1最高 self.assertEqual(result["matches"][0]["id"], "a1") self.assertAlmostEqual(result["matches"][0]["match_score"], 0.9) def test_with_doubao_failure_fallback(self): """豆包调用失败降级.""" mock_client = MagicMock() mock_client.is_available = True mock_client.chat_completion = MagicMock(return_value=None) with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = semantic_match_assets("测试", self._make_assets()) self.assertEqual(result["source"], "fallback") def test_each_match_has_required_fields(self): """每个匹配结果都有必要字段.""" mock_client = MagicMock() mock_client.is_available = False with patch("app.services.ai_service.get_doubao_client", return_value=mock_client): result = semantic_match_assets("测试", self._make_assets()) for item in result["matches"]: self.assertIn("id", item) self.assertIn("match_score", item) self.assertIn("match_reason", item) if __name__ == "__main__": unittest.main()