102 lines
4.2 KiB
Python
102 lines
4.2 KiB
Python
import asyncio
|
||
import time
|
||
from typing import Dict
|
||
from utils.logger import session_logger
|
||
from services.interrupt_handler import interrupt_handler
|
||
from utils.text_splitter import should_skip_tts
|
||
from services.tts_config import TTSConfig
|
||
from services.tts_audio_cleaner import TTSAudioCleaner
|
||
from services.interruption_helper import InterruptionHelper
|
||
|
||
class TTSSynthesizer:
|
||
def __init__(self, device_id: str, session_id: str):
|
||
self.device_id = device_id
|
||
self.session_id = session_id
|
||
self.session_key = (device_id, session_id)
|
||
self.request_timeout = TTSConfig.get_request_timeout()
|
||
self.max_tts_errors = TTSConfig.get_max_tts_errors()
|
||
|
||
def check_interruption(self):
|
||
return interrupt_handler.is_interrupted(self.session_key)
|
||
|
||
async def synthesize_audio(
|
||
self,
|
||
tts_service,
|
||
text_queue: asyncio.Queue,
|
||
url_queue: asyncio.Queue,
|
||
selected_role: Dict,
|
||
language: str = None,
|
||
):
|
||
sequence_number = 0
|
||
tts_errors_count = 0
|
||
error_notified = False
|
||
first_tts_submit_time = None
|
||
await TTSAudioCleaner.prepare_output_directory()
|
||
try:
|
||
while True:
|
||
if self.check_interruption():
|
||
session_logger.info(
|
||
self.device_id, self.session_id, "TTS合成过程收到中断请求,退出"
|
||
)
|
||
break
|
||
|
||
text_item = await text_queue.get()
|
||
if text_item is None:
|
||
session_logger.info(self.device_id, self.session_id, "TTS文本队列为空,合成结束")
|
||
break
|
||
|
||
if should_skip_tts(text_item):
|
||
session_logger.info(self.device_id, self.session_id, f"跳过不需要TTS的文本: {text_item}")
|
||
continue
|
||
|
||
if first_tts_submit_time is None:
|
||
first_tts_submit_time = time.perf_counter()
|
||
url_queue.first_tts_submit_time = first_tts_submit_time
|
||
|
||
try:
|
||
sequence_number += 1
|
||
tts_start_time = time.perf_counter()
|
||
session_logger.info(self.device_id, self.session_id, f"开始合成音频 #{sequence_number}: {text_item}")
|
||
|
||
output_file_prefix = TTSAudioCleaner.generate_filename(
|
||
self.device_id, self.session_id, sequence_number
|
||
)
|
||
|
||
urls = await tts_service.tts(
|
||
text_item,
|
||
output_file_prefix=output_file_prefix,
|
||
tts_format="mp3",
|
||
selected_role=selected_role,
|
||
language=language
|
||
)
|
||
|
||
tts_end_time = time.perf_counter()
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"音频合成完成 #{sequence_number}, 耗时: {tts_end_time - tts_start_time:.2f}秒"
|
||
)
|
||
|
||
if urls and not self.check_interruption():
|
||
for url in urls:
|
||
await url_queue.put(url)
|
||
except Exception as e:
|
||
tts_errors_count += 1
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"TTS合成失败 (第{tts_errors_count}次错误): {str(e)}",
|
||
exc_info=True
|
||
)
|
||
|
||
if tts_errors_count >= self.max_tts_errors and not error_notified:
|
||
error_notified = True
|
||
await InterruptionHelper.notify_client_tts_error(self.device_id, self.session_id)
|
||
await InterruptionHelper.send_error_prompt_sound(self.device_id, self.session_id)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id, self.session_id, f"TTS合成任务出错: {e}", exc_info=True
|
||
)
|
||
finally:
|
||
await url_queue.put(None)
|