349 lines
12 KiB
Python
Executable File
349 lines
12 KiB
Python
Executable File
"""统一 AI 服务层 — 豆包大模型接入.
|
||
|
||
提供基于字节跳动豆包大模型的 AI 能力:
|
||
- 智能标题生成(爆款/情感/信息三种风格)
|
||
- 后续扩展:智能素材匹配、AI 推荐片段编排等
|
||
|
||
设计原则:
|
||
1. 无 API Key 或调用失败时自动降级为本地模拟,不阻塞主流程
|
||
2. 统一的客户端封装,新增能力只需加方法
|
||
3. 所有模型相关配置集中在 Settings
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from typing import Any, Dict, List, Optional
|
||
|
||
from packages.domain.ai_parsing import generate_titles_fallback as _generate_titles_fallback_base
|
||
from packages.domain.ai_parsing import keyword_match_fallback as _semantic_match_fallback_base
|
||
from packages.domain.ai_parsing import parse_semantic_match_response as _parse_semantic_match_base
|
||
from packages.domain.ai_parsing import parse_titles_from_response as _parse_titles_from_response
|
||
from packages.shared.ai_client import get_doubao_client
|
||
|
||
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款热门手机深度评测",
|
||
],
|
||
},
|
||
}
|
||
|
||
|
||
# ── 智能标题生成 ─────────────────────────────────────────────────────────────
|
||
|
||
|
||
def _generate_titles_fallback(
|
||
description: str,
|
||
style: str = "viral",
|
||
count: int = 5,
|
||
) -> List[str]:
|
||
"""本地降级:基于模板规则生成标题(薄包装,转发到 ai_parsing 模块)."""
|
||
style_info = TITLE_STYLES.get(style, TITLE_STYLES["viral"])
|
||
return _generate_titles_fallback_base(description, style_info, count)
|
||
|
||
|
||
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 = get_doubao_client()
|
||
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 _semantic_match_fallback(
|
||
description: str,
|
||
assets: List[Dict[str, Any]],
|
||
) -> List[Dict[str, Any]]:
|
||
"""本地降级:基于关键词的简单匹配(薄包装,转发到 ai_parsing 模块)."""
|
||
return _semantic_match_fallback_base(description, assets)
|
||
|
||
|
||
def _parse_semantic_match_response(
|
||
content: str,
|
||
asset_ids: List[str],
|
||
) -> Optional[Dict[str, float]]:
|
||
"""从模型返回中解析素材匹配度(薄包装,转发到 ai_parsing 模块)."""
|
||
result = _parse_semantic_match_base(content, asset_ids)
|
||
if result is None:
|
||
return None
|
||
return dict(result)
|
||
|
||
|
||
def semantic_match_assets(
|
||
description: str,
|
||
assets: List[Dict[str, Any]],
|
||
top_k: int = 0,
|
||
) -> Dict[str, Any]:
|
||
"""智能素材语义匹配.
|
||
|
||
根据用户描述,评估每个素材的语义匹配度并排序。
|
||
|
||
Args:
|
||
description: 用户描述的目标视频内容
|
||
assets: 素材列表,每个素材需含 id/name/tags/description 等字段
|
||
top_k: 返回前K个,0表示返回全部
|
||
|
||
Returns:
|
||
{
|
||
"matches": [{"asset_id": ..., "match_score": ..., ...}],
|
||
"source": "doubao" | "fallback",
|
||
"description": "...",
|
||
"total": 总数,
|
||
}
|
||
"""
|
||
description = (description or "").strip()
|
||
if not assets:
|
||
return {"matches": [], "source": "fallback", "description": description, "total": 0}
|
||
|
||
client = get_doubao_client()
|
||
if not client.is_available:
|
||
logger.info("豆包API未配置,使用本地降级做素材语义匹配")
|
||
matched = _semantic_match_fallback(description, assets)
|
||
if top_k > 0:
|
||
matched = matched[:top_k]
|
||
return {
|
||
"matches": matched,
|
||
"source": "fallback",
|
||
"description": description,
|
||
"total": len(assets),
|
||
}
|
||
|
||
# 构建素材信息(控制 token 数量)
|
||
asset_summaries = []
|
||
for asset in assets[:50]: # 最多传50个素材给模型
|
||
aid = asset.get("id", "")
|
||
name = asset.get("name", "")[:50]
|
||
tags = asset.get("tags", [])
|
||
tags_str = ",".join(str(t) for t in tags[:5])
|
||
desc = str(asset.get("description", ""))[:80]
|
||
asset_summaries.append(f"ID:{aid} | 名称:{name} | 标签:[{tags_str}] | 描述:{desc}")
|
||
|
||
asset_ids = [str(a.get("id", "")) for a in assets[:50]]
|
||
|
||
system_prompt = (
|
||
"你是一个专业的视频素材匹配助手。"
|
||
"根据用户的视频目标描述,评估每个素材的匹配程度。\n"
|
||
"评分规则:\n"
|
||
"- 0.0-0.3: 完全不相关\n"
|
||
"- 0.3-0.6: 有一定关联但不够匹配\n"
|
||
"- 0.6-0.8: 比较匹配,适合使用\n"
|
||
"- 0.8-1.0: 高度匹配,非常适合\n"
|
||
"只返回JSON对象,key为素材ID,value为匹配分数(0-1之间的小数)。"
|
||
"不要其他文字说明。"
|
||
)
|
||
|
||
user_prompt = (
|
||
f"目标视频描述:{description}\n\n"
|
||
f"素材列表:\n" + "\n".join(asset_summaries) + "\n\n请返回每个素材的匹配分数JSON:"
|
||
)
|
||
|
||
messages = [
|
||
{"role": "system", "content": system_prompt},
|
||
{"role": "user", "content": user_prompt},
|
||
]
|
||
|
||
result = client.chat_completion(
|
||
messages=messages,
|
||
temperature=0.3,
|
||
max_tokens=1024,
|
||
)
|
||
|
||
if result:
|
||
scores = _parse_semantic_match_response(result, asset_ids)
|
||
if scores:
|
||
# 把评分填回素材
|
||
matched = []
|
||
for asset in assets:
|
||
aid = str(asset.get("id", ""))
|
||
score = scores.get(aid, 0.3) # 没评分的给默认偏低分
|
||
matched.append(
|
||
{
|
||
**asset,
|
||
"match_score": round(score, 3),
|
||
"match_reason": "doubao_semantic",
|
||
}
|
||
)
|
||
matched.sort(key=lambda x: x["match_score"], reverse=True)
|
||
|
||
logger.info(
|
||
"豆包语义匹配完成: assets=%d top_score=%.2f description=%s...",
|
||
len(matched),
|
||
matched[0]["match_score"] if matched else 0,
|
||
description[:20],
|
||
)
|
||
|
||
if top_k > 0:
|
||
matched = matched[:top_k]
|
||
|
||
return {
|
||
"matches": matched,
|
||
"source": "doubao",
|
||
"description": description,
|
||
"total": len(assets),
|
||
}
|
||
logger.warning("豆包语义匹配返回解析失败,降级到本地: %s", result[:100])
|
||
|
||
# 降级
|
||
matched = _semantic_match_fallback(description, assets)
|
||
if top_k > 0:
|
||
matched = matched[:top_k]
|
||
return {
|
||
"matches": matched,
|
||
"source": "fallback",
|
||
"description": description,
|
||
"total": len(assets),
|
||
}
|
||
|
||
|
||
# ── 单例入口 ─────────────────────────────────────────────────────────────────
|
||
|
||
|
||
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 = get_doubao_client()
|
||
|
||
@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)
|
||
|
||
def semantic_match(
|
||
self,
|
||
description: str,
|
||
assets: List[Dict[str, Any]],
|
||
top_k: int = 0,
|
||
) -> Dict[str, Any]:
|
||
return semantic_match_assets(description, assets, top_k)
|