diff --git a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java index 592b3c1d..f28c4e55 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/config/service/impl/ConfigServiceImpl.java @@ -251,6 +251,7 @@ public class ConfigServiceImpl implements ConfigService { String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" }; String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId }; String intentLLMModelId = null; + String memLocalShortLLMModelId = null; for (int i = 0; i < modelIds.length; i++) { if (modelIds[i] == null) { @@ -269,7 +270,7 @@ public class ConfigServiceImpl implements ConfigService { Map map = (Map) model.getConfigJson(); if ("intent_llm".equals(map.get("type"))) { intentLLMModelId = (String) map.get("llm"); - if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) { + if (StringUtils.isNotBlank(intentLLMModelId) && intentLLMModelId.equals(llmModelId)) { intentLLMModelId = null; } } @@ -281,10 +282,31 @@ public class ConfigServiceImpl implements ConfigService { } } } + if ("Memory".equals(modelTypes[i])) { + Map map = (Map) model.getConfigJson(); + if ("mem_local_short".equals(map.get("type"))) { + memLocalShortLLMModelId = (String) map.get("llm"); + if (StringUtils.isNotBlank(memLocalShortLLMModelId) + && memLocalShortLLMModelId.equals(llmModelId)) { + memLocalShortLLMModelId = null; + } + } + } // 如果是LLM类型,且intentLLMModelId不为空,则添加附加模型 - if ("LLM".equals(modelTypes[i]) && intentLLMModelId != null) { - ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache); - typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson()); + if ("LLM".equals(modelTypes[i])) { + if (StringUtils.isNotBlank(intentLLMModelId)) { + if (!typeConfig.containsKey(intentLLMModelId)) { + ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache); + typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson()); + } + } + if (StringUtils.isNotBlank(memLocalShortLLMModelId)) { + if (!typeConfig.containsKey(memLocalShortLLMModelId)) { + ModelConfigEntity memLocalShortLLM = modelConfigService + .getModelById(memLocalShortLLMModelId, isCache); + typeConfig.put(memLocalShortLLM.getId(), memLocalShortLLM.getConfigJson()); + } + } } } result.put(modelTypes[i], typeConfig); diff --git a/main/manager-api/src/main/resources/db/changelog/202505271412.sql b/main/manager-api/src/main/resources/db/changelog/202505271412.sql new file mode 100644 index 00000000..70cd2e0d --- /dev/null +++ b/main/manager-api/src/main/resources/db/changelog/202505271412.sql @@ -0,0 +1,4 @@ +-- 本地短期记忆配置可以设置独立的LLM + +update `ai_model_provider` set fields = '[{"key":"llm","label":"LLM模型","type":"string"}]' where id = 'SYSTEM_Memory_mem_local_short'; +update `ai_model_config` set config_json = '{\"type\": \"mem_local_short\", \"llm\": \"LLM_ChatGLMLLM\"}' where id = 'Memory_mem_local_short'; diff --git a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml index 9d3808da..aadd2235 100755 --- a/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml +++ b/main/manager-api/src/main/resources/db/changelog/db.changelog-master.yaml @@ -162,4 +162,11 @@ databaseChangeLog: changes: - sqlFile: encoding: utf8 - path: classpath:db/changelog/202505151451.sql \ No newline at end of file + path: classpath:db/changelog/202505151451.sql + - changeSet: + id: 202505271412 + author: hrz + changes: + - sqlFile: + encoding: utf8 + path: classpath:db/changelog/202505271412.sql \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/llm/coze/coze.py b/main/xiaozhi-server/core/providers/llm/coze/coze.py index 60a04667..19002ac8 100644 --- a/main/xiaozhi-server/core/providers/llm/coze/coze.py +++ b/main/xiaozhi-server/core/providers/llm/coze/coze.py @@ -25,7 +25,7 @@ class LLMProvider(LLMProviderBase): self.session_conversation_map = {} # 存储session_id和conversation_id的映射 check_model_key("CozeLLM", self.personal_access_token) - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): coze_api_token = self.personal_access_token coze_api_base = COZE_CN_BASE_URL diff --git a/main/xiaozhi-server/core/providers/llm/dify/dify.py b/main/xiaozhi-server/core/providers/llm/dify/dify.py index 52ca9853..8b01261f 100644 --- a/main/xiaozhi-server/core/providers/llm/dify/dify.py +++ b/main/xiaozhi-server/core/providers/llm/dify/dify.py @@ -17,7 +17,7 @@ class LLMProvider(LLMProviderBase): self.session_conversation_map = {} # 存储session_id和conversation_id的映射 check_model_key("DifyLLM", self.api_key) - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: # 取最后一条用户消息 last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") diff --git a/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py b/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py index b2e9d6b7..a5581541 100644 --- a/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py +++ b/main/xiaozhi-server/core/providers/llm/fastgpt/fastgpt.py @@ -16,7 +16,7 @@ class LLMProvider(LLMProviderBase): self.variables = config.get("variables", {}) check_model_key("FastGPTLLM", self.api_key) - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: # 取最后一条用户消息 last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") diff --git a/main/xiaozhi-server/core/providers/llm/gemini/gemini.py b/main/xiaozhi-server/core/providers/llm/gemini/gemini.py index d1cf238e..3369aa2d 100644 --- a/main/xiaozhi-server/core/providers/llm/gemini/gemini.py +++ b/main/xiaozhi-server/core/providers/llm/gemini/gemini.py @@ -112,7 +112,7 @@ class LLMProvider(LLMProviderBase): ] # Gemini文档提到,无需维护session-id,直接用dialogue拼接而成 - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): yield from self._generate(dialogue, None) def response_with_functions(self, session_id, dialogue, functions=None): diff --git a/main/xiaozhi-server/core/providers/llm/homeassistant/homeassistant.py b/main/xiaozhi-server/core/providers/llm/homeassistant/homeassistant.py index c436fc9e..f203e5f7 100644 --- a/main/xiaozhi-server/core/providers/llm/homeassistant/homeassistant.py +++ b/main/xiaozhi-server/core/providers/llm/homeassistant/homeassistant.py @@ -14,7 +14,7 @@ class LLMProvider(LLMProviderBase): self.base_url = config.get("base_url", config.get("url")) # 默认使用 base_url self.api_url = f"{self.base_url}/api/conversation/process" # 拼接完整的 API URL - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: # home assistant语音助手自带意图,无需使用xiaozhi ai自带的,只需要把用户说的话传递给home assistant即可 diff --git a/main/xiaozhi-server/core/providers/llm/ollama/ollama.py b/main/xiaozhi-server/core/providers/llm/ollama/ollama.py index aad6354f..c973cacd 100644 --- a/main/xiaozhi-server/core/providers/llm/ollama/ollama.py +++ b/main/xiaozhi-server/core/providers/llm/ollama/ollama.py @@ -18,13 +18,13 @@ class LLMProvider(LLMProviderBase): self.client = OpenAI( base_url=self.base_url, - api_key="ollama" # Ollama doesn't need an API key but OpenAI client requires one + api_key="ollama", # Ollama doesn't need an API key but OpenAI client requires one ) # 检查是否是qwen3模型 self.is_qwen3 = self.model_name and self.model_name.lower().startswith("qwen3") - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: # 如果是qwen3模型,在用户最后一条消息中添加/no_think指令 if self.is_qwen3: @@ -35,7 +35,9 @@ class LLMProvider(LLMProviderBase): for i in range(len(dialogue_copy) - 1, -1, -1): if dialogue_copy[i]["role"] == "user": # 在用户消息前添加/no_think指令 - dialogue_copy[i]["content"] = "/no_think " + dialogue_copy[i]["content"] + dialogue_copy[i]["content"] = ( + "/no_think " + dialogue_copy[i]["content"] + ) logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令") break @@ -43,9 +45,7 @@ class LLMProvider(LLMProviderBase): dialogue = dialogue_copy responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True + model=self.model_name, messages=dialogue, stream=True ) is_active = True # 用于处理跨chunk的标签 @@ -53,29 +53,33 @@ class LLMProvider(LLMProviderBase): for chunk in responses: try: - delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None - content = delta.content if hasattr(delta, 'content') else '' + delta = ( + chunk.choices[0].delta + if getattr(chunk, "choices", None) + else None + ) + content = delta.content if hasattr(delta, "content") else "" if content: # 将内容添加到缓冲区 buffer += content # 处理缓冲区中的标签 - while '' in buffer and '' in buffer: + while "" in buffer and "" in buffer: # 找到完整的标签并移除 - pre = buffer.split('', 1)[0] - post = buffer.split('', 1)[1] + pre = buffer.split("", 1)[0] + post = buffer.split("", 1)[1] buffer = pre + post # 处理只有开始标签的情况 - if '' in buffer: + if "" in buffer: is_active = False - buffer = buffer.split('', 1)[0] + buffer = buffer.split("", 1)[0] # 处理只有结束标签的情况 - if '' in buffer: + if "" in buffer: is_active = True - buffer = buffer.split('', 1)[1] + buffer = buffer.split("", 1)[1] # 如果当前处于活动状态且缓冲区有内容,则输出 if is_active and buffer: @@ -100,7 +104,9 @@ class LLMProvider(LLMProviderBase): for i in range(len(dialogue_copy) - 1, -1, -1): if dialogue_copy[i]["role"] == "user": # 在用户消息前添加/no_think指令 - dialogue_copy[i]["content"] = "/no_think " + dialogue_copy[i]["content"] + dialogue_copy[i]["content"] = ( + "/no_think " + dialogue_copy[i]["content"] + ) logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令") break @@ -119,9 +125,15 @@ class LLMProvider(LLMProviderBase): for chunk in stream: try: - delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None - content = delta.content if hasattr(delta, 'content') else None - tool_calls = delta.tool_calls if hasattr(delta, 'tool_calls') else None + delta = ( + chunk.choices[0].delta + if getattr(chunk, "choices", None) + else None + ) + content = delta.content if hasattr(delta, "content") else None + tool_calls = ( + delta.tool_calls if hasattr(delta, "tool_calls") else None + ) # 如果是工具调用,直接传递 if tool_calls: @@ -134,21 +146,21 @@ class LLMProvider(LLMProviderBase): buffer += content # 处理缓冲区中的标签 - while '' in buffer and '' in buffer: + while "" in buffer and "" in buffer: # 找到完整的标签并移除 - pre = buffer.split('', 1)[0] - post = buffer.split('', 1)[1] + pre = buffer.split("", 1)[0] + post = buffer.split("", 1)[1] buffer = pre + post # 处理只有开始标签的情况 - if '' in buffer: + if "" in buffer: is_active = False - buffer = buffer.split('', 1)[0] + buffer = buffer.split("", 1)[0] # 处理只有结束标签的情况 - if '' in buffer: + if "" in buffer: is_active = True - buffer = buffer.split('', 1)[1] + buffer = buffer.split("", 1)[1] # 如果当前处于活动状态且缓冲区有内容,则输出 if is_active and buffer: diff --git a/main/xiaozhi-server/core/providers/llm/xinference/xinference.py b/main/xiaozhi-server/core/providers/llm/xinference/xinference.py index b90b0418..a4e9a5e6 100644 --- a/main/xiaozhi-server/core/providers/llm/xinference/xinference.py +++ b/main/xiaozhi-server/core/providers/llm/xinference/xinference.py @@ -15,39 +15,45 @@ class LLMProvider(LLMProviderBase): # 如果没有v1,增加v1 if not self.base_url.endswith("/v1"): self.base_url = f"{self.base_url}/v1" - - logger.bind(tag=TAG).info(f"Initializing Xinference LLM provider with model: {self.model_name}, base_url: {self.base_url}") + + logger.bind(tag=TAG).info( + f"Initializing Xinference LLM provider with model: {self.model_name}, base_url: {self.base_url}" + ) try: self.client = OpenAI( base_url=self.base_url, - api_key="xinference" # Xinference has a similar setup to Ollama where it doesn't need an actual key + api_key="xinference", # Xinference has a similar setup to Ollama where it doesn't need an actual key ) logger.bind(tag=TAG).info("Xinference client initialized successfully") except Exception as e: logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}") raise - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: - logger.bind(tag=TAG).debug(f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}") - responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True + logger.bind(tag=TAG).debug( + f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}" ) - is_active=True + responses = self.client.chat.completions.create( + model=self.model_name, messages=dialogue, stream=True + ) + is_active = True for chunk in responses: try: - delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None - content = delta.content if hasattr(delta, 'content') else '' + delta = ( + chunk.choices[0].delta + if getattr(chunk, "choices", None) + else None + ) + content = delta.content if hasattr(delta, "content") else "" if content: - if '' in content: + if "" in content: is_active = False - content = content.split('')[0] - if '' in content: + content = content.split("")[0] + if "" in content: is_active = True - content = content.split('')[-1] + content = content.split("")[-1] if is_active: yield content except Exception as e: @@ -59,10 +65,14 @@ class LLMProvider(LLMProviderBase): def response_with_functions(self, session_id, dialogue, functions=None): try: - logger.bind(tag=TAG).debug(f"Sending function call request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}") + logger.bind(tag=TAG).debug( + f"Sending function call request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}" + ) if functions: - logger.bind(tag=TAG).debug(f"Function calls enabled with: {[f.get('function', {}).get('name') for f in functions]}") - + logger.bind(tag=TAG).debug( + f"Function calls enabled with: {[f.get('function', {}).get('name') for f in functions]}" + ) + stream = self.client.chat.completions.create( model=self.model_name, messages=dialogue, @@ -74,7 +84,7 @@ class LLMProvider(LLMProviderBase): delta = chunk.choices[0].delta content = delta.content tool_calls = delta.tool_calls - + if content: yield content, tool_calls elif tool_calls: @@ -82,4 +92,7 @@ class LLMProvider(LLMProviderBase): except Exception as e: logger.bind(tag=TAG).error(f"Error in Xinference function call: {e}") - yield {"type": "content", "content": f"【Xinference服务响应异常: {str(e)}】"} + yield { + "type": "content", + "content": f"【Xinference服务响应异常: {str(e)}】", + } diff --git a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py index 844035bd..9f855cc6 100644 --- a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py +++ b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py @@ -167,7 +167,12 @@ class MemoryProvider(MemoryProviderBase): msgStr += f"当前时间:{time_str}" if self.save_to_file: - result = self.llm.response_no_stream(short_term_memory_prompt, msgStr) + result = self.llm.response_no_stream( + short_term_memory_prompt, + msgStr, + max_tokens=2000, + temperature=0.2, + ) json_str = extract_json_data(result) try: json.loads(json_str) # 检查json格式是否正确 @@ -177,7 +182,10 @@ class MemoryProvider(MemoryProviderBase): print("Error:", e) else: result = self.llm.response_no_stream( - short_term_memory_prompt_only_content, msgStr + short_term_memory_prompt_only_content, + msgStr, + max_tokens=2000, + temperature=0.2, ) save_mem_local_short(self.role_id, result) logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}")