阶段3 任务3.06:TTS 合成 API #166

Merged
xiaoxia merged 1 commits from feature/task-306-tts-api into develop 2026-07-02 11:16:40 +08:00
9 changed files with 868 additions and 0 deletions
@@ -0,0 +1,58 @@
"""Task 3.06: Create tts_jobs table
Revision ID: 020
Revises: 019
Create Date: 2026-07-02
新增 tts_jobs 表,用于存储 TTS 合成任务。
"""
from alembic import op
import sqlalchemy as sa
revision = "020"
down_revision = "019"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"tts_jobs",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("input_text", sa.Text(), nullable=False),
sa.Column("voice_id", sa.String(100), nullable=False, server_default=""),
sa.Column("voice_model", sa.String(100), nullable=False, server_default=""),
sa.Column("project_id", sa.String(36), nullable=False, server_default=""),
sa.Column("voice_clone_profile_id", sa.String(36), nullable=False, server_default=""),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("output_audio_url", sa.Text(), nullable=False, server_default=""),
sa.Column("output_audio_key", sa.String(500), nullable=False, server_default=""),
sa.Column("duration", sa.Float(), nullable=False, server_default="0"),
sa.Column("file_size", sa.Integer(), nullable=False, server_default="0"),
sa.Column("sample_rate", sa.Integer(), nullable=False, server_default="22050"),
sa.Column("format", sa.String(20), nullable=False, server_default="mp3"),
sa.Column("error_message", sa.Text(), nullable=False, server_default=""),
sa.Column("retry_count", sa.Integer(), nullable=False, server_default="0"),
sa.Column("max_retries", sa.Integer(), nullable=False, server_default="3"),
sa.Column("metadata", sa.JSON(), nullable=False, server_default="{}"),
sa.Column("started_at", sa.DateTime(), nullable=True),
sa.Column("completed_at", sa.DateTime(), nullable=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.func.now(),
),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.func.now(),
),
)
def downgrade() -> None:
op.drop_table("tts_jobs")
+6
View File
@@ -20,6 +20,7 @@ from app.api.routes.task_center import router as task_center_router
from app.api.routes.templates import router as templates_router
from app.api.routes.titles import router as titles_router
from app.api.routes.upload import router as upload_router
from app.api.routes.tts import router as tts_router
from app.api.routes.voice_clones import router as voice_clones_router
from app.api.routes.voices import router as voices_router
from fastapi import APIRouter
@@ -139,3 +140,8 @@ api_router.include_router(
prefix="/edit-plans",
tags=["EditPlan"],
)
api_router.include_router(
tts_router,
prefix="/tts",
tags=["TTS"],
)
+168
View File
@@ -0,0 +1,168 @@
"""TTS 合成 API 路由。"""
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.tts import (
ListTTSJobResponse,
TTSSynthesizeRequest,
TTSSynthesizeResponse,
TTSJobResponse,
TTSStatusResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.tts_job_repository import (
SQLAlchemyTTSJobRepository,
)
from packages.application.tts_job.use_cases import (
CreateTTSJobUseCase,
DeleteTTSJobUseCase,
GetTTSJobStatusUseCase,
GetTTSJobUseCase,
ListTTSJobsUseCase,
TTSJobNotFoundError,
)
router = APIRouter()
def _get_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTTSJobRepository:
return SQLAlchemyTTSJobRepository(session)
def _to_response(job) -> TTSJobResponse:
return TTSJobResponse(
id=job.id,
user_id=job.user_id,
input_text=job.input_text,
voice_id=job.voice_id,
voice_model=job.voice_model,
project_id=job.project_id,
voice_clone_profile_id=job.voice_clone_profile_id,
status=job.status,
output_audio_url=job.output_audio_url,
output_audio_key=job.output_audio_key,
duration=job.duration,
file_size=job.file_size,
sample_rate=job.sample_rate,
format=job.format,
error_message=job.error_message,
retry_count=job.retry_count,
max_retries=job.max_retries,
metadata=job.metadata,
started_at=job.started_at,
completed_at=job.completed_at,
created_at=job.created_at,
updated_at=job.updated_at,
)
@router.post("/synthesize", response_model=TTSSynthesizeResponse, status_code=status.HTTP_201_CREATED)
def synthesize(
request: TTSSynthesizeRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
) -> TTSSynthesizeResponse:
"""发起 TTS 合成任务。
创建 TTS 任务,状态为 pending,等待后续 CosyVoice API 调用。
"""
user_id = authenticated_user.user.id
use_case = CreateTTSJobUseCase(repository)
job = use_case.execute(
user_id=user_id,
input_text=request.text,
voice_id=request.voice_id,
voice_model=request.voice_model,
voice_clone_profile_id=request.voice_clone_profile_id,
metadata=request.metadata_,
)
return TTSSynthesizeResponse(
job_id=job.id,
status=job.status,
message="合成任务已创建",
)
@router.get("/jobs", response_model=ListTTSJobResponse)
def list_tts_jobs(
page: int = Query(default=1, ge=1, description="页码"),
page_size: int = Query(default=20, ge=1, le=100, description="每页数量"),
status_filter: Optional[str] = Query(None, alias="status"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
) -> ListTTSJobResponse:
"""列出用户的 TTS 合成任务。"""
user_id = authenticated_user.user.id
use_case = ListTTSJobsUseCase(repository)
skip = (page - 1) * page_size
items, total = use_case.execute(
user_id, status=status_filter, skip=skip, limit=page_size
)
return ListTTSJobResponse(
items=[_to_response(j) for j in items],
total=total,
page=page,
page_size=page_size,
)
@router.get("/jobs/{job_id}", response_model=TTSJobResponse)
def get_tts_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
) -> TTSJobResponse:
"""获取 TTS 任务详情。"""
user_id = authenticated_user.user.id
use_case = GetTTSJobUseCase(repository)
try:
job = use_case.execute(job_id, user_id)
except TTSJobNotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
return _to_response(job)
@router.get("/jobs/{job_id}/status", response_model=TTSStatusResponse)
def get_tts_job_status(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
) -> TTSStatusResponse:
"""查询 TTS 合成状态(用于前端轮询)。"""
user_id = authenticated_user.user.id
use_case = GetTTSJobStatusUseCase(repository)
try:
job = use_case.execute(job_id, user_id)
except TTSJobNotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
return TTSStatusResponse(
id=job.id,
status=job.status,
output_audio_url=job.output_audio_url,
error_message=job.error_message,
duration=job.duration,
retry_count=job.retry_count,
created_at=job.created_at,
updated_at=job.updated_at,
)
@router.delete("/jobs/{job_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
def delete_tts_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
) -> Response:
"""删除 TTS 合成任务。"""
user_id = authenticated_user.user.id
use_case = DeleteTTSJobUseCase(repository)
deleted = use_case.execute(job_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
return Response(status_code=204)
+89
View File
@@ -0,0 +1,89 @@
"""TTS 合成 API Schema。"""
from __future__ import annotations
from datetime import datetime
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
class TTSSynthesizeRequest(BaseModel):
"""TTS 合成请求。"""
text: str = Field(..., min_length=1, max_length=10000, description="合成文本")
voice_id: str = Field("", description="音色 ID")
output_name: str = Field("", description="输出文件名")
language: str = Field("zh-CN", description="语言")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
voice_model: str = Field("", description="语音模型名称")
voice_clone_profile_id: str = Field("", description="关联的音色克隆档案 ID")
format: str = Field("mp3", description="输出格式(mp3/wav/pcm")
metadata_: Optional[Dict[str, Any]] = Field(
default=None, alias="metadata", description="额外元数据"
)
class Config:
populate_by_name = True
class TTSJobResponse(BaseModel):
"""TTS 任务响应。"""
id: str
user_id: str
input_text: str
voice_id: str = ""
voice_model: str = ""
project_id: str = ""
voice_clone_profile_id: str = ""
status: str
output_audio_url: str = ""
output_audio_key: str = ""
duration: float = 0.0
file_size: int = 0
sample_rate: int = 22050
format: str = "mp3"
error_message: str = ""
retry_count: int = 0
max_retries: int = 3
metadata_: Optional[Dict[str, Any]] = Field(
default=None, alias="metadata", description="额外元数据"
)
started_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
updated_at: datetime
class Config:
populate_by_name = True
class TTSStatusResponse(BaseModel):
"""TTS 任务状态响应(用于轮询)。"""
id: str
status: str
output_audio_url: str = ""
error_message: str = ""
duration: float = 0.0
retry_count: int = 0
created_at: datetime
updated_at: datetime
class TTSSynthesizeResponse(BaseModel):
"""TTS 合成创建响应。"""
job_id: str
status: str
message: str = "合成任务已创建"
class ListTTSJobResponse(BaseModel):
"""TTS 任务列表响应。"""
items: List[TTSJobResponse]
total: int
page: int
page_size: int
@@ -430,4 +430,33 @@ class JobModel(Base):
started_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
class TTSJobModel(Base):
"""TTS 合成任务 ORM 模型。"""
__tablename__ = "tts_jobs"
id = Column(String(36), primary_key=True)
user_id = Column(String(36), nullable=False, index=True)
input_text = Column(Text, nullable=False)
voice_id = Column(String(100), nullable=False, default="")
voice_model = Column(String(100), nullable=False, default="")
project_id = Column(String(36), nullable=False, default="")
voice_clone_profile_id = Column(String(36), nullable=False, default="")
status = Column(String(20), nullable=False, default="pending", index=True)
output_audio_url = Column(Text, nullable=False, default="")
output_audio_key = Column(String(500), nullable=False, default="")
duration = Column(Float, nullable=False, default=0.0)
file_size = Column(Integer, nullable=False, default=0)
sample_rate = Column(Integer, nullable=False, default=22050)
format = Column(String(20), nullable=False, default="mp3")
error_message = Column(Text, nullable=False, default="")
retry_count = Column(Integer, nullable=False, default=0)
max_retries = Column(Integer, nullable=False, default=3)
metadata_ = Column("metadata", JSON, nullable=False, default=dict)
started_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -0,0 +1,175 @@
"""SQLAlchemy implementation of TTSJobRepository."""
from __future__ import annotations
from typing import List, Optional
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import TTSJobModel
from packages.domain.tts_job import TTSJob
class SQLAlchemyTTSJobRepository:
"""SQLAlchemy TTS 任务仓储。"""
def __init__(self, session: Session) -> None:
self.session = session
def create(self, job: TTSJob) -> TTSJob:
model = TTSJobModel(
id=job.id,
user_id=job.user_id,
input_text=job.input_text,
voice_id=job.voice_id,
voice_model=job.voice_model,
project_id=job.project_id,
voice_clone_profile_id=job.voice_clone_profile_id,
status=job.status,
output_audio_url=job.output_audio_url,
output_audio_key=job.output_audio_key,
duration=job.duration,
file_size=job.file_size,
sample_rate=job.sample_rate,
format=job.format,
error_message=job.error_message,
retry_count=job.retry_count,
max_retries=job.max_retries,
metadata_=job.metadata,
started_at=job.started_at,
completed_at=job.completed_at,
)
self.session.add(model)
self.session.commit()
self.session.refresh(model)
return self._model_to_entity(model)
def get(self, job_id: str) -> Optional[TTSJob]:
model = (
self.session.query(TTSJobModel)
.filter(
TTSJobModel.id == job_id,
TTSJobModel.status != "deleted",
)
.first()
)
if model is None:
return None
return self._model_to_entity(model)
def update(self, job: TTSJob) -> TTSJob:
model = (
self.session.query(TTSJobModel)
.filter(TTSJobModel.id == job.id)
.first()
)
if model is None:
raise ValueError(f"TTSJob {job.id} not found")
model.input_text = job.input_text
model.voice_id = job.voice_id
model.voice_model = job.voice_model
model.project_id = job.project_id
model.voice_clone_profile_id = job.voice_clone_profile_id
model.status = job.status
model.output_audio_url = job.output_audio_url
model.output_audio_key = job.output_audio_key
model.duration = job.duration
model.file_size = job.file_size
model.sample_rate = job.sample_rate
model.format = job.format
model.error_message = job.error_message
model.retry_count = job.retry_count
model.max_retries = job.max_retries
model.metadata_ = job.metadata
model.started_at = job.started_at
model.completed_at = job.completed_at
self.session.commit()
self.session.refresh(model)
return self._model_to_entity(model)
def delete(self, job_id: str) -> bool:
model = (
self.session.query(TTSJobModel)
.filter(TTSJobModel.id == job_id)
.first()
)
if model is None:
return False
model.status = "deleted"
self.session.commit()
return True
def list_by_user(
self,
user_id: str,
*,
status: Optional[str] = None,
limit: int = 50,
offset: int = 0,
) -> List[TTSJob]:
query = self.session.query(TTSJobModel).filter(
TTSJobModel.user_id == user_id,
TTSJobModel.status != "deleted",
)
if status:
query = query.filter(TTSJobModel.status == status)
query = query.order_by(TTSJobModel.created_at.desc())
models = query.offset(offset).limit(limit).all()
return [self._model_to_entity(m) for m in models]
def count_by_user(self, user_id: str, *, status: Optional[str] = None) -> int:
query = (
self.session.query(TTSJobModel)
.filter(
TTSJobModel.user_id == user_id,
TTSJobModel.status != "deleted",
)
)
if status:
query = query.filter(TTSJobModel.status == status)
return query.count()
def list_by_profile(
self,
voice_clone_profile_id: str,
*,
status: Optional[str] = None,
limit: int = 50,
offset: int = 0,
) -> List[TTSJob]:
query = self.session.query(TTSJobModel).filter(
TTSJobModel.voice_clone_profile_id == voice_clone_profile_id,
TTSJobModel.status != "deleted",
)
if status:
query = query.filter(TTSJobModel.status == status)
query = query.order_by(TTSJobModel.created_at.desc())
models = query.offset(offset).limit(limit).all()
return [self._model_to_entity(m) for m in models]
@staticmethod
def _model_to_entity(model: TTSJobModel) -> TTSJob:
return TTSJob(
id=model.id,
user_id=model.user_id,
input_text=model.input_text or "",
voice_id=model.voice_id or "",
voice_model=model.voice_model or "",
project_id=model.project_id or "",
voice_clone_profile_id=model.voice_clone_profile_id or "",
status=model.status,
output_audio_url=model.output_audio_url or "",
output_audio_key=model.output_audio_key or "",
duration=model.duration or 0.0,
file_size=model.file_size or 0,
sample_rate=model.sample_rate or 22050,
format=model.format or "mp3",
error_message=model.error_message or "",
retry_count=model.retry_count or 0,
max_retries=model.max_retries or 3,
metadata=model.metadata_ or {},
started_at=model.started_at,
completed_at=model.completed_at,
created_at=model.created_at,
updated_at=model.updated_at,
)
+1
View File
@@ -0,0 +1 @@
"""TTS Job application layer."""
+110
View File
@@ -0,0 +1,110 @@
"""TTS Job use cases."""
from __future__ import annotations
from typing import List, Optional
from packages.domain.tts_job import TTSJob
from packages.ports.tts_job_repository import TTSJobRepository
class TTSJobNotFoundError(Exception):
"""TTS 任务未找到。"""
pass
class CreateTTSJobUseCase:
"""创建 TTS 合成任务。"""
def __init__(self, repository: TTSJobRepository) -> None:
self.repository = repository
def execute(
self,
user_id: str,
input_text: str,
*,
voice_id: str = "",
voice_model: str = "",
project_id: str = "",
voice_clone_profile_id: str = "",
sample_rate: int = 22050,
format: str = "mp3",
max_retries: int = 3,
metadata: Optional[dict] = None,
) -> TTSJob:
"""创建 TTS 合成任务,状态为 pending。"""
job = TTSJob.create(
user_id=user_id,
input_text=input_text,
voice_id=voice_id,
voice_model=voice_model,
project_id=project_id,
voice_clone_profile_id=voice_clone_profile_id,
sample_rate=sample_rate,
format=format,
max_retries=max_retries,
metadata=metadata,
)
return self.repository.create(job)
class ListTTSJobsUseCase:
"""列出用户的 TTS 合成任务。"""
def __init__(self, repository: TTSJobRepository) -> None:
self.repository = repository
def execute(
self,
user_id: str,
*,
status: Optional[str] = None,
skip: int = 0,
limit: int = 50,
) -> tuple[List[TTSJob], int]:
items = self.repository.list_by_user(
user_id, status=status, limit=limit, offset=skip
)
total = self.repository.count_by_user(user_id, status=status)
return items, total
class GetTTSJobUseCase:
"""获取 TTS 任务详情。"""
def __init__(self, repository: TTSJobRepository) -> None:
self.repository = repository
def execute(self, job_id: str, user_id: str) -> TTSJob:
job = self.repository.get(job_id)
if job is None or job.user_id != user_id:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
return job
class GetTTSJobStatusUseCase:
"""查询 TTS 任务状态(用于轮询)。"""
def __init__(self, repository: TTSJobRepository) -> None:
self.repository = repository
def execute(self, job_id: str, user_id: str) -> TTSJob:
job = self.repository.get(job_id)
if job is None or job.user_id != user_id:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
return job
class DeleteTTSJobUseCase:
"""删除 TTS 任务(软删除)。"""
def __init__(self, repository: TTSJobRepository) -> None:
self.repository = repository
def execute(self, job_id: str, user_id: str) -> bool:
job = self.repository.get(job_id)
if job is None or job.user_id != user_id:
return False
return self.repository.delete(job_id)
+232
View File
@@ -0,0 +1,232 @@
"""TTS 合成 API 单元测试。"""
from __future__ import annotations
import pytest
from datetime import datetime, timezone
from unittest.mock import MagicMock
from packages.domain.tts_job import TTSJob, TTSJobStatus
from packages.application.tts_job.use_cases import (
CreateTTSJobUseCase,
DeleteTTSJobUseCase,
GetTTSJobStatusUseCase,
GetTTSJobUseCase,
ListTTSJobsUseCase,
TTSJobNotFoundError,
)
def _make_job(**kwargs) -> TTSJob:
defaults = {
"id": "test_job_001",
"user_id": "user_001",
"input_text": "测试文本",
"voice_id": "voice_001",
"voice_model": "",
"project_id": "",
"voice_clone_profile_id": "",
"status": TTSJobStatus.PENDING,
"output_audio_url": "",
"output_audio_key": "",
"duration": 0.0,
"file_size": 0,
"sample_rate": 22050,
"format": "mp3",
"error_message": "",
"retry_count": 0,
"max_retries": 3,
"metadata": {},
"started_at": None,
"completed_at": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
defaults.update(kwargs)
return TTSJob(**defaults)
class TestCreateTTSJobUseCase:
"""测试创建 TTS 合成任务。"""
def test_create_success(self) -> None:
"""正常创建。"""
repo = MagicMock()
repo.create.side_effect = lambda j: j
use_case = CreateTTSJobUseCase(repo)
job = use_case.execute(
user_id="user_001",
input_text="你好世界",
voice_id="voice_001",
)
assert job.user_id == "user_001"
assert job.input_text == "你好世界"
assert job.status == TTSJobStatus.PENDING
repo.create.assert_called_once()
def test_create_with_metadata(self) -> None:
"""带元数据创建。"""
repo = MagicMock()
repo.create.side_effect = lambda j: j
use_case = CreateTTSJobUseCase(repo)
job = use_case.execute(
user_id="user_001",
input_text="测试",
metadata={"source": "api"},
)
assert job.metadata == {"source": "api"}
def test_create_empty_text_raises(self) -> None:
"""空文本应报错。"""
repo = MagicMock()
use_case = CreateTTSJobUseCase(repo)
with pytest.raises(ValueError, match="input_text"):
use_case.execute(user_id="user_001", input_text=" ")
class TestListTTSJobsUseCase:
"""测试列出 TTS 合成任务。"""
def test_list_empty(self) -> None:
"""空列表。"""
repo = MagicMock()
repo.list_by_user.return_value = []
repo.count_by_user.return_value = 0
use_case = ListTTSJobsUseCase(repo)
items, total = use_case.execute("user_001")
assert items == []
assert total == 0
repo.list_by_user.assert_called_once_with("user_001", status=None, limit=50, offset=0)
def test_list_with_pagination(self) -> None:
"""分页查询。"""
repo = MagicMock()
jobs = [_make_job(id=f"job_{i}") for i in range(3)]
repo.list_by_user.return_value = jobs
repo.count_by_user.return_value = 10
use_case = ListTTSJobsUseCase(repo)
items, total = use_case.execute("user_001", skip=5, limit=3)
assert len(items) == 3
assert total == 10
repo.list_by_user.assert_called_once_with("user_001", status=None, limit=3, offset=5)
def test_list_with_status_filter(self) -> None:
"""按状态过滤。"""
repo = MagicMock()
repo.list_by_user.return_value = []
repo.count_by_user.return_value = 0
use_case = ListTTSJobsUseCase(repo)
use_case.execute("user_001", status="completed")
repo.list_by_user.assert_called_once_with("user_001", status="completed", limit=50, offset=0)
class TestGetTTSJobUseCase:
"""测试获取 TTS 任务详情。"""
def test_get_success(self) -> None:
"""正常获取。"""
job = _make_job()
repo = MagicMock()
repo.get.return_value = job
use_case = GetTTSJobUseCase(repo)
result = use_case.execute("test_job_001", "user_001")
assert result.id == "test_job_001"
assert result.user_id == "user_001"
def test_get_not_found(self) -> None:
"""任务不存在。"""
repo = MagicMock()
repo.get.return_value = None
use_case = GetTTSJobUseCase(repo)
with pytest.raises(TTSJobNotFoundError):
use_case.execute("nonexistent", "user_001")
def test_get_wrong_user(self) -> None:
"""用户不匹配。"""
job = _make_job(user_id="other_user")
repo = MagicMock()
repo.get.return_value = job
use_case = GetTTSJobUseCase(repo)
with pytest.raises(TTSJobNotFoundError):
use_case.execute("test_job_001", "user_001")
class TestGetTTSJobStatusUseCase:
"""测试查询 TTS 任务状态。"""
def test_status_success(self) -> None:
"""正常查询状态。"""
job = _make_job(status=TTSJobStatus.COMPLETED, output_audio_url="https://example.com/audio.mp3")
repo = MagicMock()
repo.get.return_value = job
use_case = GetTTSJobStatusUseCase(repo)
result = use_case.execute("test_job_001", "user_001")
assert result.status == TTSJobStatus.COMPLETED
assert result.output_audio_url == "https://example.com/audio.mp3"
def test_status_not_found(self) -> None:
"""任务不存在。"""
repo = MagicMock()
repo.get.return_value = None
use_case = GetTTSJobStatusUseCase(repo)
with pytest.raises(TTSJobNotFoundError):
use_case.execute("nonexistent", "user_001")
class TestDeleteTTSJobUseCase:
"""测试删除 TTS 任务。"""
def test_delete_success(self) -> None:
"""正常删除。"""
job = _make_job()
repo = MagicMock()
repo.get.return_value = job
repo.delete.return_value = True
use_case = DeleteTTSJobUseCase(repo)
result = use_case.execute("test_job_001", "user_001")
assert result is True
repo.delete.assert_called_once_with("test_job_001")
def test_delete_not_found(self) -> None:
"""任务不存在。"""
repo = MagicMock()
repo.get.return_value = None
use_case = DeleteTTSJobUseCase(repo)
result = use_case.execute("nonexistent", "user_001")
assert result is False
def test_delete_wrong_user(self) -> None:
"""用户不匹配。"""
job = _make_job(user_id="other_user")
repo = MagicMock()
repo.get.return_value = job
use_case = DeleteTTSJobUseCase(repo)
result = use_case.execute("test_job_001", "user_001")
assert result is False