d6ab413dcd
CI/CD Pipeline / Deploy Staging (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Frontend Lint (push) Failing after 47h57m37s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 47h57m37s
- 新增 text_splitter.py: 长文本智能分段(句子边界 + 短段合并) - 新增 audio_merger.py: FFmpeg concat 音频合并器 - workflow.py: 分段合成完整流程(同步合并 / 异步轮询 / 失败处理) - tts_synthesis.py: 新增 process_tts_segment_synthesis Celery 任务 - tts.py: 路由层自动识别分段任务并分发到对应 Celery task - 23 个单元测试全部通过,P0 回归测试无退化
502 lines
18 KiB
Python
502 lines
18 KiB
Python
"""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}")
|