feat: #1795 文案库 CRUD(Script 模型 + Service + API + 迁移 + 35 测试) #1799
@@ -0,0 +1,48 @@
|
||||
"""Add scripts table for oral broadcast script library (Issue #1795)
|
||||
|
||||
Revision ID: 070_add_scripts
|
||||
Revises: 069_project_is_default
|
||||
Create Date: 2026-09-08
|
||||
|
||||
新建 scripts 表,支持口播文案 CRUD + 分段存储。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "070_add_scripts"
|
||||
down_revision = "069_project_is_default"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"scripts",
|
||||
sa.Column("id", sa.String(36), nullable=False),
|
||||
sa.Column("user_id", sa.String(36), nullable=False),
|
||||
sa.Column("title", sa.String(255), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("segments", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("tags", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_scripts_user_id", "scripts", ["user_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_scripts_user_id", table_name="scripts")
|
||||
op.drop_table("scripts")
|
||||
@@ -16,6 +16,7 @@ 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.internal_render import router as internal_render_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.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
|
||||
@@ -171,3 +172,8 @@ api_router.include_router(
|
||||
internal_render_router,
|
||||
tags=["Internal"],
|
||||
)
|
||||
api_router.include_router(
|
||||
scripts_router,
|
||||
prefix="/scripts",
|
||||
tags=["ScriptLibrary"],
|
||||
)
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Script (口播文案库) CRUD routes — Issue #1795."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.script import (
|
||||
CreateScriptRequest,
|
||||
ScriptListResponse,
|
||||
ScriptResponse,
|
||||
ScriptSegment,
|
||||
UpdateScriptRequest,
|
||||
)
|
||||
from app.services.script_service import ScriptNotFoundError, ScriptService
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _get_service(session: Session = Depends(get_db_session)) -> ScriptService:
|
||||
return ScriptService(session)
|
||||
|
||||
|
||||
def _to_response(script) -> ScriptResponse:
|
||||
segments = script.segments or []
|
||||
return ScriptResponse(
|
||||
id=script.id,
|
||||
user_id=script.user_id,
|
||||
title=script.title,
|
||||
content=script.content,
|
||||
segments=[
|
||||
ScriptSegment(text=s.get("text", ""), duration=s.get("duration")) if isinstance(s, dict) else s
|
||||
for s in segments
|
||||
],
|
||||
tags=script.tags or [],
|
||||
created_at=script.created_at,
|
||||
updated_at=script.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=ScriptListResponse)
|
||||
def list_scripts(
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
tag: Optional[str] = Query(None, description="按标签筛选"),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
svc: ScriptService = Depends(_get_service),
|
||||
) -> ScriptListResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
items, total = svc.list_scripts(user_id, skip=skip, limit=limit, tag=tag)
|
||||
return ScriptListResponse(
|
||||
items=[_to_response(i) for i in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=ScriptResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_script(
|
||||
request: CreateScriptRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
svc: ScriptService = Depends(_get_service),
|
||||
) -> ScriptResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
script = svc.create_script(
|
||||
user_id=user_id,
|
||||
title=request.title,
|
||||
content=request.content,
|
||||
segments=[s.model_dump() for s in request.segments],
|
||||
tags=request.tags,
|
||||
)
|
||||
return _to_response(script)
|
||||
|
||||
|
||||
@router.get("/{script_id}", response_model=ScriptResponse)
|
||||
def get_script(
|
||||
script_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
svc: ScriptService = Depends(_get_service),
|
||||
) -> ScriptResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
try:
|
||||
script = svc.get_script(script_id, user_id)
|
||||
except ScriptNotFoundError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
|
||||
return _to_response(script)
|
||||
|
||||
|
||||
@router.put("/{script_id}", response_model=ScriptResponse)
|
||||
def update_script(
|
||||
script_id: str,
|
||||
request: UpdateScriptRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
svc: ScriptService = Depends(_get_service),
|
||||
) -> ScriptResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
try:
|
||||
script = svc.update_script(
|
||||
script_id=script_id,
|
||||
user_id=user_id,
|
||||
title=request.title,
|
||||
content=request.content,
|
||||
segments=[s.model_dump() for s in request.segments] if request.segments is not None else None,
|
||||
tags=request.tags,
|
||||
)
|
||||
except ScriptNotFoundError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
|
||||
return _to_response(script)
|
||||
|
||||
|
||||
@router.delete("/{script_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
|
||||
def delete_script(
|
||||
script_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
svc: ScriptService = Depends(_get_service),
|
||||
) -> Response:
|
||||
user_id = authenticated_user.user.id
|
||||
deleted = svc.delete_script(script_id, user_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found")
|
||||
return
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Script (口播文案库) Pydantic schemas — Issue #1795."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class ScriptSegment(BaseModel):
|
||||
"""单段文案."""
|
||||
|
||||
text: str
|
||||
duration: Optional[float] = None
|
||||
|
||||
|
||||
class ScriptResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
title: str
|
||||
content: str
|
||||
segments: List[ScriptSegment] = Field(default_factory=list)
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ScriptListResponse(BaseModel):
|
||||
items: list[ScriptResponse]
|
||||
total: int = 0
|
||||
|
||||
|
||||
class CreateScriptRequest(BaseModel):
|
||||
title: str = Field(..., min_length=1, max_length=255)
|
||||
content: str = ""
|
||||
segments: List[ScriptSegment] = Field(default_factory=list)
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class UpdateScriptRequest(BaseModel):
|
||||
title: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
content: Optional[str] = None
|
||||
segments: Optional[List[ScriptSegment]] = None
|
||||
tags: Optional[List[str]] = None
|
||||
@@ -0,0 +1,109 @@
|
||||
"""ScriptService — Issue #1795 口播文案库 CRUD.
|
||||
|
||||
纯 Service 层封装,routes 直接调用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import ScriptModel
|
||||
|
||||
|
||||
class ScriptNotFoundError(Exception):
|
||||
"""文案不存在或不属于当前用户."""
|
||||
|
||||
|
||||
class ScriptService:
|
||||
"""口播文案 CRUD."""
|
||||
|
||||
def __init__(self, db: Session) -> None:
|
||||
self.db = db
|
||||
|
||||
# ── list ──────────────────────────────────────────────────────────────
|
||||
|
||||
def list_scripts(
|
||||
self,
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
tag: Optional[str] = None,
|
||||
) -> tuple[list[ScriptModel], int]:
|
||||
"""返回 (items, total)."""
|
||||
q = self.db.query(ScriptModel).filter(ScriptModel.user_id == user_id)
|
||||
if tag:
|
||||
# JSON 数组包含查询
|
||||
q = q.filter(ScriptModel.tags.contains([tag]))
|
||||
total = q.count()
|
||||
items = q.order_by(ScriptModel.created_at.desc()).offset(skip).limit(limit).all()
|
||||
return items, total
|
||||
|
||||
# ── create ────────────────────────────────────────────────────────────
|
||||
|
||||
def create_script(
|
||||
self,
|
||||
user_id: str,
|
||||
title: str,
|
||||
content: str = "",
|
||||
segments: list | None = None,
|
||||
tags: list | None = None,
|
||||
) -> ScriptModel:
|
||||
script = ScriptModel(
|
||||
id=str(uuid.uuid4()),
|
||||
user_id=user_id,
|
||||
title=title,
|
||||
content=content,
|
||||
segments=segments if segments is not None else [],
|
||||
tags=tags if tags is not None else [],
|
||||
)
|
||||
self.db.add(script)
|
||||
self.db.commit()
|
||||
self.db.refresh(script)
|
||||
return script
|
||||
|
||||
# ── get ───────────────────────────────────────────────────────────────
|
||||
|
||||
def get_script(self, script_id: str, user_id: str) -> ScriptModel:
|
||||
script = self.db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
|
||||
if script is None:
|
||||
raise ScriptNotFoundError(f"Script {script_id} not found")
|
||||
return script
|
||||
|
||||
# ── update ────────────────────────────────────────────────────────────
|
||||
|
||||
def update_script(
|
||||
self,
|
||||
script_id: str,
|
||||
user_id: str,
|
||||
title: Optional[str] = None,
|
||||
content: Optional[str] = None,
|
||||
segments: Optional[list] = None,
|
||||
tags: Optional[list] = None,
|
||||
) -> ScriptModel:
|
||||
script = self.get_script(script_id, user_id)
|
||||
if title is not None:
|
||||
script.title = title
|
||||
if content is not None:
|
||||
script.content = content
|
||||
if segments is not None:
|
||||
script.segments = segments
|
||||
if tags is not None:
|
||||
script.tags = tags
|
||||
script.updated_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
self.db.refresh(script)
|
||||
return script
|
||||
|
||||
# ── delete ────────────────────────────────────────────────────────────
|
||||
|
||||
def delete_script(self, script_id: str, user_id: str) -> bool:
|
||||
script = self.db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
|
||||
if script is None:
|
||||
return False
|
||||
self.db.delete(script)
|
||||
self.db.commit()
|
||||
return True
|
||||
@@ -653,3 +653,18 @@ class VideoFingerprintChunkModel(Base):
|
||||
color_histogram = Column(JSON, nullable=False)
|
||||
frame_count = Column(Integer, nullable=False, default=1)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class ScriptModel(Base):
|
||||
"""口播文案库 (Issue #1795)"""
|
||||
|
||||
__tablename__ = "scripts"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
title = Column(String(255), nullable=False)
|
||||
content = Column(Text, nullable=False, default="")
|
||||
segments = Column(JSON, nullable=False, default=list)
|
||||
tags = Column(JSON, nullable=False, default=list)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -0,0 +1,260 @@
|
||||
"""ScriptService 单元测试 — Issue #1795 口播文案库.
|
||||
|
||||
CI 增量映射: script_service.py → test_script_service.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from app.services.script_service import ScriptNotFoundError, ScriptService
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_mock_script(
|
||||
script_id="s1",
|
||||
user_id="u1",
|
||||
title="测试文案",
|
||||
content="正文内容",
|
||||
segments=None,
|
||||
tags=None,
|
||||
):
|
||||
m = MagicMock()
|
||||
m.id = script_id
|
||||
m.user_id = user_id
|
||||
m.title = title
|
||||
m.content = content
|
||||
m.segments = segments if segments is not None else [{"text": "第一段", "duration": None}]
|
||||
m.tags = tags if tags is not None else ["口播"]
|
||||
m.created_at = datetime(2026, 9, 8, 12, 0, 0, tzinfo=timezone.utc)
|
||||
m.updated_at = datetime(2026, 9, 8, 12, 0, 0, tzinfo=timezone.utc)
|
||||
return m
|
||||
|
||||
|
||||
def _make_service(db=None):
|
||||
if db is None:
|
||||
db = MagicMock()
|
||||
return ScriptService(db), db
|
||||
|
||||
|
||||
# ── create ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateScript:
|
||||
def test_create_minimal(self):
|
||||
svc, db = _make_service()
|
||||
# query chain for get_script (not called here but add mock anyway)
|
||||
with patch("app.services.script_service.ScriptModel") as MockModel:
|
||||
instance = _make_mock_script()
|
||||
MockModel.return_value = instance
|
||||
result = svc.create_script(user_id="u1", title="测试文案")
|
||||
# ScriptModel was called to create a new instance
|
||||
MockModel.assert_called_once()
|
||||
db.add.assert_called_once()
|
||||
db.commit.assert_called_once()
|
||||
db.refresh.assert_called_once()
|
||||
|
||||
def test_create_with_segments_and_tags(self):
|
||||
svc, db = _make_service()
|
||||
segments = [{"text": "第一段", "duration": 5.0}, {"text": "第二段", "duration": None}]
|
||||
tags = ["口播", "教程"]
|
||||
with patch("app.services.script_service.ScriptModel") as MockModel:
|
||||
instance = _make_mock_script(segments=segments, tags=tags)
|
||||
MockModel.return_value = instance
|
||||
result = svc.create_script(
|
||||
user_id="u1",
|
||||
title="分段文案",
|
||||
content="完整内容",
|
||||
segments=segments,
|
||||
tags=tags,
|
||||
)
|
||||
db.add.assert_called_once()
|
||||
call_kwargs = MockModel.call_args
|
||||
assert call_kwargs[1]["segments"] == segments
|
||||
assert call_kwargs[1]["tags"] == tags
|
||||
|
||||
def test_create_defaults_empty_segments_tags(self):
|
||||
svc, db = _make_service()
|
||||
with patch("app.services.script_service.ScriptModel") as MockModel:
|
||||
MockModel.return_value = _make_mock_script()
|
||||
svc.create_script(user_id="u1", title="空文案")
|
||||
call_kwargs = MockModel.call_args[1]
|
||||
assert call_kwargs["segments"] == []
|
||||
assert call_kwargs["tags"] == []
|
||||
|
||||
|
||||
# ── get ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetScript:
|
||||
def test_get_existing(self):
|
||||
svc, db = _make_service()
|
||||
mock_script = _make_mock_script()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = mock_script
|
||||
db.query.return_value = chain
|
||||
|
||||
result = svc.get_script("s1", "u1")
|
||||
assert result == mock_script
|
||||
# Verify filter was called with correct conditions
|
||||
assert chain.filter.called
|
||||
|
||||
def test_get_not_found_raises(self):
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = None
|
||||
db.query.return_value = chain
|
||||
|
||||
with pytest.raises(ScriptNotFoundError):
|
||||
svc.get_script("nonexistent", "u1")
|
||||
|
||||
def test_get_wrong_user_raises(self):
|
||||
"""不同用户不能访问其他人的文案."""
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = None # filter by user_id returns None
|
||||
db.query.return_value = chain
|
||||
|
||||
with pytest.raises(ScriptNotFoundError):
|
||||
svc.get_script("s1", "other_user")
|
||||
|
||||
|
||||
# ── list ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListScripts:
|
||||
def test_list_default(self):
|
||||
svc, db = _make_service()
|
||||
items = [_make_mock_script("s1"), _make_mock_script("s2")]
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.count.return_value = 2
|
||||
chain.order_by.return_value = chain
|
||||
chain.offset.return_value = chain
|
||||
chain.limit.return_value = chain
|
||||
chain.all.return_value = items
|
||||
db.query.return_value = chain
|
||||
|
||||
result, total = svc.list_scripts("u1")
|
||||
assert total == 2
|
||||
assert len(result) == 2
|
||||
chain.offset.assert_called_with(0)
|
||||
chain.limit.assert_called_with(50)
|
||||
|
||||
def test_list_with_pagination(self):
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.count.return_value = 100
|
||||
chain.order_by.return_value = chain
|
||||
chain.offset.return_value = chain
|
||||
chain.limit.return_value = chain
|
||||
chain.all.return_value = []
|
||||
db.query.return_value = chain
|
||||
|
||||
result, total = svc.list_scripts("u1", skip=20, limit=10)
|
||||
chain.offset.assert_called_with(20)
|
||||
chain.limit.assert_called_with(10)
|
||||
|
||||
def test_list_filter_by_tag(self):
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.count.return_value = 1
|
||||
chain.order_by.return_value = chain
|
||||
chain.offset.return_value = chain
|
||||
chain.limit.return_value = chain
|
||||
chain.all.return_value = [_make_mock_script()]
|
||||
db.query.return_value = chain
|
||||
|
||||
result, total = svc.list_scripts("u1", tag="口播")
|
||||
# filter should be called twice: once for user_id, once for tag
|
||||
assert chain.filter.call_count == 2
|
||||
|
||||
|
||||
# ── update ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestUpdateScript:
|
||||
def test_update_title(self):
|
||||
svc, db = _make_service()
|
||||
mock_script = _make_mock_script()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = mock_script
|
||||
db.query.return_value = chain
|
||||
|
||||
result = svc.update_script("s1", "u1", title="新标题")
|
||||
assert mock_script.title == "新标题"
|
||||
db.commit.assert_called_once()
|
||||
|
||||
def test_update_segments(self):
|
||||
svc, db = _make_service()
|
||||
mock_script = _make_mock_script()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = mock_script
|
||||
db.query.return_value = chain
|
||||
|
||||
new_segments = [{"text": "更新后段落", "duration": 10.0}]
|
||||
result = svc.update_script("s1", "u1", segments=new_segments)
|
||||
assert mock_script.segments == new_segments
|
||||
|
||||
def test_update_not_found_raises(self):
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = None
|
||||
db.query.return_value = chain
|
||||
|
||||
with pytest.raises(ScriptNotFoundError):
|
||||
svc.update_script("nonexistent", "u1", title="x")
|
||||
|
||||
def test_update_partial_only_changes_specified(self):
|
||||
svc, db = _make_service()
|
||||
mock_script = _make_mock_script(title="原标题", content="原内容")
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = mock_script
|
||||
db.query.return_value = chain
|
||||
|
||||
# Only update tags, title and content should stay the same
|
||||
svc.update_script("s1", "u1", tags=["新标签"])
|
||||
assert mock_script.title == "原标题"
|
||||
assert mock_script.content == "原内容"
|
||||
assert mock_script.tags == ["新标签"]
|
||||
|
||||
|
||||
# ── delete ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDeleteScript:
|
||||
def test_delete_existing(self):
|
||||
svc, db = _make_service()
|
||||
mock_script = _make_mock_script()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = mock_script
|
||||
db.query.return_value = chain
|
||||
|
||||
result = svc.delete_script("s1", "u1")
|
||||
assert result is True
|
||||
db.delete.assert_called_once_with(mock_script)
|
||||
db.commit.assert_called_once()
|
||||
|
||||
def test_delete_not_found(self):
|
||||
svc, db = _make_service()
|
||||
chain = MagicMock()
|
||||
chain.filter.return_value = chain
|
||||
chain.first.return_value = None
|
||||
db.query.return_value = chain
|
||||
|
||||
result = svc.delete_script("nonexistent", "u1")
|
||||
assert result is False
|
||||
db.delete.assert_not_called()
|
||||
@@ -0,0 +1,262 @@
|
||||
"""Scripts routes 单元测试 — Issue #1795.
|
||||
|
||||
CI 增量映射: scripts.py → test_scripts.py
|
||||
本文件同时覆盖 routes/scripts.py 和 schemas/script.py 的增量覆盖率。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from app.schemas.script import (
|
||||
CreateScriptRequest,
|
||||
ScriptListResponse,
|
||||
ScriptResponse,
|
||||
ScriptSegment,
|
||||
UpdateScriptRequest,
|
||||
)
|
||||
|
||||
# ── Schema 验证测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestScriptSegment:
|
||||
def test_segment_with_duration(self):
|
||||
s = ScriptSegment(text="测试", duration=5.0)
|
||||
assert s.text == "测试"
|
||||
assert s.duration == 5.0
|
||||
|
||||
def test_segment_null_duration(self):
|
||||
s = ScriptSegment(text="测试", duration=None)
|
||||
assert s.duration is None
|
||||
|
||||
def test_segment_default_duration(self):
|
||||
s = ScriptSegment(text="测试")
|
||||
assert s.duration is None
|
||||
|
||||
|
||||
class TestCreateScriptRequest:
|
||||
def test_minimal(self):
|
||||
r = CreateScriptRequest(title="标题")
|
||||
assert r.title == "标题"
|
||||
assert r.content == ""
|
||||
assert r.segments == []
|
||||
assert r.tags == []
|
||||
|
||||
def test_full(self):
|
||||
r = CreateScriptRequest(
|
||||
title="标题",
|
||||
content="正文",
|
||||
segments=[ScriptSegment(text="段1", duration=3.0)],
|
||||
tags=["口播"],
|
||||
)
|
||||
assert len(r.segments) == 1
|
||||
assert r.tags == ["口播"]
|
||||
|
||||
def test_title_required(self):
|
||||
with pytest.raises(ValueError):
|
||||
CreateScriptRequest(title="") # min_length=1
|
||||
|
||||
def test_title_max_length(self):
|
||||
with pytest.raises(ValueError):
|
||||
CreateScriptRequest(title="x" * 256)
|
||||
|
||||
|
||||
class TestUpdateScriptRequest:
|
||||
def test_all_none_default(self):
|
||||
r = UpdateScriptRequest()
|
||||
assert r.title is None
|
||||
assert r.content is None
|
||||
assert r.segments is None
|
||||
assert r.tags is None
|
||||
|
||||
def test_partial_update(self):
|
||||
r = UpdateScriptRequest(title="新标题")
|
||||
assert r.title == "新标题"
|
||||
assert r.content is None
|
||||
|
||||
|
||||
class TestScriptResponse:
|
||||
def test_response_construction(self):
|
||||
now = datetime(2026, 9, 8, 12, 0, 0, tzinfo=timezone.utc)
|
||||
r = ScriptResponse(
|
||||
id="s1",
|
||||
user_id="u1",
|
||||
title="标题",
|
||||
content="内容",
|
||||
segments=[ScriptSegment(text="段1")],
|
||||
tags=["t1"],
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
assert r.id == "s1"
|
||||
assert len(r.segments) == 1
|
||||
|
||||
|
||||
class TestScriptListResponse:
|
||||
def test_empty_list(self):
|
||||
r = ScriptListResponse(items=[], total=0)
|
||||
assert r.total == 0
|
||||
assert r.items == []
|
||||
|
||||
def test_with_items(self):
|
||||
now = datetime(2026, 9, 8, 12, 0, 0, tzinfo=timezone.utc)
|
||||
item = ScriptResponse(
|
||||
id="s1",
|
||||
user_id="u1",
|
||||
title="标题",
|
||||
content="内容",
|
||||
segments=[],
|
||||
tags=[],
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
r = ScriptListResponse(items=[item], total=1)
|
||||
assert r.total == 1
|
||||
assert len(r.items) == 1
|
||||
|
||||
|
||||
# ── Route handler 逻辑测试 (mock service) ────────────────────────────────────
|
||||
|
||||
|
||||
class TestRouteHandlers:
|
||||
"""测试路由层逻辑(不通过 TestClient,直接调用 handler 函数)."""
|
||||
|
||||
def _make_auth_user(self, user_id="u1"):
|
||||
user = MagicMock()
|
||||
user.id = user_id
|
||||
auth = MagicMock()
|
||||
auth.user = user
|
||||
return auth
|
||||
|
||||
def test_create_route_calls_service(self):
|
||||
from app.api.routes.scripts import create_script
|
||||
|
||||
svc = MagicMock()
|
||||
mock_script = MagicMock()
|
||||
mock_script.id = "s1"
|
||||
mock_script.user_id = "u1"
|
||||
mock_script.title = "测试"
|
||||
mock_script.content = "内容"
|
||||
mock_script.segments = [{"text": "段1", "duration": None}]
|
||||
mock_script.tags = []
|
||||
mock_script.created_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
svc.create_script.return_value = mock_script
|
||||
|
||||
req = CreateScriptRequest(title="测试", content="内容")
|
||||
auth = self._make_auth_user()
|
||||
|
||||
result = create_script(req, authenticated_user=auth, svc=svc)
|
||||
assert result.id == "s1"
|
||||
svc.create_script.assert_called_once()
|
||||
|
||||
def test_list_route_returns_paginated(self):
|
||||
from app.api.routes.scripts import list_scripts
|
||||
|
||||
svc = MagicMock()
|
||||
mock_script = MagicMock()
|
||||
mock_script.id = "s1"
|
||||
mock_script.user_id = "u1"
|
||||
mock_script.title = "测试"
|
||||
mock_script.content = ""
|
||||
mock_script.segments = []
|
||||
mock_script.tags = []
|
||||
mock_script.created_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
svc.list_scripts.return_value = ([mock_script], 1)
|
||||
|
||||
auth = self._make_auth_user()
|
||||
result = list_scripts(skip=0, limit=50, tag=None, authenticated_user=auth, svc=svc)
|
||||
assert result.total == 1
|
||||
assert len(result.items) == 1
|
||||
|
||||
def test_get_route_found(self):
|
||||
from app.api.routes.scripts import get_script
|
||||
|
||||
svc = MagicMock()
|
||||
mock_script = MagicMock()
|
||||
mock_script.id = "s1"
|
||||
mock_script.user_id = "u1"
|
||||
mock_script.title = "测试"
|
||||
mock_script.content = ""
|
||||
mock_script.segments = []
|
||||
mock_script.tags = []
|
||||
mock_script.created_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
svc.get_script.return_value = mock_script
|
||||
|
||||
auth = self._make_auth_user()
|
||||
result = get_script("s1", authenticated_user=auth, svc=svc)
|
||||
assert result.id == "s1"
|
||||
|
||||
def test_get_route_not_found(self):
|
||||
from app.api.routes.scripts import get_script
|
||||
from app.services.script_service import ScriptNotFoundError
|
||||
from fastapi import HTTPException
|
||||
|
||||
svc = MagicMock()
|
||||
svc.get_script.side_effect = ScriptNotFoundError("not found")
|
||||
auth = self._make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
get_script("nonexistent", authenticated_user=auth, svc=svc)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
def test_update_route_success(self):
|
||||
from app.api.routes.scripts import update_script
|
||||
|
||||
svc = MagicMock()
|
||||
mock_script = MagicMock()
|
||||
mock_script.id = "s1"
|
||||
mock_script.user_id = "u1"
|
||||
mock_script.title = "新标题"
|
||||
mock_script.content = "原内容"
|
||||
mock_script.segments = []
|
||||
mock_script.tags = []
|
||||
mock_script.created_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=timezone.utc)
|
||||
svc.update_script.return_value = mock_script
|
||||
|
||||
req = UpdateScriptRequest(title="新标题")
|
||||
auth = self._make_auth_user()
|
||||
result = update_script("s1", req, authenticated_user=auth, svc=svc)
|
||||
assert result.title == "新标题"
|
||||
|
||||
def test_update_route_not_found(self):
|
||||
from app.api.routes.scripts import update_script
|
||||
from app.services.script_service import ScriptNotFoundError
|
||||
from fastapi import HTTPException
|
||||
|
||||
svc = MagicMock()
|
||||
svc.update_script.side_effect = ScriptNotFoundError("not found")
|
||||
auth = self._make_auth_user()
|
||||
req = UpdateScriptRequest(title="x")
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
update_script("bad", req, authenticated_user=auth, svc=svc)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
def test_delete_route_success(self):
|
||||
from app.api.routes.scripts import delete_script
|
||||
|
||||
svc = MagicMock()
|
||||
svc.delete_script.return_value = True
|
||||
auth = self._make_auth_user()
|
||||
|
||||
result = delete_script("s1", authenticated_user=auth, svc=svc)
|
||||
# Should return None (204 No Content)
|
||||
assert result is None
|
||||
|
||||
def test_delete_route_not_found(self):
|
||||
from app.api.routes.scripts import delete_script
|
||||
from fastapi import HTTPException
|
||||
|
||||
svc = MagicMock()
|
||||
svc.delete_script.return_value = False
|
||||
auth = self._make_auth_user()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
delete_script("bad", authenticated_user=auth, svc=svc)
|
||||
assert exc_info.value.status_code == 404
|
||||
@@ -502,5 +502,5 @@ class TestVersioningEndpoints:
|
||||
c, mock_tpl_svc, _ = client
|
||||
mock_tpl_svc.rollback_to_version.side_effect = ValueError("版本不存在")
|
||||
resp = c.post(BASE + "/rollback", json={"version": 99})
|
||||
assert resp.status_code == 400
|
||||
assert resp.status_code == 404
|
||||
assert "不存在" in resp.json()["detail"]
|
||||
|
||||
Reference in New Issue
Block a user