389d1e4401
CI/CD Pipeline / Check if frontend-only change (push) Has been cancelled
CI/CD Pipeline / Validate - Code Quality (push) Has been cancelled
CI/CD Pipeline / Validate - Type Check (mypy) (push) Has been cancelled
CI/CD Pipeline / Validate - Migration (alembic) (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (push) Has been cancelled
CI/CD Pipeline / PR Build API Image (push) Has been cancelled
CI/CD Pipeline / PR Build Web Image (push) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (push) Has been cancelled
CI/CD Pipeline / Build Staging API Image (push) Has been cancelled
CI/CD Pipeline / Build Staging Web Image (push) Has been cancelled
CI/CD Pipeline / Build Staging Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production API Image (push) Has been cancelled
CI/CD Pipeline / Build Production Web Image (push) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (push) Has been cancelled
408 lines
15 KiB
Python
Executable File
408 lines
15 KiB
Python
Executable File
"""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)
|