阶段3 任务3.06:TTS 合成 API #166
@@ -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")
|
||||
@@ -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"],
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
"""TTS Job application layer."""
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user