From 0fc8f02c6e4887a7da0c1c2a745321a51318330d Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Thu, 12 Jun 2025 23:11:42 +0800 Subject: [PATCH] =?UTF-8?q?update:=E4=BC=98=E5=8C=96=E7=A9=BA=E5=AF=86?= =?UTF-8?q?=E9=92=A5=E7=9A=84=E6=8F=90=E7=A4=BA=E6=8F=90=E7=A4=BA=E6=96=B9?= =?UTF-8?q?=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../core/providers/llm/coze/coze.py | 4 +- .../core/providers/llm/dify/dify.py | 4 +- .../core/providers/llm/fastgpt/fastgpt.py | 4 +- .../core/providers/llm/gemini/gemini.py | 5 +- .../core/providers/llm/openai/openai.py | 29 +++++++---- .../core/providers/memory/mem0ai/mem0ai.py | 6 ++- .../core/providers/tts/doubao.py | 4 +- .../core/providers/tts/fishspeech.py | 5 +- .../providers/tts/huoshan_double_stream.py | 6 ++- .../core/providers/tts/linkerai.py | 10 ++-- .../core/providers/tts/openai.py | 4 +- .../core/providers/vllm/openai.py | 4 +- main/xiaozhi-server/core/utils/util.py | 14 +++--- .../plugins_func/functions/hass_init.py | 49 ++++++++++++------- 14 files changed, 96 insertions(+), 52 deletions(-) diff --git a/main/xiaozhi-server/core/providers/llm/coze/coze.py b/main/xiaozhi-server/core/providers/llm/coze/coze.py index 19002ac8..31a7d4ac 100644 --- a/main/xiaozhi-server/core/providers/llm/coze/coze.py +++ b/main/xiaozhi-server/core/providers/llm/coze/coze.py @@ -23,7 +23,9 @@ class LLMProvider(LLMProviderBase): self.bot_id = str(config.get("bot_id")) self.user_id = str(config.get("user_id")) self.session_conversation_map = {} # 存储session_id和conversation_id的映射 - check_model_key("CozeLLM", self.personal_access_token) + model_key_msg = check_model_key("CozeLLM", self.personal_access_token) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) def response(self, session_id, dialogue, **kwargs): coze_api_token = self.personal_access_token diff --git a/main/xiaozhi-server/core/providers/llm/dify/dify.py b/main/xiaozhi-server/core/providers/llm/dify/dify.py index 8b01261f..4ec04ec7 100644 --- a/main/xiaozhi-server/core/providers/llm/dify/dify.py +++ b/main/xiaozhi-server/core/providers/llm/dify/dify.py @@ -15,7 +15,9 @@ class LLMProvider(LLMProviderBase): self.mode = config.get("mode", "chat-messages") self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip("/") self.session_conversation_map = {} # 存储session_id和conversation_id的映射 - check_model_key("DifyLLM", self.api_key) + model_key_msg = check_model_key("DifyLLM", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) def response(self, session_id, dialogue, **kwargs): try: diff --git a/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py b/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py index a5581541..4f3f4259 100644 --- a/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py +++ b/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py @@ -14,7 +14,9 @@ class LLMProvider(LLMProviderBase): self.base_url = config.get("base_url") self.detail = config.get("detail", False) self.variables = config.get("variables", {}) - check_model_key("FastGPTLLM", self.api_key) + model_key_msg = check_model_key("FastGPTLLM", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) def response(self, session_id, dialogue, **kwargs): try: diff --git a/main/xiaozhi-server/core/providers/llm/gemini/gemini.py b/main/xiaozhi-server/core/providers/llm/gemini/gemini.py index 3369aa2d..608e8387 100644 --- a/main/xiaozhi-server/core/providers/llm/gemini/gemini.py +++ b/main/xiaozhi-server/core/providers/llm/gemini/gemini.py @@ -73,8 +73,9 @@ class LLMProvider(LLMProviderBase): http_proxy = cfg.get("http_proxy") https_proxy = cfg.get("https_proxy") - if not check_model_key("LLM", self.api_key): - raise ValueError("无效的Gemini API Key,请检查是否配置正确") + model_key_msg = check_model_key("LLM", self.api_key) + if model_key_msg: + log.bind(tag=TAG).error(model_key_msg) if http_proxy or https_proxy: log.bind(tag=TAG).info( diff --git a/main/xiaozhi-server/core/providers/llm/openai/openai.py b/main/xiaozhi-server/core/providers/llm/openai/openai.py index bc0e7f21..a5ee29b6 100644 --- a/main/xiaozhi-server/core/providers/llm/openai/openai.py +++ b/main/xiaozhi-server/core/providers/llm/openai/openai.py @@ -21,20 +21,27 @@ class LLMProvider(LLMProviderBase): "max_tokens": (500, int), "temperature": (0.7, lambda x: round(float(x), 1)), "top_p": (1.0, lambda x: round(float(x), 1)), - "frequency_penalty": (0, lambda x: round(float(x), 1)) + "frequency_penalty": (0, lambda x: round(float(x), 1)), } for param, (default, converter) in param_defaults.items(): value = config.get(param) try: - setattr(self, param, converter(value) if value not in (None, "") else default) + setattr( + self, + param, + converter(value) if value not in (None, "") else default, + ) except (ValueError, TypeError): setattr(self, param, default) logger.debug( - f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}") + f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}" + ) - check_model_key("LLM", self.api_key) + model_key_msg = check_model_key("LLM", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) def response(self, session_id, dialogue, **kwargs): @@ -46,7 +53,9 @@ class LLMProvider(LLMProviderBase): max_tokens=kwargs.get("max_tokens", self.max_tokens), temperature=kwargs.get("temperature", self.temperature), top_p=kwargs.get("top_p", self.top_p), - frequency_penalty=kwargs.get("frequency_penalty", self.frequency_penalty), + frequency_penalty=kwargs.get( + "frequency_penalty", self.frequency_penalty + ), ) is_active = True @@ -84,12 +93,14 @@ class LLMProvider(LLMProviderBase): for chunk in stream: # 检查是否存在有效的choice且content不为空 if getattr(chunk, "choices", None): - yield chunk.choices[0].delta.content, chunk.choices[0].delta.tool_calls + yield chunk.choices[0].delta.content, chunk.choices[ + 0 + ].delta.tool_calls # 存在 CompletionUsage 消息时,生成 Token 消耗 log - elif isinstance(getattr(chunk, 'usage', None), CompletionUsage): - usage_info = getattr(chunk, 'usage', None) + elif isinstance(getattr(chunk, "usage", None), CompletionUsage): + usage_info = getattr(chunk, "usage", None) logger.bind(tag=TAG).info( - f"Token 消耗:输入 {getattr(usage_info, 'prompt_tokens', '未知')}," + f"Token 消耗:输入 {getattr(usage_info, 'prompt_tokens', '未知')}," f"输出 {getattr(usage_info, 'completion_tokens', '未知')}," f"共计 {getattr(usage_info, 'total_tokens', '未知')}" ) diff --git a/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py b/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py index 6ef3bc8c..d61d73f3 100644 --- a/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py +++ b/main/xiaozhi-server/core/providers/memory/mem0ai/mem0ai.py @@ -12,12 +12,14 @@ class MemoryProvider(MemoryProviderBase): super().__init__(config) self.api_key = config.get("api_key", "") self.api_version = config.get("api_version", "v1.1") - have_key = check_model_key("Mem0ai", self.api_key) - if not have_key: + model_key_msg = check_model_key("Mem0ai", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) self.use_mem0 = False return else: self.use_mem0 = True + try: self.client = MemoryClient(api_key=self.api_key) logger.bind(tag=TAG).info("成功连接到 Mem0ai 服务") diff --git a/main/xiaozhi-server/core/providers/tts/doubao.py b/main/xiaozhi-server/core/providers/tts/doubao.py index 1a21ffa1..2eed9bcf 100644 --- a/main/xiaozhi-server/core/providers/tts/doubao.py +++ b/main/xiaozhi-server/core/providers/tts/doubao.py @@ -37,7 +37,9 @@ class TTSProvider(TTSProviderBase): self.api_url = config.get("api_url") self.authorization = config.get("authorization") self.header = {"Authorization": f"{self.authorization}{self.access_token}"} - check_model_key("TTS", self.access_token) + model_key_msg = check_model_key("TTS", self.access_token) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) async def text_to_speak(self, text, output_file): request_json = { diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index 3bcb1229..e192ff25 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -90,8 +90,9 @@ class TTSProvider(TTSProviderBase): self.format = config.get("response_format", "wav") self.audio_file_type = config.get("response_format", "wav") self.api_key = config.get("api_key", "YOUR_API_KEY") - have_key = check_model_key("FishSpeech TTS", self.api_key) - if not have_key: + model_key_msg = check_model_key("FishSpeech TTS", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) return self.normalize = str(config.get("normalize", True)).lower() in ( "true", diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py index 6e4b9eac..7f34c144 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py @@ -157,7 +157,9 @@ class TTSProvider(TTSProviderBase): self.opus_encoder = opus_encoder_utils.OpusEncoderUtils( sample_rate=16000, channels=1, frame_size_ms=60 ) - check_model_key("TTS", self.access_token) + model_key_msg = check_model_key("TTS", self.access_token) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) async def open_audio_channels(self, conn): try: @@ -267,7 +269,7 @@ class TTSProvider(TTSProviderBase): await handleAbortMessage(self.conn) logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本") return - + # 过滤Markdown filtered_text = MarkdownCleaner.clean_markdown(text) diff --git a/main/xiaozhi-server/core/providers/tts/linkerai.py b/main/xiaozhi-server/core/providers/tts/linkerai.py index 7493764f..ee108540 100644 --- a/main/xiaozhi-server/core/providers/tts/linkerai.py +++ b/main/xiaozhi-server/core/providers/tts/linkerai.py @@ -166,7 +166,9 @@ class TTSProvider(TTSProviderBase): ) as resp: if resp.status != 200: - logger.error(f"TTS请求失败: {resp.status}, {await resp.text()}") + logger.bind(tag=TAG).error( + f"TTS请求失败: {resp.status}, {await resp.text()}" + ) self.tts_audio_queue.put((SentenceType.LAST, [], None)) return @@ -229,7 +231,7 @@ class TTSProvider(TTSProviderBase): self._process_before_stop_play_files() except Exception as e: - logger.error(f"TTS请求异常: {e}") + logger.bind(tag=TAG).error(f"TTS请求异常: {e}") self.tts_audio_queue.put((SentenceType.LAST, [], None)) def to_tts(self, text: str) -> list: @@ -263,7 +265,7 @@ class TTSProvider(TTSProviderBase): self.api_url, params=params, headers=headers, timeout=5 ) as response: if response.status_code != 200: - logger.error( + logger.bind(tag=TAG).error( f"TTS请求失败: {response.status_code}, {response.text}" ) return [] @@ -299,5 +301,5 @@ class TTSProvider(TTSProviderBase): return opus_datas except Exception as e: - logger.error(f"TTS请求异常: {e}") + logger.bind(tag=TAG).error(f"TTS请求异常: {e}") return [] diff --git a/main/xiaozhi-server/core/providers/tts/openai.py b/main/xiaozhi-server/core/providers/tts/openai.py index e5155513..20bf9a96 100644 --- a/main/xiaozhi-server/core/providers/tts/openai.py +++ b/main/xiaozhi-server/core/providers/tts/openai.py @@ -25,7 +25,9 @@ class TTSProvider(TTSProviderBase): self.speed = float(speed) if speed else 1.0 self.output_file = config.get("output_dir", "tmp/") - check_model_key("TTS", self.api_key) + model_key_msg = check_model_key("TTS", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) async def text_to_speak(self, text, output_file): headers = { diff --git a/main/xiaozhi-server/core/providers/vllm/openai.py b/main/xiaozhi-server/core/providers/vllm/openai.py index 6ef89bca..12d8198d 100644 --- a/main/xiaozhi-server/core/providers/vllm/openai.py +++ b/main/xiaozhi-server/core/providers/vllm/openai.py @@ -34,7 +34,9 @@ class VLLMProvider(VLLMProviderBase): except (ValueError, TypeError): setattr(self, param, default) - check_model_key("VLLM", self.api_key) + model_key_msg = check_model_key("VLLM", self.api_key) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) def response(self, question, base64_image): diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index cf52bfc2..14f53b96 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -186,10 +186,8 @@ def remove_punctuation_and_length(text): def check_model_key(modelType, modelKey): if "你" in modelKey: - raise ValueError( - "你还没配置" + modelType + "的密钥,请检查一下所使用的LLM是否配置了密钥" - ) - return True + return f"配置错误: {modelType} 的 API key 未设置,当前值为: {modelKey}" + return None def parse_string_to_list(value, separator=";"): @@ -785,7 +783,9 @@ def audio_bytes_to_data(audio_bytes, file_type, is_opus=True): return p3.decode_opus_from_bytes(audio_bytes) else: # 其他格式用pydub - audio = AudioSegment.from_file(BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]) + audio = AudioSegment.from_file( + BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"] + ) audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) duration = len(audio) / 1000.0 raw_data = audio.raw_data @@ -838,11 +838,11 @@ def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1): pcm = decoder.decode(opus_frame, frame_size) pcm_datas.append(pcm) - pcm_bytes = b''.join(pcm_datas) + pcm_bytes = b"".join(pcm_datas) # 写入wav字节流 wav_buffer = BytesIO() - with wave.open(wav_buffer, 'wb') as wf: + with wave.open(wav_buffer, "wb") as wf: wf.setnchannels(channels) wf.setsampwidth(2) # 16bit wf.setframerate(sample_rate) diff --git a/main/xiaozhi-server/plugins_func/functions/hass_init.py b/main/xiaozhi-server/plugins_func/functions/hass_init.py index 8a28eef9..4e8660c9 100644 --- a/main/xiaozhi-server/plugins_func/functions/hass_init.py +++ b/main/xiaozhi-server/plugins_func/functions/hass_init.py @@ -10,9 +10,17 @@ def append_devices_to_prompt(conn): funcs = conn.config["Intent"][conn.config["selected_module"]["Intent"]].get( "functions", [] ) + + config_source = ( + "home_assistant" + if conn.config["plugins"].get("home_assistant") + else "hass_get_state" + ) + if "hass_get_state" in funcs or "hass_set_state" in funcs: prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n" - devices = conn.config["plugins"]["home_assistant"].get("devices", []) + # TODO 分割被控设备 + devices = conn.config["plugins"].get(config_source, {}).get("devices", []) if len(devices) == 0: return for device in devices: @@ -24,21 +32,26 @@ def append_devices_to_prompt(conn): def initialize_hass_handler(conn): ha_config = {} - if conn.load_function_plugin: - if conn.config["plugins"].get("home_assistant"): - ha_config["base_url"] = conn.config["plugins"]["home_assistant"].get( - "base_url" - ) - ha_config["api_key"] = conn.config["plugins"]["home_assistant"].get( - "api_key" - ) - check_model_key("home_assistant", ha_config.get("api_key")) - elif conn.config["plugins"].get("hass_get_state"): - ha_config["base_url"] = conn.config["plugins"]["hass_get_state"].get( - "base_url" - ) - ha_config["api_key"] = conn.config["plugins"]["hass_get_state"].get( - "api_key" - ) - check_model_key("home_assistant", ha_config.get("api_key")) + if not conn.load_function_plugin: + return ha_config + + # 确定配置来源 + config_source = ( + "home_assistant" + if conn.config["plugins"].get("home_assistant") + else "hass_get_state" + ) + if not conn.config["plugins"].get(config_source): + return ha_config + + # 统一获取配置 + plugin_config = conn.config["plugins"][config_source] + ha_config["base_url"] = plugin_config.get("base_url") + ha_config["api_key"] = plugin_config.get("api_key") + + # 统一检查API密钥 + model_key_msg = check_model_key("home_assistant", ha_config.get("api_key")) + if model_key_msg: + logger.bind(tag=TAG).error(model_key_msg) + return ha_config