mirror of
https://gitcode.com/ageerle/ruoyi-ai.git
synced 2026-09-13 00:14:59 +00:00
fix(rag): 完善向量路由批量写入与数据生命周期
This commit is contained in:
@@ -1131,7 +1131,8 @@ CREATE TABLE `knowledge_attach` (
|
||||
`id` bigint NOT NULL AUTO_INCREMENT COMMENT '主键',
|
||||
`knowledge_id` bigint NOT NULL COMMENT '知识库ID',
|
||||
`oss_id` bigint NULL DEFAULT NULL COMMENT '对象存储ID',
|
||||
`doc_id` varchar(11) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '文档ID',
|
||||
`doc_id` varchar(32) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '文档ID',
|
||||
`file_hash` varchar(64) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '文件SHA-256摘要',
|
||||
`name` varchar(500) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '附件名称',
|
||||
`type` varchar(50) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '附件类型',
|
||||
`create_dept` varchar(255) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '部门',
|
||||
@@ -1143,7 +1144,8 @@ CREATE TABLE `knowledge_attach` (
|
||||
`tenant_id` bigint NOT NULL DEFAULT 0 COMMENT '租户Id',
|
||||
`status` tinyint NULL DEFAULT 0 COMMENT '解析状态: 0待解析, 1解析中, 2已解析, 3解析失败',
|
||||
PRIMARY KEY (`id`) USING BTREE,
|
||||
UNIQUE INDEX `idx_kname`(`knowledge_id` ASC, `name` ASC) USING BTREE
|
||||
UNIQUE INDEX `idx_kname`(`knowledge_id` ASC, `name` ASC) USING BTREE,
|
||||
UNIQUE INDEX `uk_knowledge_file_hash`(`knowledge_id`, `file_hash`) USING BTREE
|
||||
) ENGINE = InnoDB AUTO_INCREMENT = 2033199209203183619 CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '知识库附件' ROW_FORMAT = DYNAMIC;
|
||||
|
||||
-- ----------------------------
|
||||
@@ -1156,8 +1158,9 @@ CREATE TABLE `knowledge_attach` (
|
||||
DROP TABLE IF EXISTS `knowledge_fragment`;
|
||||
CREATE TABLE `knowledge_fragment` (
|
||||
`id` bigint NOT NULL AUTO_INCREMENT COMMENT '主键',
|
||||
`fid` varchar(32) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '向量库片段ID',
|
||||
`idx` int NOT NULL COMMENT '片段索引下标',
|
||||
`doc_id` varchar(11) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '文档ID',
|
||||
`doc_id` varchar(32) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '文档ID',
|
||||
`content` text CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '文档内容',
|
||||
`create_dept` varchar(255) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '部门',
|
||||
`create_by` varchar(50) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '创建人',
|
||||
@@ -1168,6 +1171,9 @@ CREATE TABLE `knowledge_fragment` (
|
||||
`tenant_id` bigint NOT NULL DEFAULT 0 COMMENT '租户Id',
|
||||
`knowledge_id` bigint NULL DEFAULT NULL COMMENT '知识库ID',
|
||||
PRIMARY KEY (`id`) USING BTREE,
|
||||
UNIQUE INDEX `uk_fid`(`fid`) USING BTREE,
|
||||
INDEX `idx_doc_id`(`doc_id`) USING BTREE,
|
||||
INDEX `idx_knowledge_id`(`knowledge_id`) USING BTREE,
|
||||
FULLTEXT INDEX `ft_content`(`content`) WITH PARSER `ngram`
|
||||
) ENGINE = InnoDB AUTO_INCREMENT = 2033199209131880451 CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '知识片段' ROW_FORMAT = DYNAMIC;
|
||||
|
||||
@@ -1206,7 +1212,9 @@ CREATE TABLE `knowledge_info` (
|
||||
`enable_hybrid` tinyint(1) NULL DEFAULT 0 COMMENT '是否启用混合检索',
|
||||
`hybrid_alpha` double NULL DEFAULT 0.5 COMMENT '混合检索权重比例 (0.0=纯向量, 1.0=纯关键词)',
|
||||
`system_prompt` text CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL COMMENT '系统提示词',
|
||||
PRIMARY KEY (`id`) USING BTREE
|
||||
PRIMARY KEY (`id`) USING BTREE,
|
||||
INDEX `idx_tenant_user` (`tenant_id`, `user_id`) USING BTREE,
|
||||
INDEX `idx_tenant_share` (`tenant_id`, `share`) USING BTREE
|
||||
) ENGINE = InnoDB AUTO_INCREMENT = 2033198818050781187 CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '知识库' ROW_FORMAT = DYNAMIC;
|
||||
|
||||
-- ----------------------------
|
||||
|
||||
44
docs/script/sql/update/2026-07-20-knowledge-fragment-fid.sql
Normal file
44
docs/script/sql/update/2026-07-20-knowledge-fragment-fid.sql
Normal file
@@ -0,0 +1,44 @@
|
||||
-- RAG metadata migration (MySQL 8). Safe to execute repeatedly.
|
||||
ALTER TABLE `knowledge_attach`
|
||||
ADD COLUMN IF NOT EXISTS `file_hash` varchar(64) NULL DEFAULT NULL COMMENT '文件SHA-256摘要' AFTER `doc_id`,
|
||||
MODIFY COLUMN `doc_id` varchar(32) NULL DEFAULT NULL COMMENT '文档ID';
|
||||
|
||||
SET @add_file_hash = IF(EXISTS(
|
||||
SELECT 1 FROM information_schema.statistics WHERE table_schema = DATABASE()
|
||||
AND table_name = 'knowledge_attach' AND index_name = 'uk_knowledge_file_hash'),
|
||||
'SELECT 1', 'ALTER TABLE `knowledge_attach` ADD UNIQUE INDEX `uk_knowledge_file_hash` (`knowledge_id`, `file_hash`)');
|
||||
PREPARE stmt FROM @add_file_hash; EXECUTE stmt; DEALLOCATE PREPARE stmt;
|
||||
|
||||
ALTER TABLE `knowledge_fragment`
|
||||
ADD COLUMN IF NOT EXISTS `fid` varchar(32) NULL DEFAULT NULL COMMENT '向量库片段ID' AFTER `id`,
|
||||
MODIFY COLUMN `doc_id` varchar(32) NULL DEFAULT NULL COMMENT '文档ID';
|
||||
|
||||
UPDATE `knowledge_fragment`
|
||||
SET `fid` = LOWER(MD5(CONCAT('knowledge_fragment:', `id`)))
|
||||
WHERE `fid` IS NULL OR `fid` = '';
|
||||
|
||||
SET @drop_idx_fid = IF(EXISTS(
|
||||
SELECT 1 FROM information_schema.statistics WHERE table_schema = DATABASE()
|
||||
AND table_name = 'knowledge_fragment' AND index_name = 'idx_fid'),
|
||||
'ALTER TABLE `knowledge_fragment` DROP INDEX `idx_fid`', 'SELECT 1');
|
||||
PREPARE stmt FROM @drop_idx_fid; EXECUTE stmt; DEALLOCATE PREPARE stmt;
|
||||
|
||||
SET @add_uk_fid = IF(EXISTS(
|
||||
SELECT 1 FROM information_schema.statistics WHERE table_schema = DATABASE()
|
||||
AND table_name = 'knowledge_fragment' AND index_name = 'uk_fid'),
|
||||
'SELECT 1', 'ALTER TABLE `knowledge_fragment` ADD UNIQUE INDEX `uk_fid` (`fid`)');
|
||||
PREPARE stmt FROM @add_uk_fid; EXECUTE stmt; DEALLOCATE PREPARE stmt;
|
||||
|
||||
ALTER TABLE `knowledge_fragment` MODIFY COLUMN `fid` varchar(32) NOT NULL COMMENT '向量库片段ID';
|
||||
|
||||
SET @add_tenant_user = IF(EXISTS(
|
||||
SELECT 1 FROM information_schema.statistics WHERE table_schema = DATABASE()
|
||||
AND table_name = 'knowledge_info' AND index_name = 'idx_tenant_user'),
|
||||
'SELECT 1', 'ALTER TABLE `knowledge_info` ADD INDEX `idx_tenant_user` (`tenant_id`, `user_id`)');
|
||||
PREPARE stmt FROM @add_tenant_user; EXECUTE stmt; DEALLOCATE PREPARE stmt;
|
||||
|
||||
SET @add_tenant_share = IF(EXISTS(
|
||||
SELECT 1 FROM information_schema.statistics WHERE table_schema = DATABASE()
|
||||
AND table_name = 'knowledge_info' AND index_name = 'idx_tenant_share'),
|
||||
'SELECT 1', 'ALTER TABLE `knowledge_info` ADD INDEX `idx_tenant_share` (`tenant_id`, `share`)');
|
||||
PREPARE stmt FROM @add_tenant_share; EXECUTE stmt; DEALLOCATE PREPARE stmt;
|
||||
@@ -106,7 +106,10 @@ public class KnowledgeAttachController extends BaseController {
|
||||
|
||||
/**
|
||||
* 上传知识库附件
|
||||
* 注意:multipart 上传不能加 @RepeatSubmit(其参数序列化不支持 MultipartFile)
|
||||
*/
|
||||
@SaCheckPermission("system:attach:add")
|
||||
@Log(title = "知识库附件", businessType = BusinessType.INSERT)
|
||||
@PostMapping(value = "/upload")
|
||||
public R<String> upload(KnowledgeInfoUploadBo bo){
|
||||
knowledgeAttachService.upload(bo);
|
||||
@@ -118,7 +121,10 @@ public class KnowledgeAttachController extends BaseController {
|
||||
*
|
||||
* @param id 附件ID
|
||||
*/
|
||||
@SaCheckPermission("system:attach:edit")
|
||||
@Log(title = "知识库附件", businessType = BusinessType.UPDATE)
|
||||
@PostMapping("/parse/{id}")
|
||||
@RepeatSubmit()
|
||||
public R<Void> parse(@PathVariable Long id) {
|
||||
knowledgeAttachService.parse(id);
|
||||
return R.ok();
|
||||
|
||||
@@ -107,7 +107,9 @@ public class KnowledgeFragmentController extends BaseController {
|
||||
/**
|
||||
* 检索测试
|
||||
*/
|
||||
@SaCheckPermission("system:fragment:list")
|
||||
@PostMapping("/retrieval")
|
||||
@RepeatSubmit()
|
||||
public R<List<KnowledgeRetrievalVo>> retrieval(@RequestBody KnowledgeFragmentBo bo) {
|
||||
return R.ok(knowledgeFragmentService.retrieval(bo));
|
||||
}
|
||||
|
||||
@@ -37,6 +37,9 @@ public class KnowledgeAttach extends BaseEntity {
|
||||
*/
|
||||
private String docId;
|
||||
|
||||
/** SHA-256 content digest used for upload idempotency. */
|
||||
private String fileHash;
|
||||
|
||||
/**
|
||||
* 附件名称
|
||||
*/
|
||||
|
||||
@@ -27,6 +27,11 @@ public class KnowledgeFragment extends BaseEntity {
|
||||
@TableId(value = "id")
|
||||
private Long id;
|
||||
|
||||
/**
|
||||
* 向量库片段ID(与向量库中的 fid 元数据对应,用于向量定位与混合检索融合)
|
||||
*/
|
||||
private String fid;
|
||||
|
||||
/**
|
||||
* 文档ID-用于关联文本块信息
|
||||
*/
|
||||
|
||||
@@ -30,6 +30,11 @@ public class KnowledgeFragmentVo implements Serializable {
|
||||
@ExcelProperty(value = "主键")
|
||||
private Long id;
|
||||
|
||||
/**
|
||||
* 向量库片段ID
|
||||
*/
|
||||
private String fid;
|
||||
|
||||
/**
|
||||
* 文档ID-用于关联文本块信息
|
||||
*/
|
||||
|
||||
@@ -10,6 +10,7 @@ import org.ruoyi.service.vector.impl.MilvusVectorStoreStrategy;
|
||||
import org.ruoyi.service.vector.impl.QdrantVectorStoreStrategy;
|
||||
import org.ruoyi.service.vector.impl.WeaviateVectorStoreStrategy;
|
||||
import org.springframework.stereotype.Component;
|
||||
import org.ruoyi.common.core.exception.ServiceException;
|
||||
|
||||
import java.util.HashMap;
|
||||
import java.util.Map;
|
||||
@@ -45,14 +46,20 @@ public class VectorStoreStrategyFactory {
|
||||
* 获取当前配置的向量库策略
|
||||
*/
|
||||
public VectorStoreService getStrategy() {
|
||||
String vectorStoreType = vectorStoreProperties.getType();
|
||||
return getStrategy(null);
|
||||
}
|
||||
|
||||
public VectorStoreService getStrategy(String requestedType) {
|
||||
String vectorStoreType = requestedType;
|
||||
if (vectorStoreType == null || vectorStoreType.trim().isEmpty()) {
|
||||
vectorStoreType = vectorStoreProperties.getType();
|
||||
}
|
||||
if (vectorStoreType == null || vectorStoreType.trim().isEmpty()) {
|
||||
vectorStoreType = "weaviate"; // 默认使用weaviate
|
||||
}
|
||||
VectorStoreService strategy = strategies.get(vectorStoreType.toLowerCase());
|
||||
if (strategy == null) {
|
||||
log.warn("未找到向量库策略: {}, 使用默认策略: weaviate", vectorStoreType);
|
||||
strategy = strategies.get("weaviate");
|
||||
throw new ServiceException("不支持的向量库类型: " + vectorStoreType);
|
||||
}
|
||||
log.debug("使用向量库策略: {}", vectorStoreType);
|
||||
return strategy;
|
||||
|
||||
@@ -6,6 +6,8 @@ import org.apache.ibatis.annotations.Select;
|
||||
import org.ruoyi.domain.entity.knowledge.KnowledgeAttach;
|
||||
import org.ruoyi.domain.vo.knowledge.KnowledgeAttachVo;
|
||||
import org.ruoyi.common.mybatis.core.mapper.BaseMapperPlus;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 知识库附件Mapper接口
|
||||
@@ -21,4 +23,10 @@ public interface KnowledgeAttachMapper extends BaseMapperPlus<KnowledgeAttach, K
|
||||
*/
|
||||
@Select("SELECT COUNT(*) FROM knowledge_attach WHERE knowledge_id = #{knowledgeId}")
|
||||
int countByKnowledgeId(@Param("knowledgeId") Long knowledgeId);
|
||||
|
||||
@Select("<script>SELECT knowledge_id AS knowledgeId, COUNT(*) AS documentCount " +
|
||||
"FROM knowledge_attach WHERE knowledge_id IN " +
|
||||
"<foreach collection='knowledgeIds' item='id' open='(' separator=',' close=')'>#{id}</foreach> " +
|
||||
"GROUP BY knowledge_id</script>")
|
||||
List<Map<String, Object>> countByKnowledgeIds(@Param("knowledgeIds") List<Long> knowledgeIds);
|
||||
}
|
||||
|
||||
@@ -34,7 +34,7 @@ public interface KnowledgeFragmentMapper extends BaseMapperPlus<KnowledgeFragmen
|
||||
"</script>")
|
||||
List<DocFragmentCountVo> selectFragmentCountByDocIds(@Param("docIds") List<String> docIds);
|
||||
@Select("<script>" +
|
||||
"SELECT id, doc_id AS docId, content, idx, knowledge_id AS knowledgeId " +
|
||||
"SELECT id, fid, doc_id AS docId, content, idx, knowledge_id AS knowledgeId " +
|
||||
"FROM knowledge_fragment " +
|
||||
"WHERE knowledge_id = #{knowledgeId} " +
|
||||
"AND MATCH (content) AGAINST (#{query} IN NATURAL LANGUAGE MODE) " +
|
||||
|
||||
@@ -2,6 +2,7 @@ package org.ruoyi.service.knowledge.impl;
|
||||
|
||||
import cn.hutool.core.collection.CollUtil;
|
||||
import cn.hutool.core.util.RandomUtil;
|
||||
import cn.hutool.crypto.digest.DigestUtil;
|
||||
import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.toolkit.Wrappers;
|
||||
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
|
||||
@@ -11,6 +12,7 @@ import org.ruoyi.common.chat.domain.vo.chat.ChatModelVo;
|
||||
import org.ruoyi.common.chat.service.chat.IChatModelService;
|
||||
import org.ruoyi.enums.KnowledgeAttachStatus;
|
||||
import org.ruoyi.common.core.domain.dto.OssDTO;
|
||||
import org.ruoyi.common.core.exception.ServiceException;
|
||||
import org.ruoyi.common.core.service.OssService;
|
||||
import org.ruoyi.common.core.utils.MapstructUtils;
|
||||
import org.ruoyi.common.core.utils.SpringUtils;
|
||||
@@ -32,6 +34,7 @@ import org.ruoyi.service.knowledge.IKnowledgeAttachService;
|
||||
import org.ruoyi.service.knowledge.IKnowledgeInfoService;
|
||||
import org.ruoyi.service.knowledge.ResourceLoader;
|
||||
import org.ruoyi.service.vector.VectorStoreService;
|
||||
import org.ruoyi.service.retrieval.KnowledgeRetrievalService;
|
||||
import org.springframework.scheduling.annotation.Async;
|
||||
import org.springframework.stereotype.Service;
|
||||
import org.springframework.web.multipart.MultipartFile;
|
||||
@@ -60,6 +63,7 @@ public class KnowledgeAttachServiceImpl implements IKnowledgeAttachService {
|
||||
private final ResourceLoaderFactory resourceLoaderFactory;
|
||||
private final VectorStoreService vectorStoreService;
|
||||
private final OssService ossService;
|
||||
private final KnowledgeRetrievalService knowledgeRetrievalService;
|
||||
|
||||
@Override
|
||||
public KnowledgeAttachVo queryById(Long id) {
|
||||
@@ -126,18 +130,44 @@ public class KnowledgeAttachServiceImpl implements IKnowledgeAttachService {
|
||||
|
||||
@Override
|
||||
public Boolean deleteWithValidByIds(Collection<Long> ids, Boolean isValid) {
|
||||
// 删除附件前,同步清理其片段记录与向量库中的向量
|
||||
List<KnowledgeAttach> attaches = baseMapper.selectByIds(ids);
|
||||
for (KnowledgeAttach attach : attaches) {
|
||||
String docId = attach.getDocId();
|
||||
String kid = String.valueOf(attach.getKnowledgeId());
|
||||
vectorStoreService.removeByDocId(docId, kid);
|
||||
knowledgeFragmentMapper.delete(
|
||||
Wrappers.<KnowledgeFragment>lambdaQuery().eq(KnowledgeFragment::getDocId, docId));
|
||||
if (attach.getOssId() != null) {
|
||||
ossService.deleteFile(attach.getOssId());
|
||||
}
|
||||
knowledgeRetrievalService.invalidateKnowledge(kid);
|
||||
}
|
||||
return baseMapper.deleteByIds(ids) > 0;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void upload(KnowledgeInfoUploadBo bo) {
|
||||
MultipartFile file = bo.getFile();
|
||||
final String fileHash;
|
||||
try (InputStream input = file.getInputStream()) {
|
||||
fileHash = DigestUtil.sha256Hex(input);
|
||||
} catch (Exception e) {
|
||||
throw new ServiceException("计算文件摘要失败", e);
|
||||
}
|
||||
boolean duplicate = baseMapper.exists(Wrappers.<KnowledgeAttach>lambdaQuery()
|
||||
.eq(KnowledgeAttach::getKnowledgeId, bo.getKnowledgeId())
|
||||
.eq(KnowledgeAttach::getFileHash, fileHash));
|
||||
if (duplicate) {
|
||||
throw new ServiceException("该文件已上传,请勿重复提交");
|
||||
}
|
||||
OssDTO ossDTO = ossService.uploadFile(file);
|
||||
|
||||
KnowledgeAttach knowledgeAttach = new KnowledgeAttach();
|
||||
knowledgeAttach.setKnowledgeId(bo.getKnowledgeId());
|
||||
knowledgeAttach.setOssId(ossDTO.getOssId());
|
||||
knowledgeAttach.setDocId(RandomUtil.randomString(10));
|
||||
knowledgeAttach.setFileHash(fileHash);
|
||||
knowledgeAttach.setName(ossDTO.getOriginalName());
|
||||
knowledgeAttach.setType(ossDTO.getFileSuffix());
|
||||
knowledgeAttach.setStatus(KnowledgeAttachStatus.WAITING.getCode()); // 待解析
|
||||
@@ -166,6 +196,8 @@ public class KnowledgeAttachServiceImpl implements IKnowledgeAttachService {
|
||||
|
||||
Long knowledgeId = attach.getKnowledgeId();
|
||||
String docId = attach.getDocId();
|
||||
List<KnowledgeFragment> oldFragments = knowledgeFragmentMapper.selectList(
|
||||
Wrappers.<KnowledgeFragment>lambdaQuery().eq(KnowledgeFragment::getDocId, docId));
|
||||
|
||||
// 获取文件信息并下载
|
||||
List<OssDTO> ossDTOs = ossService.selectByIds(String.valueOf(attach.getOssId()));
|
||||
@@ -180,25 +212,25 @@ public class KnowledgeAttachServiceImpl implements IKnowledgeAttachService {
|
||||
}
|
||||
List<String> chunkList = resourceLoader.getChunkList(content, String.valueOf(knowledgeId));
|
||||
|
||||
List<String> fids = new ArrayList<>();
|
||||
List<KnowledgeFragment> knowledgeFragmentList = new ArrayList<>();
|
||||
if (CollUtil.isNotEmpty(chunkList)) {
|
||||
for (int i = 0; i < chunkList.size(); i++) {
|
||||
String fid = RandomUtil.randomString(10);
|
||||
fids.add(fid);
|
||||
KnowledgeFragment knowledgeFragment = new KnowledgeFragment();
|
||||
knowledgeFragment.setKnowledgeId(knowledgeId);
|
||||
knowledgeFragment.setDocId(docId);
|
||||
knowledgeFragment.setIdx(i);
|
||||
knowledgeFragment.setContent(chunkList.get(i));
|
||||
knowledgeFragment.setCreateTime(new Date());
|
||||
knowledgeFragmentList.add(knowledgeFragment);
|
||||
}
|
||||
knowledgeFragmentMapper.delete(Wrappers.<KnowledgeFragment>lambdaQuery().eq(KnowledgeFragment::getDocId, docId));
|
||||
knowledgeFragmentMapper.insertBatch(knowledgeFragmentList);
|
||||
log.info("文档切片并入库完成,共计 {} 个片段。id: {}", chunkList.size(), id);
|
||||
if (CollUtil.isEmpty(chunkList)) {
|
||||
throw new RuntimeException("文档分片结果为空,请检查文档内容或分片器是否支持该文件类型");
|
||||
}
|
||||
|
||||
// 重新解析前先清理旧的向量数据,避免向量重复累积
|
||||
List<String> fids = new ArrayList<>();
|
||||
List<KnowledgeFragment> knowledgeFragmentList = new ArrayList<>();
|
||||
for (int i = 0; i < chunkList.size(); i++) {
|
||||
String fid = RandomUtil.randomString(10);
|
||||
fids.add(fid);
|
||||
KnowledgeFragment knowledgeFragment = new KnowledgeFragment();
|
||||
knowledgeFragment.setKnowledgeId(knowledgeId);
|
||||
knowledgeFragment.setDocId(docId);
|
||||
knowledgeFragment.setFid(fid);
|
||||
knowledgeFragment.setIdx(i);
|
||||
knowledgeFragment.setContent(chunkList.get(i));
|
||||
knowledgeFragment.setCreateTime(new Date());
|
||||
knowledgeFragmentList.add(knowledgeFragment);
|
||||
}
|
||||
KnowledgeInfoVo knowledgeInfoVo = knowledgeInfoService.queryById(knowledgeId);
|
||||
ChatModelVo chatModelVo = chatModelService.selectModelByName(knowledgeInfoVo.getEmbeddingModel());
|
||||
|
||||
@@ -211,7 +243,27 @@ public class KnowledgeAttachServiceImpl implements IKnowledgeAttachService {
|
||||
storeEmbeddingBo.setEmbeddingModelName(knowledgeInfoVo.getEmbeddingModel());
|
||||
storeEmbeddingBo.setApiKey(chatModelVo.getApiKey());
|
||||
storeEmbeddingBo.setBaseUrl(chatModelVo.getApiHost());
|
||||
vectorStoreService.storeEmbeddings(storeEmbeddingBo);
|
||||
try {
|
||||
vectorStoreService.storeEmbeddings(storeEmbeddingBo);
|
||||
for (KnowledgeFragment old : oldFragments) {
|
||||
if (StringUtils.isNotBlank(old.getFid())) {
|
||||
vectorStoreService.removeByFid(old.getFid(), String.valueOf(knowledgeId));
|
||||
}
|
||||
}
|
||||
} catch (Exception vectorError) {
|
||||
for (String newFid : fids) {
|
||||
try {
|
||||
vectorStoreService.removeByFid(newFid, String.valueOf(knowledgeId));
|
||||
} catch (Exception cleanupError) {
|
||||
log.error("补偿删除新向量失败, kid={}, fid={}", knowledgeId, newFid, cleanupError);
|
||||
}
|
||||
}
|
||||
throw vectorError;
|
||||
}
|
||||
|
||||
knowledgeFragmentMapper.delete(Wrappers.<KnowledgeFragment>lambdaQuery().eq(KnowledgeFragment::getDocId, docId));
|
||||
knowledgeFragmentMapper.insertBatch(knowledgeFragmentList);
|
||||
knowledgeRetrievalService.invalidateKnowledge(String.valueOf(knowledgeId));
|
||||
|
||||
attach.setStatus(KnowledgeAttachStatus.COMPLETED.getCode()); // 已完成
|
||||
baseMapper.updateById(attach);
|
||||
|
||||
@@ -40,6 +40,7 @@ public class KnowledgeFragmentServiceImpl implements IKnowledgeFragmentService {
|
||||
private final IKnowledgeInfoService knowledgeInfoService;
|
||||
private final IChatModelService chatModelService;
|
||||
private final KnowledgeRetrievalService knowledgeRetrievalService;
|
||||
private final org.ruoyi.service.vector.VectorStoreService vectorStoreService;
|
||||
|
||||
/**
|
||||
* 查询知识片段
|
||||
@@ -114,7 +115,11 @@ public class KnowledgeFragmentServiceImpl implements IKnowledgeFragmentService {
|
||||
public Boolean updateByBo(KnowledgeFragmentBo bo) {
|
||||
KnowledgeFragment update = MapstructUtils.convert(bo, KnowledgeFragment.class);
|
||||
validEntityBeforeSave(update);
|
||||
return baseMapper.updateById(update) > 0;
|
||||
boolean updated = baseMapper.updateById(update) > 0;
|
||||
if (updated && update.getKnowledgeId() != null) {
|
||||
knowledgeRetrievalService.invalidateKnowledge(String.valueOf(update.getKnowledgeId()));
|
||||
}
|
||||
return updated;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -136,6 +141,14 @@ public class KnowledgeFragmentServiceImpl implements IKnowledgeFragmentService {
|
||||
if(isValid){
|
||||
//TODO 做一些业务上的校验,判断是否需要校验
|
||||
}
|
||||
// 删除 DB 片段前,同步删除向量库中对应向量
|
||||
List<KnowledgeFragment> fragments = baseMapper.selectByIds(ids);
|
||||
for (KnowledgeFragment fragment : fragments) {
|
||||
if (StringUtils.isNotBlank(fragment.getFid()) && fragment.getKnowledgeId() != null) {
|
||||
vectorStoreService.removeByFid(fragment.getFid(), String.valueOf(fragment.getKnowledgeId()));
|
||||
knowledgeRetrievalService.invalidateKnowledge(String.valueOf(fragment.getKnowledgeId()));
|
||||
}
|
||||
}
|
||||
return baseMapper.deleteByIds(ids) > 0;
|
||||
}
|
||||
|
||||
|
||||
@@ -16,6 +16,9 @@ import org.ruoyi.mapper.knowledge.KnowledgeAttachMapper;
|
||||
import org.ruoyi.mapper.knowledge.KnowledgeInfoMapper;
|
||||
import org.ruoyi.service.knowledge.IKnowledgeInfoService;
|
||||
import org.springframework.stereotype.Service;
|
||||
import org.springframework.transaction.annotation.Transactional;
|
||||
import org.ruoyi.service.retrieval.KnowledgeRetrievalService;
|
||||
import org.ruoyi.common.core.service.OssService;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
@@ -36,6 +39,12 @@ public class KnowledgeInfoServiceImpl implements IKnowledgeInfoService {
|
||||
|
||||
private final KnowledgeAttachMapper knowledgeAttachMapper;
|
||||
|
||||
private final org.ruoyi.mapper.knowledge.KnowledgeFragmentMapper knowledgeFragmentMapper;
|
||||
|
||||
private final org.ruoyi.service.vector.VectorStoreService vectorStoreService;
|
||||
private final KnowledgeRetrievalService knowledgeRetrievalService;
|
||||
private final OssService ossService;
|
||||
|
||||
/**
|
||||
* 查询知识库
|
||||
*
|
||||
@@ -97,10 +106,14 @@ public class KnowledgeInfoServiceImpl implements IKnowledgeInfoService {
|
||||
*/
|
||||
private void fillDocumentCount(List<KnowledgeInfoVo> records) {
|
||||
if (records == null || records.isEmpty()) return;
|
||||
for (KnowledgeInfoVo vo : records) {
|
||||
int count = knowledgeAttachMapper.countByKnowledgeId(vo.getId());
|
||||
vo.setDocumentCount(count);
|
||||
List<Long> ids = records.stream().map(KnowledgeInfoVo::getId).toList();
|
||||
Map<Long, Integer> counts = new java.util.HashMap<>();
|
||||
for (Map<String, Object> row : knowledgeAttachMapper.countByKnowledgeIds(ids)) {
|
||||
Number kid = (Number) (row.get("knowledgeId") != null ? row.get("knowledgeId") : row.get("knowledgeid"));
|
||||
Number count = (Number) (row.get("documentCount") != null ? row.get("documentCount") : row.get("documentcount"));
|
||||
if (kid != null && count != null) counts.put(kid.longValue(), count.intValue());
|
||||
}
|
||||
records.forEach(vo -> vo.setDocumentCount(counts.getOrDefault(vo.getId(), 0)));
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -130,7 +143,9 @@ public class KnowledgeInfoServiceImpl implements IKnowledgeInfoService {
|
||||
public Boolean updateByBo(KnowledgeInfoBo bo) {
|
||||
KnowledgeInfo update = MapstructUtils.convert(bo, KnowledgeInfo.class);
|
||||
validEntityBeforeSave(update);
|
||||
return baseMapper.updateById(update) > 0;
|
||||
boolean updated = baseMapper.updateById(update) > 0;
|
||||
if (updated) knowledgeRetrievalService.invalidateKnowledge(String.valueOf(bo.getId()));
|
||||
return updated;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -148,10 +163,33 @@ public class KnowledgeInfoServiceImpl implements IKnowledgeInfoService {
|
||||
* @return 是否删除成功
|
||||
*/
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public Boolean deleteWithValidByIds(Collection<Long> ids, Boolean isValid) {
|
||||
if(isValid){
|
||||
//TODO 做一些业务上的校验,判断是否需要校验
|
||||
}
|
||||
for (Long kid : ids) {
|
||||
KnowledgeInfo info = baseMapper.selectById(kid);
|
||||
// 1. 删除向量库中该知识库的所有向量(按文档逐个清理,三种向量库行为一致)
|
||||
List<org.ruoyi.domain.entity.knowledge.KnowledgeAttach> attaches = knowledgeAttachMapper.selectList(
|
||||
Wrappers.lambdaQuery(org.ruoyi.domain.entity.knowledge.KnowledgeAttach.class)
|
||||
.eq(org.ruoyi.domain.entity.knowledge.KnowledgeAttach::getKnowledgeId, kid));
|
||||
vectorStoreService.removeById(String.valueOf(kid), info == null ? null : info.getVectorModel());
|
||||
List<Long> ossIds = attaches.stream()
|
||||
.map(org.ruoyi.domain.entity.knowledge.KnowledgeAttach::getOssId)
|
||||
.filter(java.util.Objects::nonNull).toList();
|
||||
if (!ossIds.isEmpty()) {
|
||||
for (Long ossId : ossIds) {
|
||||
ossService.deleteFile(ossId);
|
||||
}
|
||||
}
|
||||
// 2. 删除该知识库下的附件与片段记录
|
||||
knowledgeAttachMapper.delete(Wrappers.lambdaQuery(org.ruoyi.domain.entity.knowledge.KnowledgeAttach.class)
|
||||
.eq(org.ruoyi.domain.entity.knowledge.KnowledgeAttach::getKnowledgeId, kid));
|
||||
knowledgeFragmentMapper.delete(Wrappers.lambdaQuery(org.ruoyi.domain.entity.knowledge.KnowledgeFragment.class)
|
||||
.eq(org.ruoyi.domain.entity.knowledge.KnowledgeFragment::getKnowledgeId, kid));
|
||||
knowledgeRetrievalService.invalidateKnowledge(String.valueOf(kid));
|
||||
}
|
||||
return baseMapper.deleteByIds(ids) > 0;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,9 @@ import dev.langchain4j.store.embedding.filter.MetadataFilterBuilder;
|
||||
import dev.langchain4j.store.embedding.milvus.MilvusEmbeddingStore;
|
||||
import io.milvus.param.IndexType;
|
||||
import io.milvus.param.MetricType;
|
||||
import io.milvus.v2.client.ConnectConfig;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.service.collection.request.DropCollectionReq;
|
||||
import lombok.SneakyThrows;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.ruoyi.common.chat.domain.vo.chat.ChatModelVo;
|
||||
@@ -37,13 +40,16 @@ import java.util.stream.IntStream;
|
||||
public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
|
||||
private final KnowledgeAttachMapper knowledgeAttachMapper;
|
||||
private final org.ruoyi.mapper.knowledge.KnowledgeInfoMapper knowledgeInfoMapper;
|
||||
|
||||
public MilvusVectorStoreStrategy(VectorStoreProperties vectorStoreProperties,
|
||||
IChatModelService chatModelService,
|
||||
EmbeddingModelFactory embeddingModelFactory,
|
||||
KnowledgeAttachMapper knowledgeAttachMapper) {
|
||||
KnowledgeAttachMapper knowledgeAttachMapper,
|
||||
org.ruoyi.mapper.knowledge.KnowledgeInfoMapper knowledgeInfoMapper) {
|
||||
super(vectorStoreProperties, embeddingModelFactory, chatModelService);
|
||||
this.knowledgeAttachMapper = knowledgeAttachMapper;
|
||||
this.knowledgeInfoMapper = knowledgeInfoMapper;
|
||||
}
|
||||
|
||||
// 缓存不同集合与 autoFlush 配置的 Milvus 连接
|
||||
@@ -100,9 +106,11 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
long startTime = System.currentTimeMillis();
|
||||
|
||||
// 复用连接,写入场景使用 autoFlush=false 以提升批量插入性能
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, dimension, false);
|
||||
// addAll is already batched; flush before returning so newly parsed fragments are immediately searchable.
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, dimension, true);
|
||||
|
||||
IntStream.range(0, chunkList.size()).forEach(i -> {
|
||||
List<TextSegment> segments = new ArrayList<>(chunkList.size());
|
||||
for (int i = 0; i < chunkList.size(); i++) {
|
||||
String text = chunkList.get(i);
|
||||
String fid = fidList.get(i);
|
||||
Metadata metadata = new Metadata();
|
||||
@@ -110,13 +118,18 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
metadata.put("kid", kid);
|
||||
metadata.put("docId", docId);
|
||||
|
||||
TextSegment textSegment = TextSegment.from(text, metadata);
|
||||
Embedding embedding = embeddingModel.embed(text).content();
|
||||
segments.add(TextSegment.from(text, metadata));
|
||||
}
|
||||
List<Embedding> embeddings = embeddingModel.embedAll(segments).content();
|
||||
if (embeddings.size() != segments.size()) {
|
||||
throw new org.ruoyi.common.core.exception.ServiceException("Embedding 返回数量与分片数量不一致");
|
||||
}
|
||||
for (Embedding embedding : embeddings) {
|
||||
// 单位化处理
|
||||
float[] vector = embedding.vector();
|
||||
normalize(vector);
|
||||
embeddingStore.add(Embedding.from(vector), textSegment);
|
||||
});
|
||||
}
|
||||
embeddingStore.addAll(embeddings, segments);
|
||||
long endTime = System.currentTimeMillis();
|
||||
log.info("Milvus向量存储完成消耗时间:{}秒", (endTime - startTime) / 1000);
|
||||
}
|
||||
@@ -174,6 +187,7 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
if (segment == null) continue;
|
||||
|
||||
String docId = segment.metadata().getString("docId");
|
||||
String fid = segment.metadata().getString("fid");
|
||||
String sourceName = "未知来源";
|
||||
if (docId != null) {
|
||||
KnowledgeAttach attach = knowledgeAttachMapper.selectOne(new LambdaQueryWrapper<KnowledgeAttach>()
|
||||
@@ -188,6 +202,8 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
double score = match.score();
|
||||
|
||||
resultList.add(org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo.builder()
|
||||
.id(fid)
|
||||
.docId(docId)
|
||||
.content(segment.text())
|
||||
.score(score)
|
||||
.sourceName(sourceName)
|
||||
@@ -200,16 +216,40 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
@SneakyThrows
|
||||
public void removeById(String id, String modelName) {
|
||||
// 注意:此处原逻辑使用 collectionname + id,保持现状
|
||||
int dimension = getModelDimension(modelName);
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(vectorStoreProperties.getMilvus().getCollectionname() + id, dimension, false);
|
||||
embeddingStore.remove(id);
|
||||
String collectionName = vectorStoreProperties.getMilvus().getCollectionname() + id;
|
||||
MilvusClientV2 client = new MilvusClientV2(ConnectConfig.builder()
|
||||
.uri(vectorStoreProperties.getMilvus().getUrl()).build());
|
||||
try {
|
||||
client.dropCollection(DropCollectionReq.builder().collectionName(collectionName).build());
|
||||
storeCache.keySet().removeIf(key -> key.startsWith(collectionName));
|
||||
log.info("Milvus collection deleted: {}", collectionName);
|
||||
} catch (Exception e) {
|
||||
log.error("Milvus collection delete failed: {}", collectionName, e);
|
||||
throw new org.ruoyi.common.core.exception.ServiceException("Milvus向量集合删除失败");
|
||||
} finally {
|
||||
client.close();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 根据知识库ID解析其 embedding 模型维度,失败时回退默认 1024
|
||||
*/
|
||||
private int getDimensionByKid(String kid) {
|
||||
try {
|
||||
org.ruoyi.domain.entity.knowledge.KnowledgeInfo info = knowledgeInfoMapper.selectById(Long.parseLong(kid));
|
||||
if (info != null && info.getEmbeddingModel() != null) {
|
||||
return getModelDimension(info.getEmbeddingModel());
|
||||
}
|
||||
} catch (Exception e) {
|
||||
log.warn("根据 kid={} 解析向量维度失败,使用默认 1024: {}", kid, e.getMessage());
|
||||
}
|
||||
return 1024;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void removeByDocId(String docId, String kid) {
|
||||
String collectionName = vectorStoreProperties.getMilvus().getCollectionname() + kid;
|
||||
// 使用默认维度,因为删除操作不需要精确的维度信息
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, 1024, false);
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, getDimensionByKid(kid), true);
|
||||
Filter filter = MetadataFilterBuilder.metadataKey("docId").isEqualTo(docId);
|
||||
embeddingStore.removeAll(filter);
|
||||
log.info("Milvus成功删除 docId={} 的所有向量数据", docId);
|
||||
@@ -218,8 +258,7 @@ public class MilvusVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
@Override
|
||||
public void removeByFid(String fid, String kid) {
|
||||
String collectionName = vectorStoreProperties.getMilvus().getCollectionname() + kid;
|
||||
// 使用默认维度,因为删除操作不需要精确的维度信息
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, 1024, false);
|
||||
EmbeddingStore<TextSegment> embeddingStore = getMilvusStore(collectionName, getDimensionByKid(kid), true);
|
||||
Filter filter = MetadataFilterBuilder.metadataKey("fid").isEqualTo(fid);
|
||||
embeddingStore.removeAll(filter);
|
||||
log.info("Milvus成功删除 fid={} 的所有向量数据", fid);
|
||||
|
||||
@@ -129,20 +129,26 @@ public class QdrantVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
log.info("Qdrant向量存储条数记录: {}", chunkList.size());
|
||||
long startTime = System.currentTimeMillis();
|
||||
|
||||
IntStream.range(0, chunkList.size()).forEach(i -> {
|
||||
List<TextSegment> segments = new ArrayList<>(chunkList.size());
|
||||
for (int i = 0; i < chunkList.size(); i++) {
|
||||
String text = chunkList.get(i);
|
||||
String fid = fidList.get(i);
|
||||
Metadata metadata = new Metadata();
|
||||
metadata.put(METADATA_FID_KEY, fid);
|
||||
metadata.put(METADATA_KID_KEY, kid);
|
||||
metadata.put(METADATA_DOC_ID_KEY, docId);
|
||||
TextSegment textSegment = TextSegment.from(text, metadata);
|
||||
Embedding embedding = embeddingModel.embed(text).content();
|
||||
segments.add(TextSegment.from(text, metadata));
|
||||
}
|
||||
List<Embedding> embeddings = embeddingModel.embedAll(segments).content();
|
||||
if (embeddings.size() != segments.size()) {
|
||||
throw new ServiceException("Embedding 返回数量与分片数量不一致");
|
||||
}
|
||||
for (Embedding embedding : embeddings) {
|
||||
// 单位化处理
|
||||
float[] vector = embedding.vector();
|
||||
normalize(vector);
|
||||
embeddingStore.add(Embedding.from(vector), textSegment);
|
||||
});
|
||||
}
|
||||
embeddingStore.addAll(embeddings, segments);
|
||||
|
||||
long endTime = System.currentTimeMillis();
|
||||
log.info("Qdrant向量存储完成消耗时间:{}秒", (endTime - startTime) / 1000);
|
||||
@@ -228,6 +234,12 @@ public class QdrantVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
docId = docIdValue.getStringValue();
|
||||
}
|
||||
|
||||
String fid = null;
|
||||
JsonWithInt.Value fidValue = point.getPayloadMap().get(METADATA_FID_KEY);
|
||||
if (fidValue != null && fidValue.hasStringValue()) {
|
||||
fid = fidValue.getStringValue();
|
||||
}
|
||||
|
||||
String sourceName = "未知来源";
|
||||
if (docId != null) {
|
||||
KnowledgeAttach attach = knowledgeAttachMapper.selectOne(new LambdaQueryWrapper<KnowledgeAttach>()
|
||||
@@ -239,6 +251,8 @@ public class QdrantVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
}
|
||||
|
||||
resultList.add(org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo.builder()
|
||||
.id(fid)
|
||||
.docId(docId)
|
||||
.content(content)
|
||||
.score((double) point.getScore())
|
||||
.sourceName(sourceName)
|
||||
|
||||
@@ -6,6 +6,8 @@ import org.ruoyi.domain.bo.vector.QueryVectorBo;
|
||||
import org.ruoyi.domain.bo.vector.StoreEmbeddingBo;
|
||||
import org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo;
|
||||
import org.ruoyi.factory.VectorStoreStrategyFactory;
|
||||
import org.ruoyi.mapper.knowledge.KnowledgeInfoMapper;
|
||||
import org.ruoyi.domain.entity.knowledge.KnowledgeInfo;
|
||||
import org.ruoyi.service.vector.VectorStoreService;
|
||||
import org.springframework.context.annotation.Primary;
|
||||
import org.springframework.stereotype.Service;
|
||||
@@ -24,6 +26,7 @@ import java.util.List;
|
||||
public class VectorStoreServiceImpl implements VectorStoreService {
|
||||
|
||||
private final VectorStoreStrategyFactory strategyFactory;
|
||||
private final KnowledgeInfoMapper knowledgeInfoMapper;
|
||||
|
||||
|
||||
/**
|
||||
@@ -33,9 +36,23 @@ public class VectorStoreServiceImpl implements VectorStoreService {
|
||||
return strategyFactory.getStrategy();
|
||||
}
|
||||
|
||||
private VectorStoreService getStrategy(String type) {
|
||||
return strategyFactory.getStrategy(type);
|
||||
}
|
||||
|
||||
private String vectorTypeForKid(String kid) {
|
||||
try {
|
||||
KnowledgeInfo info = knowledgeInfoMapper.selectById(Long.parseLong(kid));
|
||||
return info == null ? null : info.getVectorModel();
|
||||
} catch (Exception e) {
|
||||
log.warn("无法解析知识库向量类型, kid={}", kid, e);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void createSchema(String kid, String modelName) {
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(vectorTypeForKid(kid));
|
||||
strategy.createSchema(kid, modelName);
|
||||
}
|
||||
|
||||
@@ -43,7 +60,7 @@ public class VectorStoreServiceImpl implements VectorStoreService {
|
||||
public void storeEmbeddings(StoreEmbeddingBo storeEmbeddingBo) {
|
||||
log.info("存储向量数据: kid={}, docId={}, 数据条数={}",
|
||||
storeEmbeddingBo.getKid(), storeEmbeddingBo.getDocId(), storeEmbeddingBo.getChunkList().size());
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(storeEmbeddingBo.getVectorStoreName());
|
||||
strategy.storeEmbeddings(storeEmbeddingBo);
|
||||
}
|
||||
|
||||
@@ -51,35 +68,35 @@ public class VectorStoreServiceImpl implements VectorStoreService {
|
||||
public List<String> getQueryVector(QueryVectorBo queryVectorBo) {
|
||||
log.info("查询向量数据: kid={}, query={}, maxResults={}",
|
||||
queryVectorBo.getKid(), queryVectorBo.getQuery(), queryVectorBo.getMaxResults());
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(queryVectorBo.getVectorModelName());
|
||||
return strategy.getQueryVector(queryVectorBo);
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<KnowledgeRetrievalVo> search(QueryVectorBo queryVectorBo) {
|
||||
log.info("执行测试搜索: kid={}, query={}", queryVectorBo.getKid(), queryVectorBo.getQuery());
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(queryVectorBo.getVectorModelName());
|
||||
return strategy.search(queryVectorBo);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void removeById(String id, String modelName) {
|
||||
log.info("根据ID删除向量数据: id={}, modelName={}", id, modelName);
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(modelName);
|
||||
strategy.removeById(id, modelName);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void removeByDocId(String docId, String kid) {
|
||||
log.info("根据docId删除向量数据: docId={}, kid={}", docId, kid);
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(vectorTypeForKid(kid));
|
||||
strategy.removeByDocId(docId, kid);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void removeByFid(String fid, String kid) {
|
||||
log.info("根据fid删除向量数据: fid={}, kid={}", fid, kid);
|
||||
VectorStoreService strategy = getCurrentStrategy();
|
||||
VectorStoreService strategy = getStrategy(vectorTypeForKid(kid));
|
||||
strategy.removeByFid(fid, kid);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package org.ruoyi.service.vector.impl;
|
||||
import cn.hutool.json.JSONObject;
|
||||
import dev.langchain4j.data.embedding.Embedding;
|
||||
import dev.langchain4j.model.embedding.EmbeddingModel;
|
||||
import dev.langchain4j.data.segment.TextSegment;
|
||||
|
||||
import io.weaviate.client.WeaviateClient;
|
||||
import lombok.SneakyThrows;
|
||||
@@ -18,6 +19,8 @@ import org.springframework.stereotype.Component;
|
||||
import io.weaviate.client.Config;
|
||||
import io.weaviate.client.base.Result;
|
||||
import io.weaviate.client.v1.batch.api.ObjectsBatchDeleter;
|
||||
import io.weaviate.client.v1.batch.api.ObjectsBatcher;
|
||||
import io.weaviate.client.v1.data.model.WeaviateObject;
|
||||
import io.weaviate.client.v1.batch.model.BatchDeleteResponse;
|
||||
import io.weaviate.client.v1.filters.Operator;
|
||||
import io.weaviate.client.v1.filters.WhereFilter;
|
||||
@@ -43,8 +46,12 @@ import java.util.Map;
|
||||
@Component
|
||||
public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
|
||||
private WeaviateClient client;
|
||||
private volatile WeaviateClient client;
|
||||
private final KnowledgeAttachMapper knowledgeAttachMapper;
|
||||
/**
|
||||
* 已确认存在的 class 缓存,避免每次检索都全量拉取 schema
|
||||
*/
|
||||
private final java.util.Set<String> knownClasses = java.util.concurrent.ConcurrentHashMap.newKeySet();
|
||||
|
||||
public WeaviateVectorStoreStrategy(VectorStoreProperties vectorStoreProperties,
|
||||
IChatModelService chatModelService,
|
||||
@@ -54,6 +61,22 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
this.knowledgeAttachMapper = knowledgeAttachMapper;
|
||||
}
|
||||
|
||||
/**
|
||||
* 懒加载单例客户端,避免 remove 等方法在未调用 createSchema 时 NPE
|
||||
*/
|
||||
private WeaviateClient getClient() {
|
||||
if (client == null) {
|
||||
synchronized (this) {
|
||||
if (client == null) {
|
||||
String protocol = vectorStoreProperties.getWeaviate().getProtocol();
|
||||
String host = vectorStoreProperties.getWeaviate().getHost();
|
||||
client = new WeaviateClient(new Config(protocol, host));
|
||||
}
|
||||
}
|
||||
}
|
||||
return client;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getVectorStoreType() {
|
||||
return "weaviate";
|
||||
@@ -61,13 +84,12 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
|
||||
@Override
|
||||
public void createSchema(String kid, String embeddingModelName) {
|
||||
String protocol = vectorStoreProperties.getWeaviate().getProtocol();
|
||||
String host = vectorStoreProperties.getWeaviate().getHost();
|
||||
String className = vectorStoreProperties.getWeaviate().getClassname() + kid;
|
||||
// 创建 Weaviate 客户端
|
||||
client = new WeaviateClient(new Config(protocol, host));
|
||||
if (knownClasses.contains(className)) {
|
||||
return;
|
||||
}
|
||||
// 检查类是否存在,如果不存在就创建 schema
|
||||
Result<Schema> schemaResult = client.schema().getter().run();
|
||||
Result<Schema> schemaResult = getClient().schema().getter().run();
|
||||
Schema schema = schemaResult.getResult();
|
||||
boolean classExists = false;
|
||||
for (WeaviateClass weaviateClass : schema.getClasses()) {
|
||||
@@ -88,13 +110,15 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
Property.builder().name("docId").dataType(Collections.singletonList("text")).build())
|
||||
)
|
||||
.build();
|
||||
Result<Boolean> createResult = client.schema().classCreator().withClass(build).run();
|
||||
Result<Boolean> createResult = getClient().schema().classCreator().withClass(build).run();
|
||||
if (createResult.hasErrors()) {
|
||||
log.error("Schema 创建失败: {}", createResult.getError());
|
||||
throw new ServiceException("Weaviate Schema 创建失败: " + createResult.getError());
|
||||
} else {
|
||||
log.info("Schema 创建成功: {}", className);
|
||||
}
|
||||
}
|
||||
knownClasses.add(className);
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -107,10 +131,16 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
String docId = storeEmbeddingBo.getDocId();
|
||||
log.info("向量存储条数记录: {}", chunkList.size());
|
||||
long startTime = System.currentTimeMillis();
|
||||
List<TextSegment> segments = chunkList.stream().map(TextSegment::from).toList();
|
||||
List<Embedding> embeddings = embeddingModel.embedAll(segments).content();
|
||||
if (embeddings.size() != chunkList.size()) {
|
||||
throw new ServiceException("Embedding 返回数量与分片数量不一致");
|
||||
}
|
||||
ObjectsBatcher batcher = getClient().batch().objectsBatcher();
|
||||
for (int i = 0; i < chunkList.size(); i++) {
|
||||
String text = chunkList.get(i);
|
||||
String fid = fidList.get(i);
|
||||
Embedding embedding = embeddingModel.embed(text).content();
|
||||
Embedding embedding = embeddings.get(i);
|
||||
Map<String, Object> properties = Map.of(
|
||||
"text", text,
|
||||
"fid", fid,
|
||||
@@ -121,11 +151,13 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
normalize(vectorArray);
|
||||
Float[] vector = toObjectArray(vectorArray);
|
||||
|
||||
client.data().creator()
|
||||
.withClassName(vectorStoreProperties.getWeaviate().getClassname() + kid)
|
||||
.withProperties(properties)
|
||||
.withVector(vector)
|
||||
.run();
|
||||
batcher.withObject(WeaviateObject.builder()
|
||||
.className(vectorStoreProperties.getWeaviate().getClassname() + kid)
|
||||
.properties(properties).vector(vector).build());
|
||||
}
|
||||
Result<?> batchResult = batcher.run();
|
||||
if (batchResult.hasErrors()) {
|
||||
throw new ServiceException("Weaviate 批量写入失败: " + batchResult.getError());
|
||||
}
|
||||
long endTime = System.currentTimeMillis();
|
||||
log.info("向量存储完成消耗时间:" + (endTime - startTime) / 1000 + "秒");
|
||||
@@ -169,7 +201,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
queryVectorBo.getMaxResults()
|
||||
);
|
||||
|
||||
Result<GraphQLResponse> result = client.graphQL().raw().withQuery(graphQLQuery).run();
|
||||
Result<GraphQLResponse> result = getClient().graphQL().raw().withQuery(graphQLQuery).run();
|
||||
List<String> resultList = new ArrayList<>();
|
||||
if (result != null && !result.hasErrors()) {
|
||||
Object data = result.getResult().getData();
|
||||
@@ -211,6 +243,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
" Get {\n" +
|
||||
" %s(nearVector: {vector: [%s]} limit: %d) {\n" +
|
||||
" text\n" +
|
||||
" fid\n" +
|
||||
" docId\n" +
|
||||
" _additional {\n" +
|
||||
" distance\n" +
|
||||
@@ -223,7 +256,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
queryVectorBo.getMaxResults()
|
||||
);
|
||||
|
||||
Result<GraphQLResponse> result = client.graphQL().raw().withQuery(graphQLQuery).run();
|
||||
Result<GraphQLResponse> result = getClient().graphQL().raw().withQuery(graphQLQuery).run();
|
||||
List<org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo> resultList = new ArrayList<>();
|
||||
|
||||
if (result != null && !result.hasErrors()) {
|
||||
@@ -236,6 +269,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
Map<String, Object> map = (Map<String, Object>) obj;
|
||||
String content = (String) map.get("text");
|
||||
String docId = (String) map.get("docId");
|
||||
String fid = (String) map.get("fid");
|
||||
|
||||
Map<String, Object> additional = (Map<String, Object>) map.get("_additional");
|
||||
Double distance = Double.valueOf(String.valueOf(additional.get("distance")));
|
||||
@@ -253,6 +287,8 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
}
|
||||
|
||||
resultList.add(org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo.builder()
|
||||
.id(fid)
|
||||
.docId(docId)
|
||||
.content(content)
|
||||
.score(score)
|
||||
.sourceName(sourceName)
|
||||
@@ -265,12 +301,10 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
@Override
|
||||
@SneakyThrows
|
||||
public void removeById(String id, String modelName) {
|
||||
String protocol = vectorStoreProperties.getWeaviate().getProtocol();
|
||||
String host = vectorStoreProperties.getWeaviate().getHost();
|
||||
String className = vectorStoreProperties.getWeaviate().getClassname();
|
||||
String finalClassName = className + id;
|
||||
WeaviateClient client = new WeaviateClient(new Config(protocol, host));
|
||||
Result<Boolean> result = client.schema().classDeleter().withClassName(finalClassName).run();
|
||||
Result<Boolean> result = getClient().schema().classDeleter().withClassName(finalClassName).run();
|
||||
knownClasses.remove(finalClassName);
|
||||
if (result.hasErrors()) {
|
||||
log.error("失败删除向量: " + result.getError());
|
||||
throw new ServiceException("失败删除向量数据!");
|
||||
@@ -288,7 +322,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
.operator(Operator.Equal)
|
||||
.valueText(docId)
|
||||
.build();
|
||||
ObjectsBatchDeleter deleter = client.batch().objectsBatchDeleter();
|
||||
ObjectsBatchDeleter deleter = getClient().batch().objectsBatchDeleter();
|
||||
Result<BatchDeleteResponse> result = deleter.withClassName(className)
|
||||
.withWhere(whereFilter)
|
||||
.run();
|
||||
@@ -308,7 +342,7 @@ public class WeaviateVectorStoreStrategy extends AbstractVectorStoreStrategy {
|
||||
.operator(Operator.Equal)
|
||||
.valueText(fid)
|
||||
.build();
|
||||
ObjectsBatchDeleter deleter = client.batch().objectsBatchDeleter();
|
||||
ObjectsBatchDeleter deleter = getClient().batch().objectsBatchDeleter();
|
||||
Result<BatchDeleteResponse> result = deleter.withClassName(className)
|
||||
.withWhere(whereFilter)
|
||||
.run();
|
||||
|
||||
Reference in New Issue
Block a user