Files
xiaoxia-saas/packages/application/tts_job/workflow.py
T
xiaoxia 09d2b12ea8
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Frontend Lint (push) Successful in 57s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 3m7s
CI/CD Pipeline / Unit Tests (push) Successful in 3m13s
CI/CD Pipeline / Integration Tests (push) Successful in 1m22s
CI Build & Deploy Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 1m32s
CI Build & Deploy Pipeline / Build Staging API Image (push) Successful in 18m38s
CI Build & Deploy Pipeline / Build Staging Web Image (push) Successful in 19s
CI Build & Deploy Pipeline / Build Staging Worker Image (push) Successful in 8m7s
CI Build & Deploy Pipeline / Build Production API Image (push) Has been skipped
CI Build & Deploy Pipeline / Build Production Web Image (push) Has been skipped
CI Build & Deploy Pipeline / Build Production Worker Image (push) Has been skipped
CI Build & Deploy Pipeline / Deploy Production (push) Has been skipped
CI Build & Deploy Pipeline / Production Browser E2E (push) Has been skipped
CI Build & Deploy Pipeline / Staging E2E Tests (push) Successful in 2m16s
CI Build & Deploy Pipeline / Staging API Integration Tests (push) Successful in 4m35s
feat(ci): P1-1 Phase2 后端启用F401+F841并修复存量 (#470)
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
2026-07-17 14:06:50 +08:00

584 lines
23 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""TTS Job workflow orchestration — Phase 3.
编排 TTS 合成的完整流程:
1. 创建 TTSJobpending
2. 提交 CosyVoice 合成任务
3. 轮询处理合成结果(成功/失败)
4. 重试失败的合成
"""
from __future__ import annotations
import io
import logging
import os
import shutil
import tempfile
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Optional
from packages.application.cosyvoice_service import CosyVoiceAuthError, CosyVoiceError, CosyVoiceService
from packages.application.tts_job.audio_merger import AudioMerger
from packages.application.tts_job.text_splitter import split_text
from packages.domain.tts_job import TTSJob, TTSJobStatus
from packages.ports.tts_job_repository import TTSJobRepository
from packages.shared.storage import SharedStorageService, get_shared_storage_service
from packages.shared.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
safe_download_bytes,
safe_download_file,
)
logger = logging.getLogger(__name__)
# 长文本分段阈值:超过此字符数自动分段合成
_SEGMENT_THRESHOLD = 500
# 分段并发上限
_MAX_SEGMENT_WORKERS = 5
class TTSWorkflowError(Exception):
"""TTS 合成工作流异常。"""
pass
class TTSJobNotFoundError(Exception):
"""TTS 任务未找到。"""
pass
class TTSWorkflowService:
"""TTS 合成工作流编排服务。
协调 TTSJobRepository + CosyVoiceService
实现完整的 TTS 合成生命周期管理。
"""
def __init__(
self,
repository: TTSJobRepository,
cosyvoice_service: CosyVoiceService,
storage_service: Optional[SharedStorageService] = None,
) -> None:
self.repository = repository
self.cosyvoice_service = cosyvoice_service
self._storage_service = storage_service
@property
def _storage(self) -> SharedStorageService:
if self._storage_service is None:
self._storage_service = get_shared_storage_service()
return self._storage_service
def _transfer_audio_to_oss(
self,
temp_url: str,
user_id: str,
job_id: str,
audio_format: str = "mp3",
) -> tuple[str, str]:
"""下载 CosyVoice 临时音频并转存到 OSS。
Returns:
(permanent_url, storage_key) 元组。
转存失败时回退到原始临时 URL,storage_key 为空字符串。
"""
storage_key = f"tts-outputs/{user_id}/{job_id}.{audio_format}"
content_type_map = {
"mp3": "audio/mpeg",
"wav": "audio/wav",
"pcm": "audio/pcm",
"opus": "audio/opus",
}
content_type = content_type_map.get(audio_format, "application/octet-stream")
try:
# 安全下载临时音频(SSRF 防护 + 大小限制 + 重定向校验)
audio_data = safe_download_bytes(
temp_url,
purpose="tts_audio_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
# 上传到 OSS
file_obj = io.BytesIO(audio_data)
permanent_url = self._storage.upload_file(file_obj, storage_key, content_type=content_type)
logger.info(f"音频转存 OSS 成功: job_id={job_id}, " f"storage_key={storage_key}, size={len(audio_data)}")
return permanent_url, storage_key
except Exception as e:
logger.warning(f"音频转存 OSS 失败,使用临时 URL: " f"job_id={job_id}, error={e}")
return temp_url, ""
def start_synthesis(
self,
job_id: str,
) -> TTSJob:
"""启动 TTS 合成流程。
1. 获取 pending 状态的 TTSJob
2. 标记为 processing
3. 提交 CosyVoice 合成任务
4. 保存 task_id 到 metadata
5. 返回 jobCelery task 由调用方触发)
Args:
job_id: TTSJob ID
Returns:
TTSJob: 更新后的 job
Raises:
TTSJobNotFoundError: job 不存在
TTSWorkflowError: CosyVoice 提交失败
"""
job = self.repository.get(job_id)
if job is None:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
# 标记为 processing
job.mark_processing()
job = self.repository.update(job)
# 长文本自动分段合成
if len(job.input_text) > _SEGMENT_THRESHOLD:
return self._start_segment_synthesis(job)
try:
submit_result = self.cosyvoice_service.submit_synthesize_task(
text=job.input_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
)
# 保存 task_id / request_id 到 metadata
job_metadata = dict(job.metadata)
job_metadata["cosyvoice_task_id"] = submit_result.get("task_id", "")
job_metadata["cosyvoice_request_id"] = submit_result.get("request_id", "")
# 如果 CosyVoice 同步返回了 audio_url,转存 OSS 后标记完成
audio_url = submit_result.get("audio_url", "")
if audio_url:
permanent_url, storage_key = self._transfer_audio_to_oss(audio_url, job.user_id, job.id, job.format)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=submit_result.get("duration", 0.0),
file_size=submit_result.get("file_size", 0),
)
job.metadata = job_metadata
job = self.repository.update(job)
logger.info(f"TTS 合成同步完成: job_id={job.id}, audio_url={permanent_url}")
return job
job.metadata = job_metadata
job = self.repository.update(job)
logger.info(f"TTS 合成任务已提交: job_id={job.id}, " f"task_id={submit_result.get('task_id')}")
except (CosyVoiceError, CosyVoiceAuthError) as e:
job.mark_failed(str(e))
job = self.repository.update(job)
logger.error(f"TTS 合成提交失败: job_id={job.id}, error={e}")
except ValueError as e:
job.mark_failed(str(e))
job = self.repository.update(job)
logger.error(f"TTS 合成参数错误: job_id={job.id}, error={e}")
return job
def poll_and_process_synthesis(self, job_id: str, timeout: float = 120.0) -> TTSJob:
"""轮询/检查 CosyVoice 合成任务并处理结果.
新 CosyVoice SpeechSynthesizer 非流式接口是同步的,
start_synthesis 阶段通常已经完成. 此方法用于:
1. job 已 completed → 直接返回(同步路径已处理)
2. job 仍在 processing → 重新提交合成(兜底)
3. 分段任务 → 检查分段状态
供 Celery 后台任务调用。
"""
job = self.repository.get(job_id)
if job is None:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
# 已完成直接返回(同步路径在 start_synthesis 里已处理)
if job.status == TTSJobStatus.COMPLETED.value:
logger.info(f"TTS 任务已完成,跳过轮询: job_id={job_id}")
return job
# 检查是否为分段合成任务
segment_task_ids = (job.metadata or {}).get("segment_task_ids", [])
if segment_task_ids:
return self._poll_segment_tasks(job)
# 单段模式:同步接口下通常不会走到这里,
# 但如果因为异常导致仍在 processing,重新提交一次
task_id = (job.metadata or {}).get("cosyvoice_task_id", "")
# 新接口(同步):没有 task_id,重新合成
if not task_id:
logger.info(f"TTS 任务无 task_id,重新同步合成: job_id={job_id}")
return self._resynthesize_and_complete(job)
# 旧接口遗留的 task_id,尝试轮询(兼容过渡)
try:
result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=timeout)
return self.process_synthesis_result(
job_id,
audio_url=result["audio_url"],
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
)
except CosyVoiceError:
# 旧接口轮询失败,重新同步合成
logger.warning(f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}")
return self._resynthesize_and_complete(job)
def process_synthesis_result(
self,
job_id: str,
audio_url: str,
*,
duration: float = 0.0,
file_size: int = 0,
) -> TTSJob:
"""处理合成成功结果。
Args:
job_id: TTSJob ID
audio_url: 输出音频 URL
duration: 音频时长
file_size: 文件大小
Returns:
TTSJob: 更新后的 job
Raises:
TTSJobNotFoundError: job 不存在
"""
job = self.repository.get(job_id)
if job is None:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
# 转存音频到 OSS,获取永久 URL
permanent_url, storage_key = self._transfer_audio_to_oss(audio_url, job.user_id, job.id, job.format)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=duration,
file_size=file_size,
)
job = self.repository.update(job)
logger.info(f"TTS 合成成功: job_id={job_id}, audio_url={permanent_url}")
return job
def _resynthesize_and_complete(self, job: TTSJob) -> TTSJob:
"""重新同步合成并完成任务(兜底路径).
当 poll_and_process_synthesis 发现 job 仍在 processing 且无 task_id 时,
重新调用同步合成接口,转存 OSS 后标记完成。
"""
try:
# 从 metadata 读取合成参数(兼容旧数据,无则用默认值)
job_metadata = job.metadata or {}
speed = float(job_metadata.get("speed", 1.0))
volume = int(job_metadata.get("volume", 50))
result = self.cosyvoice_service.submit_synthesize_task(
text=job.input_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
speed=speed,
volume=volume,
)
audio_url = result.get("audio_url", "")
if not audio_url:
raise CosyVoiceError("重新合成未返回 audio_url")
return self.process_synthesis_result(
job.id,
audio_url=audio_url,
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
)
except Exception as e:
logger.error(f"重新同步合成失败: job_id={job.id}, error={e}")
return self.process_synthesis_failure(job.id, str(e))
def process_synthesis_failure(self, job_id: str, error_message: str) -> TTSJob:
"""处理合成失败结果。
Args:
job_id: TTSJob ID
error_message: 错误信息
Returns:
TTSJob: 更新后的 job
Raises:
TTSJobNotFoundError: job 不存在
"""
job = self.repository.get(job_id)
if job is None:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
job.mark_failed(error_message)
job = self.repository.update(job)
logger.error(f"TTS 合成失败: job_id={job_id}, error={error_message}")
return job
# ── P1: 长文本分段合成 ─────────────────────────────────────
def _upload_merged_to_oss(
self, merged_data: bytes, user_id: str, job_id: str, audio_format: str
) -> tuple[str, str]:
"""上传合并后的音频数据到 OSS。
Returns:
(permanent_url, storage_key) 元组。
上传失败时返回 ("", "")。
"""
storage_key = f"tts-outputs/{user_id}/{job_id}.{audio_format}"
content_type_map = {
"mp3": "audio/mpeg",
"wav": "audio/wav",
"pcm": "audio/pcm",
"opus": "audio/opus",
}
content_type = content_type_map.get(audio_format, "application/octet-stream")
try:
file_obj = io.BytesIO(merged_data)
permanent_url = self._storage.upload_file(file_obj, storage_key, content_type=content_type)
return permanent_url, storage_key
except Exception as e:
logger.warning(f"分段合并音频转存 OSS 失败: job_id={job_id}, error={e}")
return "", ""
def _start_segment_synthesis(self, job: TTSJob) -> TTSJob:
"""长文本分段合成入口。
将文本分段后并发提交到 CosyVoice,根据同步/异步结果走不同路径。
"""
segments = split_text(job.input_text, max_chars=_SEGMENT_THRESHOLD)
logger.info(f"长文本分段合成: job_id={job.id}, " f"原文={len(job.input_text)}字, 段数={len(segments)}")
# 记录分段信息到 metadata
job_metadata = dict(job.metadata)
job_metadata["segment_count"] = len(segments)
# 并发提交所有分段
results = self._submit_segments_concurrent(segments, job)
if results is None:
# 提交阶段已失败,_submit_segments_concurrent 内部已标记 failed
return self.repository.get(job.id)
# 判断同步还是异步
has_audio_urls = any(r.get("audio_url", "") for r in results)
has_task_ids = any(r.get("task_id", "") for r in results)
if has_audio_urls and not has_task_ids:
# 所有分段同步返回音频,直接合并
return self._process_segments_sync(job, results)
# 异步路径:保存各分段的 task_id 供后续轮询
segment_task_ids = [r.get("task_id", "") for r in results]
segment_audio_urls = [r.get("audio_url", "") for r in results]
job_metadata["segment_task_ids"] = segment_task_ids
job_metadata["segment_audio_urls"] = segment_audio_urls
job_metadata["segment_format"] = job.format
job.metadata = job_metadata
job = self.repository.update(job)
logger.info(f"分段合成任务已提交(异步): job_id={job.id}, " f"段数={len(segments)}")
return job
def _submit_segments_concurrent(self, segments: list[str], job: TTSJob) -> list[dict] | None:
"""并发提交分段合成任务。
Returns:
各分段的结果列表(保持顺序),提交失败时返回 None。
"""
max_workers = min(len(segments), _MAX_SEGMENT_WORKERS)
results: list[dict | None] = [None] * len(segments)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
future_to_idx = {}
for idx, segment_text in enumerate(segments):
future = executor.submit(
self.cosyvoice_service.submit_synthesize_task,
text=segment_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
)
future_to_idx[future] = idx
for future in as_completed(future_to_idx):
idx = future_to_idx[future]
try:
results[idx] = future.result()
except Exception as e:
logger.error(f"分段合成提交失败: job_id={job.id}, " f"segment={idx}, error={e}")
self._handle_segment_failure(job, f"分段 {idx + 1} 合成提交失败: {e}")
return None
return results # type: ignore[return-value]
def _process_segments_sync(self, job: TTSJob, results: list[dict]) -> TTSJob:
"""同步路径:所有分段已返回 audio_url,下载合并后转存 OSS。"""
merged_data, total_duration = self._download_and_merge_segments(results, job)
# 直接上传合并后的音频 bytes 到 OSS
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=total_duration,
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(f"分段合成完成: job_id={job.id}, " f"merged_size={len(merged_data)}, duration={total_duration:.1f}")
return job
def _download_and_merge_segments(self, results: list[dict], job: TTSJob) -> tuple[bytes, float]:
"""下载各分段音频并合并。
Returns:
(merged_audio_bytes, total_duration)
"""
temp_dir = tempfile.mkdtemp(prefix="tts_segments_")
try:
audio_paths: list[str] = []
total_duration = 0.0
for idx, result in enumerate(results):
audio_url = result.get("audio_url", "")
if not audio_url:
raise TTSWorkflowError(f"分段 {idx + 1} 没有返回 audio_url")
total_duration += result.get("duration", 0.0)
# 安全下载分段音频到临时文件(SSRF 防护 + 大小限制)
seg_path = os.path.join(temp_dir, f"seg_{idx:03d}.{job.format}")
safe_download_file(
audio_url,
seg_path,
purpose="tts_segment_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
audio_paths.append(seg_path)
# 合并
merger = AudioMerger()
merged_data = merger.merge(audio_paths, output_format=job.format)
return merged_data, total_duration
finally:
shutil.rmtree(temp_dir, ignore_errors=True)
def _poll_segment_tasks(self, job: TTSJob) -> TTSJob:
"""分段任务完成检查(适配新同步接口).
新 CosyVoice SpeechSynthesizer 非流式接口为同步接口,
分段任务在提交时应已同步返回 audio_url。
若历史任务处于 processing 且有 segment_task_ids 但缺少 audio_url
则对缺失分段重新同步合成,全部完成后合并音频。
"""
segment_task_ids: list[str] = (job.metadata or {}).get("segment_task_ids", [])
segment_audio_urls: list[str] = (job.metadata or {}).get("segment_audio_urls", [])
segment_count = len(segment_task_ids)
if segment_count == 0:
logger.warning(f"分段任务无 task_id: job_id={job.id}")
self._handle_segment_failure(job, "分段任务数据异常:无分段信息")
return self.repository.get(job.id)
# 从 metadata 读取合成参数
job_metadata = job.metadata or {}
speed = float(job_metadata.get("speed", 1.0))
volume = int(job_metadata.get("volume", 50))
# 分段文本(用于缺失段重新合成)
segments = split_text(job.input_text, max_chars=_SEGMENT_THRESHOLD)
results: list[dict | None] = [None] * segment_count
# 已有音频的分段直接用
for idx in range(segment_count):
if idx < len(segment_audio_urls) and segment_audio_urls[idx]:
results[idx] = {
"audio_url": segment_audio_urls[idx],
"duration": 0.0,
"file_size": 0,
}
# 找出缺失音频的分段索引
missing_indices = [i for i in range(segment_count) if results[i] is None]
if missing_indices:
logger.info(f"分段任务重新合成缺失段: job_id={job.id}, " f"缺失={len(missing_indices)}/{segment_count}")
# 并发重新合成缺失分段
max_workers = min(len(missing_indices), _MAX_SEGMENT_WORKERS)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
future_to_idx = {}
for idx in missing_indices:
segment_text = segments[idx] if idx < len(segments) else ""
future = executor.submit(
self.cosyvoice_service.submit_synthesize_task,
text=segment_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
speed=speed,
volume=volume,
)
future_to_idx[future] = idx
for future in as_completed(future_to_idx):
idx = future_to_idx[future]
try:
results[idx] = future.result()
except Exception as e:
logger.error(f"分段重新合成失败: job_id={job.id}, " f"segment={idx}, error={e}")
self._handle_segment_failure(job, f"分段 {idx + 1} 重新合成失败: {e}")
return self.repository.get(job.id)
# 所有分段完成,下载合并
if all(r is not None for r in results):
try:
merged_data, total_duration = self._download_and_merge_segments(results, job)
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=total_duration,
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(f"分段合成完成(重新合成路径): job_id={job.id}, " f"merged_size={len(merged_data)}")
return job
except Exception as e:
self._handle_segment_failure(job, f"分段合并失败: {e}")
return self.repository.get(job.id)
# 理论上不会到这里(全部重新合成要么成功要么失败)
self._handle_segment_failure(job, "分段合成结果不完整")
return self.repository.get(job.id)
def _handle_segment_failure(self, job: TTSJob, error_message: str) -> None:
"""分段合成失败处理。"""
job.mark_failed(error_message)
self.repository.update(job)
logger.error(f"分段合成失败: job_id={job.id}, error={error_message}")