update:智控台添加独立记忆模型配置

This commit is contained in:
hrz
2025-05-27 15:28:17 +08:00
parent 626692df29
commit be7ef08f40
11 changed files with 125 additions and 59 deletions
@@ -251,6 +251,7 @@ public class ConfigServiceImpl implements ConfigService {
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" }; String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" };
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId }; String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId };
String intentLLMModelId = null; String intentLLMModelId = null;
String memLocalShortLLMModelId = null;
for (int i = 0; i < modelIds.length; i++) { for (int i = 0; i < modelIds.length; i++) {
if (modelIds[i] == null) { if (modelIds[i] == null) {
@@ -269,7 +270,7 @@ public class ConfigServiceImpl implements ConfigService {
Map<String, Object> map = (Map<String, Object>) model.getConfigJson(); Map<String, Object> map = (Map<String, Object>) model.getConfigJson();
if ("intent_llm".equals(map.get("type"))) { if ("intent_llm".equals(map.get("type"))) {
intentLLMModelId = (String) map.get("llm"); intentLLMModelId = (String) map.get("llm");
if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) { if (StringUtils.isNotBlank(intentLLMModelId) && intentLLMModelId.equals(llmModelId)) {
intentLLMModelId = null; intentLLMModelId = null;
} }
} }
@@ -281,10 +282,31 @@ public class ConfigServiceImpl implements ConfigService {
} }
} }
} }
if ("Memory".equals(modelTypes[i])) {
Map<String, Object> map = (Map<String, Object>) 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不为空,则添加附加模型 // 如果是LLM类型,且intentLLMModelId不为空,则添加附加模型
if ("LLM".equals(modelTypes[i]) && intentLLMModelId != null) { if ("LLM".equals(modelTypes[i])) {
ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache); if (StringUtils.isNotBlank(intentLLMModelId)) {
typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson()); 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); result.put(modelTypes[i], typeConfig);
@@ -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';
@@ -162,4 +162,11 @@ databaseChangeLog:
changes: changes:
- sqlFile: - sqlFile:
encoding: utf8 encoding: utf8
path: classpath:db/changelog/202505151451.sql path: classpath:db/changelog/202505151451.sql
- changeSet:
id: 202505271412
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202505271412.sql
@@ -25,7 +25,7 @@ class LLMProvider(LLMProviderBase):
self.session_conversation_map = {} # 存储session_id和conversation_id的映射 self.session_conversation_map = {} # 存储session_id和conversation_id的映射
check_model_key("CozeLLM", self.personal_access_token) 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_token = self.personal_access_token
coze_api_base = COZE_CN_BASE_URL coze_api_base = COZE_CN_BASE_URL
@@ -17,7 +17,7 @@ class LLMProvider(LLMProviderBase):
self.session_conversation_map = {} # 存储session_id和conversation_id的映射 self.session_conversation_map = {} # 存储session_id和conversation_id的映射
check_model_key("DifyLLM", self.api_key) check_model_key("DifyLLM", self.api_key)
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# 取最后一条用户消息 # 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
@@ -16,7 +16,7 @@ class LLMProvider(LLMProviderBase):
self.variables = config.get("variables", {}) self.variables = config.get("variables", {})
check_model_key("FastGPTLLM", self.api_key) check_model_key("FastGPTLLM", self.api_key)
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# 取最后一条用户消息 # 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
@@ -112,7 +112,7 @@ class LLMProvider(LLMProviderBase):
] ]
# Gemini文档提到,无需维护session-id,直接用dialogue拼接而成 # Gemini文档提到,无需维护session-id,直接用dialogue拼接而成
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
yield from self._generate(dialogue, None) yield from self._generate(dialogue, None)
def response_with_functions(self, session_id, dialogue, functions=None): def response_with_functions(self, session_id, dialogue, functions=None):
@@ -14,7 +14,7 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url", config.get("url")) # 默认使用 base_url self.base_url = config.get("base_url", config.get("url")) # 默认使用 base_url
self.api_url = f"{self.base_url}/api/conversation/process" # 拼接完整的 API 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: try:
# home assistant语音助手自带意图,无需使用xiaozhi ai自带的,只需要把用户说的话传递给home assistant即可 # home assistant语音助手自带意图,无需使用xiaozhi ai自带的,只需要把用户说的话传递给home assistant即可
@@ -18,13 +18,13 @@ class LLMProvider(LLMProviderBase):
self.client = OpenAI( self.client = OpenAI(
base_url=self.base_url, 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模型 # 检查是否是qwen3模型
self.is_qwen3 = self.model_name and self.model_name.lower().startswith("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: try:
# 如果是qwen3模型,在用户最后一条消息中添加/no_think指令 # 如果是qwen3模型,在用户最后一条消息中添加/no_think指令
if self.is_qwen3: if self.is_qwen3:
@@ -35,7 +35,9 @@ class LLMProvider(LLMProviderBase):
for i in range(len(dialogue_copy) - 1, -1, -1): for i in range(len(dialogue_copy) - 1, -1, -1):
if dialogue_copy[i]["role"] == "user": if dialogue_copy[i]["role"] == "user":
# 在用户消息前添加/no_think指令 # 在用户消息前添加/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指令") logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令")
break break
@@ -43,9 +45,7 @@ class LLMProvider(LLMProviderBase):
dialogue = dialogue_copy dialogue = dialogue_copy
responses = self.client.chat.completions.create( responses = self.client.chat.completions.create(
model=self.model_name, model=self.model_name, messages=dialogue, stream=True
messages=dialogue,
stream=True
) )
is_active = True is_active = True
# 用于处理跨chunk的标签 # 用于处理跨chunk的标签
@@ -53,29 +53,33 @@ class LLMProvider(LLMProviderBase):
for chunk in responses: for chunk in responses:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else '' chunk.choices[0].delta
if getattr(chunk, "choices", None)
else None
)
content = delta.content if hasattr(delta, "content") else ""
if content: if content:
# 将内容添加到缓冲区 # 将内容添加到缓冲区
buffer += content buffer += content
# 处理缓冲区中的标签 # 处理缓冲区中的标签
while '<think>' in buffer and '</think>' in buffer: while "<think>" in buffer and "</think>" in buffer:
# 找到完整的<think></think>标签并移除 # 找到完整的<think></think>标签并移除
pre = buffer.split('<think>', 1)[0] pre = buffer.split("<think>", 1)[0]
post = buffer.split('</think>', 1)[1] post = buffer.split("</think>", 1)[1]
buffer = pre + post buffer = pre + post
# 处理只有开始标签的情况 # 处理只有开始标签的情况
if '<think>' in buffer: if "<think>" in buffer:
is_active = False is_active = False
buffer = buffer.split('<think>', 1)[0] buffer = buffer.split("<think>", 1)[0]
# 处理只有结束标签的情况 # 处理只有结束标签的情况
if '</think>' in buffer: if "</think>" in buffer:
is_active = True is_active = True
buffer = buffer.split('</think>', 1)[1] buffer = buffer.split("</think>", 1)[1]
# 如果当前处于活动状态且缓冲区有内容,则输出 # 如果当前处于活动状态且缓冲区有内容,则输出
if is_active and buffer: if is_active and buffer:
@@ -100,7 +104,9 @@ class LLMProvider(LLMProviderBase):
for i in range(len(dialogue_copy) - 1, -1, -1): for i in range(len(dialogue_copy) - 1, -1, -1):
if dialogue_copy[i]["role"] == "user": if dialogue_copy[i]["role"] == "user":
# 在用户消息前添加/no_think指令 # 在用户消息前添加/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指令") logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令")
break break
@@ -119,9 +125,15 @@ class LLMProvider(LLMProviderBase):
for chunk in stream: for chunk in stream:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else None chunk.choices[0].delta
tool_calls = delta.tool_calls if hasattr(delta, 'tool_calls') else None 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: if tool_calls:
@@ -134,21 +146,21 @@ class LLMProvider(LLMProviderBase):
buffer += content buffer += content
# 处理缓冲区中的标签 # 处理缓冲区中的标签
while '<think>' in buffer and '</think>' in buffer: while "<think>" in buffer and "</think>" in buffer:
# 找到完整的<think></think>标签并移除 # 找到完整的<think></think>标签并移除
pre = buffer.split('<think>', 1)[0] pre = buffer.split("<think>", 1)[0]
post = buffer.split('</think>', 1)[1] post = buffer.split("</think>", 1)[1]
buffer = pre + post buffer = pre + post
# 处理只有开始标签的情况 # 处理只有开始标签的情况
if '<think>' in buffer: if "<think>" in buffer:
is_active = False is_active = False
buffer = buffer.split('<think>', 1)[0] buffer = buffer.split("<think>", 1)[0]
# 处理只有结束标签的情况 # 处理只有结束标签的情况
if '</think>' in buffer: if "</think>" in buffer:
is_active = True is_active = True
buffer = buffer.split('</think>', 1)[1] buffer = buffer.split("</think>", 1)[1]
# 如果当前处于活动状态且缓冲区有内容,则输出 # 如果当前处于活动状态且缓冲区有内容,则输出
if is_active and buffer: if is_active and buffer:
@@ -15,39 +15,45 @@ class LLMProvider(LLMProviderBase):
# 如果没有v1,增加v1 # 如果没有v1,增加v1
if not self.base_url.endswith("/v1"): if not self.base_url.endswith("/v1"):
self.base_url = f"{self.base_url}/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: try:
self.client = OpenAI( self.client = OpenAI(
base_url=self.base_url, 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") logger.bind(tag=TAG).info("Xinference client initialized successfully")
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}") logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}")
raise raise
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
logger.bind(tag=TAG).debug(f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}") logger.bind(tag=TAG).debug(
responses = self.client.chat.completions.create( f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}"
model=self.model_name,
messages=dialogue,
stream=True
) )
is_active=True responses = self.client.chat.completions.create(
model=self.model_name, messages=dialogue, stream=True
)
is_active = True
for chunk in responses: for chunk in responses:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else '' chunk.choices[0].delta
if getattr(chunk, "choices", None)
else None
)
content = delta.content if hasattr(delta, "content") else ""
if content: if content:
if '<think>' in content: if "<think>" in content:
is_active = False is_active = False
content = content.split('<think>')[0] content = content.split("<think>")[0]
if '</think>' in content: if "</think>" in content:
is_active = True is_active = True
content = content.split('</think>')[-1] content = content.split("</think>")[-1]
if is_active: if is_active:
yield content yield content
except Exception as e: except Exception as e:
@@ -59,10 +65,14 @@ class LLMProvider(LLMProviderBase):
def response_with_functions(self, session_id, dialogue, functions=None): def response_with_functions(self, session_id, dialogue, functions=None):
try: 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: 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( stream = self.client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=dialogue, messages=dialogue,
@@ -74,7 +84,7 @@ class LLMProvider(LLMProviderBase):
delta = chunk.choices[0].delta delta = chunk.choices[0].delta
content = delta.content content = delta.content
tool_calls = delta.tool_calls tool_calls = delta.tool_calls
if content: if content:
yield content, tool_calls yield content, tool_calls
elif tool_calls: elif tool_calls:
@@ -82,4 +92,7 @@ class LLMProvider(LLMProviderBase):
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error in Xinference function call: {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)}",
}
@@ -167,7 +167,12 @@ class MemoryProvider(MemoryProviderBase):
msgStr += f"当前时间:{time_str}" msgStr += f"当前时间:{time_str}"
if self.save_to_file: 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) json_str = extract_json_data(result)
try: try:
json.loads(json_str) # 检查json格式是否正确 json.loads(json_str) # 检查json格式是否正确
@@ -177,7 +182,10 @@ class MemoryProvider(MemoryProviderBase):
print("Error:", e) print("Error:", e)
else: else:
result = self.llm.response_no_stream( 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) save_mem_local_short(self.role_id, result)
logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}") logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}")