diff --git a/.gitignore b/.gitignore index 61cea7bc..ec89322a 100644 --- a/.gitignore +++ b/.gitignore @@ -156,3 +156,6 @@ main/xiaozhi-server/models/SenseVoiceSmall/model.pt main/xiaozhi-server/models/sherpa-onnx* /main/xiaozhi-server/audio_ref/ /audio_ref/ +/asr-models/iic/SenseVoiceSmall/ +/main/xiaozhi-server/asr-models/iic/SenseVoiceSmall/ +/models/SenseVoiceSmall/model.pt diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 27a2f6c1..d4076a3a 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -256,6 +256,15 @@ TTS: appid: 你的火山引擎语音合成服务appid access_token: 你的火山引擎语音合成服务access_token cluster: volcano_tts + #火山tts,支持双向流式tts + HuoshanTTS: + type: huoshan + # 如果是机智云 wss://bytedance.gizwitsapi.com/api/v3/tts/bidirection + # 机智云不需要天填 appid + ws_url: wss://openspeech.bytedance.com/api/v3/tts/bidirection + appid: 你的火山引擎语音合成服务appid + access_token: 你的火山引擎语音合成服务access_token + speaker: zh_female_meilinvyou_moon_bigtts CosyVoiceSiliconflow: type: siliconflow # 硅基流动TTS diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 1e4d0fac..5a5daaf0 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -11,11 +11,12 @@ import websockets from typing import Dict, Any import plugins_func.loadplugins from config.logger import setup_logging +from core.providers.tts.dto.dto import TTSMessageDTO, MsgType from core.utils.dialogue import Message, Dialogue from core.handle.textHandle import handleTextMessage from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string, get_ip_info from concurrent.futures import ThreadPoolExecutor, TimeoutError -from core.handle.sendAudioHandle import sendAudioMessage, sendAudioMessageStream +from core.handle.sendAudioHandle import sendAudioMessage from core.handle.receiveAudioHandle import handleAudioMessage from core.handle.functionHandler import FunctionHandler from plugins_func.register import Action @@ -45,6 +46,7 @@ class ConnectionHandler: self.session_id = None self.prompt = None self.welcome_msg = None + self.u_id = None # 客户端状态相关 self.client_abort = False @@ -58,6 +60,7 @@ class ConnectionHandler: self.audio_play_queue = queue.Queue() max_workers = self.config.get("TTS_SET", {}).get("MAX_WORKERS", 10) self.executor = ThreadPoolExecutor(max_workers=max_workers) + self.start_tts_request_flag = False # 依赖的组件 self.vad = _vad @@ -157,14 +160,14 @@ class ConnectionHandler: # 异步初始化 await self.loop.run_in_executor(None, self._initialize_components) - # tts 消化线程 - tts_priority = threading.Thread(target=self._tts_priority_thread, daemon=True) - tts_priority.start() - # 音频播放 消化线程 + self.stop_event.clear() audio_play_priority = threading.Thread(target=self._audio_play_priority_thread, daemon=True) audio_play_priority.start() + # 打开音频通道 + await self.tts.open_audio_channels() + try: async for message in self.websocket: await self._route_message(message) @@ -196,9 +199,9 @@ class ConnectionHandler: if self.private_config: self.prompt = self.private_config.private_config.get("prompt", self.prompt) - self.client_ip_info = get_ip_info(self.client_ip) - self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}") - self.prompt = self.prompt + f"\n我在:{self.client_ip_info}" + # self.client_ip_info = get_ip_info(self.client_ip) + # self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}") + # self.prompt = self.prompt + f"\n我在:{self.client_ip_info}" self.dialogue.put(Message(role="system", content=self.prompt)) self.func_handler = FunctionHandler(self.config) @@ -258,6 +261,8 @@ class ConnectionHandler: self.llm_finish_task = False text_index = 0 + uuid_str = str(uuid.uuid4()) + msg_type = None for content in llm_responses: response_message.append(content) if self.client_abort: @@ -265,62 +270,14 @@ class ConnectionHandler: end_time = time.time() self.logger.bind(tag=TAG).debug(f"大模型返回时间: {end_time - start_time} 秒, 生成token={content}") - - # 合并当前全部文本并处理未分割部分 - full_text = "".join(response_message) - current_text = full_text[processed_chars:] # 从未处理的位置开始 - - # 查找最后一个有效标点 - punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ") - last_punct_pos = -1 - for punct in punctuations: - pos = current_text.rfind(punct) - if pos > last_punct_pos: - last_punct_pos = pos - - # 找到分割点则处理 - if last_punct_pos != -1: - segment_text_raw = current_text[:last_punct_pos + 1] - segment_text = get_string_no_punctuation_or_emoji(segment_text_raw) - if segment_text: - # 强制设置空字符,测试TTS出错返回语音的健壮性 - # if text_index % 2 == 0: - # segment_text = " " - text_index += 1 - self.recode_first_last_text(segment_text, text_index) - if self.tts_stream: - stream_queue = queue.Queue() - self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue) - self.tts_queue_stream.put({ - "text": segment_text, - "chunk_queque": stream_queue, - "text_index": text_index - }) - else: - future = self.executor.submit(self.speak_and_play, segment_text, text_index) - self.tts_queue.put(future) - processed_chars += len(segment_text_raw) # 更新已处理字符位置 - - # 处理最后剩余的文本 - full_text = "".join(response_message) - remaining_text = full_text[processed_chars:] - if remaining_text: - segment_text = get_string_no_punctuation_or_emoji(remaining_text) - if segment_text: - text_index += 1 - self.recode_first_last_text(segment_text, text_index) - if self.tts_stream: - stream_queue = queue.Queue() - self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue, text_index) - self.tts_queue_stream.put({ - "text": segment_text, - "chunk_queque": stream_queue, - "text_index": text_index - }) - else: - future = self.executor.submit(self.speak_and_play, segment_text, text_index) - self.tts_queue.put(future) - + if text_index == 0: + msg_type = MsgType.START_TTS_REQUEST + else: + msg_type = MsgType.TTS_TEXT_REQUEST + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=msg_type, content=content)) + text_index += 1 + msg_type = MsgType.STOP_TTS_REQUEST + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=msg_type, content="")) self.llm_finish_task = True self.dialogue.put(Message(role="assistant", content="".join(response_message))) self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False)) @@ -372,6 +329,9 @@ class ConnectionHandler: function_id = None function_arguments = "" content_arguments = "" + uuid_str = str(uuid.uuid4()).replace("-", "") + self.u_id = uuid_str + msg_type = None for response in llm_responses: content, tools_call = response if content is not None and len(content) > 0: @@ -398,39 +358,16 @@ class ConnectionHandler: end_time = time.time() self.logger.bind(tag=TAG).debug(f"大模型返回时间: {end_time - start_time} 秒, 生成token={content}") - - # 处理文本分段和TTS逻辑 - # 合并当前全部文本并处理未分割部分 - full_text = "".join(response_message) - current_text = full_text[processed_chars:] # 从未处理的位置开始 - - # 查找最后一个有效标点 - punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ") - last_punct_pos = -1 - for punct in punctuations: - pos = current_text.rfind(punct) - if pos > last_punct_pos: - last_punct_pos = pos - - # 找到分割点则处理 - if last_punct_pos != -1: - segment_text_raw = current_text[:last_punct_pos + 1] - segment_text = get_string_no_punctuation_or_emoji(segment_text_raw) - if segment_text: - text_index += 1 - self.recode_first_last_text(segment_text, text_index) - if self.tts_stream: - stream_queue = queue.Queue() - self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue) - self.tts_queue_stream.put({ - "text": segment_text, - "chunk_queque": stream_queue, - "text_index": text_index - }) - else: - future = self.executor.submit(self.speak_and_play, segment_text, text_index) - self.tts_queue.put(future) - processed_chars += len(segment_text_raw) # 更新已处理字符位置 + if text_index == 0: + self.tts.tts_text_queue.put( + TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.START_TTS_REQUEST, content='')) + self.start_tts_request_flag = True + self.tts.tts_text_queue.put( + TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.TTS_TEXT_REQUEST, content=content)) + text_index += 1 + if self.start_tts_request_flag: + self.start_tts_request_flag = False + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.STOP_TTS_REQUEST, content='')) # 处理function call if tool_call_flag: @@ -464,26 +401,6 @@ class ConnectionHandler: result = self.func_handler.handle_llm_function_call(self, function_call_data) self._handle_function_result(result, function_call_data, text_index + 1) - # 处理最后剩余的文本 - full_text = "".join(response_message) - remaining_text = full_text[processed_chars:] - if remaining_text: - segment_text = get_string_no_punctuation_or_emoji(remaining_text) - if segment_text: - text_index += 1 - self.recode_first_last_text(segment_text, text_index) - if self.tts_stream: - stream_queue = queue.Queue() - self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue, text_index) - self.tts_queue_stream.put({ - "text": segment_text, - "chunk_queque": stream_queue, - "text_index": text_index - }) - else: - future = self.executor.submit(self.speak_and_play, segment_text, text_index) - self.tts_queue.put(future) - # 存储对话内容 if len(response_message) > 0: self.dialogue.put(Message(role="assistant", content="".join(response_message))) @@ -497,20 +414,9 @@ class ConnectionHandler: if result.action == Action.RESPONSE: # 直接回复前端 text = result.response self.recode_first_last_text(text, text_index) - if self.tts_stream: - stream_queue = queue.Queue() - self.executor.submit(self.speak_and_play_stream, text, stream_queue, text_index) - self.tts_queue_stream.put({ - "text": text, - "chunk_queque": stream_queue, - "text_index": text_index - }) - else: - future = self.executor.submit(self.speak_and_play, text, text_index) - self.tts_queue.put(future) + asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop) self.dialogue.put(Message(role="assistant", content=text)) elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复 - text = result.result if text is not None and len(text) > 0: function_id = function_call_data["id"] @@ -528,14 +434,12 @@ class ConnectionHandler: elif result.action == Action.NOTFOUND: text = result.result self.recode_first_last_text(text, text_index) - future = self.executor.submit(self.speak_and_play, text, text_index) - self.tts_queue.put(future) + asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop) self.dialogue.put(Message(role="assistant", content=text)) else: text = result.result self.recode_first_last_text(text, text_index) - future = self.executor.submit(self.speak_and_play, text, text_index) - self.tts_queue.put(future) + asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop) self.dialogue.put(Message(role="assistant", content=text)) def _tts_priority_thread(self): @@ -615,22 +519,10 @@ class ConnectionHandler: while not self.stop_event.is_set(): text = None try: - if self.tts_stream: - data, text, text_index = self.audio_play_queue.get() - if isinstance(data, list): - future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, data, text, text_index), - self.loop) - future.result() - else: - future = asyncio.run_coroutine_threadsafe( - sendAudioMessageStream(self, data, text, text_index), - self.loop) - future.result() - else: - opus_datas, text, text_index = self.audio_play_queue.get() - future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, opus_datas, text, text_index), - self.loop) - future.result() + ttsMessageDTO = self.tts.tts_audio_queue.get() + future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, ttsMessageDTO), + self.loop) + future.result() except Exception as e: self.logger.bind(tag=TAG).error(f"audio_play_priority priority_thread: {text} {e}") @@ -670,12 +562,12 @@ class ConnectionHandler: async def close(self): """资源清理方法""" - # 清理其他资源 self.stop_event.set() self.executor.shutdown(wait=False) if self.websocket: await self.websocket.close() + await self.tts.close() self.logger.bind(tag=TAG).info("连接资源已释放") def reset_vad_states(self): diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index 014c22ec..51010fc5 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -4,94 +4,33 @@ from config.logger import setup_logging import json import asyncio import time + +from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, MsgType from core.utils.util import remove_punctuation_and_length, get_string_no_punctuation_or_emoji TAG = __name__ logger = setup_logging() -async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_finish_task=False): +async def sendAudioMessage(conn, ttsMessageDTO: TTSMessageDTO): + if ttsMessageDTO.u_id != conn.u_id: + logger.bind(tag=TAG).info(f"msg id:{ttsMessageDTO.u_id},不是当前对话,当前对话id:{conn.u_id}") + return # 发送句子开始消息 - if text_index == conn.tts_first_text_index: - logger.bind(tag=TAG).info(f"发送第一段语音: {text}") - await send_tts_message(conn, "sentence_start", text) - - # 初始化流控参数 - frame_duration = 60 # 毫秒 - start_time = time.time() # 使用高精度计时器 - # 初始化流控参数 - frame_duration = 60 # 毫秒 - start_time_chunk = time.perf_counter() # 使用高精度计时器 - play_position = 0 # 已播放的时长(毫秒) - while True: - try: - start_get_queue = time.time() - # 尝试获取数据,如果没有数据,则等待一小段时间再试 - audio_data_chunke = None - try: - audio_data_chunke = audios_queue.get(timeout=5) # 设置超时为1秒 - except Exception as e: - # 如果超时,继续等待 - logger.bind(tag=TAG).error(f"获取队列超时~{e}") - - audio_opus_datas = audio_data_chunke.get('data') if audio_data_chunke else None - duration = audio_data_chunke.get('duration') if audio_data_chunke else 0 - - if audio_data_chunke: - start_time = time.time() - # 检查是否超过 5 秒没有数据 - if time.time() - start_time > 15: - logger.bind(tag=TAG).error("超过15秒没有数据,退出。") - break - - if audio_data_chunke and audio_data_chunke.get("end", True): - break - - if audio_opus_datas: - for opus_packet in audio_opus_datas: - if conn.client_abort: - return - logger.bind(tag=TAG).info(f'发送数据长度:{len(opus_packet)}') - await conn.websocket.send(opus_packet) - play_position += frame_duration # 更新播放位置 - start_time = time.time() # 更新获取数据的时间 - except Exception as e: - logger.bind(tag=TAG).error(f"发生错误: {e}") - traceback.print_exc() # 打印错误堆栈 - await send_tts_message(conn, "sentence_end", text) - - print(f'{text_index}-{conn.tts_last_text_index}') - expected_time = start_time_chunk + (play_position / 1000) - current_time = time.perf_counter() - # 等待直到预期时间 - delay = expected_time - current_time - if delay > 0: - await asyncio.sleep(delay) - # 发送结束消息(如果是最后一个文本) - logger.bind(tag=TAG).info(f"{conn.llm_finish_task},{text_index},{conn.tts_last_text_index}") - if conn.llm_finish_task and text_index == conn.tts_last_text_index: - await send_tts_message(conn, 'stop', None) - if conn.close_after_chat or "拜拜" in text or "再见" in text: - await conn.close() - - - -async def sendAudioMessage(conn, audios, text, text_index=0): - # 发送句子开始消息 - if text_index == conn.tts_first_text_index: - logger.bind(tag=TAG).info(f"发送第一段语音: {text}") - await send_tts_message(conn, "sentence_start", text) + if SentenceType.SENTENCE_START == ttsMessageDTO.sentence_type: + logger.bind(tag=TAG).info(f"发送第一段语音: {ttsMessageDTO.tts_finish_text}") + await send_tts_message(conn, "sentence_start", ttsMessageDTO.tts_finish_text) # 流控参数优化 original_frame_duration = 60 # 原始帧时长(毫秒) adjusted_frame_duration = int(original_frame_duration * 0.8) # 缩短20% - total_frames = len(audios) # 获取总帧数 + total_frames = len(ttsMessageDTO.content) # 获取总帧数 compensation = total_frames * (original_frame_duration - adjusted_frame_duration) / 1000 # 补偿时间(秒) start_time = time.perf_counter() play_position = 0 # 已播放时长(毫秒) - for opus_packet in audios: + for opus_packet in ttsMessageDTO.content: if conn.client_abort: return @@ -110,15 +49,17 @@ async def sendAudioMessage(conn, audios, text, text_index=0): # 补偿因加速损失的时长 if compensation > 0: await asyncio.sleep(compensation) - - await send_tts_message(conn, "sentence_end", text) + if SentenceType.SENTENCE_END == ttsMessageDTO.sentence_type: + logger.bind(tag=TAG).info(f"发送最后一段语音: {ttsMessageDTO.tts_finish_text}") + await send_tts_message(conn, "sentence_end", ttsMessageDTO.tts_finish_text) # 发送结束消息(如果是最后一个文本) - if conn.llm_finish_task and text_index == conn.tts_last_text_index: + if conn.llm_finish_task and MsgType.STOP_TTS_RESPONSE == ttsMessageDTO.msg_type: await send_tts_message(conn, 'stop', None) if conn.close_after_chat: await conn.close() + async def send_tts_message(conn, state, text=None): """发送 TTS 状态消息""" message = { diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 5dec4b4d..a363c3d7 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -1,20 +1,219 @@ import asyncio +import gc +import io +import threading +import traceback +import uuid +from concurrent.futures import ThreadPoolExecutor + +import torch +import torchaudio + from config.logger import setup_logging import os import numpy as np import opuslib_next from pydub import AudioSegment from abc import ABC, abstractmethod +from core.utils import textUtils import queue +from core.providers.tts.dto.dto import MsgType, TTSMessageDTO, SentenceType + TAG = __name__ logger = setup_logging() class TTSProviderBase(ABC): def __init__(self, config, delete_audio_file): + self.config = config self.delete_audio_file = delete_audio_file self.output_file = config.get("output_dir") + self.tts_text_queue = queue.Queue() + self.tts_audio_queue = queue.Queue() + self.enable_two_way = False + self.stop_event = threading.Event() + + self.tts_text_buff = [] + self.punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ", ",", ",") + self.tts_request = False + self.processed_chars = 0 + self.stream = False + self.last_to_opus_raw = b'' + + # 启动tts_text_queue监听线程 + # 线程任务相关 + self.loop = asyncio.get_event_loop() + self.process_tasks_loop = asyncio.get_event_loop() + self.max_workers = self.config.get("TTS_SET", {}).get("MAX_WORKERS", 3) + self.active_tasks = set() # 追踪当前运行的任务 + self.executor = ThreadPoolExecutor(max_workers=self.max_workers) + + async def open_audio_channels(self): + # 启动tts_text_queue监听线程 + tts_priority = threading.Thread(target=self._tts_text_priority_thread, daemon=True) + tts_priority.start() + + async def stop_listen_resource(self): + """资源清理方法""" + self.stop_event.set() + self.tts_text_queue = None + self.tts_audio_queue = None + gc.collect() # 强制执行垃圾回收 + + async def close(self): + pass + + def _get_segment_text(self): + # 合并当前全部文本并处理未分割部分 + full_text = "".join(self.tts_text_buff) + current_text = full_text[self.processed_chars:] # 从未处理的位置开始 + last_punct_pos = -1 + for punct in self.punctuations: + pos = current_text.rfind(punct) + if (pos != -1 and last_punct_pos == -1) or (pos != -1 and pos < last_punct_pos): + last_punct_pos = pos + if last_punct_pos != -1: + segment_text_raw = current_text[:last_punct_pos + 1] + segment_text = textUtils.get_string_no_punctuation_or_emoji(segment_text_raw) + self.processed_chars += len(segment_text_raw) # 更新已处理字符位置 + return segment_text + else: + return None + + async def process_generator(self, generator): + async for tts_data in generator: + self.tts_audio_queue.put(tts_data) + + def _tts_text_priority_thread(self): + logger.bind(tag=TAG).info("开始监听tts文本") + if self.enable_two_way: + self._enable_two_way_tts() + else: + self._no_enable_two_way_tts() + + async def start_session(self, session_id): + pass + + async def finish_session(self, session_id): + pass + + async def tts_one_sentence(self,text): + uuid_str = str(uuid.uuid4()).replace("-", "") + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.START_TTS_REQUEST, content='')) + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.TTS_TEXT_REQUEST, content=text)) + self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.STOP_TTS_REQUEST, content=text)) + + def _enable_two_way_tts(self): + while not self.stop_event.is_set(): + try: + ttsMessageDTO = self.tts_text_queue.get() + msg_type = ttsMessageDTO.msg_type + if msg_type == MsgType.START_TTS_REQUEST: + # 开始传输tts文本 + self.tts_request = True + self.u_id = ttsMessageDTO.u_id + # 开启session + future = asyncio.run_coroutine_threadsafe(self.start_session(ttsMessageDTO.u_id), loop=self.loop) + future.result() + # await self.start_session(ttsMessageDTO.u_id) + elif self.tts_request and msg_type == MsgType.TTS_TEXT_REQUEST: + future = asyncio.run_coroutine_threadsafe( + self.text_to_speak(u_id=ttsMessageDTO.u_id, text=ttsMessageDTO.content), loop=self.loop) + future.result() + elif msg_type == MsgType.STOP_TTS_REQUEST: + self.tts_request = False + future = asyncio.run_coroutine_threadsafe(self.finish_session(ttsMessageDTO.u_id), loop=self.loop) + future.result() + + + except Exception as e: + logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}") + # 报错了。要关闭说话 + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.STOP_TTS_RESPONSE, content=[], + tts_finish_text='', + sentence_type=None + ) + ) + traceback.print_exc() + + def _no_enable_two_way_tts(self): + # 为这个线程创建一个新的事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + while not self.stop_event.is_set(): + try: + ttsMessageDTO = self.tts_text_queue.get() + msg_type = ttsMessageDTO.msg_type + if not self.enable_two_way: + if msg_type == MsgType.START_TTS_REQUEST: + # 开始传输tts文本 + self.tts_request = True + self.processed_chars = 0 + self.tts_text_buff = [] + elif self.tts_request and msg_type == MsgType.TTS_TEXT_REQUEST: + self.tts_text_buff.append(ttsMessageDTO.content) + elif msg_type == MsgType.STOP_TTS_REQUEST: + # 结束传输tts文本,处理最尾巴的数据 + self.tts_request = False + segment_text = self._get_segment_text() + if segment_text: + # 修改部分:创建协程对象 + # 修改部分:创建协程对象 + tts_generator = self.text_to_speak(ttsMessageDTO.u_id, segment_text, + True if msg_type == MsgType.STOP_TTS_REQUEST else False, + True if msg_type == MsgType.START_TTS_REQUEST else False) + future = asyncio.run_coroutine_threadsafe(self.process_generator(tts_generator), self.loop) + self.active_tasks.add(future) + if self.active_tasks: + async def wrap_future(future): + return await asyncio.wrap_future(future) + + wrapped_tasks = [wrap_future(task) for task in self.active_tasks] + done, _ = loop.run_until_complete(asyncio.wait(wrapped_tasks)) + self.active_tasks -= done + + # 发送合成结束 + self.tts_audio_queue.put(TTSMessageDTO(u_id=ttsMessageDTO.u_id, + msg_type=MsgType.STOP_TTS_RESPONSE, + content=[], + tts_finish_text='', + sentence_type=SentenceType.SENTENCE_END)) + + segment_text = self._get_segment_text() + if segment_text: + # 确保这里得到的是协程对象 + tts_generator = self.text_to_speak( + ttsMessageDTO.u_id, + segment_text, + msg_type == MsgType.STOP_TTS_REQUEST, + msg_type == MsgType.START_TTS_REQUEST + ) + # 提交协程到事件循环 + tts_generator_future = asyncio.run_coroutine_threadsafe( + self.process_generator(tts_generator), + loop + ) + self.active_tasks.add(tts_generator_future) + if len(self.active_tasks) >= self.max_workers: + # 等待所有任务完成 + try: + async def wrap_future(future): + return await asyncio.wrap_future(future) + + wrapped_tasks = [wrap_future(task) for task in self.active_tasks] + done, _ = loop.run_until_complete(asyncio.wait(wrapped_tasks)) + self.active_tasks -= done + except Exception as e: + logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}") + traceback.print_exc() + else: + pass + except Exception as e: + logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}") + traceback.print_exc() @abstractmethod def generate_filename(self): @@ -38,7 +237,7 @@ class TTSProviderBase(ABC): logger.bind(tag=TAG).info(f"Failed to generate TTS file: {e}") return None - def to_tts_stream(self, text, queue: queue.Queue, text_index=0): + def to_tts_stream(self, u_id, text, queue: queue.Queue, text_index=0): try: asyncio.run(self.text_to_speak_stream(text, queue, text_index)) except Exception as e: @@ -46,7 +245,7 @@ class TTSProviderBase(ABC): return None @abstractmethod - async def text_to_speak(self, text, output_file): + async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): pass async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0): @@ -95,7 +294,17 @@ class TTSProviderBase(ABC): return opus_datas, duration - def wav_to_opus_data_audio_raw(self, raw_data): + def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000): + tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0) + with io.BytesIO() as bf: + torchaudio.save(bf, tts_speech, src_rate, format="wav") + audio = AudioSegment.from_file(bf, format="wav") + audio = audio.set_channels(1).set_frame_rate(to_rate) + return audio + + def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False): + raw_data = self.last_to_opus_raw + raw_data_var + self.last_to_opus_raw = b'' # 初始化Opus编码器 encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) @@ -110,8 +319,13 @@ class TTSProviderBase(ABC): chunk = raw_data[i:i + frame_size * 2] # 如果最后一帧不足,补零 - if len(chunk) < frame_size * 2: - # logger.bind(tag=TAG).info("开始补0") + # 缓存记录一下 + if len(chunk) < frame_size * 2 and not is_end: + logger.bind(tag=TAG).info("如果最后一帧不足,缓存记录一下") + self.last_to_opus_raw = chunk + break + if len(chunk) < frame_size * 2 and is_end: + logger.bind(tag=TAG).info("是最后一句了,补零") chunk += b'\x00' * (frame_size * 2 - len(chunk)) # 转换为numpy数组处理 diff --git a/main/xiaozhi-server/core/providers/tts/dto/__init__.py b/main/xiaozhi-server/core/providers/tts/dto/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/main/xiaozhi-server/core/providers/tts/dto/dto.py b/main/xiaozhi-server/core/providers/tts/dto/dto.py new file mode 100644 index 00000000..706e7d35 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tts/dto/dto.py @@ -0,0 +1,42 @@ +from enum import Enum +from typing import Union + + +class MsgType(Enum): + # 请求类型 + START_TTS_REQUEST = "START_TTS_REQUEST" + TTS_TEXT_REQUEST = "TTS_TEXT_REQUEST" + STOP_TTS_REQUEST = "STOP_TTS_REQUEST" + + # 返回类型 + START_TTS_RESPONSE = "START_TTS_RESPONSE" + TTS_TEXT_RESPONSE = "TTS_TEXT_RESPONSE" + STOP_TTS_RESPONSE = "STOP_TTS_RESPONSE" + + +class SentenceType(Enum): + # 句子开始 + SENTENCE_START = "SENTENCE_START" + # 句子结束 + SENTENCE_END = "SENTENCE_END" + + +class TTSMessageDTO: + def __init__(self, u_id: str, msg_type: MsgType, content: Union[str, bytes], tts_finish_text=None, + sentence_type: SentenceType = None, duration=0): + if not isinstance(msg_type, MsgType): + raise ValueError("msg_type must be an instance of MsgType Enum") + if not isinstance(content, (str, list, bytes)): + raise ValueError("content must be of type str or bytes") + + # 唯一id,每个合成到合成结束,使用同一个id + self.u_id = u_id + self.msg_type = msg_type + self.sentence_type = sentence_type + self.content = content + self.tts_finish_text = tts_finish_text + self.duration = duration + + def __repr__(self): + content_preview = self.content if isinstance(self.content, str) else "" + return f"MessageDTO(msg_type={self.msg_type}, content={content_preview})" diff --git a/main/xiaozhi-server/core/providers/tts/edge.py b/main/xiaozhi-server/core/providers/tts/edge.py index 3c02597d..a9fda52a 100644 --- a/main/xiaozhi-server/core/providers/tts/edge.py +++ b/main/xiaozhi-server/core/providers/tts/edge.py @@ -1,9 +1,17 @@ +import io import os import uuid import edge_tts from datetime import datetime -from core.providers.tts.base import TTSProviderBase +from pydub import AudioSegment + +from config.logger import setup_logging +from core.providers.tts.base import TTSProviderBase +from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType + +TAG = __name__ +logger = setup_logging() class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): @@ -13,6 +21,25 @@ class TTSProvider(TTSProviderBase): def generate_filename(self, extension=".mp3"): return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") - async def text_to_speak(self, text, output_file): - communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice - await communicate.save(output_file) + async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): + try: + communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice + tmp_file = self.generate_filename() + await communicate.save(tmp_file) + + # 使用 pydub 读取临时文件 + audio = AudioSegment.from_file(tmp_file, format="mp3") + audio = audio.set_channels(1).set_frame_rate(16000) + opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,sentence_type=SentenceType.SENTENCE_START) + # 用完后删除临时文件 + try: + os.remove(tmp_file) + except FileNotFoundError: + # 若文件不存在,忽略该异常 + pass + except Exception as e: + logger.bind(tag=TAG).error(f"TTSProvider text_to_speak error: {e}") + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], tts_finish_text=text,sentence_type=SentenceType.SENTENCE_START) + + diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index 37f63095..4acac43a 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -17,6 +17,8 @@ from pydub import AudioSegment from typing_extensions import Annotated from datetime import datetime from typing import Literal + +from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType from core.utils.util import check_model_key from core.providers.tts.base import TTSProviderBase from config.logger import setup_logging @@ -118,55 +120,6 @@ class TTSProvider(TTSProviderBase): def generate_filename(self, extension=".wav"): return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") - async def text_to_speak(self, text, output_file): - # Prepare reference data - byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio] - ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text] - - data = { - "text": text, - "references": [ - ServeReferenceAudio( - audio=audio if audio else b"", text=text - ) - for text, audio in zip(ref_texts, byte_audios) - ], - "reference_id": self.reference_id, - "normalize": self.normalize, - "format": self.format, - "max_new_tokens": self.max_new_tokens, - "chunk_length": self.chunk_length, - "top_p": self.top_p, - "repetition_penalty": self.repetition_penalty, - "temperature": self.temperature, - "streaming": self.streaming, - "use_memory_cache": self.use_memory_cache, - "seed": self.seed, - } - - pydantic_data = ServeTTSRequest(**data) - - response = requests.post( - self.api_url, - data=ormsgpack.packb(pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC), - headers={ - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/msgpack", - }, - ) - - if response.status_code == 200: - audio_content = response.content - - with open(output_file, "wb") as audio_file: - audio_file.write(audio_content) - - - - else: - print(f"Request failed with status code {response.status_code}") - print(response.json()) - def _get_audio_from_tts(self, data_bytes): tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0) with io.BytesIO() as bf: @@ -175,7 +128,7 @@ class TTSProvider(TTSProviderBase): audio = audio.set_channels(1).set_frame_rate(16000) return audio - async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0): + async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): try: data = { "text": text, @@ -219,6 +172,7 @@ class TTSProvider(TTSProviderBase): }, ) as response: if response.status_code == 200: + index = 0 for chunk in response.iter_content(): # 拼接当前块和上一块数据 chunk_total += chunk @@ -226,44 +180,31 @@ class TTSProvider(TTSProviderBase): if len(chunk_total) % 2 == 0 and chunk_total[-2:] == b'\x00\x00': audio = self._get_audio_from_tts(chunk_total) audio_raw = audio_raw + audio.raw_data - # 长度凑够2贞开始发送,60ms*4=240ms - if len(audio_raw) >= 7680: - duration = 60 * len(audio_raw) // 1920 - if (len(audio_raw) % 1920) > 0: - duration += 60 - duration = duration / 1000.0 - # logger.bind(tag=TAG).info(f'发送数据长度:{len(audio_raw)}') + # 长度凑够2贞开始发送,60ms*2=120ms + if len(audio_raw) >= 3840: opus_datas = self.wav_to_opus_data_audio_raw(audio_raw) - queue.put({ - "data": opus_datas, - "duration": duration, - "end": False, - "text_index": text_index - }) + if index == 0: + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, sentence_type=SentenceType.SENTENCE_START) + else: + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, sentence_type=None) audio_raw = b'' chunk_total = b'' if len(chunk_total) > 0: audio = self._get_audio_from_tts(chunk_total) audio_raw = audio_raw + audio.raw_data - duration = 60 * len(audio_raw) // 1920 - if (len(audio_raw) % 1920) > 0: - duration += 60 - duration = duration / 1000.0 - # 把 audio 转成 opus - # logger.bind(tag=TAG).info(f'发送数据长度:{len(audio_raw)}') opus_datas = self.wav_to_opus_data_audio_raw(audio_raw) - queue.put({ - "data": opus_datas, - "duration": duration, - "end": False - }) + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, + tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END) + else: + yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], + tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END) else: print('请求失败:', response.status_code, response.text) - queue.put({ - "data": None, - "end": True - }) except Exception as e: logger.bind(tag=TAG).error("tts发生错误") traceback.print_exc() diff --git a/main/xiaozhi-server/core/providers/tts/huoshan.py b/main/xiaozhi-server/core/providers/tts/huoshan.py new file mode 100644 index 00000000..6489cfa9 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tts/huoshan.py @@ -0,0 +1,391 @@ +import asyncio +import io +import os +import subprocess +import threading +import traceback +import uuid +import json +import base64 +import requests +from datetime import datetime +from mutagen.oggopus import OggOpus + +import websockets + +from config.logger import setup_logging +from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType +from core.utils.util import check_model_key +from core.providers.tts.base import TTSProviderBase + +TAG = __name__ +logger = setup_logging() + +PROTOCOL_VERSION = 0b0001 +DEFAULT_HEADER_SIZE = 0b0001 + +# Message Type: +FULL_CLIENT_REQUEST = 0b0001 +AUDIO_ONLY_RESPONSE = 0b1011 +FULL_SERVER_RESPONSE = 0b1001 +ERROR_INFORMATION = 0b1111 + +# Message Type Specific Flags +MsgTypeFlagNoSeq = 0b0000 # Non-terminal packet with no sequence +MsgTypeFlagPositiveSeq = 0b1 # Non-terminal packet with sequence > 0 +MsgTypeFlagLastNoSeq = 0b10 # last packet with no sequence +MsgTypeFlagNegativeSeq = 0b11 # Payload contains event number (int32) +MsgTypeFlagWithEvent = 0b100 +# Message Serialization +NO_SERIALIZATION = 0b0000 +JSON = 0b0001 +# Message Compression +COMPRESSION_NO = 0b0000 +COMPRESSION_GZIP = 0b0001 + +EVENT_NONE = 0 +EVENT_Start_Connection = 1 + +EVENT_FinishConnection = 2 + +EVENT_ConnectionStarted = 50 # 成功建连 + +EVENT_ConnectionFailed = 51 # 建连失败(可能是无法通过权限认证) + +EVENT_ConnectionFinished = 52 # 连接结束 + +# 上行Session事件 +EVENT_StartSession = 100 + +EVENT_FinishSession = 102 +# 下行Session事件 +EVENT_SessionStarted = 150 +EVENT_SessionFinished = 152 + +EVENT_SessionFailed = 153 + +# 上行通用事件 +EVENT_TaskRequest = 200 + +# 下行TTS事件 +EVENT_TTSSentenceStart = 350 + +EVENT_TTSSentenceEnd = 351 + +EVENT_TTSResponse = 352 + + +class Header: + def __init__(self, + protocol_version=PROTOCOL_VERSION, + header_size=DEFAULT_HEADER_SIZE, + message_type: int = 0, + message_type_specific_flags: int = 0, + serial_method: int = NO_SERIALIZATION, + compression_type: int = COMPRESSION_NO, + reserved_data=0): + self.header_size = header_size + self.protocol_version = protocol_version + self.message_type = message_type + self.message_type_specific_flags = message_type_specific_flags + self.serial_method = serial_method + self.compression_type = compression_type + self.reserved_data = reserved_data + + def as_bytes(self) -> bytes: + return bytes([ + (self.protocol_version << 4) | self.header_size, + (self.message_type << 4) | self.message_type_specific_flags, + (self.serial_method << 4) | self.compression_type, + self.reserved_data + ]) + + +class Optional: + def __init__(self, event: int = EVENT_NONE, sessionId: str = None, sequence: int = None): + self.event = event + self.sessionId = sessionId + self.errorCode: int = 0 + self.connectionId: str | None = None + self.response_meta_json: str | None = None + self.sequence = sequence + + # 转成 byte 序列 + def as_bytes(self) -> bytes: + option_bytes = bytearray() + if self.event != EVENT_NONE: + option_bytes.extend(self.event.to_bytes(4, "big", signed=True)) + if self.sessionId is not None: + session_id_bytes = str.encode(self.sessionId) + size = len(session_id_bytes).to_bytes(4, "big", signed=True) + option_bytes.extend(size) + option_bytes.extend(session_id_bytes) + if self.sequence is not None: + option_bytes.extend(self.sequence.to_bytes(4, "big", signed=True)) + return option_bytes + + +class Response: + def __init__(self, header: Header, optional: Optional): + self.optional = optional + self.header = header + self.payload: bytes | None = None + + def __str__(self): + return super().__str__() + + +class TTSProvider(TTSProviderBase): + def __init__(self, config, delete_audio_file): + super().__init__(config, delete_audio_file) + self.appId = config.get("appid") + self.access_token = config.get("access_token") + self.cluster = config.get("cluster") + self.voice = config.get("voice") + self.ws_url = config.get("ws_url") + self.authorization = config.get("authorization") + self.speaker = config.get("speaker") + self.header = {"Authorization": f"{self.authorization}{self.access_token}"} + self.stop_event_response = threading.Event() + self.enable_two_way = True + self.start_connection_flag = False + self.tts_text = "" + + async def open_audio_channels(self): + await super().open_audio_channels() + ws_header = { + "X-Api-App-Key": self.appId, + "X-Api-Access-Key": self.access_token, + "X-Api-Resource-Id": 'volc.service_type.10029', + "X-Api-Connect-Id": uuid.uuid4(), + } + self.ws = await websockets.connect(self.ws_url, additional_headers=ws_header, max_size=1000000000) + tts_priority = threading.Thread(target=self._start_monitor_tts_response_thread(), daemon=True) + tts_priority.start() + + def generate_filename(self, extension=".wav"): + return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + + async def send_event(self, header: bytes, optional: bytes | None = None, + payload: bytes = None): + full_client_request = bytearray(header) + if optional is not None: + full_client_request.extend(optional) + if payload is not None: + payload_size = len(payload).to_bytes(4, 'big', signed=True) + full_client_request.extend(payload_size) + full_client_request.extend(payload) + await self.ws.send(full_client_request) + + async def send_text(self, speaker: str, text: str, session_id): + header = Header(message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON).as_bytes() + optional = Optional(event=EVENT_TaskRequest, sessionId=session_id).as_bytes() + payload = self.get_payload_bytes(event=EVENT_TaskRequest, text=text, speaker=speaker) + return await self.send_event(header, optional, payload) + + # 读取 res 数组某段 字符串内容 + def read_res_content(self, res: bytes, offset: int): + content_size = int.from_bytes(res[offset: offset + 4], "big", signed=True) + offset += 4 + content = str(res[offset: offset + content_size]) + offset += content_size + return content, offset + + # 读取 payload + def read_res_payload(self, res: bytes, offset: int): + payload_size = int.from_bytes(res[offset: offset + 4], "big", signed=True) + offset += 4 + payload = res[offset: offset + payload_size] + offset += payload_size + return payload, offset + + def parser_response(self, res) -> Response: + if isinstance(res, str): + raise RuntimeError(res) + response = Response(Header(), Optional()) + # 解析结果 + # header + header = response.header + num = 0b00001111 + header.protocol_version = res[0] >> 4 & num + header.header_size = res[0] & 0x0f + header.message_type = (res[1] >> 4) & num + header.message_type_specific_flags = res[1] & 0x0f + header.serialization_method = res[2] >> num + header.message_compression = res[2] & 0x0f + header.reserved = res[3] + # + offset = 4 + optional = response.optional + if header.message_type == FULL_SERVER_RESPONSE or AUDIO_ONLY_RESPONSE: + # read event + if header.message_type_specific_flags == MsgTypeFlagWithEvent: + optional.event = int.from_bytes(res[offset:8], "big", signed=True) + offset += 4 + if optional.event == EVENT_NONE: + return response + # read connectionId + elif optional.event == EVENT_ConnectionStarted: + optional.connectionId, offset = self.read_res_content(res, offset) + elif optional.event == EVENT_ConnectionFailed: + optional.response_meta_json, offset = self.read_res_content(res, offset) + elif (optional.event == EVENT_SessionStarted + or optional.event == EVENT_SessionFailed + or optional.event == EVENT_SessionFinished): + optional.sessionId, offset = self.read_res_content(res, offset) + optional.response_meta_json, offset = self.read_res_content(res, offset) + else: + optional.sessionId, offset = self.read_res_content(res, offset) + response.payload, offset = self.read_res_payload(res, offset) + + elif header.message_type == ERROR_INFORMATION: + optional.errorCode = int.from_bytes(res[offset:offset + 4], "big", signed=True) + offset += 4 + response.payload, offset = self.read_res_payload(res, offset) + return response + + async def start_connection(self): + header = Header(message_type=FULL_CLIENT_REQUEST, message_type_specific_flags=MsgTypeFlagWithEvent).as_bytes() + optional = Optional(event=EVENT_Start_Connection).as_bytes() + payload = str.encode("{}") + return await self.send_event(header, optional, payload) + + def print_response(self, res, tag_msg: str): + logger.bind(tag=TAG).info(f'===>{tag_msg} header:{res.header.__dict__}') + logger.bind(tag=TAG).info(f'===>{tag_msg} optional:{res.optional.__dict__}') + + def get_payload_bytes(self, uid='1234', event=EVENT_NONE, text='', speaker='', audio_format='pcm', + audio_sample_rate=16000): + return str.encode(json.dumps( + { + "user": {"uid": uid}, + "event": event, + "namespace": "BidirectionalTTS", + "req_params": { + "text": text, + "speaker": speaker, + "audio_params": { + "format": audio_format, + "sample_rate": audio_sample_rate + } + } + } + )) + + async def finish_connection(self): + header = Header(message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON + ).as_bytes() + optional = Optional(event=EVENT_FinishConnection).as_bytes() + payload = str.encode('{}') + await self.send_event(header, optional, payload) + return + + async def start_session(self, session_id): + self.stop_event_response.clear() + header = Header(message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON + ).as_bytes() + optional = Optional(event=EVENT_StartSession, sessionId=session_id).as_bytes() + payload = self.get_payload_bytes(event=EVENT_StartSession, speaker=self.speaker) + await self.send_event(header, optional, payload) + + async def finish_session(self, session_id): + self.stop_event_response.set() + header = Header(message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON + ).as_bytes() + optional = Optional(event=EVENT_FinishSession, sessionId=session_id).as_bytes() + payload = str.encode('{}') + await self.send_event(header, optional, payload) + return + + async def reset(self): + # 关闭之前的对话 + if self.start_connection_flag: + await self.finish_connection() + self.start_connection_flag = False + await self.start_connection() + self.start_connection_flag = True + await super().reset() + + async def close(self): + """资源清理方法""" + await self.finish_connection() + await self.ws.close() + + async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): + # 发送文本 + await self.send_text(self.speaker, text, u_id) + return + + def _start_monitor_tts_response_thread(self): + # 初始化链接 + asyncio.run_coroutine_threadsafe(self._start_monitor_tts_response(), loop=self.loop) + + async def _start_monitor_tts_response(self): + chunk_total = b'' + while True: + try: + msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop + res = self.parser_response(msg) + self.print_response(res, 'send_text res:') + + if res.optional.event == EVENT_TTSResponse and res.header.message_type == AUDIO_ONLY_RESPONSE: + logger.bind(tag=TAG).info(f'推送数据到队列里面~~') + opus_datas = self.wav_to_opus_data_audio_raw(res.payload) + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, + tts_finish_text="", sentence_type=None, duration=0 + ) + ) + elif res.optional.event == EVENT_TTSSentenceStart: + json_data = json.loads(res.payload.decode('utf-8')) + self.tts_text = json_data.get("text", "") + logger.bind(tag=TAG).info(f'句子开始~~{self.tts_text}') + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], + tts_finish_text=self.tts_text, + sentence_type=SentenceType.SENTENCE_START + ) + ) + elif res.optional.event == EVENT_TTSSentenceEnd: + logger.bind(tag=TAG).info(f'句子结束~~{self.tts_text}') + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], + tts_finish_text=self.tts_text, + sentence_type=SentenceType.SENTENCE_END + ) + ) + elif res.optional.event == EVENT_SessionFinished: + logger.bind(tag=TAG).info(f'会话结束~~,最后一句补零') + opus_datas = self.wav_to_opus_data_audio_raw(b'', is_end=True) + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, + tts_finish_text="", sentence_type=None, duration=0 + ) + ) + self.tts_audio_queue.put( + TTSMessageDTO( + u_id=self.u_id, msg_type=MsgType.STOP_TTS_RESPONSE, content=[], + tts_finish_text=self.tts_text, + sentence_type=SentenceType.SENTENCE_END + ) + ) + else: + continue + except websockets.ConnectionClosed: + break # 连接关闭时退出监听 + except Exception as e: + logger.bind(tag=TAG).error(f"Error in _start_monitor_tts_response: {e}") + traceback.print_exc() + continue diff --git a/main/xiaozhi-server/core/utils/textUtils.py b/main/xiaozhi-server/core/utils/textUtils.py new file mode 100644 index 00000000..7fadd3ff --- /dev/null +++ b/main/xiaozhi-server/core/utils/textUtils.py @@ -0,0 +1,34 @@ +def get_string_no_punctuation_or_emoji(s): + """去除字符串首尾的空格、标点符号和表情符号""" + chars = list(s) + # 处理开头的字符 + start = 0 + while start < len(chars) and is_punctuation_or_emoji(chars[start]): + start += 1 + # 处理结尾的字符 + end = len(chars) - 1 + while end >= start and is_punctuation_or_emoji(chars[end]): + end -= 1 + return ''.join(chars[start:end + 1]) + +def is_punctuation_or_emoji(char): + """检查字符是否为空格、指定标点或表情符号""" + # 定义需要去除的中英文标点(包括全角/半角) + punctuation_set = { + ',', ',', # 中文逗号 + 英文逗号 + '。', '.', # 中文句号 + 英文句号 + '!', '!', # 中文感叹号 + 英文感叹号 + '-', '-', # 英文连字符 + 中文全角横线 + '、' # 中文顿号 + } + if char.isspace() or char in punctuation_set: + return True + # 检查表情符号(保留原有逻辑) + code_point = ord(char) + emoji_ranges = [ + (0x1F600, 0x1F64F), (0x1F300, 0x1F5FF), + (0x1F680, 0x1F6FF), (0x1F900, 0x1F9FF), + (0x1FA70, 0x1FAFF), (0x2600, 0x26FF), + (0x2700, 0x27BF) + ] + return any(start <= code_point <= end for start, end in emoji_ranges) \ No newline at end of file diff --git a/main/xiaozhi-server/core/websocket_server.py b/main/xiaozhi-server/core/websocket_server.py index 41f14515..78ac230c 100644 --- a/main/xiaozhi-server/core/websocket_server.py +++ b/main/xiaozhi-server/core/websocket_server.py @@ -78,6 +78,7 @@ class WebSocketServer: async def _handle_connection(self, websocket): """处理新连接,每次创建独立的ConnectionHandler""" # 创建ConnectionHandler时传入当前server实例 + # tts 变成链接的时候创建,避免并非问题 handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent) self.active_connections.add(handler) try: diff --git a/main/xiaozhi-server/plugins_func/functions/play_music.py b/main/xiaozhi-server/plugins_func/functions/play_music.py index 75f3834d..f89ff8ef 100644 --- a/main/xiaozhi-server/plugins_func/functions/play_music.py +++ b/main/xiaozhi-server/plugins_func/functions/play_music.py @@ -7,10 +7,11 @@ import asyncio import difflib import traceback from pathlib import Path + +from core.providers.tts.dto.dto import TTSMessageDTO, MsgType from core.utils import p3 from core.handle.sendAudioHandle import send_stt_message -from plugins_func.register import register_function,ToolType, ActionResponse, Action - +from plugins_func.register import register_function, ToolType, ActionResponse, Action TAG = __name__ logger = setup_logging() @@ -18,22 +19,22 @@ logger = setup_logging() MUSIC_CACHE = {} play_music_function_desc = { - "type": "function", - "function": { - "name": "play_music", - "description": "唱歌、听歌、播放音乐方法。比如用户说播放音乐,参数为:random,比如用户说播放两只老虎,参数为:两只老虎", - "parameters": { - "type": "object", - "properties": { - "song_name": { - "type": "string", - "description": "歌曲名称,如果没有指定具体歌名则为'random'" - } - }, - "required": ["song_name"] - } + "type": "function", + "function": { + "name": "play_music", + "description": "唱歌、听歌、播放音乐方法。比如用户说播放音乐,参数为:random,比如用户说播放两只老虎,参数为:两只老虎", + "parameters": { + "type": "object", + "properties": { + "song_name": { + "type": "string", + "description": "歌曲名称,如果没有指定具体歌名则为'random'" } - } + }, + "required": ["song_name"] + } + } +} @register_function('play_music', play_music_function_desc, ToolType.SYSTEM_CTL) @@ -189,7 +190,12 @@ async def play_local_music(conn, specific_file=None): opus_packets, duration = p3.decode_opus_from_file(music_path) else: opus_packets, duration = conn.tts.audio_to_opus_data(music_path) - conn.audio_play_queue.put((opus_packets, selected_music, 0)) + conn.tts.tts_audio_queue.put( + TTSMessageDTO( + u_id=conn.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_packets, + tts_finish_text="", sentence_type=None, duration=0 + ) + ) except Exception as e: logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}") diff --git a/main/xiaozhi-server/requirements.txt b/main/xiaozhi-server/requirements.txt index 6afec6b5..2c040b8c 100755 --- a/main/xiaozhi-server/requirements.txt +++ b/main/xiaozhi-server/requirements.txt @@ -22,3 +22,4 @@ mem0ai==0.1.62 bs4==0.0.2 modelscope==1.23.2 sherpa_onnx==1.11.0 +mutagen==1.47.0