From 3d23b94c6efa6c76dc04604ad94eda5d9b14cf18 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Thu, 23 Jul 2026 11:39:19 +0800 Subject: [PATCH] =?UTF-8?q?feat(#674):=20=E6=8E=A5=E5=85=A5=E8=B1=86?= =?UTF-8?q?=E5=8C=85=E5=A4=A7=E6=A8=A1=E5=9E=8B=20Phase=201=20-=20?= =?UTF-8?q?=E6=99=BA=E8=83=BD=E6=A0=87=E9=A2=98=E7=94=9F=E6=88=90=20+=20?= =?UTF-8?q?=E7=BB=9F=E4=B8=80AI=E6=9C=8D=E5=8A=A1=E5=B1=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 DoubaoAIClient 封装豆包 OpenAI 兼容接口,支持重试 - 新增智能标题生成,支持 viral/emotional/informative 三种风格 - 无 API Key 或调用失败时自动降级为本地规则生成 - 新增 /api/v1/ai/titles/generate 和 /styles 两个接口 - 24个单元测试覆盖客户端、解析、降级、边界等场景 - 配置项:DOUBAO_API_KEY/MODEL/BASE_URL/TIMEOUT/MAX_RETRIES --- apps/api/app/api/router.py | 6 + apps/api/app/api/routes/ai.py | 72 ++++++ apps/api/app/config.py | 8 + apps/api/app/services/ai_service.py | 347 ++++++++++++++++++++++++++++ tests/unit/test_ai_service.py | 299 ++++++++++++++++++++++++ 5 files changed, 732 insertions(+) create mode 100755 apps/api/app/api/routes/ai.py create mode 100755 apps/api/app/services/ai_service.py create mode 100755 tests/unit/test_ai_service.py diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 15cb50fe2..754b1a7b5 100755 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -9,6 +9,7 @@ from app.api.routes.feature_flags import router as feature_flags_router from app.api.routes.generation_tasks import router as generation_tasks_router from app.api.routes.health import router as health_check_router from app.api.routes.ingest_jobs import router as ingest_jobs_router +from app.api.routes.ai import router as ai_router from app.api.routes.internal_render import router as internal_render_router from app.api.routes.projects import router as projects_router from app.api.routes.share import router as share_router @@ -134,6 +135,11 @@ api_router.include_router( prefix="/tts", tags=["TTS"], ) +api_router.include_router( + ai_router, + prefix="/ai", + tags=["AI"], +) api_router.include_router( feature_flags_router, tags=["Internal"], diff --git a/apps/api/app/api/routes/ai.py b/apps/api/app/api/routes/ai.py new file mode 100755 index 000000000..f315d3126 --- /dev/null +++ b/apps/api/app/api/routes/ai.py @@ -0,0 +1,72 @@ +"""AI 相关接口 — 智能标题、智能素材匹配等. + +基于豆包大模型的 AI 能力接口,未配置 API Key 时自动降级为本地模拟。 +""" + +from __future__ import annotations + +from typing import List, Literal, Optional + +from app.services.ai_service import TITLE_STYLES, generate_smart_titles +from fastapi import APIRouter +from pydantic import BaseModel, Field + +router = APIRouter() + + +# ── 请求/响应模型 ──────────────────────────────────────────────────────────── + + +class GenerateTitlesRequest(BaseModel): + """智能标题生成请求.""" + + description: str = Field(..., min_length=1, max_length=500, description="视频内容描述") + style: Literal["viral", "emotional", "informative"] = Field( + default="viral", + description="标题风格:viral爆款 / emotional情感 / informative信息", + ) + count: int = Field(default=5, ge=3, le=10, description="生成数量,3-10个") + + +class GenerateTitlesResponse(BaseModel): + """智能标题生成响应.""" + + titles: List[str] = Field(..., description="生成的标题列表") + style: str = Field(..., description="实际使用的风格") + source: str = Field(..., description="来源:doubao 或 fallback") + description: str = Field(..., description="原始描述") + + +class TitleStyleInfo(BaseModel): + """标题风格信息.""" + + key: str + name: str + description: str + + +# ── 路由 ──────────────────────────────────────────────────────────────────── + + +@router.post("/titles/generate", response_model=GenerateTitlesResponse) +def generate_titles(request: GenerateTitlesRequest): + """生成智能标题. + + 根据视频描述生成指定风格的标题,支持爆款、情感、信息三种风格。 + 未配置豆包 API Key 时自动降级为本地规则生成。 + """ + result = generate_smart_titles( + description=request.description, + style=request.style, + count=request.count, + ) + return GenerateTitlesResponse(**result) + + +@router.get("/titles/styles", response_model=List[TitleStyleInfo]) +def list_title_styles(): + """获取支持的标题风格列表.""" + return [ + TitleStyleInfo(key=key, name=info["name"], description=info["description"]) + for key, info in TITLE_STYLES.items() + ] diff --git a/apps/api/app/config.py b/apps/api/app/config.py index 68ec61b5d..893e50ff0 100755 --- a/apps/api/app/config.py +++ b/apps/api/app/config.py @@ -109,6 +109,14 @@ class Settings(BaseSettings): # 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService RENDER_ENGINE: str = "legacy" + # 豆包大模型配置(火山引擎方舟平台) + # 未配置 API Key 时自动降级为本地模拟生成 + DOUBAO_API_KEY: str = "" + DOUBAO_MODEL: str = "doubao-seed-1-6-250615" + DOUBAO_BASE_URL: str = "https://ark.cn-beijing.volces.com/api/v3" + DOUBAO_TIMEOUT: int = 30 + DOUBAO_MAX_RETRIES: int = 2 + model_config = SettingsConfigDict( env_file=".env", env_file_encoding="utf-8", diff --git a/apps/api/app/services/ai_service.py b/apps/api/app/services/ai_service.py new file mode 100755 index 000000000..4f82b972b --- /dev/null +++ b/apps/api/app/services/ai_service.py @@ -0,0 +1,347 @@ +"""统一 AI 服务层 — 豆包大模型接入. + +提供基于字节跳动豆包大模型的 AI 能力: +- 智能标题生成(爆款/情感/信息三种风格) +- 后续扩展:智能素材匹配、AI 推荐片段编排等 + +设计原则: +1. 无 API Key 或调用失败时自动降级为本地模拟,不阻塞主流程 +2. 统一的客户端封装,新增能力只需加方法 +3. 所有模型相关配置集中在 Settings +""" + +from __future__ import annotations + +import json +import logging +import random +import time +from typing import Any, Dict, List, Optional + +import httpx +from app.config import get_settings + +logger = logging.getLogger(__name__) + + +# ── 智能标题风格定义 ───────────────────────────────────────────────────────── + +TITLE_STYLES = { + "viral": { + "name": "爆款", + "description": "吸引点击、引发好奇的爆款标题,带有数字、疑问或反差感", + "examples": [ + "3个方法让你效率翻倍,第2个最绝", + "为什么越努力越穷?真相扎心了", + "看完这个,我删掉了手机里一半的APP", + ], + }, + "emotional": { + "name": "情感", + "description": "触动人心、引发共鸣的情感向标题", + "examples": [ + "那些年我们一起追过的梦想", + "生活不易,但请相信光", + "致每一个在城市里打拼的你", + ], + }, + "informative": { + "name": "信息", + "description": "清晰直白、传递核心信息的干货标题", + "examples": [ + "2026年最新个税政策解读,一文讲透", + "新手剪辑入门:从0到1完整指南", + "产品对比:10款热门手机深度评测", + ], + }, +} + + +# ── 豆包 AI 客户端 ────────────────────────────────────────────────────────── + + +class DoubaoAIClient: + """豆包大模型 API 客户端. + + 使用火山引擎方舟平台的 OpenAI 兼容接口。 + 未配置 API Key 时,is_available 返回 False,调用方应降级处理。 + """ + + def __init__(self) -> None: + settings = get_settings() + self.api_key: str = settings.DOUBAO_API_KEY + self.model: str = settings.DOUBAO_MODEL + self.base_url: str = settings.DOUBAO_BASE_URL.rstrip("/") + self.timeout: int = settings.DOUBAO_TIMEOUT + self.max_retries: int = settings.DOUBAO_MAX_RETRIES + + @property + def is_available(self) -> bool: + """是否可用(配置了 API Key).""" + return bool(self.api_key) + + def _chat_completion( + self, + messages: List[Dict[str, str]], + temperature: float = 0.7, + max_tokens: int = 1024, + ) -> Optional[str]: + """调用豆包 Chat Completion 接口. + + Returns: + 模型返回的文本内容,失败返回 None + """ + if not self.is_available: + return None + + url = f"{self.base_url}/chat/completions" + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + payload = { + "model": self.model, + "messages": messages, + "temperature": temperature, + "max_tokens": max_tokens, + } + + last_error: Optional[Exception] = None + for attempt in range(self.max_retries + 1): + try: + response = httpx.post( + url, + headers=headers, + json=payload, + timeout=self.timeout, + ) + response.raise_for_status() + data = response.json() + content = data["choices"][0]["message"]["content"] + return content.strip() + except Exception as e: + last_error = e + if attempt < self.max_retries: + wait = 0.5 * (2**attempt) + logger.warning( + "豆包API调用失败,%s秒后重试 (第%d/%d次): %s", + wait, + attempt + 1, + self.max_retries + 1, + e, + ) + time.sleep(wait) + + logger.error("豆包API调用最终失败: %s", last_error) + return None + + +# ── 智能标题生成 ───────────────────────────────────────────────────────────── + + +def _generate_titles_fallback( + description: str, + style: str = "viral", + count: int = 5, +) -> List[str]: + """本地降级:基于模板规则生成标题. + + 当豆包 API 不可用或调用失败时使用,保证接口始终有返回。 + """ + style_info = TITLE_STYLES.get(style, TITLE_STYLES["viral"]) + examples = style_info["examples"] + + # 从描述中提取关键词(取前几个词) + keywords = [w for w in description.strip().split() if len(w) > 1][:3] + keyword = keywords[0] if keywords else "精彩内容" + + # 基于模板生成 + templates = [ + f"「{keyword}」{examples[0][:10]}...", + f"{keyword}:{examples[1]}", + f"关于{keyword},你不知道的3件事", + f"{keyword}入门指南,新手必看", + f"深度解析:{keyword}背后的秘密", + f"{keyword}怎么做?手把手教你", + f"干货分享 | {keyword}全攻略", + f"建议收藏:{keyword}实用技巧", + f"{keyword}避坑指南,别再踩雷了", + f"一分钟搞懂{keyword}", + ] + + random.shuffle(templates) + return templates[: min(count, len(templates))] + + +def _parse_titles_from_response(content: str) -> List[str]: + """从模型返回中解析标题列表. + + 支持多种返回格式: + - JSON 数组: ["标题1", "标题2"] + - 编号列表: 1. 标题1 / 2. 标题2 + - 换行分隔: 标题1\n标题2 + - 带破折号: - 标题1 + """ + if not content: + return [] + + # 尝试解析 JSON + try: + # 清理可能的 markdown 代码块标记 + cleaned = content.strip() + if cleaned.startswith("```"): + cleaned = cleaned.strip("`") + if cleaned.lower().startswith("json"): + cleaned = cleaned[4:] + cleaned = cleaned.strip() + + data = json.loads(cleaned) + if isinstance(data, list): + return [str(item).strip() for item in data if str(item).strip()] + if isinstance(data, dict) and "titles" in data: + titles = data["titles"] + if isinstance(titles, list): + return [str(t).strip() for t in titles if str(t).strip()] + except (json.JSONDecodeError, ValueError): + pass + + # 尝试按行解析 + titles: List[str] = [] + for line in content.strip().split("\n"): + line = line.strip() + if not line: + continue + # 去掉编号前缀 "1. " "1、" "(1)" + import re + + line = re.sub(r"^[\d]+[\.、\))]\s*", "", line) + # 去掉破折号前缀 "- " "• " + line = re.sub(r"^[-•·]\s*", "", line) + # 去掉引号 + line = line.strip('"').strip("'").strip("「」") + if line and len(line) < 100: # 过滤过长的行 + titles.append(line) + + return titles + + +def generate_smart_titles( + description: str, + style: str = "viral", + count: int = 5, +) -> Dict[str, Any]: + """生成智能标题. + + Args: + description: 视频内容描述 + style: 标题风格 viral/emotional/informative + count: 生成数量(5-10) + + Returns: + { + "titles": [...], + "style": "viral", + "source": "doubao" | "fallback", # 实际来源 + "description": "...", + } + """ + # 参数校验与边界处理 + if style not in TITLE_STYLES: + style = "viral" + count = max(3, min(10, count)) # 3-10 个 + description = (description or "").strip() + + client = DoubaoAIClient() + if not client.is_available: + logger.info("豆包API未配置,使用本地降级生成标题") + titles = _generate_titles_fallback(description, style, count) + return { + "titles": titles, + "style": style, + "source": "fallback", + "description": description, + } + + style_info = TITLE_STYLES[style] + system_prompt = ( + f"你是一个专业的短视频标题创作专家,擅长根据视频内容生成吸引人的标题。\n" + f"请根据以下视频描述,生成{count}个{style_info['name']}风格的标题。\n" + f"风格说明:{style_info['description']}\n" + f"要求:\n" + f"1. 每个标题控制在8-25字之间\n" + f"2. 直接返回JSON数组格式,不要其他文字\n" + f"3. 标题要贴合内容,有吸引力" + ) + + user_prompt = f"视频描述:{description}\n\n请生成标题:" + + messages = [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt}, + ] + + result = client._chat_completion( + messages=messages, + temperature=0.8, + max_tokens=512, + ) + + if result: + titles = _parse_titles_from_response(result) + if len(titles) >= 2: # 至少解析出2个才算成功 + titles = titles[:count] + logger.info( + "豆包智能标题生成成功: style=%s count=%d description=%s...", + style, + len(titles), + description[:20], + ) + return { + "titles": titles, + "style": style, + "source": "doubao", + "description": description, + } + logger.warning("豆包返回内容解析失败,降级到本地生成: %s", result[:100]) + + # 降级到本地生成 + titles = _generate_titles_fallback(description, style, count) + return { + "titles": titles, + "style": style, + "source": "fallback", + "description": description, + } + + +# ── 单例入口 ───────────────────────────────────────────────────────────────── + + +def get_ai_service() -> "AIService": + """获取 AI 服务单例.""" + global _ai_service + if _ai_service is None: + _ai_service = AIService() + return _ai_service + + +_ai_service: Optional["AIService"] = None + + +class AIService: + """AI 服务统一入口,便于后续扩展更多能力.""" + + def __init__(self) -> None: + self._client = DoubaoAIClient() + + @property + def is_available(self) -> bool: + return self._client.is_available + + def generate_titles( + self, + description: str, + style: str = "viral", + count: int = 5, + ) -> Dict[str, Any]: + return generate_smart_titles(description, style, count) diff --git a/tests/unit/test_ai_service.py b/tests/unit/test_ai_service.py new file mode 100755 index 000000000..738737659 --- /dev/null +++ b/tests/unit/test_ai_service.py @@ -0,0 +1,299 @@ +"""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 + DoubaoAIClient, + TITLE_STYLES, + _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() -- 2.54.0