feat: #1209 AI智能选片段接入MediaKit视频理解 #1224
@@ -24,6 +24,91 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
def _build_asset_analyses(
|
||||
asset_ids: list[str],
|
||||
db: Session,
|
||||
) -> dict[str, str]:
|
||||
"""调用 MediaKit 视频理解,返回 {asset_id: 分析文本}.
|
||||
|
||||
如果 MediaKit 不可用或分析失败,返回空 dict(调用方降级处理)。
|
||||
"""
|
||||
if not asset_ids:
|
||||
return {}
|
||||
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
client = get_mediakit_client()
|
||||
if not client.is_available:
|
||||
logger.info("MediaKit 未配置,跳过视频理解分析")
|
||||
return {}
|
||||
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
storage_svc = get_shared_storage_service()
|
||||
|
||||
# 查找素材并获取下载 URL
|
||||
video_urls: list[str] = []
|
||||
url_to_asset_id: dict[str, str] = {}
|
||||
|
||||
for aid in asset_ids[:10]: # MediaKit 单次最多 10 个视频
|
||||
asset = asset_repo.get(aid)
|
||||
if not asset or not asset.storage_key:
|
||||
continue
|
||||
# 只处理视频素材
|
||||
mime = getattr(asset, "mime_type", "")
|
||||
if not mime.startswith("video/"):
|
||||
continue
|
||||
try:
|
||||
url = storage_svc.get_download_url(asset.storage_key)
|
||||
if url:
|
||||
video_urls.append(url)
|
||||
url_to_asset_id[url] = aid
|
||||
except Exception as e:
|
||||
logger.warning("获取素材URL失败: asset_id=%s error=%s", aid, str(e))
|
||||
|
||||
if not video_urls:
|
||||
logger.info("无可用视频素材,跳过视频理解分析")
|
||||
return {}
|
||||
|
||||
# 调用 MediaKit 视频理解
|
||||
prompt = (
|
||||
"请简要描述这段视频的主要内容,包括:场景(室内/室外/具体场所)、"
|
||||
"主体(人物/物体/动物)、动作/活动、氛围/情绪、主要色调。"
|
||||
"控制在100字以内。"
|
||||
)
|
||||
|
||||
contents = client.analyze_videos(
|
||||
video_urls=video_urls,
|
||||
prompt=prompt,
|
||||
level="Economy",
|
||||
)
|
||||
|
||||
if not contents:
|
||||
logger.warning("MediaKit 视频理解未返回结果")
|
||||
return {}
|
||||
|
||||
# 将结果映射回 asset_id
|
||||
analyses: dict[str, str] = {}
|
||||
for i, content in enumerate(contents):
|
||||
if i < len(video_urls) and content:
|
||||
asset_id = url_to_asset_id.get(video_urls[i])
|
||||
if asset_id:
|
||||
analyses[asset_id] = content
|
||||
|
||||
logger.info(
|
||||
"MediaKit 视频理解完成: total=%d analyzed=%d",
|
||||
len(video_urls),
|
||||
len(analyses),
|
||||
)
|
||||
return analyses
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("MediaKit 视频理解异常,将降级到无分析模式: %s", str(e))
|
||||
return {}
|
||||
|
||||
|
||||
@router.post("/ai-recommend", response_model=AIRecommendResponse)
|
||||
def editor_ai_recommend(
|
||||
template_id: str,
|
||||
@@ -46,12 +131,16 @@ def editor_ai_recommend(
|
||||
|
||||
from packages.shared.ai_service import run_ai_recommend
|
||||
|
||||
# 调用 MediaKit 视频理解,获取素材内容分析
|
||||
asset_analyses = _build_asset_analyses(body.asset_ids, db)
|
||||
|
||||
result = run_ai_recommend(
|
||||
plan_id=plan_id,
|
||||
template_id=plan.template_id,
|
||||
asset_ids=body.asset_ids,
|
||||
editing_mode=body.editing_mode,
|
||||
target_duration=body.target_duration,
|
||||
asset_analyses=asset_analyses,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -196,20 +196,54 @@ def _call_ai_recommend_service(
|
||||
asset_ids: List[str],
|
||||
editing_mode: str,
|
||||
target_duration: float,
|
||||
asset_analyses: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""调用 AI 推荐服务生成片段编排方案.
|
||||
|
||||
优先使用豆包大模型生成,失败或未配置时降级为本地规则生成。
|
||||
当提供 asset_analyses 时,会将每个素材的视频理解结果注入 prompt,
|
||||
让 LLM 能基于视频实际内容做智能编排。
|
||||
|
||||
Args:
|
||||
plan_id: 剪辑计划 ID
|
||||
template_id: 模板 ID
|
||||
asset_ids: 素材 ID 列表
|
||||
editing_mode: 剪辑模式
|
||||
target_duration: 目标时长(秒)
|
||||
asset_analyses: 可选,{asset_id: 视频理解文本} 映射
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.info("豆包API未配置,使用本地降级生成AI推荐方案")
|
||||
return _fallback_recommend_clips(plan_id, template_id, asset_ids, editing_mode, target_duration)
|
||||
|
||||
# 构建素材描述(含视频理解结果)
|
||||
asset_analyses = asset_analyses or {}
|
||||
asset_lines = []
|
||||
for aid in asset_ids[:30]:
|
||||
analysis = asset_analyses.get(aid, "")
|
||||
if analysis:
|
||||
# 截断过长的分析结果,避免 token 爆炸
|
||||
analysis_truncated = analysis[:300] + ("..." if len(analysis) > 300 else "")
|
||||
asset_lines.append(f" - 素材ID: {aid}\n 内容描述: {analysis_truncated}")
|
||||
else:
|
||||
asset_lines.append(f" - 素材ID: {aid}")
|
||||
|
||||
assets_desc = "\n".join(asset_lines)
|
||||
has_analysis = any(aid in asset_analyses for aid in asset_ids[:30])
|
||||
|
||||
# 构建 prompt
|
||||
system_prompt = (
|
||||
"你是一个专业的视频剪辑导演助手。"
|
||||
"根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
"你是一个专业的视频剪辑导演助手。" "根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
)
|
||||
if has_analysis:
|
||||
system_prompt += (
|
||||
"每个素材附带了 AI 视频理解的内容描述,请根据素材的实际内容来决策编排:\n"
|
||||
"- 将内容相关的素材放在一起,保持叙事连贯\n"
|
||||
"- 根据素材内容合理安排片段顺序(如开场用吸引人的画面、高潮部分紧凑切换等)\n"
|
||||
"- 为每个片段选择最匹配的素材,并在 text_content 中体现素材主题\n"
|
||||
)
|
||||
system_prompt += (
|
||||
"要求:\n"
|
||||
"1. 片段类型分为三类:intro(开场)、showcase(展示)、outro(结尾)\n"
|
||||
"2. 每个片段包含:clip_type、order、text_content(字幕/标题文字)、"
|
||||
@@ -221,7 +255,6 @@ def _call_ai_recommend_service(
|
||||
'返回格式:{"clips": [...], "title": "视频标题", "confidence": 0.85}'
|
||||
)
|
||||
|
||||
assets_desc = "\n".join([f" - 素材ID: {aid}" for i, aid in enumerate(asset_ids[:30])])
|
||||
user_prompt = (
|
||||
f"剪辑计划ID: {plan_id}\n"
|
||||
f"模板ID: {template_id}\n"
|
||||
@@ -246,11 +279,12 @@ def _call_ai_recommend_service(
|
||||
parsed = _parse_recommend_response(result, asset_ids, target_duration)
|
||||
if parsed and len(parsed["clips"]) >= 2:
|
||||
logger.info(
|
||||
"豆包AI推荐生成成功: plan_id=%s clips=%d duration=%.1f confidence=%.2f",
|
||||
"豆包AI推荐生成成功: plan_id=%s clips=%d duration=%.1f confidence=%.2f has_analysis=%s",
|
||||
plan_id,
|
||||
len(parsed["clips"]),
|
||||
parsed["total_duration"],
|
||||
parsed["confidence"],
|
||||
has_analysis,
|
||||
)
|
||||
return parsed
|
||||
logger.warning("豆包AI推荐返回解析失败,降级到本地方案: %s", result[:100])
|
||||
@@ -354,6 +388,7 @@ def run_ai_recommend(
|
||||
asset_ids: List[str],
|
||||
editing_mode: str = "one_take",
|
||||
target_duration: float = 30.0,
|
||||
asset_analyses: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""执行 AI 推荐片段方案
|
||||
|
||||
@@ -363,17 +398,19 @@ def run_ai_recommend(
|
||||
asset_ids: 素材 ID 列表
|
||||
editing_mode: 剪辑模式 (one_take / pip / voice_over / voice_pip)
|
||||
target_duration: 目标时长(秒)
|
||||
asset_analyses: 可选,{asset_id: 视频理解文本} 映射
|
||||
|
||||
Returns:
|
||||
推荐方案 dict,包含 clips / config / total_duration / confidence
|
||||
"""
|
||||
logger.info(
|
||||
"AI 推荐片段方案: plan_id=%s template_id=%s assets=%d mode=%s duration=%.1f",
|
||||
"AI 推荐片段方案: plan_id=%s template_id=%s assets=%d mode=%s duration=%.1f has_analysis=%s",
|
||||
plan_id,
|
||||
template_id,
|
||||
len(asset_ids),
|
||||
editing_mode,
|
||||
target_duration,
|
||||
bool(asset_analyses),
|
||||
)
|
||||
result = _call_ai_recommend_service(
|
||||
plan_id=plan_id,
|
||||
@@ -381,6 +418,7 @@ def run_ai_recommend(
|
||||
asset_ids=asset_ids,
|
||||
editing_mode=editing_mode,
|
||||
target_duration=target_duration,
|
||||
asset_analyses=asset_analyses,
|
||||
)
|
||||
logger.info(
|
||||
"AI 推荐完成: plan_id=%s clips=%d duration=%.1f confidence=%.2f",
|
||||
|
||||
@@ -115,13 +115,153 @@ class MediaKitClient:
|
||||
logger.exception("MediaKit 抽帧任务提交异常: %s", str(e))
|
||||
return None
|
||||
|
||||
def analyze_videos(
|
||||
self,
|
||||
video_urls: List[str],
|
||||
prompt: str,
|
||||
level: str = "Economy",
|
||||
poll_interval: float = 3.0,
|
||||
max_poll_attempts: int = 60,
|
||||
) -> Optional[List[str]]:
|
||||
"""调用 MediaKit 视频理解智能策略 API.
|
||||
|
||||
基于火山方舟视觉大模型,对输入的视频 URL 列表进行内容分析,
|
||||
返回每个视频的自然语言描述(用于智能选片段等场景)。
|
||||
|
||||
Args:
|
||||
video_urls: 视频 URL 列表(最多 10 个,需公网可访问)
|
||||
prompt: 指导大模型分析的自然语言指令
|
||||
level: 分析档位 Economy / Balanced / Quality
|
||||
poll_interval: 轮询间隔(秒)
|
||||
max_poll_attempts: 最大轮询次数
|
||||
|
||||
Returns:
|
||||
分析结果列表,每个元素对应 video_urls 中同索引视频的分析文本。
|
||||
失败返回 None。
|
||||
"""
|
||||
if not self.is_available:
|
||||
logger.warning("MediaKit 未配置,跳过视频理解")
|
||||
return None
|
||||
|
||||
if not video_urls:
|
||||
return None
|
||||
|
||||
task_id = self._submit_video_understand_task(video_urls, prompt, level)
|
||||
if not task_id:
|
||||
return None
|
||||
|
||||
return self._poll_video_understand_result(task_id, poll_interval, max_poll_attempts)
|
||||
|
||||
def _submit_video_understand_task(
|
||||
self,
|
||||
video_urls: List[str],
|
||||
prompt: str,
|
||||
level: str,
|
||||
) -> Optional[str]:
|
||||
"""提交视频理解任务,返回 task_id."""
|
||||
url = f"{self.base_url}/tools/video-understand-router"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
payload = {
|
||||
"video_urls": video_urls,
|
||||
"prompt": prompt,
|
||||
"level": level,
|
||||
}
|
||||
|
||||
try:
|
||||
response = httpx.post(url, headers=headers, json=payload, timeout=self.timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
task_id = data.get("task_id")
|
||||
if not task_id:
|
||||
logger.error("MediaKit 视频理解任务提交失败: 无 task_id. response=%s", data)
|
||||
return None
|
||||
|
||||
logger.info(
|
||||
"MediaKit 视频理解任务已提交: task_id=%s videos=%d level=%s",
|
||||
task_id,
|
||||
len(video_urls),
|
||||
level,
|
||||
)
|
||||
return task_id
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("MediaKit 视频理解任务提交异常: %s", str(e))
|
||||
return None
|
||||
|
||||
def _poll_video_understand_result(
|
||||
self,
|
||||
task_id: str,
|
||||
poll_interval: float,
|
||||
max_poll_attempts: int,
|
||||
) -> Optional[List[str]]:
|
||||
"""轮询视频理解任务结果,返回 contents 列表."""
|
||||
url = f"{self.base_url}/tasks/{task_id}"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
}
|
||||
|
||||
for attempt in range(max_poll_attempts):
|
||||
try:
|
||||
response = httpx.get(url, headers=headers, timeout=self.timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
status = data.get("status")
|
||||
if status in ("completed", "success"):
|
||||
result = data.get("result", {})
|
||||
contents = result.get("contents", [])
|
||||
logger.info(
|
||||
"MediaKit 视频理解完成: task_id=%s videos=%d",
|
||||
task_id,
|
||||
len(contents),
|
||||
)
|
||||
return contents if contents else None
|
||||
|
||||
elif status == "failed":
|
||||
error_msg = data.get("error", "unknown error")
|
||||
logger.error(
|
||||
"MediaKit 视频理解任务失败: task_id=%s error=%s",
|
||||
task_id,
|
||||
error_msg,
|
||||
)
|
||||
return None
|
||||
|
||||
logger.debug(
|
||||
"MediaKit 视频理解进行中: task_id=%s status=%s attempt=%d/%d",
|
||||
task_id,
|
||||
status,
|
||||
attempt + 1,
|
||||
max_poll_attempts,
|
||||
)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
"MediaKit 视频理解轮询异常: task_id=%s error=%s",
|
||||
task_id,
|
||||
str(e),
|
||||
)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
logger.error(
|
||||
"MediaKit 视频理解超时: task_id=%s max_attempts=%d",
|
||||
task_id,
|
||||
max_poll_attempts,
|
||||
)
|
||||
return None
|
||||
|
||||
def _poll_task_result(
|
||||
self,
|
||||
task_id: str,
|
||||
poll_interval: float,
|
||||
max_poll_attempts: int,
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""轮询任务状态,返回结果."""
|
||||
"""轮询抽帧任务状态,返回结果."""
|
||||
url = f"{self.base_url}/tasks/{task_id}"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
|
||||
Executable
+436
@@ -0,0 +1,436 @@
|
||||
"""#1209 MediaKit 视频理解 + AI 智能选片段集成测试.
|
||||
|
||||
覆盖范围:
|
||||
1. MediaKitClient.analyze_videos — 正常流程、不可用降级、空输入
|
||||
2. _call_ai_recommend_service 带 asset_analyses — prompt 注入验证
|
||||
3. _build_asset_analyses — API 路由层集成逻辑
|
||||
4. run_ai_recommend — asset_analyses 透传
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
# ── MediaKitClient.analyze_videos ─────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestMediaKitClientAnalyzeVideos:
|
||||
"""MediaKitClient.analyze_videos 单元测试."""
|
||||
|
||||
def _make_client(self, api_key: str = "test-key") -> object:
|
||||
from packages.shared.mediakit_client import MediaKitClient
|
||||
|
||||
with patch("packages.shared.mediakit_client.get_shared_settings") as mock_settings:
|
||||
mock_settings.return_value = MagicMock(
|
||||
mediakit_api_key=api_key,
|
||||
mediakit_base_url="https://mock.mediakit.com/api/v1",
|
||||
mediakit_timeout=30,
|
||||
)
|
||||
return MediaKitClient()
|
||||
|
||||
def test_analyze_videos_not_available(self):
|
||||
"""未配置 API Key 时返回 None."""
|
||||
client = self._make_client(api_key="")
|
||||
result = client.analyze_videos(
|
||||
video_urls=["https://example.com/video.mp4"],
|
||||
prompt="describe this video",
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_analyze_videos_empty_urls(self):
|
||||
"""空 URL 列表返回 None."""
|
||||
client = self._make_client()
|
||||
result = client.analyze_videos(video_urls=[], prompt="describe")
|
||||
assert result is None
|
||||
|
||||
@patch("packages.shared.mediakit_client.httpx.post")
|
||||
@patch("packages.shared.mediakit_client.httpx.get")
|
||||
def test_analyze_videos_success(self, mock_get, mock_post):
|
||||
"""正常提交任务并获取结果."""
|
||||
# 提交任务返回 task_id
|
||||
mock_post.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"success": True, "task_id": "test-task-123"},
|
||||
)
|
||||
mock_post.return_value.raise_for_status = MagicMock()
|
||||
|
||||
# 轮询返回 completed
|
||||
mock_get.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {
|
||||
"success": True,
|
||||
"task_id": "test-task-123",
|
||||
"status": "completed",
|
||||
"result": {
|
||||
"duration": 30.5,
|
||||
"contents": [
|
||||
"视频展示了城市日景,包含车流和行人,氛围繁忙。",
|
||||
"视频展示了夜晚霓虹灯特写,色调偏暖。",
|
||||
],
|
||||
},
|
||||
},
|
||||
)
|
||||
mock_get.return_value.raise_for_status = MagicMock()
|
||||
|
||||
client = self._make_client()
|
||||
result = client.analyze_videos(
|
||||
video_urls=["https://example.com/v1.mp4", "https://example.com/v2.mp4"],
|
||||
prompt="描述视频内容",
|
||||
level="Economy",
|
||||
poll_interval=0.01, # 测试用快速轮询
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert len(result) == 2
|
||||
assert "城市日景" in result[0]
|
||||
assert "霓虹灯" in result[1]
|
||||
|
||||
# 验证提交参数
|
||||
call_args = mock_post.call_args
|
||||
assert "/tools/video-understand-router" in call_args.args[0]
|
||||
payload = call_args.kwargs["json"]
|
||||
assert payload["video_urls"] == ["https://example.com/v1.mp4", "https://example.com/v2.mp4"]
|
||||
assert payload["prompt"] == "描述视频内容"
|
||||
assert payload["level"] == "Economy"
|
||||
|
||||
@patch("packages.shared.mediakit_client.httpx.post")
|
||||
def test_analyze_videos_submit_failure(self, mock_post):
|
||||
"""提交任务失败返回 None."""
|
||||
mock_post.side_effect = Exception("Network error")
|
||||
|
||||
client = self._make_client()
|
||||
result = client.analyze_videos(
|
||||
video_urls=["https://example.com/v1.mp4"],
|
||||
prompt="describe",
|
||||
poll_interval=0.01,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@patch("packages.shared.mediakit_client.httpx.post")
|
||||
@patch("packages.shared.mediakit_client.httpx.get")
|
||||
def test_analyze_videos_task_failed(self, mock_get, mock_post):
|
||||
"""任务状态为 failed 时返回 None."""
|
||||
mock_post.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"task_id": "task-fail"},
|
||||
)
|
||||
mock_post.return_value.raise_for_status = MagicMock()
|
||||
|
||||
mock_get.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"status": "failed", "error": "model timeout"},
|
||||
)
|
||||
mock_get.return_value.raise_for_status = MagicMock()
|
||||
|
||||
client = self._make_client()
|
||||
result = client.analyze_videos(
|
||||
video_urls=["https://example.com/v.mp4"],
|
||||
prompt="describe",
|
||||
poll_interval=0.01,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@patch("packages.shared.mediakit_client.httpx.post")
|
||||
@patch("packages.shared.mediakit_client.httpx.get")
|
||||
def test_analyze_videos_timeout(self, mock_get, mock_post):
|
||||
"""超过最大轮询次数返回 None."""
|
||||
mock_post.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"task_id": "task-slow"},
|
||||
)
|
||||
mock_post.return_value.raise_for_status = MagicMock()
|
||||
|
||||
# 一直返回 processing
|
||||
mock_get.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"status": "processing"},
|
||||
)
|
||||
mock_get.return_value.raise_for_status = MagicMock()
|
||||
|
||||
client = self._make_client()
|
||||
result = client.analyze_videos(
|
||||
video_urls=["https://example.com/v.mp4"],
|
||||
prompt="describe",
|
||||
poll_interval=0.001,
|
||||
max_poll_attempts=3,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
# ── _call_ai_recommend_service with asset_analyses ────────────────────────────
|
||||
|
||||
|
||||
class TestCallAiRecommendWithAnalysis:
|
||||
"""测试 _call_ai_recommend_service 带 asset_analyses 的行为."""
|
||||
|
||||
@patch("packages.shared.ai_service.get_doubao_client")
|
||||
def test_analysis_included_in_prompt(self, mock_get_client):
|
||||
"""有分析结果时,prompt 包含视频内容描述."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
# 模拟豆包返回一个有效 JSON
|
||||
mock_client.chat_completion.return_value = (
|
||||
'{"clips": [{"clip_type": "intro", "order": 0, "text_content": "开场",'
|
||||
'"duration": 3.0, "transition_effect": "fade", "asset_id": "asset1",'
|
||||
'"start_time": 0.0, "config": {}}, {"clip_type": "outro", "order": 1,'
|
||||
'"text_content": "结尾", "duration": 3.0, "transition_effect": "fade",'
|
||||
'"asset_id": "", "start_time": 0.0, "config": {}}],'
|
||||
'"title": "测试视频", "confidence": 0.9}'
|
||||
)
|
||||
|
||||
from packages.shared.ai_service import _call_ai_recommend_service
|
||||
|
||||
result = _call_ai_recommend_service(
|
||||
plan_id="plan-1",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["asset1", "asset2"],
|
||||
editing_mode="one_take",
|
||||
target_duration=30.0,
|
||||
asset_analyses={
|
||||
"asset1": "室内场景,一位女性在桌前讲解产品,氛围轻松专业",
|
||||
"asset2": "室外公园,阳光充足,有孩子在玩耍",
|
||||
},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 2
|
||||
|
||||
# 验证 prompt 中包含了视频分析内容
|
||||
messages = mock_client.chat_completion.call_args.kwargs["messages"]
|
||||
system_prompt = messages[0]["content"]
|
||||
assert "叙事连贯" in system_prompt # 有分析时的额外指导
|
||||
|
||||
user_prompt = messages[1]["content"]
|
||||
assert "室内场景" in user_prompt
|
||||
assert "室外公园" in user_prompt
|
||||
|
||||
@patch("packages.shared.ai_service.get_doubao_client")
|
||||
def test_no_analysis_basic_prompt(self, mock_get_client):
|
||||
"""无分析结果时,prompt 保持基本格式."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
mock_client.chat_completion.return_value = (
|
||||
'{"clips": [{"clip_type": "showcase", "order": 0, "text_content": "展示",'
|
||||
'"duration": 5.0, "transition_effect": "cut", "asset_id": "a1",'
|
||||
'"start_time": 0.0, "config": {}}],'
|
||||
'"title": "简单视频", "confidence": 0.8}'
|
||||
)
|
||||
|
||||
from packages.shared.ai_service import _call_ai_recommend_service
|
||||
|
||||
result = _call_ai_recommend_service(
|
||||
plan_id="plan-2",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["a1"],
|
||||
editing_mode="one_take",
|
||||
target_duration=10.0,
|
||||
# 不传 asset_analyses
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
|
||||
# 验证 prompt 中不含叙事连贯等指导
|
||||
messages = mock_client.chat_completion.call_args.kwargs["messages"]
|
||||
system_prompt = messages[0]["content"]
|
||||
assert "叙事连贯" not in system_prompt
|
||||
|
||||
@patch("packages.shared.ai_service.get_doubao_client")
|
||||
def test_partial_analysis_only_some_assets(self, mock_get_client):
|
||||
"""部分素材有分析结果时,只有被分析的素材包含内容描述."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
mock_client.chat_completion.return_value = (
|
||||
'{"clips": [{"clip_type": "showcase", "order": 0, "text_content": "展示",'
|
||||
'"duration": 5.0, "transition_effect": "cut", "asset_id": "a1",'
|
||||
'"start_time": 0.0, "config": {}}],'
|
||||
'"title": "部分分析", "confidence": 0.85}'
|
||||
)
|
||||
|
||||
from packages.shared.ai_service import _call_ai_recommend_service
|
||||
|
||||
_call_ai_recommend_service(
|
||||
plan_id="plan-3",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["a1", "a2", "a3"],
|
||||
editing_mode="one_take",
|
||||
target_duration=15.0,
|
||||
asset_analyses={
|
||||
"a1": "海边日落,金色阳光", # 只有 a1 有分析
|
||||
},
|
||||
)
|
||||
|
||||
messages = mock_client.chat_completion.call_args.kwargs["messages"]
|
||||
user_prompt = messages[1]["content"]
|
||||
assert "海边日落" in user_prompt
|
||||
# a2 和 a3 应该只有 ID,没有内容描述
|
||||
assert "素材ID: a2" in user_prompt
|
||||
assert "素材ID: a3" in user_prompt
|
||||
|
||||
|
||||
# ── run_ai_recommend with asset_analyses ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestRunAiRecommendWithAnalysis:
|
||||
"""run_ai_recommend 透传 asset_analyses."""
|
||||
|
||||
@patch("packages.shared.ai_service._call_ai_recommend_service")
|
||||
def test_passes_through_analysis(self, mock_call):
|
||||
"""asset_analyses 正确传递给底层服务."""
|
||||
from packages.shared.ai_service import run_ai_recommend
|
||||
|
||||
mock_call.return_value = {
|
||||
"clips": [
|
||||
{
|
||||
"clip_type": "showcase",
|
||||
"order": 0,
|
||||
"text_content": "t",
|
||||
"duration": 3.0,
|
||||
"transition_effect": "cut",
|
||||
"asset_id": "a1",
|
||||
"start_time": 0.0,
|
||||
"config": {},
|
||||
}
|
||||
],
|
||||
"config": {"title": {"text": "test", "ai_auto": True}},
|
||||
"total_duration": 3.0,
|
||||
"confidence": 0.85,
|
||||
}
|
||||
|
||||
analyses = {"a1": "海边日落场景"}
|
||||
run_ai_recommend(
|
||||
plan_id="plan-1",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["a1"],
|
||||
editing_mode="one_take",
|
||||
target_duration=10.0,
|
||||
asset_analyses=analyses,
|
||||
)
|
||||
|
||||
# 验证传递
|
||||
call_kwargs = mock_call.call_args.kwargs
|
||||
assert call_kwargs["asset_analyses"] == analyses
|
||||
|
||||
@patch("packages.shared.ai_service._call_ai_recommend_service")
|
||||
def test_default_no_analysis(self, mock_call):
|
||||
"""默认不传 asset_analyses 时为 None."""
|
||||
from packages.shared.ai_service import run_ai_recommend
|
||||
|
||||
mock_call.return_value = {
|
||||
"clips": [
|
||||
{
|
||||
"clip_type": "showcase",
|
||||
"order": 0,
|
||||
"text_content": "t",
|
||||
"duration": 3.0,
|
||||
"transition_effect": "cut",
|
||||
"asset_id": "a1",
|
||||
"start_time": 0.0,
|
||||
"config": {},
|
||||
}
|
||||
],
|
||||
"config": {"title": {"text": "test", "ai_auto": True}},
|
||||
"total_duration": 3.0,
|
||||
"confidence": 0.85,
|
||||
}
|
||||
|
||||
run_ai_recommend(
|
||||
plan_id="plan-1",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["a1"],
|
||||
)
|
||||
|
||||
call_kwargs = mock_call.call_args.kwargs
|
||||
assert call_kwargs["asset_analyses"] is None
|
||||
|
||||
|
||||
# ── _build_asset_analyses (API route helper) ─────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildAssetAnalyses:
|
||||
"""_build_asset_analyses 集成逻辑测试."""
|
||||
|
||||
def test_empty_asset_ids(self):
|
||||
"""空素材列表返回空 dict."""
|
||||
from apps.api.app.api.routes.templates_editor.ai_features import (
|
||||
_build_asset_analyses,
|
||||
)
|
||||
|
||||
result = _build_asset_analyses([], MagicMock())
|
||||
assert result == {}
|
||||
|
||||
@patch("packages.shared.mediakit_client.get_mediakit_client")
|
||||
def test_mediakit_not_available(self, mock_get_client):
|
||||
"""MediaKit 不可用时返回空 dict."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
# 需要 patch 在 packages.shared.mediakit_client 层面,因为 _build_asset_analyses 内部 from import
|
||||
from apps.api.app.api.routes.templates_editor.ai_features import (
|
||||
_build_asset_analyses,
|
||||
)
|
||||
|
||||
result = _build_asset_analyses(["asset1"], MagicMock())
|
||||
assert result == {}
|
||||
|
||||
def test_exception_returns_empty(self):
|
||||
"""任何异常都返回空 dict,不阻塞主流程."""
|
||||
from apps.api.app.api.routes.templates_editor.ai_features import (
|
||||
_build_asset_analyses,
|
||||
)
|
||||
|
||||
# 传一个 mock db,让内部自然失败
|
||||
mock_db = MagicMock()
|
||||
mock_db.side_effect = None # db 本身不抛异常,但内部操作会失败
|
||||
|
||||
result = _build_asset_analyses(["nonexistent-asset"], mock_db)
|
||||
assert result == {}
|
||||
|
||||
|
||||
# ── 长文本截断验证 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAnalysisTruncation:
|
||||
"""验证过长分析文本被截断以避免 token 爆炸."""
|
||||
|
||||
@patch("packages.shared.ai_service.get_doubao_client")
|
||||
def test_long_analysis_truncated(self, mock_get_client):
|
||||
"""超过 300 字的分析文本被截断并加 ...."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
mock_client.chat_completion.return_value = (
|
||||
'{"clips": [{"clip_type": "showcase", "order": 0, "text_content": "t",'
|
||||
'"duration": 3.0, "transition_effect": "cut", "asset_id": "a1",'
|
||||
'"start_time": 0.0, "config": {}}],'
|
||||
'"title": "test", "confidence": 0.8}'
|
||||
)
|
||||
|
||||
from packages.shared.ai_service import _call_ai_recommend_service
|
||||
|
||||
long_analysis = "A" * 500 # 500 字符
|
||||
|
||||
_call_ai_recommend_service(
|
||||
plan_id="plan-1",
|
||||
template_id="tmpl-1",
|
||||
asset_ids=["a1"],
|
||||
editing_mode="one_take",
|
||||
target_duration=10.0,
|
||||
asset_analyses={"a1": long_analysis},
|
||||
)
|
||||
|
||||
messages = mock_client.chat_completion.call_args.kwargs["messages"]
|
||||
user_prompt = messages[1]["content"]
|
||||
# 截断后应该包含 ...
|
||||
assert "..." in user_prompt
|
||||
# 原始 500 字符不应完整出现
|
||||
assert "A" * 500 not in user_prompt
|
||||
Reference in New Issue
Block a user