diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java index 00d48486..25b03034 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java @@ -3,6 +3,7 @@ package xiaozhi.modules.model.service.impl; import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.stream.Collectors; import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; @@ -19,6 +20,10 @@ import xiaozhi.common.redis.RedisKeys; import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.service.impl.BaseServiceImpl; import xiaozhi.common.utils.ConvertUtils; +import xiaozhi.modules.agent.dao.AgentDao; +import xiaozhi.modules.agent.dao.AgentTemplateDao; +import xiaozhi.modules.agent.entity.AgentEntity; +import xiaozhi.modules.agent.service.AgentTemplateService; import xiaozhi.modules.model.dao.ModelConfigDao; import xiaozhi.modules.model.dto.ModelBasicInfoDTO; import xiaozhi.modules.model.dto.ModelConfigBodyDTO; @@ -36,6 +41,9 @@ public class ModelConfigServiceImpl extends BaseServiceImpl getModelCodeList(String modelType, String modelName) { @@ -75,6 +83,7 @@ public class ModelConfigServiceImpl extends BaseServiceImpl agents = agentDao.selectList( + new QueryWrapper() + .eq("vad_model_id", modelId) + .or() + .eq("asr_model_id", modelId) + .or() + .eq("llm_model_id", modelId) + .or() + .eq("tts_model_id", modelId) + .or() + .eq("mem_model_id", modelId) + .or() + .eq("intent_model_id", modelId)); + if (!agents.isEmpty()) { + String agentNames = agents.stream() + .map(AgentEntity::getAgentName) + .collect(Collectors.joining("、")); + throw new RenException(String.format("该模型配置已被智能体[%s]引用,无法删除", agentNames)); + } + } + + /** + * 检查意图识别配置是否有引用 + * + * @param modelId 模型ID + */ + private void checkIntentConfigReference(String modelId) { + ModelConfigEntity modelConfig = modelConfigDao.selectById(modelId); + if (modelConfig != null + && "LLM".equals(modelConfig.getModelType() == null ? null : modelConfig.getModelType().toUpperCase())) { + List intentConfigs = modelConfigDao.selectList( + new QueryWrapper() + .eq("model_type", "Intent") + .like("config_json", "%" + modelId + "%")); + if (!intentConfigs.isEmpty()) { + throw new RenException("该LLM模型已被意图识别配置引用,无法删除"); + } + } + } + @Override public String getModelNameById(String id) { if (StringUtils.isBlank(id)) { diff --git a/main/xiaozhi-server/plugins_func/register.py b/main/xiaozhi-server/plugins_func/register.py index fd990760..b96f2f29 100644 --- a/main/xiaozhi-server/plugins_func/register.py +++ b/main/xiaozhi-server/plugins_func/register.py @@ -10,7 +10,10 @@ class ToolType(Enum): NONE = (1, "调用完工具后,不做其他操作") WAIT = (2, "调用工具,等待函数返回") CHANGE_SYS_PROMPT = (3, "修改系统提示词,切换角色性格或职责") - SYSTEM_CTL = (4, "系统控制,影响正常的对话流程,如退出、播放音乐等,需要传递conn参数") + SYSTEM_CTL = ( + 4, + "系统控制,影响正常的对话流程,如退出、播放音乐等,需要传递conn参数", + ) IOT_CTL = (5, "IOT设备控制,需要传递conn参数") MCP_CLIENT = (6, "MCP客户端") @@ -30,12 +33,14 @@ class Action(Enum): self.code = code self.message = message + class ActionResponse: def __init__(self, action: Action, result, response): self.action = action # 动作类型 self.result = result # 动作产生的结果 self.response = response # 直接回复的内容 + class FunctionItem: def __init__(self, name, description, func, type): self.name = name @@ -43,40 +48,49 @@ class FunctionItem: self.func = func self.type = type + class DeviceTypeRegistry: """设备类型注册表,用于管理IOT设备类型及其函数""" + def __init__(self): self.type_functions = {} # type_signature -> {func_name: FunctionItem} - + def generate_device_type_id(self, descriptor): """通过设备能力描述生成类型ID""" properties = sorted(descriptor["properties"].keys()) methods = sorted(descriptor["methods"].keys()) # 使用属性和方法的组合作为设备类型的唯一标识 - type_signature = f"{descriptor['name']}:{','.join(properties)}:{','.join(methods)}" + type_signature = ( + f"{descriptor['name']}:{','.join(properties)}:{','.join(methods)}" + ) return type_signature - + def get_device_functions(self, type_id): """获取设备类型对应的所有函数""" return self.type_functions.get(type_id, {}) - + def register_device_type(self, type_id, functions): """注册设备类型及其函数""" if type_id not in self.type_functions: self.type_functions[type_id] = functions + # 初始化函数注册字典 all_function_registry = {} device_type_registry = DeviceTypeRegistry() + def register_function(name, desc, type=None): """注册函数到函数注册字典的装饰器""" + def decorator(func): all_function_registry[name] = FunctionItem(name, desc, func, type) logger.bind(tag=TAG).debug(f"函数 '{name}' 已加载,可以注册使用") return func + return decorator + class FunctionRegistry: def __init__(self): self.function_registry = {} @@ -89,9 +103,9 @@ class FunctionRegistry: self.logger.bind(tag=TAG).error(f"函数 '{name}' 未找到") return None self.function_registry[name] = func - self.logger.bind(tag=TAG).info(f"函数 '{name}' 注册成功") + self.logger.bind(tag=TAG).debug(f"函数 '{name}' 注册成功") return func - + def unregister_function(self, name): # 注销函数,检测是否存在 if name not in self.function_registry: @@ -103,9 +117,9 @@ class FunctionRegistry: def get_function(self, name): return self.function_registry.get(name) - + def get_all_functions(self): return self.function_registry - + def get_all_function_desc(self): - return [func.description for _, func in self.function_registry.items()] \ No newline at end of file + return [func.description for _, func in self.function_registry.items()]