diff --git a/alembic/versions/020_add_tts_jobs_table.py b/alembic/versions/020_add_tts_jobs_table.py new file mode 100644 index 000000000..5b854e955 --- /dev/null +++ b/alembic/versions/020_add_tts_jobs_table.py @@ -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") diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 1b814ac7c..39a41ba47 100755 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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"], +) diff --git a/apps/api/app/api/routes/tts.py b/apps/api/app/api/routes/tts.py new file mode 100644 index 000000000..00d6a886b --- /dev/null +++ b/apps/api/app/api/routes/tts.py @@ -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) diff --git a/apps/api/app/schemas/tts.py b/apps/api/app/schemas/tts.py new file mode 100644 index 000000000..b681c310c --- /dev/null +++ b/apps/api/app/schemas/tts.py @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 99497f8fb..c1adac288 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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)) \ No newline at end of file diff --git a/packages/adapters/sqlalchemy_impl/tts_job_repository.py b/packages/adapters/sqlalchemy_impl/tts_job_repository.py new file mode 100644 index 000000000..be4cba6bf --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/tts_job_repository.py @@ -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, + ) diff --git a/packages/application/tts_job/__init__.py b/packages/application/tts_job/__init__.py new file mode 100644 index 000000000..439072c1b --- /dev/null +++ b/packages/application/tts_job/__init__.py @@ -0,0 +1 @@ +"""TTS Job application layer.""" diff --git a/packages/application/tts_job/use_cases.py b/packages/application/tts_job/use_cases.py new file mode 100644 index 000000000..063acbc40 --- /dev/null +++ b/packages/application/tts_job/use_cases.py @@ -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) diff --git a/tests/unit/test_tts_api.py b/tests/unit/test_tts_api.py new file mode 100644 index 000000000..373a41107 --- /dev/null +++ b/tests/unit/test_tts_api.py @@ -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