"""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 import time from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Optional import httpx from packages.application.cosyvoice_service import ( CosyVoiceAuthError, CosyVoiceError, CosyVoiceService, ) from packages.application.tts_job.audio_merger import AudioMergeError, 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 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: # 下载临时音频 resp = httpx.get(temp_url, timeout=60.0, follow_redirects=True) resp.raise_for_status() audio_data = resp.content # 上传到 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 合成任务并处理结果。 从 job.metadata 获取 task_id,调用 CosyVoiceService.poll_synthesize_task() 轮询状态,然后通过 process_synthesis_result / process_synthesis_failure 更新 job。 供 Celery 后台任务调用。 """ job = self.repository.get(job_id) if job is None: raise TTSJobNotFoundError(f"TTS job {job_id} not found") # 检查是否为分段合成任务 segment_task_ids = (job.metadata or {}).get("segment_task_ids", []) if segment_task_ids: return self._poll_segment_tasks(job) task_id = (job.metadata or {}).get("cosyvoice_task_id", "") if not task_id: raise ValueError(f"TTSJob {job_id} has no cosyvoice_task_id in metadata") 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), ) 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 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) # 下载分段音频到临时文件 resp = httpx.get(audio_url, timeout=60.0, follow_redirects=True) resp.raise_for_status() seg_path = os.path.join(temp_dir, f"seg_{idx:03d}.{job.format}") with open(seg_path, "wb") as f: f.write(resp.content) 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: """轮询所有分段异步任务,全部完成后合并音频。""" 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) poll_start = time.monotonic() poll_timeout = 300.0 # 分段任务超时更长 poll_interval = 2.0 while time.monotonic() - poll_start < poll_timeout: all_done = True results: list[dict | None] = [None] * segment_count for idx, task_id in enumerate(segment_task_ids): # 已经有音频的分段跳过轮询 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, } continue try: result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=poll_timeout) results[idx] = 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 results[idx] is None: all_done = False if all_done and all(r is not None for r in results): # 所有分段完成,下载合并 try: merged_data, total_duration = self._download_and_merge_segments(results, job) # 转存 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)}") return job except Exception as e: self._handle_segment_failure(job, f"分段合并失败: {e}") return self.repository.get(job.id) # 等待后重试 time.sleep(poll_interval) # 超时 self._handle_segment_failure(job, "分段合成轮询超时(300 秒)") 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}")