diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 61c5c09c3..ab8df9da4 100755 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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", diff --git a/apps/api/app/api/routes/scripts_ai.py b/apps/api/app/api/routes/scripts_ai.py new file mode 100644 index 000000000..495551e7a --- /dev/null +++ b/apps/api/app/api/routes/scripts_ai.py @@ -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) diff --git a/apps/api/app/schemas/scripts_ai.py b/apps/api/app/schemas/scripts_ai.py new file mode 100644 index 000000000..916632f04 --- /dev/null +++ b/apps/api/app/schemas/scripts_ai.py @@ -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="生成的标题列表") diff --git a/apps/api/app/services/script_asr_service.py b/apps/api/app/services/script_asr_service.py new file mode 100644 index 000000000..cdba1cc7f --- /dev/null +++ b/apps/api/app/services/script_asr_service.py @@ -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 diff --git a/requirements.txt b/requirements.txt index 1cbb54e25..f246a07be 100755 --- a/requirements.txt +++ b/requirements.txt @@ -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 diff --git a/tests/unit/test_scripts_ai.py b/tests/unit/test_scripts_ai.py new file mode 100644 index 000000000..3cf946e2b --- /dev/null +++ b/tests/unit/test_scripts_ai.py @@ -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