fix(rag): 完善向量路由批量写入与数据生命周期

This commit is contained in:
evo
2026-07-21 09:35:18 +08:00
parent c4fc1e6fcf
commit 42bc8fea95
17 changed files with 373 additions and 78 deletions

View File

@@ -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();

View File

@@ -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));
}

View File

@@ -37,6 +37,9 @@ public class KnowledgeAttach extends BaseEntity {
*/
private String docId;
/** SHA-256 content digest used for upload idempotency. */
private String fileHash;
/**
* 附件名称
*/

View File

@@ -27,6 +27,11 @@ public class KnowledgeFragment extends BaseEntity {
@TableId(value = "id")
private Long id;
/**
* 向量库片段ID与向量库中的 fid 元数据对应,用于向量定位与混合检索融合)
*/
private String fid;
/**
* 文档ID-用于关联文本块信息
*/

View File

@@ -30,6 +30,11 @@ public class KnowledgeFragmentVo implements Serializable {
@ExcelProperty(value = "主键")
private Long id;
/**
* 向量库片段ID
*/
private String fid;
/**
* 文档ID-用于关联文本块信息
*/

View File

@@ -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;

View File

@@ -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);
}

View File

@@ -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) " +

View File

@@ -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);

View File

@@ -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;
}

View File

@@ -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;
}
}

View File

@@ -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);

View File

@@ -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)

View File

@@ -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);
}
}

View File

@@ -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();