diff --git a/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java b/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java index d25c72b6..3a969713 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java +++ b/main/manager-api/src/main/java/xiaozhi/common/redis/RedisKeys.java @@ -152,4 +152,11 @@ public class RedisKeys { public static String getVoiceCloneAudioIdKey(String uuid) { return "voiceClone:audio:id:" + uuid; } + + /** + * 获取知识库缓存key + */ + public static String getKnowledgeBaseCacheKey(String datasetId) { + return "knowledge:base:" + datasetId; + } } diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentPluginMappingServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentPluginMappingServiceImpl.java index 2a53d474..450c4ce0 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentPluginMappingServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentPluginMappingServiceImpl.java @@ -1,9 +1,13 @@ package xiaozhi.modules.agent.service.impl; +import java.util.HashMap; import java.util.List; +import java.util.Map; +import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; +import com.alibaba.druid.support.json.JSONUtils; import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper; import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; @@ -11,6 +15,10 @@ import lombok.RequiredArgsConstructor; import xiaozhi.modules.agent.dao.AgentPluginMappingMapper; import xiaozhi.modules.agent.entity.AgentPluginMapping; import xiaozhi.modules.agent.service.AgentPluginMappingService; +import xiaozhi.modules.knowledge.entity.KnowledgeBaseEntity; +import xiaozhi.modules.knowledge.service.KnowledgeBaseService; +import xiaozhi.modules.model.entity.ModelConfigEntity; +import xiaozhi.modules.model.service.ModelConfigService; /** * @description 针对表【ai_agent_plugin_mapping(Agent与插件的唯一映射表)】的数据库操作Service实现 @@ -21,10 +29,39 @@ import xiaozhi.modules.agent.service.AgentPluginMappingService; public class AgentPluginMappingServiceImpl extends ServiceImpl implements AgentPluginMappingService { private final AgentPluginMappingMapper agentPluginMappingMapper; + private final KnowledgeBaseService knowledgeBaseService; + private final ModelConfigService modelConfigService; @Override public List agentPluginParamsByAgentId(String agentId) { - return agentPluginMappingMapper.selectPluginsByAgentId(agentId); + List list = agentPluginMappingMapper.selectPluginsByAgentId(agentId); + int index = 0; + for (int i = list.size() - 1; i >= 0; i--) { + AgentPluginMapping mapping = list.get(i); + if (StringUtils.isBlank(mapping.getProviderCode())) { + // 查询知识库插件参数 + KnowledgeBaseEntity knowledgeBaseEntity = knowledgeBaseService.selectById(mapping.getPluginId()); + if (knowledgeBaseEntity == null) { + list.remove(i); + continue; + } + ModelConfigEntity modelConfigEntity = modelConfigService + .getModelByIdFromCache(knowledgeBaseEntity.getRagModelId()); + if (modelConfigEntity == null) { + list.remove(i); + continue; + } + Map paramInfo = new HashMap<>(2); + paramInfo.put("name", knowledgeBaseEntity.getName()); + paramInfo.put("description", knowledgeBaseEntity.getDescription()); + mapping.setParamInfo(JSONUtils.toJSONString(paramInfo)); + String providerCode = "xzmcp_search_from_knowledgeBase_" + modelConfigEntity.getModelCode() + "_" + + index; + index++; + mapping.setProviderCode(providerCode); + } + } + return list; } @Override diff --git a/main/manager-api/src/main/java/xiaozhi/modules/knowledge/service/impl/KnowledgeBaseServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/knowledge/service/impl/KnowledgeBaseServiceImpl.java index 75edd60b..34d4cf47 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/knowledge/service/impl/KnowledgeBaseServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/knowledge/service/impl/KnowledgeBaseServiceImpl.java @@ -1,6 +1,7 @@ package xiaozhi.modules.knowledge.service.impl; import java.io.IOException; +import java.io.Serializable; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -28,6 +29,8 @@ import xiaozhi.common.constant.Constant; import xiaozhi.common.exception.ErrorCode; import xiaozhi.common.exception.RenException; import xiaozhi.common.page.PageData; +import xiaozhi.common.redis.RedisKeys; +import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.service.impl.BaseServiceImpl; import xiaozhi.common.utils.ConvertUtils; import xiaozhi.modules.knowledge.dao.KnowledgeBaseDao; @@ -47,8 +50,34 @@ public class KnowledgeBaseServiceImpl extends BaseServiceImpl getPageList(KnowledgeBaseDTO knowledgeBaseDTO, Integer page, Integer limit) { long curPage = page; @@ -169,6 +198,11 @@ public class KnowledgeBaseServiceImpl extends BaseServiceImpl getPluginList() { + // 1. 获取插件列表 LambdaQueryWrapper queryWrapper = new LambdaQueryWrapper<>(); queryWrapper.eq(ModelProviderEntity::getModelType, "Plugin"); List providerEntities = modelProviderDao.selectList(queryWrapper); - return ConvertUtils.sourceToTarget(providerEntities, ModelProviderDTO.class); + List resultList = ConvertUtils.sourceToTarget(providerEntities, ModelProviderDTO.class); + + // 2. 获取当前用户的知识库列表并追加到结果中 + UserDetail userDetail = SecurityUser.getUser(); + if (userDetail != null && userDetail.getId() != null) { + // 查询当前用户的知识库 + LambdaQueryWrapper kbQueryWrapper = new LambdaQueryWrapper<>(); + kbQueryWrapper.eq(KnowledgeBaseEntity::getCreator, userDetail.getId()); + kbQueryWrapper.eq(KnowledgeBaseEntity::getStatus, 1); // 只获取启用状态的知识库 + List knowledgeBases = knowledgeBaseDao.selectList(kbQueryWrapper); + + // 将知识库转换为ModelProviderDTO格式并添加到结果列表 + for (KnowledgeBaseEntity kb : knowledgeBases) { + ModelProviderDTO dto = new ModelProviderDTO(); + dto.setId(kb.getId()); + dto.setModelType("Rag"); + dto.setName("[知识库]" + kb.getName()); + dto.setProviderCode("ragflow"); // 假设所有RAG都使用ragflow + dto.setFields("[]"); + dto.setSort(0); + dto.setCreateDate(kb.getCreatedAt()); + dto.setUpdateDate(kb.getUpdatedAt()); + dto.setCreator(0L); + dto.setUpdater(0L); + resultList.add(dto); + } + } + + return resultList; } @Override