"""TTS Job workflow orchestration — Phase 3. 编排 TTS 合成的完整流程: 1. 创建 TTSJob(pending) 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. 返回 job(Celery 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}")