feat(#1893): scripts AI capability (douyin extract/rewrite/titles) #1930
@@ -21,6 +21,7 @@ from app.api.routes.lipsync import router as lipsync_router
|
||||
from app.api.routes.points import points_router, usage_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
from app.api.routes.scripts import router as scripts_router
|
||||
from app.api.routes.scripts_ai import router as scripts_ai_router
|
||||
from app.api.routes.share import router as share_router
|
||||
from app.api.routes.subscription import router as subscription_router
|
||||
from app.api.routes.tags import router as tags_router
|
||||
@@ -190,6 +191,11 @@ api_router.include_router(
|
||||
prefix="/scripts",
|
||||
tags=["ScriptLibrary"],
|
||||
)
|
||||
api_router.include_router(
|
||||
scripts_ai_router,
|
||||
prefix="/scripts",
|
||||
tags=["ScriptLibrary AI"],
|
||||
)
|
||||
api_router.include_router(
|
||||
ai_avatar_render_router,
|
||||
prefix="/ai-avatar/render",
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
"""Scripts AI 能力路由 — Issue #1893.
|
||||
|
||||
三个 AI 工具接口(均挂载在 /api/v1/scripts 前缀下):
|
||||
- POST /extract-from-douyin 从抖音视频提取文案(yt-dlp 下载 + ASR 转写)
|
||||
- POST /ai-rewrite AI 文案改写(复用豆包 LLM)
|
||||
- POST /ai-generate-titles AI 标题生成(复用 generate_smart_titles)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import tempfile
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.schemas.scripts_ai import (
|
||||
AiGenerateTitlesRequest,
|
||||
AiGenerateTitlesResponse,
|
||||
AiRewriteRequest,
|
||||
AiRewriteResponse,
|
||||
ExtractFromDouyinRequest,
|
||||
ExtractFromDouyinResponse,
|
||||
)
|
||||
from app.services.script_asr_service import (
|
||||
ASRNotConfiguredError,
|
||||
ASRTranscriptionError,
|
||||
transcribe_to_text,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 抖音 URL 校验:支持短链 v.douyin.com 和长链 www.douyin.com/video/
|
||||
_DOUYIN_URL_RE = re.compile(
|
||||
r"^(https?://)?(v\.douyin\.com/\S+|www\.douyin\.com/video/\S+)$",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _validate_douyin_url(url: str) -> None:
|
||||
"""校验抖音 URL 格式,不合法时抛 HTTPException(400)."""
|
||||
if not url or not url.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="链接不能为空",
|
||||
)
|
||||
if not _DOUYIN_URL_RE.match(url.strip()):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无效的抖音链接,仅支持 v.douyin.com 短链或 www.douyin.com/video/ 长链",
|
||||
)
|
||||
|
||||
|
||||
# ── 1. 从抖音视频提取文案 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/extract-from-douyin",
|
||||
response_model=ExtractFromDouyinResponse,
|
||||
)
|
||||
def extract_from_douyin(
|
||||
request: ExtractFromDouyinRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ExtractFromDouyinResponse:
|
||||
"""从抖音视频下载无水印视频并通过 ASR 提取文案."""
|
||||
source_url = request.url.strip()
|
||||
_validate_douyin_url(source_url)
|
||||
|
||||
# 确保 URL 有 scheme(yt-dlp 需要完整 URL)
|
||||
url_for_download = source_url
|
||||
if not re.match(r"^https?://", url_for_download, re.IGNORECASE):
|
||||
url_for_download = "https://" + url_for_download
|
||||
|
||||
# 使用临时目录下载视频,退出时自动清理
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="douyin_extract_") as temp_dir:
|
||||
import yt_dlp
|
||||
|
||||
ydl_opts = {
|
||||
"format": "best[ext=mp4]/best",
|
||||
"outtmpl": f"{temp_dir}/%(id)s.%(ext)s",
|
||||
"quiet": True,
|
||||
"no_warnings": True,
|
||||
"noplaylist": True,
|
||||
}
|
||||
|
||||
try:
|
||||
ydl = yt_dlp.YoutubeDL(ydl_opts)
|
||||
info = ydl.extract_info(url_for_download, download=True)
|
||||
except Exception as exc:
|
||||
logger.error("抖音视频下载失败: url=%s error=%s", source_url, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"视频下载失败: {exc}",
|
||||
) from exc
|
||||
|
||||
if info is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无法解析该抖音链接",
|
||||
)
|
||||
|
||||
video_path = ydl.prepare_filename(info)
|
||||
duration = float(info.get("duration") or 0)
|
||||
|
||||
# ASR 转写
|
||||
try:
|
||||
text = transcribe_to_text(video_path)
|
||||
except ASRNotConfiguredError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
except ASRTranscriptionError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
|
||||
return ExtractFromDouyinResponse(
|
||||
text=text,
|
||||
duration_seconds=duration,
|
||||
source_url=source_url,
|
||||
)
|
||||
|
||||
|
||||
# ── 2. AI 文案改写 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/ai-rewrite",
|
||||
response_model=AiRewriteResponse,
|
||||
)
|
||||
def ai_rewrite(
|
||||
request: AiRewriteRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> AiRewriteResponse:
|
||||
"""使用豆包大模型改写文案."""
|
||||
content = (request.content or "").strip()
|
||||
if not content:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文案内容不能为空",
|
||||
)
|
||||
|
||||
style = request.style or "口语化"
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="AI 服务不可用,请联系管理员配置豆包大模型 API Key",
|
||||
)
|
||||
|
||||
system_prompt = (
|
||||
"你是一个专业的短视频文案改写专家。请对以下文案进行改写,"
|
||||
"要求:保留原意、口语化、适合短视频口播、调整语序避免查重。"
|
||||
)
|
||||
if style:
|
||||
system_prompt += f"\n风格要求:{style}"
|
||||
|
||||
user_prompt = f"请改写以下文案:\n\n{content}"
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
]
|
||||
|
||||
try:
|
||||
rewritten = client.chat_completion(
|
||||
messages=messages,
|
||||
temperature=0.8,
|
||||
max_tokens=2048,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("AI 改写调用失败: %s", exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"AI 改写失败: {exc}",
|
||||
) from exc
|
||||
|
||||
if not rewritten:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="AI 改写未返回有效结果",
|
||||
)
|
||||
|
||||
return AiRewriteResponse(
|
||||
original=content,
|
||||
rewritten=rewritten.strip(),
|
||||
style=style,
|
||||
)
|
||||
|
||||
|
||||
# ── 3. AI 标题生成 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/ai-generate-titles",
|
||||
response_model=AiGenerateTitlesResponse,
|
||||
)
|
||||
def ai_generate_titles(
|
||||
request: AiGenerateTitlesRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> AiGenerateTitlesResponse:
|
||||
"""使用现有 generate_smart_titles 生成标题."""
|
||||
content = (request.content or "").strip()
|
||||
if not content:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文案内容不能为空",
|
||||
)
|
||||
|
||||
# count 限制在 1-5(Pydantic ge=1 le=5 已校验),但为兼容直接调用场景截断
|
||||
count = max(1, min(5, request.count))
|
||||
|
||||
from app.services.ai_service import generate_smart_titles
|
||||
|
||||
result = generate_smart_titles(
|
||||
description=content,
|
||||
style="viral",
|
||||
count=count,
|
||||
)
|
||||
|
||||
titles = result.get("titles", [])[:count]
|
||||
|
||||
return AiGenerateTitlesResponse(titles=titles)
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Scripts AI 能力 Pydantic schemas — Issue #1893.
|
||||
|
||||
抖音文案提取、AI 改写、AI 标题生成的请求/响应模型。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ── 抖音文案提取 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ExtractFromDouyinRequest(BaseModel):
|
||||
"""从抖音视频提取文案请求."""
|
||||
|
||||
url: str = Field(..., description="抖音视频链接(短链或长链)")
|
||||
|
||||
|
||||
class ExtractFromDouyinResponse(BaseModel):
|
||||
"""从抖音视频提取文案响应."""
|
||||
|
||||
text: str = Field(..., description="ASR 识别出的文案文本")
|
||||
duration_seconds: float = Field(..., description="视频时长(秒)")
|
||||
source_url: str = Field(..., description="原始视频链接")
|
||||
|
||||
|
||||
# ── AI 改写 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AiRewriteRequest(BaseModel):
|
||||
"""AI 文案改写请求."""
|
||||
|
||||
content: str = Field(..., description="原文内容")
|
||||
style: Optional[str] = Field("口语化", description="改写风格,如 口语化/正式/活泼")
|
||||
|
||||
|
||||
class AiRewriteResponse(BaseModel):
|
||||
"""AI 文案改写响应."""
|
||||
|
||||
original: str = Field(..., description="原文")
|
||||
rewritten: str = Field(..., description="改写后的文案")
|
||||
style: str = Field(..., description="使用的改写风格")
|
||||
|
||||
|
||||
# ── AI 标题生成 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AiGenerateTitlesRequest(BaseModel):
|
||||
"""AI 标题生成请求."""
|
||||
|
||||
content: str = Field(..., description="文案内容")
|
||||
count: int = Field(3, ge=1, le=5, description="生成标题数量(1-5,默认3)")
|
||||
|
||||
|
||||
class AiGenerateTitlesResponse(BaseModel):
|
||||
"""AI 标题生成响应."""
|
||||
|
||||
titles: List[str] = Field(..., description="生成的标题列表")
|
||||
@@ -0,0 +1,59 @@
|
||||
"""文案提取 ASR 服务封装 — Issue #1893.
|
||||
|
||||
将已有的 ASR 服务工厂封装为面向文案提取场景的简单接口:
|
||||
- transcribe_to_text(video_path) -> str:将视频/音频转写为纯文本
|
||||
- 未配置 ASR 时抛 ASRNotConfiguredError(路由层映射为 503)
|
||||
- ASR 调用失败时抛 ASRTranscriptionError(路由层映射为 502)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from packages.ports.asr_service import ASRServiceError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ASRNotConfiguredError(Exception):
|
||||
"""ASR 服务未配置."""
|
||||
|
||||
|
||||
class ASRTranscriptionError(Exception):
|
||||
"""ASR 转写失败."""
|
||||
|
||||
|
||||
def transcribe_to_text(media_path: str | Path) -> str:
|
||||
"""将视频/音频文件转写为纯文本.
|
||||
|
||||
Args:
|
||||
media_path: 媒体文件路径
|
||||
|
||||
Returns:
|
||||
转写出的文本
|
||||
|
||||
Raises:
|
||||
ASRNotConfiguredError: ASR 服务未配置
|
||||
ASRTranscriptionError: ASR 调用失败
|
||||
"""
|
||||
# 延迟导入,避免循环依赖和启动时副作用
|
||||
from apps.worker.services.asr_service_factory import get_asr_service
|
||||
|
||||
asr = get_asr_service()
|
||||
if asr is None:
|
||||
raise ASRNotConfiguredError("ASR 服务未配置,请联系管理员配置火山 MediaKit 或阿里云 ASR 密钥")
|
||||
|
||||
try:
|
||||
timeline = asr.transcribe(Path(media_path))
|
||||
# 拼接所有分段的文本
|
||||
text = "".join(seg.text for seg in timeline.segments)
|
||||
return text.strip()
|
||||
except ASRNotConfiguredError:
|
||||
raise
|
||||
except ASRServiceError as exc:
|
||||
logger.error("ASR 转写失败: %s", exc)
|
||||
raise ASRTranscriptionError(f"语音识别失败: {exc}") from exc
|
||||
except Exception as exc:
|
||||
logger.error("ASR 转写异常: %s", exc)
|
||||
raise ASRTranscriptionError(f"语音识别失败: {exc}") from exc
|
||||
@@ -17,3 +17,6 @@ python-dotenv==1.0.1
|
||||
# AI 数字人封面智能选帧(cover_frame_scorer 用 cv2/numpy 做清晰度/亮度/色彩评分)
|
||||
numpy==1.26.4
|
||||
opencv-python-headless==4.10.0.84
|
||||
|
||||
# yt-dlp: 抖音视频下载(#1893 文案提取)
|
||||
yt-dlp>=2024.1.0
|
||||
|
||||
@@ -0,0 +1,422 @@
|
||||
"""Scripts AI 能力单元测试 — Issue #1893.
|
||||
|
||||
测试覆盖:
|
||||
- extract_from_douyin: 成功/非法URL/下载失败/ASR未配置/ASR失败
|
||||
- ai_rewrite: 成功/空内容/LLM失败
|
||||
- ai_generate_titles: 成功/count截断/空内容
|
||||
|
||||
所有外部调用(yt_dlp、ASR、LLM)均通过 unittest.mock.patch 隔离。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pydantic
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, "apps/api")
|
||||
|
||||
|
||||
def _make_auth_user(user_id: str = "u1"):
|
||||
"""构造 mock AuthenticatedUser."""
|
||||
user = MagicMock()
|
||||
user.id = user_id
|
||||
auth = MagicMock()
|
||||
auth.user = user
|
||||
return auth
|
||||
|
||||
|
||||
def _mock_youtube_dl(
|
||||
extract_info_return=None,
|
||||
extract_info_side_effect=None,
|
||||
prepare_filename_return="/tmp/douyin_extract_abc/abc123.mp4",
|
||||
):
|
||||
"""构造 yt_dlp.YoutubeDL 的 mock.
|
||||
|
||||
路由中用法: ydl = yt_dlp.YoutubeDL(opts); info = ydl.extract_info(...)
|
||||
所以 mock_ydl_cls.return_value 就是 ydl 实例.
|
||||
"""
|
||||
mock_ydl_instance = MagicMock()
|
||||
if extract_info_side_effect is not None:
|
||||
mock_ydl_instance.extract_info.side_effect = extract_info_side_effect
|
||||
else:
|
||||
mock_ydl_instance.extract_info.return_value = extract_info_return or {
|
||||
"id": "abc123",
|
||||
"duration": 120.5,
|
||||
}
|
||||
mock_ydl_instance.prepare_filename.return_value = prepare_filename_return
|
||||
return mock_ydl_instance
|
||||
|
||||
|
||||
# ── extract_from_douyin ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestExtractFromDouyin:
|
||||
"""POST /extract-from-douyin 测试."""
|
||||
|
||||
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
||||
@patch("tempfile.TemporaryDirectory")
|
||||
@patch("yt_dlp.YoutubeDL")
|
||||
def test_extract_from_douyin_success(
|
||||
self,
|
||||
mock_ydl_cls,
|
||||
mock_tempdir,
|
||||
mock_transcribe,
|
||||
):
|
||||
"""正常流程:下载视频 + ASR 转写成功."""
|
||||
from app.api.routes.scripts_ai import extract_from_douyin
|
||||
from app.schemas.scripts_ai import ExtractFromDouyinRequest
|
||||
|
||||
mock_ydl_cls.return_value = _mock_youtube_dl(
|
||||
extract_info_return={"id": "abc123", "duration": 120.5},
|
||||
)
|
||||
|
||||
mock_td = MagicMock()
|
||||
mock_td.__enter__ = MagicMock(return_value="/tmp/douyin_extract_abc")
|
||||
mock_td.__exit__ = MagicMock(return_value=False)
|
||||
mock_tempdir.return_value = mock_td
|
||||
|
||||
mock_transcribe.return_value = "这是一段测试文案内容"
|
||||
|
||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
||||
auth = _make_auth_user()
|
||||
result = extract_from_douyin(request=req, authenticated_user=auth)
|
||||
|
||||
assert result.text == "这是一段测试文案内容"
|
||||
assert result.duration_seconds == 120.5
|
||||
assert result.source_url == "https://v.douyin.com/xxxxx/"
|
||||
mock_transcribe.assert_called_once()
|
||||
mock_tempdir.assert_called_once()
|
||||
mock_td.__exit__.assert_called_once()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_url",
|
||||
[
|
||||
"",
|
||||
"not-a-url",
|
||||
"https://www.youtube.com/watch?v=abc",
|
||||
"https://www.bilibili.com/video/BV123",
|
||||
"https://douyin.com/something",
|
||||
"ftp://v.douyin.com/xxx/",
|
||||
],
|
||||
)
|
||||
def test_extract_from_douyin_invalid_url(self, bad_url):
|
||||
"""非法 URL 返回 400."""
|
||||
from app.api.routes.scripts_ai import extract_from_douyin
|
||||
from app.schemas.scripts_ai import ExtractFromDouyinRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
req = ExtractFromDouyinRequest(url=bad_url)
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
extract_from_douyin(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@patch("tempfile.TemporaryDirectory")
|
||||
@patch("yt_dlp.YoutubeDL")
|
||||
def test_extract_from_douyin_download_failure(self, mock_ydl_cls, mock_tempdir):
|
||||
"""下载失败返回 502."""
|
||||
from app.api.routes.scripts_ai import extract_from_douyin
|
||||
from app.schemas.scripts_ai import ExtractFromDouyinRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_ydl_cls.return_value = _mock_youtube_dl(
|
||||
extract_info_side_effect=Exception("Video unavailable"),
|
||||
)
|
||||
|
||||
mock_td = MagicMock()
|
||||
mock_td.__enter__ = MagicMock(return_value="/tmp/douyin_extract_abc")
|
||||
mock_td.__exit__ = MagicMock(return_value=False)
|
||||
mock_tempdir.return_value = mock_td
|
||||
|
||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
extract_from_douyin(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
||||
@patch("tempfile.TemporaryDirectory")
|
||||
@patch("yt_dlp.YoutubeDL")
|
||||
def test_extract_from_douyin_asr_not_configured(
|
||||
self,
|
||||
mock_ydl_cls,
|
||||
mock_tempdir,
|
||||
mock_transcribe,
|
||||
):
|
||||
"""ASR 未配置返回 503."""
|
||||
from app.api.routes.scripts_ai import extract_from_douyin
|
||||
from app.schemas.scripts_ai import ExtractFromDouyinRequest
|
||||
from app.services.script_asr_service import ASRNotConfiguredError
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_ydl_cls.return_value = _mock_youtube_dl(
|
||||
extract_info_return={"id": "abc123", "duration": 60},
|
||||
)
|
||||
|
||||
mock_td = MagicMock()
|
||||
mock_td.__enter__ = MagicMock(return_value="/tmp/douyin_extract_abc")
|
||||
mock_td.__exit__ = MagicMock(return_value=False)
|
||||
mock_tempdir.return_value = mock_td
|
||||
|
||||
mock_transcribe.side_effect = ASRNotConfiguredError("ASR 服务未配置")
|
||||
|
||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
extract_from_douyin(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
||||
@patch("tempfile.TemporaryDirectory")
|
||||
@patch("yt_dlp.YoutubeDL")
|
||||
def test_extract_from_douyin_asr_failure(
|
||||
self,
|
||||
mock_ydl_cls,
|
||||
mock_tempdir,
|
||||
mock_transcribe,
|
||||
):
|
||||
"""ASR 调用失败返回 502."""
|
||||
from app.api.routes.scripts_ai import extract_from_douyin
|
||||
from app.schemas.scripts_ai import ExtractFromDouyinRequest
|
||||
from app.services.script_asr_service import ASRTranscriptionError
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_ydl_cls.return_value = _mock_youtube_dl(
|
||||
extract_info_return={"id": "abc123", "duration": 60},
|
||||
)
|
||||
|
||||
mock_td = MagicMock()
|
||||
mock_td.__enter__ = MagicMock(return_value="/tmp/douyin_extract_abc")
|
||||
mock_td.__exit__ = MagicMock(return_value=False)
|
||||
mock_tempdir.return_value = mock_td
|
||||
|
||||
mock_transcribe.side_effect = ASRTranscriptionError("语音识别失败: timeout")
|
||||
|
||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
extract_from_douyin(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
|
||||
# ── ai_rewrite ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAiRewrite:
|
||||
"""POST /ai-rewrite 测试."""
|
||||
|
||||
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
||||
def test_ai_rewrite_success(self, mock_get_client):
|
||||
"""正常改写成功."""
|
||||
from app.api.routes.scripts_ai import ai_rewrite
|
||||
from app.schemas.scripts_ai import AiRewriteRequest
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.chat_completion.return_value = "改写后的文案内容,口语化风格"
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
req = AiRewriteRequest(content="原始文案内容", style="口语化")
|
||||
auth = _make_auth_user()
|
||||
result = ai_rewrite(request=req, authenticated_user=auth)
|
||||
|
||||
assert result.original == "原始文案内容"
|
||||
assert result.rewritten == "改写后的文案内容,口语化风格"
|
||||
assert result.style == "口语化"
|
||||
mock_client.chat_completion.assert_called_once()
|
||||
|
||||
def test_ai_rewrite_empty_content(self):
|
||||
"""空内容返回 400."""
|
||||
from app.api.routes.scripts_ai import ai_rewrite
|
||||
from app.schemas.scripts_ai import AiRewriteRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
req = AiRewriteRequest(content=" ", style="口语化")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
ai_rewrite(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
||||
def test_ai_rewrite_llm_failure(self, mock_get_client):
|
||||
"""LLM 调用失败返回 502."""
|
||||
from app.api.routes.scripts_ai import ai_rewrite
|
||||
from app.schemas.scripts_ai import AiRewriteRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.chat_completion.side_effect = Exception("API timeout")
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
req = AiRewriteRequest(content="测试内容", style="口语化")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
ai_rewrite(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
||||
def test_ai_rewrite_client_unavailable(self, mock_get_client):
|
||||
"""客户端不可用返回 502."""
|
||||
from app.api.routes.scripts_ai import ai_rewrite
|
||||
from app.schemas.scripts_ai import AiRewriteRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
req = AiRewriteRequest(content="测试内容", style="口语化")
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
ai_rewrite(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
|
||||
# ── ai_generate_titles ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAiGenerateTitles:
|
||||
"""POST /ai-generate-titles 测试."""
|
||||
|
||||
@patch("app.services.ai_service.get_doubao_client")
|
||||
def test_generate_titles_success(self, mock_get_client):
|
||||
"""正常生成标题."""
|
||||
from app.api.routes.scripts_ai import ai_generate_titles
|
||||
from app.schemas.scripts_ai import AiGenerateTitlesRequest
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False # 走 fallback 路径
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
req = AiGenerateTitlesRequest(content="这是一段关于美食的文案", count=3)
|
||||
auth = _make_auth_user()
|
||||
result = ai_generate_titles(request=req, authenticated_user=auth)
|
||||
|
||||
assert len(result.titles) == 3
|
||||
assert all(isinstance(t, str) for t in result.titles)
|
||||
|
||||
@patch("app.services.ai_service.get_doubao_client")
|
||||
def test_generate_titles_count_clamp(self, mock_get_client):
|
||||
"""count 超出范围时 Pydantic 校验拦截."""
|
||||
from app.schemas.scripts_ai import AiGenerateTitlesRequest
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
auth = _make_auth_user()
|
||||
|
||||
# count=10 被 Pydantic le=5 校验拦截 → ValidationError
|
||||
with pytest.raises(pydantic.ValidationError):
|
||||
AiGenerateTitlesRequest(content="测试", count=10)
|
||||
|
||||
# count=0 被 Pydantic ge=1 校验拦截 → ValidationError
|
||||
with pytest.raises(pydantic.ValidationError):
|
||||
AiGenerateTitlesRequest(content="测试", count=0)
|
||||
|
||||
@patch("app.services.ai_service.get_doubao_client")
|
||||
def test_generate_titles_count_valid_range(self, mock_get_client):
|
||||
"""count=1 和 count=5 正常工作."""
|
||||
from app.api.routes.scripts_ai import ai_generate_titles
|
||||
from app.schemas.scripts_ai import AiGenerateTitlesRequest
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
auth = _make_auth_user()
|
||||
|
||||
# count=5
|
||||
req = AiGenerateTitlesRequest(content="测试内容", count=5)
|
||||
result = ai_generate_titles(request=req, authenticated_user=auth)
|
||||
assert len(result.titles) <= 5
|
||||
|
||||
# count=1
|
||||
req = AiGenerateTitlesRequest(content="测试内容", count=1)
|
||||
result = ai_generate_titles(request=req, authenticated_user=auth)
|
||||
assert len(result.titles) >= 1
|
||||
|
||||
def test_generate_titles_empty_content(self):
|
||||
"""空内容返回 400."""
|
||||
from app.api.routes.scripts_ai import ai_generate_titles
|
||||
from app.schemas.scripts_ai import AiGenerateTitlesRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
req = AiGenerateTitlesRequest(content="", count=3)
|
||||
auth = _make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
ai_generate_titles(request=req, authenticated_user=auth)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@patch("app.services.ai_service.get_doubao_client")
|
||||
def test_generate_titles_default_count(self, mock_get_client):
|
||||
"""不传 count 时默认 3."""
|
||||
from app.api.routes.scripts_ai import ai_generate_titles
|
||||
from app.schemas.scripts_ai import AiGenerateTitlesRequest
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get_client.return_value = mock_client
|
||||
|
||||
req = AiGenerateTitlesRequest(content="测试文案内容")
|
||||
auth = _make_auth_user()
|
||||
result = ai_generate_titles(request=req, authenticated_user=auth)
|
||||
|
||||
assert len(result.titles) == 3
|
||||
|
||||
|
||||
# ── URL 校验辅助函数 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateDouyinUrl:
|
||||
"""URL 校验逻辑单元测试."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"valid_url",
|
||||
[
|
||||
"https://v.douyin.com/abc123/",
|
||||
"http://v.douyin.com/abc123/",
|
||||
"v.douyin.com/abc123/",
|
||||
"https://www.douyin.com/video/1234567890",
|
||||
"http://www.douyin.com/video/1234567890",
|
||||
"www.douyin.com/video/1234567890",
|
||||
],
|
||||
)
|
||||
def test_valid_urls(self, valid_url):
|
||||
"""合法 URL 不抛异常."""
|
||||
from app.api.routes.scripts_ai import _validate_douyin_url
|
||||
|
||||
_validate_douyin_url(valid_url)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"invalid_url",
|
||||
[
|
||||
"",
|
||||
" ",
|
||||
"https://www.youtube.com/watch?v=abc",
|
||||
"https://www.bilibili.com/video/BV123",
|
||||
"https://douyin.com/something",
|
||||
"ftp://v.douyin.com/xxx/",
|
||||
"not-a-url",
|
||||
],
|
||||
)
|
||||
def test_invalid_urls(self, invalid_url):
|
||||
"""非法 URL 抛 400."""
|
||||
from app.api.routes.scripts_ai import _validate_douyin_url
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_validate_douyin_url(invalid_url)
|
||||
assert exc_info.value.status_code == 400
|
||||
Reference in New Issue
Block a user