feat(#674): 接入豆包大模型 Phase 1 - 智能标题生成 + 统一AI服务层 #752

Merged
xiaoxia merged 1 commits from feat/doubao-ai-integration-phase1 into develop 2026-07-23 13:17:39 +08:00
5 changed files with 732 additions and 0 deletions
+6
View File
@@ -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"],
+72
View File
@@ -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()
]
+8
View File
@@ -109,6 +109,14 @@ class Settings(BaseSettings):
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新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",
+347
View File
@@ -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)
+299
View File
@@ -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()