mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-25 08:33:53 +08:00
fix:manager意图识别bug (#762)
* fix:连接manager后无法使用functioncallbug * fix:意图识别使用llm无法播放音乐bug * fix:manager第一图识别使用独立llm无法初始化llm的bug
This commit is contained in:
+1
-1
@@ -44,7 +44,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
|||||||
@Override
|
@Override
|
||||||
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
|
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
|
||||||
IPage<AgentEntity> page = agentDao.selectPage(
|
IPage<AgentEntity> page = agentDao.selectPage(
|
||||||
getPage(params, "sort", true),
|
getPage(params, "agent_name", true),
|
||||||
new QueryWrapper<>());
|
new QueryWrapper<>());
|
||||||
return new PageData<>(page.getRecords(), page.getTotal());
|
return new PageData<>(page.getRecords(), page.getTotal());
|
||||||
}
|
}
|
||||||
|
|||||||
+20
-2
@@ -231,8 +231,9 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
boolean isCache) {
|
boolean isCache) {
|
||||||
Map<String, String> selectedModule = new HashMap<>();
|
Map<String, String> selectedModule = new HashMap<>();
|
||||||
|
|
||||||
String[] modelTypes = { "VAD", "ASR", "LLM", "TTS", "Memory", "Intent" };
|
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" };
|
||||||
String[] modelIds = { vadModelId, asrModelId, llmModelId, ttsModelId, memModelId, intentModelId };
|
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId };
|
||||||
|
String intentLLMModelId = 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) {
|
||||||
@@ -246,10 +247,27 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
if ("TTS".equals(modelTypes[i]) && voice != null) {
|
if ("TTS".equals(modelTypes[i]) && voice != null) {
|
||||||
((Map<String, Object>) model.getConfigJson()).put("private_voice", voice);
|
((Map<String, Object>) model.getConfigJson()).put("private_voice", voice);
|
||||||
}
|
}
|
||||||
|
// 如果是Intent类型,且type=intent_llm,则给他添加附加模型
|
||||||
|
if ("Intent".equals(modelTypes[i])) {
|
||||||
|
Map<String, Object> map = (Map<String, Object>) model.getConfigJson();
|
||||||
|
if ("intent_llm".equals(map.get("type"))) {
|
||||||
|
intentLLMModelId = (String) map.get("llm");
|
||||||
|
if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) {
|
||||||
|
intentLLMModelId = null;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// 如果是LLM类型,且intentLLMModelId不为空,则添加附加模型
|
||||||
|
if ("LLM".equals(modelTypes[i]) && intentLLMModelId != null) {
|
||||||
|
ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache);
|
||||||
|
typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson());
|
||||||
|
}
|
||||||
}
|
}
|
||||||
result.put(modelTypes[i], typeConfig);
|
result.put(modelTypes[i], typeConfig);
|
||||||
|
|
||||||
selectedModule.put(modelTypes[i], model.getId());
|
selectedModule.put(modelTypes[i], model.getId());
|
||||||
}
|
}
|
||||||
|
|
||||||
result.put("selected_module", selectedModule);
|
result.put("selected_module", selectedModule);
|
||||||
if (StringUtils.isNotBlank(prompt)) {
|
if (StringUtils.isNotBlank(prompt)) {
|
||||||
prompt = prompt.replace("{{assistant_name}}", "小智");
|
prompt = prompt.replace("{{assistant_name}}", "小智");
|
||||||
|
|||||||
+1
-1
@@ -221,7 +221,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
|
|||||||
params.put(Constant.PAGE, dto.getPage());
|
params.put(Constant.PAGE, dto.getPage());
|
||||||
params.put(Constant.LIMIT, dto.getLimit());
|
params.put(Constant.LIMIT, dto.getLimit());
|
||||||
IPage<DeviceEntity> page = baseDao.selectPage(
|
IPage<DeviceEntity> page = baseDao.selectPage(
|
||||||
getPage(params, "sort", true),
|
getPage(params, "mac_address", true),
|
||||||
// 定义查询条件
|
// 定义查询条件
|
||||||
new QueryWrapper<DeviceEntity>()
|
new QueryWrapper<DeviceEntity>()
|
||||||
// 必须设备关键词查找
|
// 必须设备关键词查找
|
||||||
|
|||||||
+1
-1
@@ -37,7 +37,7 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
|
|||||||
@Override
|
@Override
|
||||||
public PageData<SysParamsDTO> page(Map<String, Object> params) {
|
public PageData<SysParamsDTO> page(Map<String, Object> params) {
|
||||||
IPage<SysParamsEntity> page = baseDao.selectPage(
|
IPage<SysParamsEntity> page = baseDao.selectPage(
|
||||||
getPage(params, Constant.CREATE_DATE, false),
|
getPage(params, null, false),
|
||||||
getWrapper(params));
|
getWrapper(params));
|
||||||
|
|
||||||
return getPageData(page, SysParamsDTO.class);
|
return getPageData(page, SysParamsDTO.class);
|
||||||
|
|||||||
+1
-1
@@ -47,7 +47,7 @@ public class TimbreServiceImpl extends BaseServiceImpl<TimbreDao, TimbreEntity>
|
|||||||
params.put(Constant.PAGE, dto.getPage());
|
params.put(Constant.PAGE, dto.getPage());
|
||||||
params.put(Constant.LIMIT, dto.getLimit());
|
params.put(Constant.LIMIT, dto.getLimit());
|
||||||
IPage<TimbreEntity> page = baseDao.selectPage(
|
IPage<TimbreEntity> page = baseDao.selectPage(
|
||||||
getPage(params, "sort", true),
|
getPage(params, null, true),
|
||||||
// 定义查询条件
|
// 定义查询条件
|
||||||
new QueryWrapper<TimbreEntity>()
|
new QueryWrapper<TimbreEntity>()
|
||||||
// 必须按照ttsID查找
|
// 必须按照ttsID查找
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
-- 对0.3.0版本之前的参数进行修改
|
||||||
|
update `sys_params` set param_value = '.mp3;.wav;.p3' where param_code = 'plugins.play_music.music_ext';
|
||||||
|
update `ai_model_config` set config_json = '{\"type\": \"intent_llm\", \"llm\": \"LLM_ChatGLMLLM\"}' where id = 'Intent_intent_llm';
|
||||||
@@ -50,4 +50,11 @@ databaseChangeLog:
|
|||||||
changes:
|
changes:
|
||||||
- sqlFile:
|
- sqlFile:
|
||||||
encoding: utf8
|
encoding: utf8
|
||||||
path: classpath:db/changelog/202504112058.sql
|
path: classpath:db/changelog/202504112058.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202504131542
|
||||||
|
author: John
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202504131542.sql
|
||||||
@@ -104,8 +104,6 @@ class ConnectionHandler:
|
|||||||
|
|
||||||
self.close_after_chat = False # 是否在聊天结束后关闭连接
|
self.close_after_chat = False # 是否在聊天结束后关闭连接
|
||||||
self.use_function_call_mode = False
|
self.use_function_call_mode = False
|
||||||
if self.config["selected_module"]["Intent"] == "function_call":
|
|
||||||
self.use_function_call_mode = True
|
|
||||||
|
|
||||||
async def handle_connection(self, ws):
|
async def handle_connection(self, ws):
|
||||||
try:
|
try:
|
||||||
@@ -222,37 +220,37 @@ class ConnectionHandler:
|
|||||||
)
|
)
|
||||||
if private_config.get("VAD", None) is not None:
|
if private_config.get("VAD", None) is not None:
|
||||||
init_vad = True
|
init_vad = True
|
||||||
self.config["vad"] = private_config["VAD"]
|
self.config["VAD"] = private_config["VAD"]
|
||||||
self.config["selected_module"]["VAD"] = private_config["selected_module"][
|
self.config["selected_module"]["VAD"] = private_config["selected_module"][
|
||||||
"VAD"
|
"VAD"
|
||||||
]
|
]
|
||||||
if private_config.get("ASR", None) is not None:
|
if private_config.get("ASR", None) is not None:
|
||||||
init_asr = True
|
init_asr = True
|
||||||
self.config["asr"] = private_config["ASR"]
|
self.config["ASR"] = private_config["ASR"]
|
||||||
self.config["selected_module"]["ASR"] = private_config["selected_module"][
|
self.config["selected_module"]["ASR"] = private_config["selected_module"][
|
||||||
"ASR"
|
"ASR"
|
||||||
]
|
]
|
||||||
if private_config.get("LLM", None) is not None:
|
if private_config.get("LLM", None) is not None:
|
||||||
init_llm = True
|
init_llm = True
|
||||||
self.config["llm"] = private_config["LLM"]
|
self.config["LLM"] = private_config["LLM"]
|
||||||
self.config["selected_module"]["LLM"] = private_config["selected_module"][
|
self.config["selected_module"]["LLM"] = private_config["selected_module"][
|
||||||
"LLM"
|
"LLM"
|
||||||
]
|
]
|
||||||
if private_config.get("TTS", None) is not None:
|
if private_config.get("TTS", None) is not None:
|
||||||
init_tts = True
|
init_tts = True
|
||||||
self.config["tts"] = private_config["TTS"]
|
self.config["TTS"] = private_config["TTS"]
|
||||||
self.config["selected_module"]["TTS"] = private_config["selected_module"][
|
self.config["selected_module"]["TTS"] = private_config["selected_module"][
|
||||||
"TTS"
|
"TTS"
|
||||||
]
|
]
|
||||||
if private_config.get("Memory", None) is not None:
|
if private_config.get("Memory", None) is not None:
|
||||||
init_memory = True
|
init_memory = True
|
||||||
self.config["memory"] = private_config["Memory"]
|
self.config["Memory"] = private_config["Memory"]
|
||||||
self.config["selected_module"]["Memory"] = private_config[
|
self.config["selected_module"]["Memory"] = private_config[
|
||||||
"selected_module"
|
"selected_module"
|
||||||
]["Memory"]
|
]["Memory"]
|
||||||
if private_config.get("Intent", None) is not None:
|
if private_config.get("Intent", None) is not None:
|
||||||
init_intent = True
|
init_intent = True
|
||||||
self.config["intent"] = private_config["Intent"]
|
self.config["Intent"] = private_config["Intent"]
|
||||||
self.config["selected_module"]["Intent"] = private_config[
|
self.config["selected_module"]["Intent"] = private_config[
|
||||||
"selected_module"
|
"selected_module"
|
||||||
]["Intent"]
|
]["Intent"]
|
||||||
@@ -287,17 +285,26 @@ class ConnectionHandler:
|
|||||||
self.memory.init_memory(device_id, self.llm)
|
self.memory.init_memory(device_id, self.llm)
|
||||||
|
|
||||||
def _initialize_intent(self):
|
def _initialize_intent(self):
|
||||||
|
if (
|
||||||
|
self.config["Intent"][self.config["selected_module"]["Intent"]]["type"]
|
||||||
|
== "function_call"
|
||||||
|
):
|
||||||
|
self.use_function_call_mode = True
|
||||||
"""初始化意图识别模块"""
|
"""初始化意图识别模块"""
|
||||||
# 获取意图识别配置
|
# 获取意图识别配置
|
||||||
intent_config = self.config["Intent"]
|
intent_config = self.config["Intent"]
|
||||||
intent_type = self.config["selected_module"]["Intent"]
|
intent_type = self.config["Intent"][self.config["selected_module"]["Intent"]][
|
||||||
|
"type"
|
||||||
|
]
|
||||||
|
|
||||||
# 如果使用 nointent,直接返回
|
# 如果使用 nointent,直接返回
|
||||||
if intent_type == "nointent":
|
if intent_type == "nointent":
|
||||||
return
|
return
|
||||||
# 使用 intent_llm 模式
|
# 使用 intent_llm 模式
|
||||||
elif intent_type == "intent_llm":
|
elif intent_type == "intent_llm":
|
||||||
intent_llm_name = intent_config["intent_llm"]["llm"]
|
intent_llm_name = intent_config[self.config["selected_module"]["Intent"]][
|
||||||
|
"llm"
|
||||||
|
]
|
||||||
|
|
||||||
if intent_llm_name and intent_llm_name in self.config["LLM"]:
|
if intent_llm_name and intent_llm_name in self.config["LLM"]:
|
||||||
# 如果配置了专用LLM,则创建独立的LLM实例
|
# 如果配置了专用LLM,则创建独立的LLM实例
|
||||||
|
|||||||
@@ -57,7 +57,9 @@ class FunctionHandler:
|
|||||||
|
|
||||||
def register_config_functions(self):
|
def register_config_functions(self):
|
||||||
"""注册配置中的函数,可以不同客户端使用不同的配置"""
|
"""注册配置中的函数,可以不同客户端使用不同的配置"""
|
||||||
for func in self.config["Intent"]["function_call"].get("functions", []):
|
for func in self.config["Intent"][self.config["selected_module"]["Intent"]].get(
|
||||||
|
"functions", []
|
||||||
|
):
|
||||||
self.function_registry.register_function(func)
|
self.function_registry.register_function(func)
|
||||||
|
|
||||||
"""home assistant需要初始化提示词"""
|
"""home assistant需要初始化提示词"""
|
||||||
|
|||||||
@@ -76,6 +76,11 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
if function_name == "continue_chat":
|
if function_name == "continue_chat":
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
if function_name == "play_music":
|
||||||
|
funcItem = conn.func_handler.get_function(function_name)
|
||||||
|
if not funcItem:
|
||||||
|
conn.func_handler.function_registry.register_function("play_music")
|
||||||
|
|
||||||
function_args = None
|
function_args = None
|
||||||
if "arguments" in intent_data["function_call"]:
|
if "arguments" in intent_data["function_call"]:
|
||||||
function_args = intent_data["function_call"]["arguments"]
|
function_args = intent_data["function_call"]["arguments"]
|
||||||
|
|||||||
@@ -269,14 +269,12 @@ def register_device_type(descriptor):
|
|||||||
|
|
||||||
# 用于接受前端设备推送的搜索iot描述
|
# 用于接受前端设备推送的搜索iot描述
|
||||||
async def handleIotDescriptors(conn, descriptors):
|
async def handleIotDescriptors(conn, descriptors):
|
||||||
if not conn.use_function_call_mode:
|
|
||||||
return
|
|
||||||
wait_max_time = 5
|
wait_max_time = 5
|
||||||
while conn.func_handler is None or not conn.func_handler.finish_init:
|
while conn.func_handler is None or not conn.func_handler.finish_init:
|
||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
wait_max_time -= 1
|
wait_max_time -= 1
|
||||||
if wait_max_time <= 0:
|
if wait_max_time <= 0:
|
||||||
logger.bind(tag=TAG).error("连接对象没有func_handler")
|
logger.bind(tag=TAG).debug("连接对象没有func_handler")
|
||||||
return
|
return
|
||||||
"""处理物联网描述"""
|
"""处理物联网描述"""
|
||||||
functions_changed = False
|
functions_changed = False
|
||||||
|
|||||||
@@ -15,7 +15,15 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.base_url = config.get("base_url")
|
self.base_url = config.get("base_url")
|
||||||
else:
|
else:
|
||||||
self.base_url = config.get("url")
|
self.base_url = config.get("url")
|
||||||
self.max_tokens = config.get("max_tokens", 500)
|
max_tokens = config.get("max_tokens")
|
||||||
|
if max_tokens is None or max_tokens == "":
|
||||||
|
max_tokens = 500
|
||||||
|
|
||||||
|
try:
|
||||||
|
max_tokens = int(max_tokens)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
max_tokens = 500
|
||||||
|
self.max_tokens = max_tokens
|
||||||
|
|
||||||
check_model_key("LLM", self.api_key)
|
check_model_key("LLM", self.api_key)
|
||||||
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|||||||
@@ -42,7 +42,8 @@ class VADProvider(VADProviderBase):
|
|||||||
audio_tensor = torch.from_numpy(audio_float32)
|
audio_tensor = torch.from_numpy(audio_float32)
|
||||||
|
|
||||||
# 检测语音活动
|
# 检测语音活动
|
||||||
speech_prob = self.model(audio_tensor, 16000).item()
|
with torch.no_grad():
|
||||||
|
speech_prob = self.model(audio_tensor, 16000).item()
|
||||||
client_have_voice = speech_prob >= self.vad_threshold
|
client_have_voice = speech_prob >= self.vad_threshold
|
||||||
|
|
||||||
# 如果之前有声音,但本次没有声音,且与上次有声音的时间查已经超过了静默阈值,则认为已经说完一句话
|
# 如果之前有声音,但本次没有声音,且与上次有声音的时间查已经超过了静默阈值,则认为已经说完一句话
|
||||||
|
|||||||
@@ -1,34 +1,43 @@
|
|||||||
from plugins_func.register import register_function,ToolType, ActionResponse, Action
|
from plugins_func.register import register_function, ToolType, ActionResponse, Action
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
handle_exit_intent_function_desc = {
|
handle_exit_intent_function_desc = {
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {
|
"function": {
|
||||||
"name": "handle_exit_intent",
|
"name": "handle_exit_intent",
|
||||||
"description": "当用户想结束对话或需要退出系统时调用",
|
"description": "当用户想结束对话或需要退出系统时调用",
|
||||||
"parameters": {
|
"parameters": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"say_goodbye": {
|
"say_goodbye": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "和用户友好结束对话的告别语"
|
"description": "和用户友好结束对话的告别语",
|
||||||
}
|
|
||||||
},
|
|
||||||
"required": ["say_goodbye"]
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
},
|
||||||
|
"required": ["say_goodbye"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
@register_function('handle_exit_intent', handle_exit_intent_function_desc, ToolType.SYSTEM_CTL)
|
|
||||||
def handle_exit_intent(conn, say_goodbye: str):
|
@register_function(
|
||||||
|
"handle_exit_intent", handle_exit_intent_function_desc, ToolType.SYSTEM_CTL
|
||||||
|
)
|
||||||
|
def handle_exit_intent(conn, say_goodbye: str | None = None):
|
||||||
# 处理退出意图
|
# 处理退出意图
|
||||||
try:
|
try:
|
||||||
|
if say_goodbye is None:
|
||||||
|
say_goodbye = "再见,祝您生活愉快!"
|
||||||
conn.close_after_chat = True
|
conn.close_after_chat = True
|
||||||
logger.bind(tag=TAG).info(f"退出意图已处理:{say_goodbye}")
|
logger.bind(tag=TAG).info(f"退出意图已处理:{say_goodbye}")
|
||||||
return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response=say_goodbye)
|
return ActionResponse(
|
||||||
|
action=Action.RESPONSE, result="退出意图已处理", response=say_goodbye
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"处理退出意图错误: {e}")
|
logger.bind(tag=TAG).error(f"处理退出意图错误: {e}")
|
||||||
return ActionResponse(action=Action.NONE, result="退出意图处理失败", response="")
|
return ActionResponse(
|
||||||
|
action=Action.NONE, result="退出意图处理失败", response=""
|
||||||
|
)
|
||||||
|
|||||||
@@ -9,7 +9,9 @@ HASS_CACHE = {}
|
|||||||
|
|
||||||
def append_devices_to_prompt(conn):
|
def append_devices_to_prompt(conn):
|
||||||
if conn.use_function_call_mode:
|
if conn.use_function_call_mode:
|
||||||
funcs = conn.config["Intent"]["function_call"].get("functions", [])
|
funcs = conn.config["Intent"][conn.config["selected_module"]["Intent"]].get(
|
||||||
|
"functions", []
|
||||||
|
)
|
||||||
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
||||||
prompt = "下面是我家智能设备,可以通过homeassistant控制\n"
|
prompt = "下面是我家智能设备,可以通过homeassistant控制\n"
|
||||||
devices = conn.config["plugins"]["home_assistant"].get("devices", [])
|
devices = conn.config["plugins"]["home_assistant"].get("devices", [])
|
||||||
@@ -26,7 +28,9 @@ def initialize_hass_handler(conn):
|
|||||||
global HASS_CACHE
|
global HASS_CACHE
|
||||||
if HASS_CACHE == {}:
|
if HASS_CACHE == {}:
|
||||||
if conn.use_function_call_mode:
|
if conn.use_function_call_mode:
|
||||||
funcs = conn.config["Intent"]["function_call"].get("functions", [])
|
funcs = conn.config["Intent"][conn.config["selected_module"]["Intent"]].get(
|
||||||
|
"functions", []
|
||||||
|
)
|
||||||
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
||||||
HASS_CACHE["base_url"] = conn.config["plugins"]["home_assistant"].get(
|
HASS_CACHE["base_url"] = conn.config["plugins"]["home_assistant"].get(
|
||||||
"base_url"
|
"base_url"
|
||||||
|
|||||||
Reference in New Issue
Block a user