Merge commit '98011272ffb66c1a4e314e496bc301e0a975bd1e' into Knowledge-Base

This commit is contained in:
rainv123
2025-11-04 16:22:09 +08:00
36 changed files with 7799 additions and 32 deletions
@@ -150,6 +150,16 @@ public interface Constant {
* 火山引擎双声道语音克隆
*/
String VOICE_CLONE_HUOSHAN_DOUBLE_STREAM = "huoshan_double_stream";
/**
* RAG配置类型
*/
String RAG_CONFIG_TYPE = "RAG";
/**
* 默认RAG模型配置ID
*/
String DEFAULT_RAG_MODEL_ID = "RAG_DEFAULT_MODEL";
enum SysBaseParam {
/**
@@ -119,10 +119,6 @@ public interface ErrorCode {
int VOICEPRINT_UNREGISTER_PROCESS_ERROR = 10090; // 声纹注销处理失败
int VOICEPRINT_IDENTIFY_REQUEST_ERROR = 10091; // 声纹识别请求失败
// 设备相关错误码
int MAC_ADDRESS_ALREADY_EXISTS = 10161; // Mac地址已存在
// 模型相关错误码
int MODEL_PROVIDER_NOT_EXIST = 10162; // 供应器不存在
int LLM_NOT_EXIST = 10092; // 设置的LLM不存在
int MODEL_REFERENCED_BY_AGENT = 10093; // 该模型配置已被智能体引用,无法删除
int LLM_REFERENCED_BY_INTENT = 10094; // 该LLM模型已被意图识别配置引用,无法删除
@@ -198,4 +194,28 @@ public interface ErrorCode {
int VOICE_CLONE_PREFIX = 10158; // 复刻音色前缀
int VOICE_ID_ALREADY_EXISTS = 10159; // 音色ID已存在
int VOICE_CLONE_HUOSHAN_VOICE_ID_ERROR = 10160; // 火山引擎音色ID格式错误
// 设备相关错误码
int MAC_ADDRESS_ALREADY_EXISTS = 10161; // Mac地址已存在
// 模型相关错误码
int MODEL_PROVIDER_NOT_EXIST = 10162; // 供应器不存在
// 知识库数据集相关错误码
int Knowledge_Base_RECORD_NOT_EXISTS = 10163; // 知识库记录不存在
// RAG配置相关错误码
int RAG_CONFIG_NOT_FOUND = 10164; // RAG配置未找到
int RAG_CONFIG_TYPE_ERROR = 10165; // RAG配置类型错误
int RAG_DEFAULT_CONFIG_NOT_FOUND = 10166; // 默认RAG配置未找到
int RAG_CONFIG_MISSING_PARAMS = 10167; // RAG配置缺少必要参数
// RAG API调用相关错误码
int RAG_API_CREATE_FAILED = 10168; // RAG API创建数据集失败
int RAG_API_UPDATE_FAILED = 10169; // RAG API更新数据集失败
int RAG_API_DELETE_FAILED = 10170; // RAG API删除数据集失败
int UPLOAD_FILE_ERROR = 10171; // 上传文件失败
int RAG_API_QUERY_FAILED = 10172; // RAG API查询失败
int RAG_API_PARSE_FAILED = 10173; // RAG API解析失败
int RAG_API_OPERATION_FAILED = 10174; // RAG API操作失败
}
@@ -32,13 +32,23 @@ public class FieldMetaObjectHandler implements MetaObjectHandler {
// 创建者
strictInsertFill(metaObject, CREATOR, Long.class, user.getId());
// 创建时间
strictInsertFill(metaObject, CREATE_DATE, Date.class, date);
// 创建时间 - 支持createDate和createdAt两种字段名
if (metaObject.hasSetter(CREATE_DATE)) {
strictInsertFill(metaObject, CREATE_DATE, Date.class, date);
}
if (metaObject.hasSetter("createdAt")) {
strictInsertFill(metaObject, "createdAt", Date.class, date);
}
// 更新者
strictInsertFill(metaObject, UPDATER, Long.class, user.getId());
// 更新时间
strictInsertFill(metaObject, UPDATE_DATE, Date.class, date);
// 更新时间 - 支持updateDate和updatedAt两种字段名
if (metaObject.hasSetter(UPDATE_DATE)) {
strictInsertFill(metaObject, UPDATE_DATE, Date.class, date);
}
if (metaObject.hasSetter("updatedAt")) {
strictInsertFill(metaObject, "updatedAt", Date.class, date);
}
// 数据标识
strictInsertFill(metaObject, DATA_OPERATION, String.class, Constant.DataOperation.INSERT.getValue());
@@ -46,10 +56,17 @@ public class FieldMetaObjectHandler implements MetaObjectHandler {
@Override
public void updateFill(MetaObject metaObject) {
Date date = new Date();
// 更新者
strictUpdateFill(metaObject, UPDATER, Long.class, SecurityUser.getUserId());
// 更新时间
strictUpdateFill(metaObject, UPDATE_DATE, Date.class, new Date());
// 更新时间 - 支持updateDate和updatedAt两种字段名
if (metaObject.hasSetter(UPDATE_DATE)) {
strictUpdateFill(metaObject, UPDATE_DATE, Date.class, date);
}
if (metaObject.hasSetter("updatedAt")) {
strictUpdateFill(metaObject, "updatedAt", Date.class, date);
}
// 数据标识
strictInsertFill(metaObject, DATA_OPERATION, String.class, Constant.DataOperation.UPDATE.getValue());
@@ -45,6 +45,10 @@ public class AgentTemplateServiceImpl extends ServiceImpl<AgentTemplateDao, Agen
@Override
public void updateDefaultTemplateModelId(String modelType, String modelId) {
modelType = modelType.toUpperCase();
// 如果是rag模型,不需要更新
if (modelType.equals("RAG")) {
return;
}
UpdateWrapper<AgentTemplateEntity> wrapper = new UpdateWrapper<>();
switch (modelType) {
@@ -91,6 +91,7 @@ public class ConfigServiceImpl implements ConfigService {
null,
null,
null,
null,
result,
isCache);
@@ -195,6 +196,7 @@ public class ConfigServiceImpl implements ConfigService {
agent.getTtsModelId(),
agent.getMemModelId(),
agent.getIntentModelId(),
null,
result,
true);
@@ -371,12 +373,13 @@ public class ConfigServiceImpl implements ConfigService {
String ttsModelId,
String memModelId,
String intentModelId,
String ragModelId,
Map<String, Object> result,
boolean isCache) {
Map<String, String> selectedModule = new HashMap<>();
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM" };
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId, vllmModelId };
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM","RAG" };
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId, vllmModelId, ragModelId };
String intentLLMModelId = null;
String memLocalShortLLMModelId = null;
@@ -0,0 +1,114 @@
package xiaozhi.modules.knowledge.controller;
import org.apache.shiro.authz.annotation.RequiresPermissions;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.DeleteMapping;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.PathVariable;
import org.springframework.web.bind.annotation.PostMapping;
import org.springframework.web.bind.annotation.PutMapping;
import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.apache.commons.lang3.StringUtils;
import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.exception.RenException;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.Parameter;
import io.swagger.v3.oas.annotations.tags.Tag;
import lombok.AllArgsConstructor;
import xiaozhi.common.page.PageData;
import xiaozhi.common.utils.Result;
import xiaozhi.modules.knowledge.dto.KnowledgeBaseDTO;
import xiaozhi.modules.knowledge.service.KnowledgeBaseService;
import java.util.Map;
@AllArgsConstructor
@RestController
@RequestMapping("/api/v1")
@Tag(name = "知识库管理")
public class KnowledgeBaseController {
private final KnowledgeBaseService knowledgeBaseService;
@GetMapping("/datasets")
@Operation(summary = "分页查询知识库列表")
@RequiresPermissions("sys:role:normal")
public Result<PageData<KnowledgeBaseDTO>> getPageList(
@RequestParam(required = false) String name,
@RequestParam(required = false) String id,
@RequestParam(required = false, defaultValue = "1") Integer page,
@RequestParam(required = false, defaultValue = "10") Integer page_size,
@RequestParam(required = false) String orderby,
@RequestParam(required = false) Boolean desc) {
KnowledgeBaseDTO knowledgeBaseDTO = new KnowledgeBaseDTO();
knowledgeBaseDTO.setName(name);
knowledgeBaseDTO.setDatasetId(id);
PageData<KnowledgeBaseDTO> pageData = knowledgeBaseService.getPageList(knowledgeBaseDTO, String.valueOf(page), String.valueOf(page_size));
return new Result<PageData<KnowledgeBaseDTO>>().ok(pageData);
}
@GetMapping("/datasets/{dataset_id}")
@Operation(summary = "根据知识库ID获取知识库详情")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeBaseDTO> getByDatasetId(@PathVariable("dataset_id") String datasetId) {
KnowledgeBaseDTO knowledgeBaseDTO = knowledgeBaseService.getByDatasetId(datasetId);
return new Result<KnowledgeBaseDTO>().ok(knowledgeBaseDTO);
}
@PostMapping("/datasets")
@Operation(summary = "创建知识库")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeBaseDTO> save(@RequestBody @Validated KnowledgeBaseDTO knowledgeBaseDTO) {
KnowledgeBaseDTO resp = knowledgeBaseService.save(knowledgeBaseDTO);
return new Result<KnowledgeBaseDTO>().ok(resp);
}
@PutMapping("/datasets/{dataset_id}")
@Operation(summary = "更新知识库")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeBaseDTO> update(@PathVariable("dataset_id") String datasetId,
@RequestBody @Validated KnowledgeBaseDTO knowledgeBaseDTO) {
knowledgeBaseDTO.setDatasetId(datasetId);
KnowledgeBaseDTO resp = knowledgeBaseService.update(knowledgeBaseDTO);
return new Result<KnowledgeBaseDTO>().ok(resp);
}
@DeleteMapping("/datasets/{dataset_id}")
@Operation(summary = "删除单个知识库")
@Parameter(name = "dataset_id", description = "知识库ID", required = true)
@RequiresPermissions("sys:role:normal")
public Result<Void> delete(@PathVariable("dataset_id") String datasetId) {
knowledgeBaseService.deleteByDatasetId(datasetId);
return new Result<>();
}
@DeleteMapping("/datasets/batch")
@Operation(summary = "批量删除知识库")
@Parameter(name = "ids", description = "知识库ID列表,用逗号分隔", required = true)
@RequiresPermissions("sys:role:normal")
public Result<Void> deleteBatch(@RequestParam("ids") String ids) {
if (StringUtils.isBlank(ids)) {
throw new RenException(ErrorCode.PARAMS_GET_ERROR);
}
String[] idArray = ids.split(",");
for (String datasetId : idArray) {
if (StringUtils.isNotBlank(datasetId)) {
knowledgeBaseService.deleteByDatasetId(datasetId.trim());
}
}
return new Result<>();
}
@GetMapping("/rag-config/default")
@Operation(summary = "获取默认RAG配置")
@RequiresPermissions("sys:role:normal")
public Result<Map<String, Object>> getDefaultRAGConfig() {
Map<String, Object> config = knowledgeBaseService.getDefaultRAGConfig();
return new Result<Map<String, Object>>().ok(config);
}
}
@@ -0,0 +1,222 @@
package xiaozhi.modules.knowledge.controller;
import java.util.List;
import java.util.Map;
import org.apache.shiro.authz.annotation.RequiresPermissions;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import org.springframework.web.multipart.MultipartFile;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.Parameter;
import io.swagger.v3.oas.annotations.tags.Tag;
import lombok.AllArgsConstructor;
import xiaozhi.common.page.PageData;
import xiaozhi.common.utils.Result;
import xiaozhi.modules.knowledge.dto.KnowledgeFilesDTO;
import xiaozhi.modules.knowledge.service.KnowledgeFilesService;
@AllArgsConstructor
@RestController
@RequestMapping("/api/v1/datasets/{dataset_id}")
@Tag(name = "知识库文档管理")
public class KnowledgeFilesController {
private final KnowledgeFilesService knowledgeFilesService;
@GetMapping("/documents")
@Operation(summary = "分页查询文档列表")
@RequiresPermissions("sys:role:normal")
public Result<PageData<KnowledgeFilesDTO>> getPageList(
@PathVariable("dataset_id") String datasetId,
@RequestParam(required = false) String name,
@RequestParam(required = false, defaultValue = "1") Integer page,
@RequestParam(required = false, defaultValue = "10") Integer page_size) {
KnowledgeFilesDTO knowledgeFilesDTO = new KnowledgeFilesDTO();
knowledgeFilesDTO.setDatasetId(datasetId);
knowledgeFilesDTO.setName(name);
PageData<KnowledgeFilesDTO> pageData = knowledgeFilesService.getPageList(knowledgeFilesDTO, page, page_size);
return new Result<PageData<KnowledgeFilesDTO>>().ok(pageData);
}
@GetMapping("/documents/{document_id}")
@Operation(summary = "根据文档ID获取文档详情")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeFilesDTO> getByDocumentId(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId) {
KnowledgeFilesDTO knowledgeFilesDTO = knowledgeFilesService.getByDocumentId(documentId);
return new Result<KnowledgeFilesDTO>().ok(knowledgeFilesDTO);
}
@PostMapping("/documents")
@Operation(summary = "上传文档到知识库")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeFilesDTO> uploadDocument(
@PathVariable("dataset_id") String datasetId,
@RequestParam("file") MultipartFile file,
@RequestParam(required = false) String name,
@RequestParam(required = false) String chunkMethod,
@RequestParam(required = false) String metaFields,
@RequestParam(required = false) String parserConfig) {
KnowledgeFilesDTO resp = knowledgeFilesService.uploadDocument(datasetId, file, name,
metaFields != null ? parseJsonMap(metaFields) : null,
chunkMethod,
parserConfig != null ? parseJsonMap(parserConfig) : null);
return new Result<KnowledgeFilesDTO>().ok(resp);
}
@PutMapping("/documents/{document_id}")
@Operation(summary = "更新文档配置")
@RequiresPermissions("sys:role:normal")
public Result<KnowledgeFilesDTO> update(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId,
@RequestBody @Validated KnowledgeFilesDTO knowledgeFilesDTO) {
knowledgeFilesDTO.setDatasetId(datasetId);
knowledgeFilesDTO.setDocumentId(documentId);
KnowledgeFilesDTO resp = knowledgeFilesService.update(knowledgeFilesDTO);
return new Result<KnowledgeFilesDTO>().ok(resp);
}
@DeleteMapping("/documents/{document_id}")
@Operation(summary = "删除单个文档")
@Parameter(name = "document_id", description = "文档ID", required = true)
@RequiresPermissions("sys:role:normal")
public Result<Void> delete(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId) {
knowledgeFilesService.deleteByDocumentId(documentId, datasetId);
return new Result<>();
}
@DeleteMapping("/documents")
@Operation(summary = "批量删除文档")
@RequiresPermissions("sys:role:normal")
public Result<Void> deleteBatch(@PathVariable("dataset_id") String datasetId,
@RequestBody Map<String, List<String>> requestBody) {
List<String> ids = requestBody.get("ids");
if (ids != null && !ids.isEmpty()) {
knowledgeFilesService.deleteBatch(ids);
}
return new Result<>();
}
@PostMapping("/chunks")
@Operation(summary = "批量解析文档(切块)")
@RequiresPermissions("sys:role:normal")
public Result<Void> parseDocuments(@PathVariable("dataset_id") String datasetId,
@RequestBody Map<String, List<String>> requestBody) {
List<String> documentIds = requestBody.get("document_ids");
if (documentIds == null || documentIds.isEmpty()) {
return new Result<Void>().error("document_ids参数不能为空");
}
boolean success = knowledgeFilesService.parseDocuments(datasetId, documentIds);
if (success) {
return new Result<Void>();
} else {
return new Result<Void>().error("文档解析失败,文档可能正在处理中");
}
}
@PostMapping("/documents/{document_id}/parse")
@Operation(summary = "解析单个文档(切块)")
@RequiresPermissions("sys:role:normal")
public Result<Void> parseDocument(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId) {
List<String> documentIds = java.util.Arrays.asList(documentId);
boolean success = knowledgeFilesService.parseDocuments(datasetId, documentIds);
if (success) {
return new Result<Void>();
} else {
return new Result<Void>().error("文档解析失败,文档可能正在处理中");
}
}
@PostMapping("/documents/{document_id}/chunks")
@Operation(summary = "添加切片到指定文档")
@RequiresPermissions("sys:role:normal")
public Result<Map<String, Object>> addChunk(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId,
@RequestBody Map<String, Object> requestBody) {
String content = (String) requestBody.get("content");
List<String> importantKeywords = (List<String>) requestBody.get("important_keywords");
List<String> questions = (List<String>) requestBody.get("questions");
Map<String, Object> result = knowledgeFilesService.addChunk(datasetId, documentId, content, importantKeywords, questions);
return new Result<Map<String, Object>>().ok(result);
}
@GetMapping("/documents/{document_id}/chunks")
@Operation(summary = "列出指定文档的切片")
@RequiresPermissions("sys:role:normal")
public Result<Map<String, Object>> listChunks(@PathVariable("dataset_id") String datasetId,
@PathVariable("document_id") String documentId,
@RequestParam(required = false) String keywords,
@RequestParam(required = false, defaultValue = "1") Integer page,
@RequestParam(required = false, defaultValue = "1024") Integer page_size,
@RequestParam(required = false) String id) {
Map<String, Object> result = knowledgeFilesService.listChunks(datasetId, documentId, keywords, page, page_size, id);
return new Result<Map<String, Object>>().ok(result);
}
/**
* 召回测试
*/
@PostMapping("/retrieval-test")
@Operation(summary = "召回测试")
@RequiresPermissions("sys:role:normal")
public Result<Map<String, Object>> retrievalTest(@PathVariable("dataset_id") String datasetId,
@RequestBody Map<String, Object> params) {
try {
// 提取参数
String question = (String) params.get("question");
if (question == null || question.trim().isEmpty()) {
return new Result<Map<String, Object>>().error("问题不能为空");
}
List<String> datasetIds = (List<String>) params.get("dataset_ids");
List<String> documentIds = (List<String>) params.get("document_ids");
Integer page = (Integer) params.get("page");
Integer pageSize = (Integer) params.get("page_size");
Float similarityThreshold = (Float) params.get("similarity_threshold");
Float vectorSimilarityWeight = (Float) params.get("vector_similarity_weight");
Integer topK = (Integer) params.get("top_k");
String rerankId = (String) params.get("rerank_id");
Boolean keyword = (Boolean) params.get("keyword");
Boolean highlight = (Boolean) params.get("highlight");
List<String> crossLanguages = (List<String>) params.get("cross_languages");
Map<String, Object> metadataCondition = (Map<String, Object>) params.get("metadata_condition");
// 如果未指定数据集ID,使用当前数据集
if (datasetIds == null || datasetIds.isEmpty()) {
datasetIds = java.util.Arrays.asList(datasetId);
}
Map<String, Object> result = knowledgeFilesService.retrievalTest(
question, datasetIds, documentIds, page, pageSize, similarityThreshold,
vectorSimilarityWeight, topK, rerankId, keyword, highlight, crossLanguages, metadataCondition
);
return new Result<Map<String, Object>>().ok(result);
} catch (Exception e) {
return new Result<Map<String, Object>>().error("召回测试失败: " + e.getMessage());
}
}
/**
* 解析JSON字符串为Map对象
*/
private Map<String, Object> parseJsonMap(String jsonString) {
try {
ObjectMapper objectMapper = new ObjectMapper();
return objectMapper.readValue(jsonString, new TypeReference<Map<String, Object>>() {});
} catch (Exception e) {
throw new RuntimeException("解析JSON字符串失败: " + jsonString, e);
}
}
}
@@ -0,0 +1,14 @@
package xiaozhi.modules.knowledge.dao;
import org.apache.ibatis.annotations.Mapper;
import xiaozhi.common.dao.BaseDao;
import xiaozhi.modules.knowledge.entity.KnowledgeBaseEntity;
/**
* 知识库知识库
*/
@Mapper
public interface KnowledgeBaseDao extends BaseDao<KnowledgeBaseEntity> {
}
@@ -0,0 +1,49 @@
package xiaozhi.modules.knowledge.dto;
import java.io.Serial;
import java.io.Serializable;
import java.util.Date;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
@Data
@Schema(description = "知识库知识库")
public class KnowledgeBaseDTO implements Serializable {
@Serial
private static final long serialVersionUID = 1L;
@Schema(description = "唯一标识")
private String id;
@Schema(description = "知识库ID")
private String datasetId;
@Schema(description = "RAG模型配置ID")
private String ragModelId;
@Schema(description = "知识库名称")
private String name;
@Schema(description = "知识库描述")
private String description;
@Schema(description = "状态(0:禁用 1:启用)")
private Integer status;
@Schema(description = "创建者")
private Long creator;
@Schema(description = "创建时间")
private Date createdAt;
@Schema(description = "更新者")
private Long updater;
@Schema(description = "更新时间")
private Date updatedAt;
@Schema(description = "文档数量")
private Integer documentCount;
}
@@ -0,0 +1,61 @@
package xiaozhi.modules.knowledge.dto;
import java.io.Serial;
import java.io.Serializable;
import java.util.Date;
import java.util.Map;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
@Data
@Schema(description = "知识库文档")
public class KnowledgeFilesDTO implements Serializable {
@Serial
private static final long serialVersionUID = 1L;
@Schema(description = "唯一标识")
private String id;
@Schema(description = "文档ID")
private String documentId;
@Schema(description = "知识库ID")
private String datasetId;
@Schema(description = "文档名称")
private String name;
@Schema(description = "文档类型")
private String fileType;
@Schema(description = "文件大小(字节)")
private Long fileSize;
@Schema(description = "文件路径")
private String filePath;
@Schema(description = "元数据字段")
private Map<String, Object> metaFields;
@Schema(description = "分块方法")
private String chunkMethod;
@Schema(description = "解析器配置")
private Map<String, Object> parserConfig;
@Schema(description = "状态(0:待解析 1:解析中 2:解析成功 3:解析失败)")
private Integer status;
@Schema(description = "创建者")
private Long creator;
@Schema(description = "创建时间")
private Date createdAt;
@Schema(description = "更新者")
private Long updater;
@Schema(description = "更新时间")
private Date updatedAt;
}
@@ -0,0 +1,53 @@
package xiaozhi.modules.knowledge.entity;
import java.util.Date;
import com.baomidou.mybatisplus.annotation.FieldFill;
import com.baomidou.mybatisplus.annotation.IdType;
import com.baomidou.mybatisplus.annotation.TableField;
import com.baomidou.mybatisplus.annotation.TableId;
import com.baomidou.mybatisplus.annotation.TableName;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
@Data
@TableName(value = "ai_rag_dataset", autoResultMap = true)
@Schema(description = "知识库知识库表")
public class KnowledgeBaseEntity {
@TableId(type = IdType.ASSIGN_UUID)
@Schema(description = "唯一标识")
private String id;
@Schema(description = "知识库ID")
private String datasetId;
@Schema(description = "RAG模型配置ID")
private String ragModelId;
@Schema(description = "知识库名称")
private String name;
@Schema(description = "知识库描述")
private String description;
@Schema(description = "状态(0:禁用 1:启用)")
private Integer status;
@Schema(description = "创建者")
@TableField(fill = FieldFill.INSERT)
private Long creator;
@Schema(description = "创建时间")
@TableField(fill = FieldFill.INSERT)
private Date createdAt;
@Schema(description = "更新者")
@TableField(fill = FieldFill.UPDATE)
private Long updater;
@Schema(description = "更新时间")
@TableField(fill = FieldFill.UPDATE)
private Date updatedAt;
}
@@ -0,0 +1,85 @@
package xiaozhi.modules.knowledge.service;
import java.util.Map;
import xiaozhi.common.page.PageData;
import xiaozhi.common.service.BaseService;
import xiaozhi.modules.knowledge.dto.KnowledgeBaseDTO;
import xiaozhi.modules.knowledge.entity.KnowledgeBaseEntity;
/**
* 知识库知识库服务接口
*/
public interface KnowledgeBaseService extends BaseService<KnowledgeBaseEntity> {
/**
* 分页查询知识库列表
*
* @param knowledgeBaseDTO 查询条件
* @param page 页码
* @param limit 每页数量
* @return 分页数据
*/
PageData<KnowledgeBaseDTO> getPageList(KnowledgeBaseDTO knowledgeBaseDTO, String page, String limit);
/**
* 根据ID获取知识库详情
*
* @param id 知识库ID
* @return 知识库详情
*/
KnowledgeBaseDTO getById(String id);
/**
* 新增知识库
*
* @param knowledgeBaseDTO 知识库信息
* @return 新增的知识库
*/
KnowledgeBaseDTO save(KnowledgeBaseDTO knowledgeBaseDTO);
/**
* 更新知识库
*
* @param knowledgeBaseDTO 知识库信息
* @return 更新的知识库
*/
KnowledgeBaseDTO update(KnowledgeBaseDTO knowledgeBaseDTO);
/**
* 根据ID删除知识库
*
* @param id 知识库ID
*/
void delete(String id);
/**
* 根据知识库ID查询知识库
*
* @param datasetId 知识库ID
* @return 知识库详情
*/
KnowledgeBaseDTO getByDatasetId(String datasetId);
/**
* 根据知识库ID删除知识库
*
* @param datasetId 知识库ID
*/
void deleteByDatasetId(String datasetId);
/**
* 获取RAG配置信息
*
* @param ragModelId RAG模型配置ID
* @return RAG配置信息
*/
Map<String, Object> getRAGConfig(String ragModelId);
/**
* 获取默认RAG配置信息
*
* @return 默认RAG配置信息
*/
Map<String, Object> getDefaultRAGConfig();
}
@@ -0,0 +1,178 @@
package xiaozhi.modules.knowledge.service;
import java.util.List;
import java.util.Map;
import org.springframework.web.multipart.MultipartFile;
import xiaozhi.common.page.PageData;
import xiaozhi.modules.knowledge.dto.KnowledgeFilesDTO;
/**
* 知识库文档服务接口
*/
public interface KnowledgeFilesService {
/**
* 分页查询文档列表
*
* @param knowledgeFilesDTO 查询条件
* @param page 页码
* @param limit 每页数量
* @return 分页数据
*/
PageData<KnowledgeFilesDTO> getPageList(KnowledgeFilesDTO knowledgeFilesDTO, Integer page, Integer limit);
/**
* 根据ID获取文档详情
*
* @param id 文档ID
* @return 文档详情
*/
KnowledgeFilesDTO getById(String id);
/**
* 根据文档ID获取文档详情
*
* @param documentId 文档ID
* @return 文档详情
*/
KnowledgeFilesDTO getByDocumentId(String documentId);
/**
* 根据文档ID和知识库ID获取文档详情
*
* @param documentId 文档ID
* @param datasetId 知识库ID
* @return 文档详情
*/
KnowledgeFilesDTO getByDocumentId(String documentId, String datasetId);
/**
* 上传文档到知识库
*
* @param datasetId 知识库ID
* @param file 上传的文件
* @param name 文档名称
* @param metaFields 元数据字段
* @param chunkMethod 分块方法
* @param parserConfig 解析器配置
* @return 上传的文档信息
*/
KnowledgeFilesDTO uploadDocument(String datasetId, MultipartFile file, String name,
Map<String, Object> metaFields, String chunkMethod,
Map<String, Object> parserConfig);
/**
* 更新文档配置
*
* @param knowledgeFilesDTO 文档信息
* @return 更新的文档信息
*/
KnowledgeFilesDTO update(KnowledgeFilesDTO knowledgeFilesDTO);
/**
* 根据ID删除文档
*
* @param id 文档ID
*/
void delete(String id);
/**
* 根据文档ID删除文档
*
* @param documentId 文档ID
*/
void deleteByDocumentId(String documentId);
/**
* 根据文档ID和知识库ID删除文档
*
* @param documentId 文档ID
* @param datasetId 知识库ID
*/
void deleteByDocumentId(String documentId, String datasetId);
/**
* 批量删除文档
*
* @param ids 文档ID列表
*/
void deleteBatch(List<String> ids);
/**
* 获取RAG配置信息
*
* @param ragModelId RAG模型配置ID
* @return RAG配置信息
*/
Map<String, Object> getRAGConfig(String ragModelId);
/**
* 获取默认RAG配置信息
*
* @return 默认RAG配置信息
*/
Map<String, Object> getDefaultRAGConfig();
/**
* 解析文档(切块)
*
* @param datasetId 知识库ID
* @param documentIds 文档ID列表
* @return 解析结果
*/
boolean parseDocuments(String datasetId, List<String> documentIds);
/**
* 添加切片到指定文档
*
* @param datasetId 知识库ID
* @param documentId 文档ID
* @param content 切片内容
* @param importantKeywords 重要关键词列表
* @param questions 问题列表
* @return 添加的切片信息
*/
Map<String, Object> addChunk(String datasetId, String documentId, String content,
List<String> importantKeywords, List<String> questions);
/**
* 列出指定文档的切片
*
* @param datasetId 知识库ID
* @param documentId 文档ID
* @param keywords 关键词过滤
* @param page 页码
* @param pageSize 每页数量
* @param chunkId 切片ID
* @return 切片列表信息
*/
Map<String, Object> listChunks(String datasetId, String documentId, String keywords,
Integer page, Integer pageSize, String chunkId);
/**
* 召回测试 - 从指定数据集或文档中检索相关切片
*
* @param question 用户查询或查询关键词
* @param datasetIds 数据集ID列表
* @param documentIds 文档ID列表
* @param page 页码
* @param pageSize 每页数量
* @param similarityThreshold 最小相似度阈值
* @param vectorSimilarityWeight 向量相似度权重
* @param topK 参与向量余弦计算的切片数量
* @param rerankId 重排模型ID
* @param keyword 是否启用关键词匹配
* @param highlight 是否启用高亮显示
* @param crossLanguages 跨语言翻译列表
* @param metadataCondition 元数据过滤条件
* @return 召回测试结果
*/
Map<String, Object> retrievalTest(String question, List<String> datasetIds, List<String> documentIds,
Integer page, Integer pageSize, Float similarityThreshold,
Float vectorSimilarityWeight, Integer topK, String rerankId,
Boolean keyword, Boolean highlight, List<String> crossLanguages,
Map<String, Object> metadataCondition);
}
@@ -0,0 +1,771 @@
package xiaozhi.modules.knowledge.service.impl;
import java.io.IOException;
import java.util.List;
import java.util.Map;
import java.util.HashMap;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.apache.commons.lang3.StringUtils;
import org.springframework.http.*;
import org.springframework.stereotype.Service;
import org.springframework.web.client.RestTemplate;
import org.springframework.web.client.HttpClientErrorException;
import org.springframework.web.client.HttpServerErrorException;
import org.springframework.web.client.ResourceAccessException;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.metadata.IPage;
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
import lombok.AllArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.constant.Constant;
import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.exception.RenException;
import xiaozhi.common.page.PageData;
import xiaozhi.common.service.impl.BaseServiceImpl;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.modules.knowledge.dao.KnowledgeBaseDao;
import xiaozhi.modules.knowledge.dto.KnowledgeBaseDTO;
import xiaozhi.modules.knowledge.entity.KnowledgeBaseEntity;
import xiaozhi.modules.knowledge.service.KnowledgeBaseService;
import xiaozhi.modules.model.dao.ModelConfigDao;
import xiaozhi.modules.model.entity.ModelConfigEntity;
import xiaozhi.modules.model.service.ModelConfigService;
@Service
@AllArgsConstructor
@Slf4j
public class KnowledgeBaseServiceImpl extends BaseServiceImpl<KnowledgeBaseDao, KnowledgeBaseEntity> implements KnowledgeBaseService {
private final KnowledgeBaseDao knowledgeBaseDao;
private final ModelConfigService modelConfigService;
private final ModelConfigDao modelConfigDao;
private RestTemplate restTemplate = new RestTemplate();
@Override
public PageData<KnowledgeBaseDTO> getPageList(KnowledgeBaseDTO knowledgeBaseDTO, String page, String limit) {
long curPage = Long.parseLong(page);
long pageSize = Long.parseLong(limit);
Page<KnowledgeBaseEntity> pageInfo = new Page<>(curPage, pageSize);
QueryWrapper<KnowledgeBaseEntity> queryWrapper = new QueryWrapper<>();
// 添加查询条件
if (knowledgeBaseDTO != null) {
queryWrapper.like(StringUtils.isNotBlank(knowledgeBaseDTO.getName()), "name", knowledgeBaseDTO.getName())
.eq(knowledgeBaseDTO.getStatus() != null, "status", knowledgeBaseDTO.getStatus());
}
// 添加排序规则:按创建时间降序
queryWrapper.orderByDesc("created_at");
IPage<KnowledgeBaseEntity> knowledgeBaseEntityIPage = knowledgeBaseDao.selectPage(pageInfo, queryWrapper);
// 同步RAGFlow API的数据集状态(可选功能,可根据需要开启)
syncRAGFlowDatasetStatus(knowledgeBaseEntityIPage.getRecords());
// 获取分页数据
PageData<KnowledgeBaseDTO> pageData = getPageData(knowledgeBaseEntityIPage, KnowledgeBaseDTO.class);
// 为每个知识库获取文档数量
if (pageData != null && pageData.getList() != null) {
for (KnowledgeBaseDTO knowledgeBase : pageData.getList()) {
try {
Integer documentCount = getDocumentCountFromRAGFlow(knowledgeBase.getDatasetId(), knowledgeBase.getRagModelId());
knowledgeBase.setDocumentCount(documentCount);
} catch (Exception e) {
log.warn("获取知识库 {} 的文档数量失败: {}", knowledgeBase.getDatasetId(), e.getMessage());
knowledgeBase.setDocumentCount(0); // 设置默认值
}
}
}
return pageData;
}
@Override
public KnowledgeBaseDTO getById(String id) {
if (StringUtils.isBlank(id)) {
throw new RenException(ErrorCode.IDENTIFIER_NOT_NULL);
}
KnowledgeBaseEntity entity = knowledgeBaseDao.selectById(id);
if (entity == null) {
throw new RenException(ErrorCode.Knowledge_Base_RECORD_NOT_EXISTS);
}
return ConvertUtils.sourceToTarget(entity, KnowledgeBaseDTO.class);
}
@Override
public KnowledgeBaseDTO save(KnowledgeBaseDTO knowledgeBaseDTO) {
if (knowledgeBaseDTO == null) {
throw new RenException(ErrorCode.PARAMS_GET_ERROR);
}
String datasetId = null;
// 调用RAGFlow API创建数据集
try {
Map<String, Object> ragConfig = getValidatedRAGConfig(knowledgeBaseDTO.getRagModelId());
datasetId = createDatasetInRAGFlow(
knowledgeBaseDTO.getName(),
knowledgeBaseDTO.getDescription(),
ragConfig
);
} catch (Exception e) {
// 如果RAG API调用失败,直接抛出异常,无需回滚(因为还没有插入本地数据库)
throw e;
}
// 验证数据集ID是否已存在
KnowledgeBaseEntity existingEntity = knowledgeBaseDao.selectOne(
new QueryWrapper<KnowledgeBaseEntity>().eq("dataset_id", datasetId)
);
if (existingEntity != null) {
// 如果datasetId已存在,删除RAGFlow中的数据集并抛出异常
try {
Map<String, Object> ragConfig = getValidatedRAGConfig(knowledgeBaseDTO.getRagModelId());
deleteDatasetInRAGFlow(datasetId, ragConfig);
} catch (Exception deleteException) {
log.warn("删除重复datasetId的RAGFlow数据集失败: {}", deleteException.getMessage());
}
throw new RenException(ErrorCode.DB_RECORD_EXISTS);
}
// 创建本地实体并保存
KnowledgeBaseEntity entity = ConvertUtils.sourceToTarget(knowledgeBaseDTO, KnowledgeBaseEntity.class);
entity.setDatasetId(datasetId);
knowledgeBaseDao.insert(entity);
return ConvertUtils.sourceToTarget(entity, KnowledgeBaseDTO.class);
}
@Override
public KnowledgeBaseDTO update(KnowledgeBaseDTO knowledgeBaseDTO) {
if (knowledgeBaseDTO == null || StringUtils.isBlank(knowledgeBaseDTO.getId())) {
throw new RenException(ErrorCode.IDENTIFIER_NOT_NULL);
}
// 检查记录是否存在
KnowledgeBaseEntity existingEntity = knowledgeBaseDao.selectById(knowledgeBaseDTO.getId());
if (existingEntity == null) {
throw new RenException(ErrorCode.Knowledge_Base_RECORD_NOT_EXISTS);
}
// 验证数据集ID是否与其他记录冲突
if (StringUtils.isNotBlank(knowledgeBaseDTO.getDatasetId())) {
KnowledgeBaseEntity conflictEntity = knowledgeBaseDao.selectOne(
new QueryWrapper<KnowledgeBaseEntity>()
.eq("dataset_id", knowledgeBaseDTO.getDatasetId())
.ne("id", knowledgeBaseDTO.getId())
);
if (conflictEntity != null) {
throw new RenException(ErrorCode.DB_RECORD_EXISTS);
}
}
KnowledgeBaseEntity entity = ConvertUtils.sourceToTarget(knowledgeBaseDTO, KnowledgeBaseEntity.class);
knowledgeBaseDao.updateById(entity);
// 调用RAGFlow API更新数据集
if (StringUtils.isNotBlank(knowledgeBaseDTO.getDatasetId())) {
try {
Map<String, Object> ragConfig = getValidatedRAGConfig(knowledgeBaseDTO.getRagModelId());
updateDatasetInRAGFlow(
knowledgeBaseDTO.getDatasetId(),
knowledgeBaseDTO.getName(),
knowledgeBaseDTO.getDescription(),
ragConfig
);
} catch (Exception e) {
// 如果RAG API调用失败,回滚本地数据库操作
knowledgeBaseDao.updateById(existingEntity);
throw e;
}
}
return ConvertUtils.sourceToTarget(entity, KnowledgeBaseDTO.class);
}
@Override
public void delete(String id) {
if (StringUtils.isBlank(id)) {
throw new RenException(ErrorCode.IDENTIFIER_NOT_NULL);
}
log.info("=== 开始删除操作 ===");
log.info("删除ID: {}", id);
KnowledgeBaseEntity entity = knowledgeBaseDao.selectById(id);
if (entity == null) {
log.warn("记录不存在,ID: {}", id);
throw new RenException(ErrorCode.Knowledge_Base_RECORD_NOT_EXISTS);
}
log.info("找到记录: ID={}, datasetId={}, ragModelId={}",
entity.getId(), entity.getDatasetId(), entity.getRagModelId());
// 先调用RAGFlow API删除数据集
boolean apiDeleteSuccess = false;
if (StringUtils.isNotBlank(entity.getDatasetId()) && StringUtils.isNotBlank(entity.getRagModelId())) {
try {
log.info("开始调用RAGFlow API删除数据集");
Map<String, Object> ragConfig = getRAGConfig(entity.getRagModelId());
validateRagConfig(ragConfig);
deleteDatasetInRAGFlow(entity.getDatasetId(), ragConfig);
log.info("RAGFlow API删除调用完成");
apiDeleteSuccess = true;
} catch (Exception e) {
log.error("删除RAGFlow数据集失败: {}", e.getMessage());
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED, "删除RAGFlow数据集失败: " + e.getMessage());
}
} else {
log.warn("datasetId或ragModelId为空,跳过RAGFlow删除");
apiDeleteSuccess = true; // 没有RAG数据集,视为成功
}
// API删除成功后再删除本地记录
if (apiDeleteSuccess) {
int deleteCount = knowledgeBaseDao.deleteById(id);
log.info("本地数据库删除结果: {}", deleteCount > 0 ? "成功" : "失败");
}
log.info("=== 删除操作结束 ===");
}
@Override
public KnowledgeBaseDTO getByDatasetId(String datasetId) {
if (StringUtils.isBlank(datasetId)) {
throw new RenException(ErrorCode.PARAMS_GET_ERROR);
}
KnowledgeBaseEntity entity = knowledgeBaseDao.selectOne(
new QueryWrapper<KnowledgeBaseEntity>().eq("dataset_id", datasetId)
);
if (entity == null) {
throw new RenException(ErrorCode.Knowledge_Base_RECORD_NOT_EXISTS);
}
return ConvertUtils.sourceToTarget(entity, KnowledgeBaseDTO.class);
}
@Override
public void deleteByDatasetId(String datasetId) {
if (StringUtils.isBlank(datasetId)) {
throw new RenException(ErrorCode.PARAMS_GET_ERROR);
}
log.info("=== 开始通过datasetId删除操作 ===");
log.info("删除datasetId: {}", datasetId);
KnowledgeBaseEntity entity = knowledgeBaseDao.selectOne(
new QueryWrapper<KnowledgeBaseEntity>().eq("dataset_id", datasetId)
);
if (entity == null) {
log.warn("记录不存在,datasetId: {}", datasetId);
throw new RenException(ErrorCode.Knowledge_Base_RECORD_NOT_EXISTS);
}
log.info("找到记录: ID={}, datasetId={}, ragModelId={}",
entity.getId(), entity.getDatasetId(), entity.getRagModelId());
// 先删除本地数据库记录
int deleteCount = knowledgeBaseDao.deleteById(entity.getId());
log.info("本地数据库删除结果: {}", deleteCount > 0 ? "成功" : "失败");
// 调用RAGFlow API删除数据集
if (StringUtils.isNotBlank(entity.getDatasetId()) && StringUtils.isNotBlank(entity.getRagModelId())) {
try {
log.info("开始调用RAGFlow API删除数据集");
Map<String, Object> ragConfig = getValidatedRAGConfig(entity.getRagModelId());
deleteDatasetInRAGFlow(entity.getDatasetId(), ragConfig);
log.info("RAGFlow API删除调用完成");
} catch (Exception e) {
log.warn("删除RAGFlow数据集失败: {}", e.getMessage());
}
} else {
log.warn("datasetId或ragModelId为空,跳过RAGFlow删除");
}
log.info("=== 通过datasetId删除操作结束 ===");
}
@Override
public Map<String, Object> getRAGConfig(String ragModelId) {
if (StringUtils.isBlank(ragModelId)) {
throw new RenException(ErrorCode.PARAMS_GET_ERROR);
}
// 从缓存获取模型配置
ModelConfigEntity modelConfig = modelConfigService.getModelByIdFromCache(ragModelId);
if (modelConfig == null || modelConfig.getConfigJson() == null) {
throw new RenException(ErrorCode.RAG_CONFIG_NOT_FOUND);
}
// 验证是否为RAG类型配置
if (!Constant.RAG_CONFIG_TYPE.equals(modelConfig.getModelType().toUpperCase())) {
throw new RenException(ErrorCode.RAG_CONFIG_TYPE_ERROR);
}
Map<String, Object> config = modelConfig.getConfigJson();
// 验证必要的配置参数
validateRagConfig(config);
// 返回配置信息
return config;
}
@Override
public Map<String, Object> getDefaultRAGConfig() {
// 获取默认RAG模型配置
QueryWrapper<ModelConfigEntity> queryWrapper = new QueryWrapper<>();
queryWrapper.eq("model_type", Constant.RAG_CONFIG_TYPE)
.eq("is_default", 1)
.eq("is_enabled", 1);
List<ModelConfigEntity> modelConfigs = modelConfigDao.selectList(queryWrapper);
if (modelConfigs == null || modelConfigs.isEmpty()) {
throw new RenException(ErrorCode.RAG_DEFAULT_CONFIG_NOT_FOUND);
}
ModelConfigEntity defaultConfig = modelConfigs.get(0);
if (defaultConfig.getConfigJson() == null) {
throw new RenException(ErrorCode.RAG_CONFIG_NOT_FOUND);
}
Map<String, Object> config = defaultConfig.getConfigJson();
// 验证必要的配置参数
validateRagConfig(config);
return config;
}
/**
* 验证RAG配置中是否包含必要的参数
*/
private void validateRagConfig(Map<String, Object> config) {
if (config == null) {
throw new RenException(ErrorCode.RAG_CONFIG_NOT_FOUND);
}
// 从配置中提取必要的参数
String baseUrl = (String) config.get("base_url");
String apiKey = (String) config.get("api_key");
// 验证base_url是否存在且非空
if (StringUtils.isBlank(baseUrl)) {
throw new RenException(ErrorCode.RAG_CONFIG_MISSING_PARAMS);
}
}
/**
* 调用RAGFlow API创建数据集
*/
private String createDatasetInRAGFlow(String name, String description, Map<String, Object> ragConfig) {
String datasetId = null;
try {
String baseUrl = (String) ragConfig.get("base_url");
String apiKey = (String) ragConfig.get("api_key");
log.info("开始调用RAGFlow API创建数据集, name: {}", name);
log.debug("RAGFlow配置 - baseUrl: {}, apiKey: {}", baseUrl, StringUtils.isBlank(apiKey) ? "未配置" : "已配置");
// 构建请求URL
String url = baseUrl + "/api/v1/datasets";
log.debug("请求URL: {}", url);
// 构建请求体
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("name", name);
if (StringUtils.isNotBlank(description)) {
requestBody.put("description", description);
}
log.debug("请求体: {}", requestBody);
// 设置请求头
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.set("Authorization", "Bearer " + apiKey);
HttpEntity<Map<String, Object>> requestEntity = new HttpEntity<>(requestBody, headers);
// 发送POST请求
log.info("发送POST请求到RAGFlow API...");
ResponseEntity<String> response = restTemplate.exchange(url, HttpMethod.POST, requestEntity, String.class);
log.info("RAGFlow API响应状态码: {}", response.getStatusCode());
log.debug("RAGFlow API响应内容: {}", response.getBody());
if (!response.getStatusCode().is2xxSuccessful()) {
log.error("RAGFlow API调用失败,状态码: {}, 响应内容: {}", response.getStatusCode(), response.getBody());
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED);
}
// 解析响应体,提取datasetId
String responseBody = response.getBody();
if (StringUtils.isNotBlank(responseBody)) {
try {
// 解析RAGFlow API响应,支持多种可能的字段名
ObjectMapper objectMapper = new ObjectMapper();
Map<String, Object> responseMap = objectMapper.readValue(responseBody, Map.class);
log.debug("RAGFlow API响应解析结果: {}", responseMap);
// 首先检查响应码
Integer code = (Integer) responseMap.get("code");
if (code != null && code == 0) {
// 响应码为0表示成功,从data字段中获取datasetId
Object dataObj = responseMap.get("data");
if (dataObj instanceof Map) {
Map<String, Object> dataMap = (Map<String, Object>) dataObj;
datasetId = (String) dataMap.get("id");
if (StringUtils.isBlank(datasetId)) {
// 如果id字段为空,尝试其他可能的字段名
datasetId = (String) dataMap.get("dataset_id");
datasetId = (String) dataMap.get("datasetId");
}
}
} else {
// 如果响应码不为0,说明API调用失败
log.error("RAGFlow API调用失败,响应码: {}, 响应内容: {}", code, responseBody);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "RAGFlow API调用失败,响应码: " + code);
}
log.info("从RAGFlow API响应中解析出datasetId: {}", datasetId);
log.debug("完整响应内容: {}", responseBody);
} catch (Exception e) {
log.error("解析RAGFlow API响应失败: {}, 响应内容: {}", e.getMessage(), responseBody);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "解析RAGFlow响应失败: " + e.getMessage());
}
}
if (StringUtils.isBlank(datasetId)) {
log.error("无法从RAGFlow API响应中获取datasetId,响应内容: {}", responseBody);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "RAGFlow API响应中未包含datasetId");
}
log.info("RAGFlow数据集创建成功,datasetId: {}", datasetId);
} catch (HttpClientErrorException e) {
log.error("RAGFlow API调用失败 - HTTP错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "创建RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (HttpServerErrorException e) {
log.error("RAGFlow API调用失败 - 服务器错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "创建RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (ResourceAccessException e) {
log.error("RAGFlow API调用失败 - 网络连接错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "创建RAGFlow数据集失败: 网络连接错误 - " + e.getMessage());
} catch (Exception e) {
log.error("RAGFlow API调用失败 - 未知错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_CREATE_FAILED, "创建RAGFlow数据集失败: " + e.getMessage());
}
return datasetId;
}
/**
* 调用RAGFlow API更新数据集
*/
private void updateDatasetInRAGFlow(String datasetId, String name, String description, Map<String, Object> ragConfig) {
try {
String baseUrl = (String) ragConfig.get("base_url");
String apiKey = (String) ragConfig.get("api_key");
log.info("开始调用RAGFlow API更新数据集,datasetId: {}, name: {}", datasetId, name);
log.debug("RAGFlow配置 - baseUrl: {}, apiKey: {}", baseUrl, StringUtils.isBlank(apiKey) ? "未配置" : "已配置");
// 构建请求URL
String url = baseUrl + "/api/v1/datasets/" + datasetId;
log.debug("请求URL: {}", url);
// 构建请求体
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("dataset_id", datasetId);
requestBody.put("name", name);
if (StringUtils.isNotBlank(description)) {
requestBody.put("description", description);
}
log.debug("请求体: {}", requestBody);
// 设置请求头
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.set("Authorization", "Bearer " + apiKey);
HttpEntity<Map<String, Object>> requestEntity = new HttpEntity<>(requestBody, headers);
// 发送PUT请求
log.info("发送PUT请求到RAGFlow API...");
ResponseEntity<String> response = restTemplate.exchange(url, HttpMethod.PUT, requestEntity, String.class);
log.info("RAGFlow API响应状态码: {}", response.getStatusCode());
log.debug("RAGFlow API响应内容: {}", response.getBody());
if (!response.getStatusCode().is2xxSuccessful()) {
log.error("RAGFlow API调用失败,状态码: {}, 响应内容: {}", response.getStatusCode(), response.getBody());
throw new RenException(ErrorCode.RAG_API_UPDATE_FAILED);
}
log.info("RAGFlow数据集更新成功,datasetId: {}", datasetId);
} catch (HttpClientErrorException e) {
log.error("RAGFlow API调用失败 - HTTP错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_UPDATE_FAILED, "更新RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (HttpServerErrorException e) {
log.error("RAGFlow API调用失败 - 服务器错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_UPDATE_FAILED, "更新RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (ResourceAccessException e) {
log.error("RAGFlow API调用失败 - 网络连接错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_UPDATE_FAILED, "更新RAGFlow数据集失败: 网络连接错误 - " + e.getMessage());
} catch (Exception e) {
log.error("RAGFlow API调用失败 - 未知错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_UPDATE_FAILED, "更新RAGFlow数据集失败: " + e.getMessage());
}
}
/**
* 调用RAGFlow API删除数据集
*/
private void deleteDatasetInRAGFlow(String datasetId, Map<String, Object> ragConfig) {
try {
String baseUrl = (String) ragConfig.get("base_url");
String apiKey = (String) ragConfig.get("api_key");
log.info("开始调用RAGFlow API删除数据集,datasetId: {}", datasetId);
log.debug("RAGFlow配置 - baseUrl: {}, apiKey: {}", baseUrl, StringUtils.isBlank(apiKey) ? "未配置" : "已配置");
// 构建请求URL
String url = baseUrl + "/api/v1/datasets";
log.debug("请求URL: {}", url);
// 构建请求体
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("ids", List.of(datasetId));
log.debug("请求体: {}", requestBody);
// 设置请求头
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.set("Authorization", "Bearer " + apiKey);
HttpEntity<Map<String, Object>> requestEntity = new HttpEntity<>(requestBody, headers);
// 发送DELETE请求
log.info("发送DELETE请求到RAGFlow API...");
ResponseEntity<String> response = restTemplate.exchange(url, HttpMethod.DELETE, requestEntity, String.class);
log.info("RAGFlow API响应状态码: {}", response.getStatusCode());
log.debug("RAGFlow API响应内容: {}", response.getBody());
if (!response.getStatusCode().is2xxSuccessful()) {
log.error("RAGFlow API调用失败,状态码: {}, 响应内容: {}", response.getStatusCode(), response.getBody());
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED);
}
log.info("RAGFlow数据集删除成功,datasetId: {}", datasetId);
} catch (HttpClientErrorException e) {
log.error("RAGFlow API调用失败 - HTTP错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED, "删除RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (HttpServerErrorException e) {
log.error("RAGFlow API调用失败 - 服务器错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED, "删除RAGFlow数据集失败: " + e.getMessage() + ", 响应: " + e.getResponseBodyAsString());
} catch (ResourceAccessException e) {
log.error("RAGFlow API调用失败 - 网络连接错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED, "删除RAGFlow数据集失败: 网络连接错误 - " + e.getMessage());
} catch (Exception e) {
log.error("RAGFlow API调用失败 - 未知错误: {}", e.getMessage(), e);
throw new RenException(ErrorCode.RAG_API_DELETE_FAILED, "删除RAGFlow数据集失败: " + e.getMessage());
}
}
/**
* 获取RAG配置并验证
*/
private Map<String, Object> getValidatedRAGConfig(String ragModelId) {
Map<String, Object> ragConfig;
if (StringUtils.isNotBlank(ragModelId)) {
ragConfig = getRAGConfig(ragModelId);
} else {
ragConfig = getDefaultRAGConfig();
}
// 验证配置
validateRagConfig(ragConfig);
return ragConfig;
}
/**
* 从RAGFlow API获取知识库的文档数量
*/
private Integer getDocumentCountFromRAGFlow(String datasetId, String ragModelId) {
if (StringUtils.isBlank(datasetId) || StringUtils.isBlank(ragModelId)) {
log.warn("datasetId或ragModelId为空,无法获取文档数量");
return 0;
}
try {
log.info("开始获取知识库 {} 的文档数量", datasetId);
// 获取RAG配置
Map<String, Object> ragConfig = getValidatedRAGConfig(ragModelId);
String baseUrl = (String) ragConfig.get("base_url");
String apiKey = (String) ragConfig.get("api_key");
// 构建请求URL - 调用RAGFlow API获取文档列表,但不返回文档详情,只获取总数
String url = baseUrl + "/api/v1/datasets/" + datasetId + "/documents?page=1&size=1";
log.debug("请求URL: {}", url);
// 设置请求头
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.set("Authorization", "Bearer " + apiKey);
HttpEntity<String> requestEntity = new HttpEntity<>(headers);
// 发送GET请求
log.info("发送GET请求到RAGFlow API获取文档数量...");
ResponseEntity<String> response = restTemplate.exchange(url, HttpMethod.GET, requestEntity, String.class);
log.info("RAGFlow API响应状态码: {}", response.getStatusCode());
if (!response.getStatusCode().is2xxSuccessful()) {
log.error("RAGFlow API调用失败,状态码: {}, 响应内容: {}", response.getStatusCode(), response.getBody());
return 0;
}
String responseBody = response.getBody();
log.debug("RAGFlow API响应内容: {}", responseBody);
// 解析响应
ObjectMapper objectMapper = new ObjectMapper();
Map<String, Object> responseMap = objectMapper.readValue(responseBody, Map.class);
Integer code = (Integer) responseMap.get("code");
if (code != null && code == 0) {
Object dataObj = responseMap.get("data");
if (dataObj instanceof Map) {
Map<String, Object> dataMap = (Map<String, Object>) dataObj;
Object totalObj = dataMap.get("total");
if (totalObj instanceof Integer) {
Integer documentCount = (Integer) totalObj;
log.info("获取知识库 {} 的文档数量成功: {}", datasetId, documentCount);
return documentCount;
} else if (totalObj instanceof Long) {
Long documentCount = (Long) totalObj;
log.info("获取知识库 {} 的文档数量成功: {}", datasetId, documentCount);
return documentCount.intValue();
}
}
} else {
log.error("RAGFlow API调用失败,响应码: {}, 响应内容: {}", code, responseBody);
}
} catch (IOException e) {
log.error("解析RAGFlow API响应失败: {}", e.getMessage(), e);
} catch (HttpClientErrorException e) {
log.error("RAGFlow API调用失败 - HTTP错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
} catch (HttpServerErrorException e) {
log.error("RAGFlow API调用失败 - 服务器错误: {}, 状态码: {}, 响应内容: {}",
e.getMessage(), e.getStatusCode(), e.getResponseBodyAsString(), e);
} catch (ResourceAccessException e) {
log.error("RAGFlow API调用失败 - 网络连接错误: {}", e.getMessage(), e);
} catch (Exception e) {
log.error("获取文档数量失败: {}", e.getMessage(), e);
}
return 0;
}
/**
* 同步RAGFlow API的数据集状态(可选功能)
*/
private void syncRAGFlowDatasetStatus(List<KnowledgeBaseEntity> entities) {
if (entities == null || entities.isEmpty()) {
log.debug("没有需要同步状态的数据集");
return;
}
log.info("开始同步RAGFlow数据集状态,共{}个数据集", entities.size());
for (KnowledgeBaseEntity entity : entities) {
if (StringUtils.isNotBlank(entity.getDatasetId()) && StringUtils.isNotBlank(entity.getRagModelId())) {
try {
log.debug("开始同步数据集 {} 的状态", entity.getDatasetId());
Map<String, Object> ragConfig = getValidatedRAGConfig(entity.getRagModelId());
String baseUrl = (String) ragConfig.get("base_url");
String apiKey = (String) ragConfig.get("api_key");
// 使用正确的API端点获取数据集列表,然后过滤
String url = baseUrl + "/api/v1/datasets";
log.debug("请求URL: {}", url);
HttpHeaders headers = new HttpHeaders();
headers.set("Authorization", "Bearer " + apiKey);
HttpEntity<Void> requestEntity = new HttpEntity<>(headers);
// 发送GET请求获取所有数据集
ResponseEntity<Map> response = restTemplate.exchange(url, HttpMethod.GET, requestEntity, Map.class);
if (response.getStatusCode().is2xxSuccessful() && response.getBody() != null) {
Map<String, Object> responseBody = response.getBody();
// 根据响应结构判断数据集是否存在
boolean exists = checkDatasetExists(responseBody, entity.getDatasetId());
if (exists) {
log.info("数据集 {} 在RAGFlow中状态正常", entity.getDatasetId());
} else {
log.warn("数据集 {} 在RAGFlow中不存在", entity.getDatasetId());
}
} else {
log.warn("获取数据集列表失败,状态码: {}", response.getStatusCode());
}
} catch (Exception e) {
log.error("同步数据集 {} 状态失败: {}", entity.getDatasetId(), e.getMessage());
}
}
}
log.info("RAGFlow数据集状态同步完成");
}
/**
* 检查数据集是否存在
*/
private boolean checkDatasetExists(Map<String, Object> responseBody, String datasetId) {
try {
// 根据RAGFlow API的实际响应结构来解析
if (responseBody.containsKey("data")) {
Object data = responseBody.get("data");
if (data instanceof List) {
List<Map<String, Object>> datasets = (List<Map<String, Object>>) data;
return datasets.stream()
.anyMatch(dataset -> datasetId.equals(dataset.get("id")) ||
datasetId.equals(dataset.get("dataset_id")));
}
}
return false;
} catch (Exception e) {
log.error("检查数据集存在性失败: {}", e.getMessage());
return false;
}
}
}
@@ -0,0 +1,24 @@
-- 添加RAG模型供应器和配置
-- -------------------------------------------------------
-- 添加RAG模型供应器
delete from `ai_model_provider` where id = 'SYSTEM_RAG_ragflow';
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
('SYSTEM_RAG_ragflow', 'RAG', 'ragflow', 'RAGFlow', '[{"key": "base_url", "type": "string", "label": "服务地址"}, {"key": "api_key", "type": "string", "label": "API密钥"}]', 1, 1, NOW(), 1, NOW());
-- 添加RAG模型配置
delete from `ai_model_config` where id = 'RAG_RAGFlow';
INSERT INTO `ai_model_config` VALUES ('RAG_RAGFlow', 'RAG', 'ragflow', 'RAGFlow', 1, 1, '{"type": "ragflow", "base_url": "http://localhost", "api_key": "your_api_key_here"}', 'https://github.com/infiniflow/ragflow/blob/main/README_zh.md', 'RAGFlow配置说明:
一、快速部署教程(docker部署)
1.$ sysctl vm.max_map_count
2.$ sysctl -w vm.max_map_count=262144
3.$ git clone https://github.com/infiniflow/ragflow.git
4.docker compose -f docker-compose.yml up -d
5.$ docker logs -f docker-ragflow-cpu-1
6.注冊登录后,点击右上角头像,获得RAGFlow的API KEY和API服务器地址。使用RAGFlow前请在Model Provider中添加模型和设置默认模型。
二、如果您希望关掉注册功能
1.停止服务 docker compose down
2. sed -i ''s/REGISTER_ENABLED=1/REGISTER_ENABLED=0/g'' .env
3.cat .env | grep -i register
4.看到REGISTER_ENABLED=0 重启服务即可。', 1, NULL, NULL, NULL, NULL);
@@ -0,0 +1,19 @@
-- 知识库表
DROP TABLE IF EXISTS `ai_rag_dataset`;
CREATE TABLE `ai_rag_dataset` (
`id` VARCHAR(32) NOT NULL COMMENT '唯一标识',
`dataset_id` VARCHAR(64) NOT NULL COMMENT '知识库ID',
`rag_model_id` VARCHAR(64) COMMENT 'RAG模型配置ID',
`name` VARCHAR(100) NOT NULL COMMENT '知识库名称',
`description` TEXT COMMENT '知识库描述',
`status` TINYINT(1) DEFAULT 1 COMMENT '状态:0停用 1启用',
`creator` BIGINT COMMENT '创建者',
`created_at` DATETIME COMMENT '创建时间',
`updater` BIGINT COMMENT '更新者',
`updated_at` DATETIME COMMENT '更新时间',
PRIMARY KEY (`id`),
UNIQUE KEY `uk_dataset_id` (`dataset_id`),
INDEX `idx_ai_rag_dataset_status` (`status`),
INDEX `idx_ai_rag_dataset_creator` (`creator`),
INDEX `idx_ai_rag_dataset_created_at` (`created_at`)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='知识库表';
@@ -402,3 +402,17 @@ databaseChangeLog:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202510191042.sql
- changeSet:
id: 202510250955
author: rainv123
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202510250955.sql
- changeSet:
id: 202510251150
author: rainv123
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202510251150.sql
@@ -168,4 +168,17 @@
10159=Voice ID already exists
10160=Huoshan Engine voice ID format error, must start with S_
10161=Mac address already exists
10162=Model provider does not exist
10162=Model provider does not exist
10163=Knowledge base record does not exist
10164=RAG configuration not found
10165=RAG configuration type error
10166=Default RAG configuration not found
10167=RAG configuration missing required parameters
10168=RAG API create dataset failed
10169=RAG API update dataset failed
10170=RAG API delete dataset failed
10171=Upload file failed
10172=RAG API query failed
10173=RAG API parse failed
10174=RAG API operation failed
@@ -168,4 +168,16 @@
10159=\u97F3\u8272ID\u5DF2\u5B58\u5728
10160=\u706B\u5C71\u5F15\u64CE\u97F3\u8272ID\u683C\u5F0F\u9519\u8BEF\uFF0C\u5FC5\u987B\u4EE5S_\u5F00\u5934
10161=Mac\u5730\u5740\u5DF2\u5B58\u5728
10162=\u6A21\u578B\u4F9B\u5E94\u5668\u4E0D\u5B58\u5728
10162=\u6A21\u578B\u4F9B\u5E94\u5668\u4E0D\u5B58\u5728
10163=\u77E5\u8BC6\u5E93\u8BB0\u5F55\u4E0D\u5B58\u5728
10164=RAG\u914D\u7F6E\u672A\u627E\u5230
10165=RAG\u914D\u7F6E\u7C7B\u578B\u9519\u8BEF
10166=\u9ED8\u8BA4RAG\u914D\u7F6E\u672A\u627E\u5230
10167=RAG\u914D\u7F6E\u7F3A\u5C11\u5FC5\u8981\u53C2\u6570
10168=RAG API\u521B\u5EFA\u6570\u636E\u96C6\u5931\u8D25
10169=RAG API\u66F4\u65B0\u6570\u636E\u96C6\u5931\u8D25
10170=RAG API\u5220\u9664\u6570\u636E\u96C6\u5931\u8D25
10171=\u4E0A\u4F20\u6587\u4EF6\u5931\u8D25
10172=RAG API\u67E5\u8BE2\u5931\u8D25
10173=RAG API\u89E3\u6790\u5931\u8D25
10174=RAG API\u64CD\u4F5C\u5931\u8D25
@@ -169,4 +169,15 @@
10160=\u706B\u5C71\u5F15\u64CE\u97F3\u8272ID\u683C\u5F0F\u932F\u8AA4\uFF0C\u5FC5\u9808\u4EE5S_\u958B\u982D
10161=Mac\u5730\u5740\u5DF2\u5B58\u5728
10162=\u6A21\u578B\u63D0\u4F9B\u5546\u4E0D\u5B58\u5728
10163=\u77E5\u8B58\u5EAB\u8A18\u9304\u4E0D\u5B58\u5728
10164=RAG\u914D\u7F6E\u672A\u627E\u5230
10165=RAG\u914D\u7F6E\u985E\u578B\u932F\u8AA4
10166=\u9810\u8A2DRAG\u914D\u7F6E\u672A\u627E\u5230
10167=RAG\u914D\u7F6E\u7F3A\u5C11\u5FC5\u8981\u53C3\u6578
10168=RAG API\u5275\u5EFA\u6578\u64DA\u96C6\u5931\u6557
10169=RAG API\u66F4\u65B0\u6578\u64DA\u96C6\u5931\u6557
10170=RAG API\u522A\u9664\u6578\u64DA\u96C6\u5931\u6557
10171=\u4E0A\u50B3\u6587\u4EF6\u5931\u6557
10172=RAG API\u67E5\u8A62\u5931\u6557
10173=RAG API\u5256\u6790\u5931\u6557
10174=RAG API\u64CD\u4F5C\u5931\u6557
@@ -0,0 +1,19 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="xiaozhi.modules.knowledge.dao.KnowledgeBaseDao">
<resultMap id="KnowledgeBaseResultMap" type="xiaozhi.modules.knowledge.entity.KnowledgeBaseEntity">
<id column="id" property="id"/>
<result column="dataset_id" property="datasetId"/>
<result column="dataset_name" property="datasetName"/>
<result column="dataset_type" property="datasetType"/>
<result column="description" property="description"/>
<result column="status" property="status"/>
<result column="config_json" property="configJson" typeHandler="com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler"/>
<result column="remark" property="remark"/>
<result column="updater" property="updater"/>
<result column="updated_at" property="updatedAt"/>
<result column="creator" property="creator"/>
<result column="created_at" property="createdAt"/>
</resultMap>
</mapper>