Files
xiaoxia-saas/packages/application/tts_job/workflow.py
T
灵应 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
feat: P1 长文本分段合成
- 新增 text_splitter.py: 长文本智能分段(句子边界 + 短段合并)
- 新增 audio_merger.py: FFmpeg concat 音频合并器
- workflow.py: 分段合成完整流程(同步合并 / 异步轮询 / 失败处理)
- tts_synthesis.py: 新增 process_tts_segment_synthesis Celery 任务
- tts.py: 路由层自动识别分段任务并分发到对应 Celery task
- 23 个单元测试全部通过,P0 回归测试无退化
2026-07-07 12:56:52 +08:00

502 lines
18 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
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. 返回 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 合成任务并处理结果。
从 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}")