"""TTS 合成 API 路由。""" from __future__ import annotations import logging from typing import Optional from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app from app.dependencies import ( get_audio_url_signer, get_cosyvoice_service, get_db_session, get_user_repository, get_voice_clone_profile_repository, get_voice_library_repository, ) from app.schemas.tts import ( ListTTSJobResponse, SaveToLibraryRequest, SaveToLibraryResponse, TTSJobResponse, TTSStatusResponse, TTSSynthesizeRequest, TTSSynthesizeResponse, ) from fastapi import APIRouter, Depends, HTTPException, Query, Response, WebSocket, WebSocketDisconnect, status from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.tts_job_repository import ( SQLAlchemyTTSJobRepository, ) from packages.adapters.sqlalchemy_impl.voice_library_repository import SQLAlchemyVoiceLibraryRepository from packages.application.cosyvoice_service import CosyVoiceService from packages.application.tts_job.streaming_service import TTSStreamingService from packages.application.tts_job.use_cases import ( CreateTTSJobUseCase, DeleteTTSJobUseCase, GetTTSJobStatusUseCase, GetTTSJobUseCase, ListTTSJobsUseCase, TTSJobNotFoundError, ) from packages.application.tts_job.workflow import TTSWorkflowService from packages.application.voice_library.commands import CreateVoiceLibraryCommand from packages.application.voice_library.use_cases import ( CreateVoiceLibraryUseCase, QuotaExceededError, ) from packages.domain.voice_presets import list_voices from packages.ports.user_repository import UserRepository logger = logging.getLogger(__name__) router = APIRouter() @router.get("/presets", summary="获取预设音色列表") def list_preset_voices( gender: Optional[str] = Query(None, description="按性别筛选: male/female/child"), style: Optional[str] = Query(None, description="按风格筛选: stable/lively/customer_service/narration/news/story"), keyword: Optional[str] = Query(None, description="按关键词搜索"), _user: AuthenticatedUser = Depends(get_current_user), ) -> list[dict]: """获取可用的预设音色列表。 用于配音功能的音色选择。 """ voices = list_voices(gender=gender, style=style, keyword=keyword) return [ { "voice_id": v.voice_id, "name": v.name, "gender": v.gender.value, "style": v.style.value, "description": v.description, "default_speed": v.default_speed, "default_pitch": v.default_pitch, "sample_rate": v.sample_rate, "language": v.language, } for v in voices ] def _get_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTTSJobRepository: return SQLAlchemyTTSJobRepository(session) def _to_response(job, sign_url=None) -> TTSJobResponse: output_url = job.output_audio_url if sign_url and output_url: output_url = sign_url(output_url) 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=output_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), cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service), voice_clone_repo=Depends(get_voice_clone_profile_repository), ) -> TTSSynthesizeResponse: """发起 TTS 合成任务。 创建 TTS 任务 → 提交 CosyVoice 合成 → 触发 Celery 异步轮询。 与音色克隆接口保持一致:CosyVoice 失败时不抛 500,而是返回 201 + failed 状态任务记录。 """ user_id = authenticated_user.user.id # 校验 voice_clone_profile_id 归属(防止越权使用他人克隆音色) if request.voice_clone_profile_id: profile = voice_clone_repo.get(request.voice_clone_profile_id) if profile is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail="Voice clone profile not found", ) if profile.user_id != user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to voice clone profile", ) 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_, ) # 提交 CosyVoice 合成任务 workflow = TTSWorkflowService( repository=repository, cosyvoice_service=cosyvoice_service, ) try: job = workflow.start_synthesis(job.id) except Exception as e: # 兜底:workflow 内部已捕获 CosyVoiceError / ValueError, # 但 DB 异常、网络异常等意外错误可能逃逸。 # 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。 logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True) try: job = workflow.process_synthesis_failure(job.id, str(e)) except Exception as inner_e: logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}") # 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询 if job.status.value == "processing": # 分段合成任务 vs 普通单段任务 segment_task_ids = (job.metadata or {}).get("segment_task_ids", []) is_segment = len(segment_task_ids) > 0 try: if is_segment: celery_app.send_task("worker.process_tts_segment_synthesis", args=[job.id]) else: celery_app.send_task("worker.process_tts_synthesis", args=[job.id]) except Exception as e: # Celery 调度失败,标记 job 为 failed try: workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}") except Exception as inner_e: logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}") 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), sign_url=Depends(get_audio_url_signer), ) -> 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, sign_url) 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), sign_url=Depends(get_audio_url_signer), ) -> TTSJobResponse: """获取 TTS 任务详情。""" user_id = authenticated_user.user.id use_case = GetTTSJobUseCase(repository) try: job = use_case.execute(job_id, user_id) except TTSJobNotFoundError as _e: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found") from _e return _to_response(job, sign_url) @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), sign_url=Depends(get_audio_url_signer), ) -> TTSStatusResponse: """查询 TTS 合成状态(用于前端轮询)。""" user_id = authenticated_user.user.id use_case = GetTTSJobStatusUseCase(repository) try: job = use_case.execute(job_id, user_id) except TTSJobNotFoundError as _e: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found") from _e output_url = job.output_audio_url if output_url: output_url = sign_url(output_url) return TTSStatusResponse( id=job.id, status=job.status, output_audio_url=output_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_model=None, 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 @router.post( "/jobs/{job_id}/save-to-library", response_model=SaveToLibraryResponse, status_code=status.HTTP_201_CREATED, ) def save_tts_job_to_library( job_id: str, request: SaveToLibraryRequest = SaveToLibraryRequest(), authenticated_user: AuthenticatedUser = Depends(get_current_user), tts_repository: SQLAlchemyTTSJobRepository = Depends(_get_repository), voice_library_repository: SQLAlchemyVoiceLibraryRepository = Depends(get_voice_library_repository), user_repository: UserRepository = Depends(get_user_repository), sign_url=Depends(get_audio_url_signer), ) -> SaveToLibraryResponse: """将已完成的 TTS 合成结果保存到配音库。 自动携带音色名、时长、语速等元信息。 """ user_id = authenticated_user.user.id # 获取 TTS job get_use_case = GetTTSJobUseCase(tts_repository) try: job = get_use_case.execute(job_id, user_id) except TTSJobNotFoundError as _e: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found") from _e # 校验已完成 if not job.is_completed: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail="TTS job is not completed yet", ) # 构建配音素材名称 name = request.name or f"TTS-{job.id[:8]}" # 构建元信息 metadata_ = { "source": "tts_job", "tts_job_id": job.id, "format": job.format, "sample_rate": job.sample_rate, } if job.metadata: # 保留原始 job 的有用元信息 for key in ("speed", "language"): if key in job.metadata: metadata_[key] = job.metadata[key] # 获取用户套餐(用于配额检查) user = user_repository.find_by_id(user_id) plan_name = getattr(user, "subscription_plan", "free") if user else "free" # 构建命令并执行 command = CreateVoiceLibraryCommand( user_id=user_id, name=name, text=job.input_text, voice_provider="cosyvoice", voice_id=job.voice_id, voice_name=job.voice_model or "", audio_url=job.output_audio_url, duration=job.duration, file_size=job.file_size, status="completed", project_id=job.project_id or "", tags=[], metadata_=metadata_, ) use_case = CreateVoiceLibraryUseCase(voice_library_repository) try: item = use_case.execute(command, plan_name=plan_name or "free") except QuotaExceededError as exc: raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail=f"配音库配额已满({exc.used}/{exc.limit}),请升级套餐", ) from exc return SaveToLibraryResponse( id=item.id, name=item.name, audio_url=sign_url(item.audio_url) if item.audio_url else "", duration=item.duration, voice_id=item.voice_id, voice_name=item.voice_name, status=item.status, ) @router.websocket("/ws/tts/stream") async def tts_websocket_stream( websocket: WebSocket, cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service), ) -> None: """WebSocket 流式 TTS 合成。 协议: 1. 客户端发送 JSON 文本帧: {"text": "...", "voice_id": "...", ...} 2. 服务端发送 JSON 状态帧 + 二进制音频帧 3. 完成时发送 JSON 结束帧 """ await websocket.accept() try: message = await websocket.receive_json() params = { "text": message.get("text", ""), "voice_id": message.get("voice_id", ""), "sample_rate": message.get("sample_rate", 0), "format": message.get("format", "mp3"), "speed": message.get("speed", 1.0), } streaming_service = TTSStreamingService(cosyvoice_service) await streaming_service.synthesize_and_stream(websocket, params) except WebSocketDisconnect: logger.info("WebSocket 客户端主动断开连接") except Exception as e: logger.error(f"WebSocket 流式合成异常: {e}", exc_info=True) try: await websocket.send_json({"type": "error", "message": f"服务异常: {e}"}) except Exception as send_err: logger.warning("WebSocket 错误消息发送失败(连接可能已断开): %s", send_err)