diff --git a/config.yaml b/config.yaml index 68b4b978..de65c3ab 100644 --- a/config.yaml +++ b/config.yaml @@ -27,7 +27,9 @@ delete_audio: true selected_module: ASR: FunASR VAD: SileroVAD + # 将根据配置名称对应的type调用实际的LLM适配器 LLM: ChatGLMLLM + # TTS将根据配置名称对应的type调用实际的TTS适配器 TTS: EdgeTTS ASR: @@ -42,31 +44,45 @@ VAD: min_silence_duration_ms: 700 # 如果说话停顿比较长,可以把这个值设置大一些 LLM: + # 当前支持的type为openai、dify,可自行适配 AliLLM: + # 定义LLM API类型 + type: openai # 可在这里找到你的 api_key https://bailian.console.aliyun.com/?apiKey=1#/api-key base_url: https://dashscope.aliyuncs.com/compatible-mode/v1 model_name: qwen-turbo api_key: 你的阿里云dashscope api key DeepSeekLLM: + # 定义LLM API类型 + type: openai # 可在这里找到你的api key https://platform.deepseek.com/ model_name: deepseek-chat url: https://api.deepseek.com api_key: 你的deepseek api key ChatGLMLLM: + # 定义LLM API类型 + type: openai # 可在这里找到你的api key https://bigmodel.cn/usercenter/proj-mgmt/apikeys model_name: glm-4-flash url: https://open.bigmodel.cn/api/paas/v4/ api_key: 你的bigmodel api key DifyLLM: + # 定义LLM API类型 + type: dify # 建议使用本地部署的dify接口,国内部分区域访问dify公有云接口可能会受限 # 如果使用DifyLLM,配置文件里prompt(提示词)是无效的,需要在dify控制台设置提示词 base_url: 你的私有化部署的dify接口地址 api_key: 你的dify api key TTS: + # 当前支持的type为edge、doubao,可自行适配 EdgeTTS: + # 定义TTS API类型 + type: edge voice: zh-CN-XiaoxiaoNeural output_file: tmp/ DoubaoTTS: + # 定义TTS API类型 + type: doubao # 火山引擎语音合成服务,需要先在火山引擎控制台创建应用并获取appid和access_token # 山引擎语音一定要购买花钱,起步价30元,就有100并发了。如果用免费的只有2个并发,会经常报tts错误 # 购买服务后,购买免费的音色后,可能要等半小时左右,才能使用。 diff --git a/core/providers/llm/base.py b/core/providers/llm/base.py new file mode 100644 index 00000000..ac8a807c --- /dev/null +++ b/core/providers/llm/base.py @@ -0,0 +1,7 @@ +from abc import ABC, abstractmethod + +class LLMProviderBase(ABC): + @abstractmethod + def response(self, session_id, dialogue): + """LLM response generator""" + pass \ No newline at end of file diff --git a/core/providers/llm/dify/dify.py b/core/providers/llm/dify/dify.py new file mode 100644 index 00000000..8fe7fdde --- /dev/null +++ b/core/providers/llm/dify/dify.py @@ -0,0 +1,36 @@ +import logging +import requests +from core.providers.llm.base import LLMProviderBase + +logger = logging.getLogger(__name__) + +class LLMProvider(LLMProviderBase): + def __init__(self, config): + self.api_key = config["api_key"] + self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/') + + def response(self, session_id, dialogue): + try: + # 取最后一条用户消息 + last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") + + # 发起流式请求 + with requests.post( + f"{self.base_url}/chat-messages", + headers={"Authorization": f"Bearer {self.api_key}"}, + json={ + "query": last_msg["content"], + "response_mode": "streaming", + "user": session_id, + "inputs": {} + }, + stream=True + ) as r: + for line in r.iter_lines(): + if line.startswith(b'data: '): + event = json.loads(line[6:]) + if event.get('answer'): + yield event['answer'] + + except Exception: + yield "【服务响应异常】" \ No newline at end of file diff --git a/core/providers/llm/openai/openai.py b/core/providers/llm/openai/openai.py new file mode 100644 index 00000000..ef2074b6 --- /dev/null +++ b/core/providers/llm/openai/openai.py @@ -0,0 +1,32 @@ +import logging +import openai +from core.providers.llm.base import LLMProviderBase + +logger = logging.getLogger(__name__) + +class LLMProvider(LLMProviderBase): + def __init__(self, config): + self.model_name = config.get("model_name") + self.api_key = config.get("api_key") + if 'base_url' in config: + self.base_url = config.get("base_url") + else: + self.base_url = config.get("url") + self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) + + def response(self, session_id, dialogue): + try: + responses = self.client.chat.completions.create( + model=self.model_name, + messages=dialogue, + stream=True + ) + for chunk in responses: + # 检查是否存在有效的choice且content不为空 + if chunk.choices and len(chunk.choices) > 0: + delta = chunk.choices[0].delta + content = getattr(delta, 'content', '') + if content: # 仅在content非空时生成 + yield content + except Exception as e: + logger.error(f"Error in response generation: {e}") \ No newline at end of file diff --git a/core/providers/tts/base.py b/core/providers/tts/base.py new file mode 100644 index 00000000..62389f02 --- /dev/null +++ b/core/providers/tts/base.py @@ -0,0 +1,89 @@ +import asyncio +import logging +import os +import json +import uuid +import base64 +# from datetime import datetime +# import edge_tts +import numpy as np +import opuslib +# import requests +# from core.utils.util import read_config, get_project_dir +from pydub import AudioSegment +from abc import ABC, abstractmethod + +logger = logging.getLogger(__name__) + +class TTSProviderBase(ABC): + def __init__(self, config, delete_audio_file): + self.delete_audio_file = delete_audio_file + self.output_file = config.get("output_file") + + @abstractmethod + def generate_filename(self): + pass + + def to_tts(self, text): + tmp_file = self.generate_filename() + try: + max_repeat_time = 5 + while not os.path.exists(tmp_file) and max_repeat_time > 0: + asyncio.run(self.text_to_speak(text, tmp_file)) + if not os.path.exists(tmp_file): + max_repeat_time = max_repeat_time - 1 + logger.error(f"语音生成失败: {text}:{tmp_file},再试{max_repeat_time}次") + + if max_repeat_time>0: + logger.info(f"语音生成成功: {text}:{tmp_file},重试{5-max_repeat_time}次") + + return tmp_file + except Exception as e: + logger.info(f"Failed to generate TTS file: {e}") + return None + + @abstractmethod + async def text_to_speak(self, text, output_file): + pass + + def wav_to_opus_data(self, wav_file_path): + # 使用pydub加载PCM文件 + # 获取文件后缀名 + file_type = os.path.splitext(wav_file_path)[1] + if file_type: + file_type = file_type.lstrip('.') + audio = AudioSegment.from_file(wav_file_path, format=file_type) + + duration = len(audio) / 1000.0 + + # 转换为单声道和16kHz采样率(确保与编码器匹配) + audio = audio.set_channels(1).set_frame_rate(16000) + + # 获取原始PCM数据(16位小端) + raw_data = audio.raw_data + + # 初始化Opus编码器 + encoder = opuslib.Encoder(16000, 1, opuslib.APPLICATION_AUDIO) + + # 编码参数 + frame_duration = 60 # 60ms per frame + frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame + + opus_datas = [] + # 按帧处理所有音频数据(包括最后一帧可能补零) + for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample + # 获取当前帧的二进制数据 + chunk = raw_data[i:i + frame_size * 2] + + # 如果最后一帧不足,补零 + if len(chunk) < frame_size * 2: + chunk += b'\x00' * (frame_size * 2 - len(chunk)) + + # 转换为numpy数组处理 + np_frame = np.frombuffer(chunk, dtype=np.int16) + + # 编码Opus数据 + opus_data = encoder.encode(np_frame.tobytes(), frame_size) + opus_datas.append(opus_data) + + return opus_datas, duration \ No newline at end of file diff --git a/core/providers/tts/doubao.py b/core/providers/tts/doubao.py new file mode 100644 index 00000000..073f2291 --- /dev/null +++ b/core/providers/tts/doubao.py @@ -0,0 +1,55 @@ +import os +import uuid +import json +import base64 +import requests +from datetime import datetime +from core.providers.tts.base import TTSProviderBase + +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.host = "openspeech.bytedance.com" + self.api_url = f"https://{self.host}/api/v1/tts" + self.header = {"Authorization": f"Bearer;{self.access_token}"} + + 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): + request_json = { + "app": { + "appid": self.appid, + "token": "access_token", + "cluster": self.cluster + }, + "user": { + "uid": "1" + }, + "audio": { + "voice_type": self.voice, + "encoding": "wav", + "speed_ratio": 1.0, + "volume_ratio": 1.0, + "pitch_ratio": 1.0, + }, + "request": { + "reqid": str(uuid.uuid4()), + "text": text, + "text_type": "plain", + "operation": "query", + "with_frontend": 1, + "frontend_type": "unitTson" + } + } + + resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header) + if "data" in resp.json(): + data = resp.json()["data"] + file_to_save = open(output_file, "wb") + file_to_save.write(base64.b64decode(data)) \ No newline at end of file diff --git a/core/providers/tts/edge.py b/core/providers/tts/edge.py new file mode 100644 index 00000000..a0e28864 --- /dev/null +++ b/core/providers/tts/edge.py @@ -0,0 +1,17 @@ +import os +import uuid +import edge_tts +from datetime import datetime +from core.providers.tts.base import TTSProviderBase + +class TTSProvider(TTSProviderBase): + def __init__(self, config, delete_audio_file): + super().__init__(config, delete_audio_file) + self.voice = config.get("voice") + + 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) \ No newline at end of file diff --git a/core/server.py b/core/server.py index a4454e61..c28c3151 100644 --- a/core/server.py +++ b/core/server.py @@ -25,11 +25,17 @@ class WebSocketServer: self.config["delete_audio"] ), llm.create_instance( - self.config["selected_module"]["LLM"], + self.config["selected_module"]["LLM"] + if not 'type' in self.config["LLM"][self.config["selected_module"]["LLM"]] + else + self.config["LLM"][self.config["selected_module"]["LLM"]]['type'], self.config["LLM"][self.config["selected_module"]["LLM"]], ), tts.create_instance( - self.config["selected_module"]["TTS"], + self.config["selected_module"]["TTS"] + if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]] + else + self.config["TTS"][self.config["selected_module"]["TTS"]]["type"], self.config["TTS"][self.config["selected_module"]["TTS"]], self.config["delete_audio"] ) diff --git a/core/utils/llm.py b/core/utils/llm.py index 0b1e8ccc..79b48dbb 100644 --- a/core/utils/llm.py +++ b/core/utils/llm.py @@ -1,7 +1,10 @@ +import os +import sys import json import logging import openai import requests +import importlib from datetime import datetime from core.utils.util import is_segment from core.utils.util import get_string_no_punctuation_or_emoji @@ -11,132 +14,15 @@ from abc import ABC, abstractmethod logger = logging.getLogger(__name__) -class LLM(ABC): - @abstractmethod - def response(self, session_id, dialogue): - """LLM response generator""" - pass - - -class DeepSeekLLM(LLM): - def __init__(self, config): - self.model_name = config.get("model_name") - self.api_key = config.get("api_key") - self.base_url = config.get("url") - self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) - - def response(self, session_id, dialogue): - logger.info(f"Generating response using {dialogue}") - try: - responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True - ) - for chunk in responses: - # 检查是否存在有效的choice且content不为空 - if chunk.choices and len(chunk.choices) > 0: - delta = chunk.choices[0].delta - content = getattr(delta, 'content', '') - if content: # 仅在content非空时生成 - yield content - except Exception as e: - logger.error(f"Error in response generation: {e}") - - -class ChatGLMLLM(LLM): - def __init__(self, config): - self.model_name = config.get("model_name") - self.api_key = config.get("api_key") - self.base_url = config.get("url") - self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) - - def response(self, session_id, dialogue): - try: - responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True - ) - for chunk in responses: - # 检查是否存在有效的choice且content不为空 - if chunk.choices and len(chunk.choices) > 0: - delta = chunk.choices[0].delta - content = getattr(delta, 'content', '') - if content: # 仅在content非空时生成 - yield content - except Exception as e: - logger.error(f"Error in response generation: {e}") - -class AliLLM(LLM): - def __init__(self, config): - self.model_name = config.get("model_name") - self.api_key = config.get("api_key") - self.base_url = config.get("base_url") - self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) - - def response(self, session_id, dialogue): - try: - responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True - ) - for chunk in responses: - # 检查是否存在有效的choice且content不为空 - if chunk.choices and len(chunk.choices) > 0: - delta = chunk.choices[0].delta - content = getattr(delta, 'content', '') - if content: # 仅在content非空时生成 - yield content - except Exception as e: - logger.error(f"Error in response generation: {e}") - -class DifyLLM(LLM): - def __init__(self, config): - self.api_key = config["api_key"] - self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/') - - def response(self, session_id, dialogue): - try: - # 取最后一条用户消息 - last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") - - # 发起流式请求 - with requests.post( - f"{self.base_url}/chat-messages", - headers={"Authorization": f"Bearer {self.api_key}"}, - json={ - "query": last_msg["content"], - "response_mode": "streaming", - "user": session_id, - "inputs": {} - }, - stream=True - ) as r: - for line in r.iter_lines(): - if line.startswith(b'data: '): - event = json.loads(line[6:]) - if event.get('answer'): - yield event['answer'] - - except Exception: - yield "【服务响应异常】" - - def create_instance(class_name, *args, **kwargs): - # 获取类对象 - cls_map = { - "DeepSeekLLM": DeepSeekLLM, - "ChatGLMLLM": ChatGLMLLM, - "DifyLLM": DifyLLM, - "AliLLM": AliLLM, - # 可扩展其他LLM实现 - } + # 创建LLM实例 + if os.path.exists(os.path.join('core', 'providers', 'llm', class_name, f'{class_name}.py')): + lib_name = f'core.providers.llm.{class_name}.{class_name}' + if lib_name not in sys.modules: + sys.modules[lib_name] = importlib.import_module(f'{lib_name}') + return sys.modules[lib_name].LLMProvider(*args, **kwargs) - if cls := cls_map.get(class_name): - return cls(*args, **kwargs) - raise ValueError(f"不支持的LLM类型: {class_name}") + raise ValueError(f"不支持的LLM类型: {class_name},请检查该配置的type是否设置正确") if __name__ == "__main__": @@ -145,7 +31,10 @@ if __name__ == "__main__": """ config = read_config(get_project_dir() + "config.yaml") llm = create_instance( - config["selected_module"]["LLM"], + config["selected_module"]["LLM"] + if not "type" in config["LLM"][config["selected_module"]["LLM"]] + else + config["LLM"][config["selected_module"]["LLM"]]["type"], config["LLM"][config["selected_module"]["LLM"]] ) diff --git a/core/utils/tts.py b/core/utils/tts.py index a18a0eb9..25914cde 100644 --- a/core/utils/tts.py +++ b/core/utils/tts.py @@ -1,9 +1,11 @@ import asyncio import logging import os +import sys import json import uuid import base64 +import importlib from datetime import datetime import edge_tts import numpy as np @@ -16,150 +18,15 @@ from abc import ABC, abstractmethod logger = logging.getLogger(__name__) -class TTS(ABC): - def __init__(self, config, delete_audio_file): - self.delete_audio_file = delete_audio_file - self.output_file = config.get("output_file") - - @abstractmethod - def generate_filename(self): - pass - - def to_tts(self, text): - tmp_file = self.generate_filename() - try: - max_repeat_time = 5 - while not os.path.exists(tmp_file) and max_repeat_time > 0: - asyncio.run(self.text_to_speak(text, tmp_file)) - if not os.path.exists(tmp_file): - max_repeat_time = max_repeat_time - 1 - logger.error(f"语音生成失败: {text}:{tmp_file},再试{max_repeat_time}次") - - return tmp_file - except Exception as e: - logger.info(f"Failed to generate TTS file: {e}") - return None - - @abstractmethod - async def text_to_speak(self, text, output_file): - pass - - def wav_to_opus_data(self, wav_file_path): - # 使用pydub加载PCM文件 - # 获取文件后缀名 - file_type = os.path.splitext(wav_file_path)[1] - if file_type: - file_type = file_type.lstrip('.') - audio = AudioSegment.from_file(wav_file_path, format=file_type) - - duration = len(audio) / 1000.0 - - # 转换为单声道和16kHz采样率(确保与编码器匹配) - audio = audio.set_channels(1).set_frame_rate(16000) - - # 获取原始PCM数据(16位小端) - raw_data = audio.raw_data - - # 初始化Opus编码器 - encoder = opuslib.Encoder(16000, 1, opuslib.APPLICATION_AUDIO) - - # 编码参数 - frame_duration = 60 # 60ms per frame - frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame - - opus_datas = [] - # 按帧处理所有音频数据(包括最后一帧可能补零) - for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample - # 获取当前帧的二进制数据 - chunk = raw_data[i:i + frame_size * 2] - - # 如果最后一帧不足,补零 - if len(chunk) < frame_size * 2: - chunk += b'\x00' * (frame_size * 2 - len(chunk)) - - # 转换为numpy数组处理 - np_frame = np.frombuffer(chunk, dtype=np.int16) - - # 编码Opus数据 - opus_data = encoder.encode(np_frame.tobytes(), frame_size) - opus_datas.append(opus_data) - - return opus_datas, duration - - -class EdgeTTS(TTS): - def __init__(self, config, delete_audio_file): - super().__init__(config, delete_audio_file) - self.voice = config.get("voice") - - 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) - - -class DoubaoTTS(TTS): - 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.host = "openspeech.bytedance.com" - self.api_url = f"https://{self.host}/api/v1/tts" - self.header = {"Authorization": f"Bearer;{self.access_token}"} - - 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): - request_json = { - "app": { - "appid": self.appid, - "token": "access_token", - "cluster": self.cluster - }, - "user": { - "uid": "1" - }, - "audio": { - "voice_type": self.voice, - "encoding": "wav", - "speed_ratio": 1.0, - "volume_ratio": 1.0, - "pitch_ratio": 1.0, - }, - "request": { - "reqid": str(uuid.uuid4()), - "text": text, - "text_type": "plain", - "operation": "query", - "with_frontend": 1, - "frontend_type": "unitTson" - } - } - - resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header) - if "data" in resp.json(): - data = resp.json()["data"] - file_to_save = open(output_file, "wb") - file_to_save.write(base64.b64decode(data)) - - def create_instance(class_name, *args, **kwargs): - # 获取类对象 - cls_map = { - "DoubaoTTS": DoubaoTTS, - "EdgeTTS": EdgeTTS, - # 可扩展其他TTS实现 - } + # 创建TTS实例 + if os.path.exists(os.path.join('core', 'providers', 'tts', f'{class_name}.py')): + lib_name = f'core.providers.tts.{class_name}' + if lib_name not in sys.modules: + sys.modules[lib_name] = importlib.import_module(f'{lib_name}') + return sys.modules[lib_name].TTSProvider(*args, **kwargs) - if cls := cls_map.get(class_name): - return cls(*args, **kwargs) - raise ValueError(f"不支持的TTS类型: {class_name}") + raise ValueError(f"不支持的TTS类型: {class_name},请检查该配置的type是否设置正确") if __name__ == "__main__": @@ -168,7 +35,10 @@ if __name__ == "__main__": """ config = read_config(get_project_dir() + "config.yaml") tts = create_instance( - config["selected_module"]["TTS"], + config["selected_module"]["TTS"] + if not 'type' in config["TTS"][config["selected_module"]["TTS"]] + else + config["TTS"][config["selected_module"]["TTS"]]["type"], config["TTS"][config["selected_module"]["TTS"]], config["delete_audio"] )