package cn.iocoder.yudao.module.ai.service.knowledge;
|
|
import cn.hutool.core.collection.CollUtil;
|
import cn.hutool.core.map.MapUtil;
|
import cn.hutool.core.util.StrUtil;
|
import cn.iocoder.yudao.framework.common.enums.CommonStatusEnum;
|
import cn.iocoder.yudao.framework.common.pojo.PageResult;
|
import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.*;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDO;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeSegmentDO;
|
import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeDocumentMapper;
|
import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeSegmentMapper;
|
import cn.iocoder.yudao.module.ai.service.knowledge.bo.AiKnowledgeSegmentSearchReqBO;
|
import cn.iocoder.yudao.module.ai.service.knowledge.bo.AiKnowledgeSegmentSearchRespBO;
|
import cn.iocoder.yudao.module.ai.service.model.AiModelService;
|
import jakarta.annotation.Resource;
|
import lombok.extern.slf4j.Slf4j;
|
import org.springframework.ai.document.Document;
|
import org.springframework.ai.vectorstore.SearchRequest;
|
import org.springframework.ai.vectorstore.VectorStore;
|
import org.springframework.ai.vectorstore.filter.Filter;
|
import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder;
|
import org.springframework.stereotype.Service;
|
import org.springframework.transaction.annotation.Transactional;
|
|
import java.util.*;
|
import java.util.stream.Collectors;
|
|
import static cn.iocoder.yudao.framework.common.exception.util.ServiceExceptionUtil.exception;
|
import static cn.iocoder.yudao.module.ai.enums.ErrorCodeConstants.*;
|
|
@Slf4j
|
@Service
|
public class AiKnowledgeSegmentServiceImpl implements AiKnowledgeSegmentService {
|
|
private static final String METADATA_KNOWLEDGE_ID = "knowledgeId";
|
private static final String METADATA_DOCUMENT_ID = "documentId";
|
private static final String METADATA_SEGMENT_ID = "segmentId";
|
|
@Resource
|
private AiKnowledgeSegmentMapper segmentMapper;
|
@Resource
|
private AiKnowledgeDocumentMapper documentMapper;
|
@Resource
|
private AiKnowledgeService knowledgeService;
|
@Resource
|
private AiModelService modelService;
|
|
@Override
|
public AiKnowledgeSegmentDO getSegment(Long id) {
|
return segmentMapper.selectById(id);
|
}
|
|
@Override
|
public AiKnowledgeSegmentDO validateSegmentExists(Long id) {
|
AiKnowledgeSegmentDO segment = segmentMapper.selectById(id);
|
if (segment == null) throw exception(KNOWLEDGE_SEGMENT_NOT_EXISTS);
|
return segment;
|
}
|
|
@Override
|
public PageResult<AiKnowledgeSegmentDO> getSegmentPage(AiKnowledgeSegmentPageReqVO pageReqVO) {
|
return segmentMapper.selectPage(pageReqVO);
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public Long createSegment(AiKnowledgeSegmentSaveReqVO saveReqVO) {
|
if (saveReqVO.getKnowledgeId() == null) throw exception(KNOWLEDGE_NOT_EXISTS);
|
if (saveReqVO.getDocumentId() == null) throw exception(KNOWLEDGE_DOCUMENT_NOT_EXISTS);
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(saveReqVO.getKnowledgeId());
|
AiKnowledgeSegmentDO segment = new AiKnowledgeSegmentDO();
|
segment.setKnowledgeId(saveReqVO.getKnowledgeId());
|
segment.setDocumentId(saveReqVO.getDocumentId());
|
segment.setContent(saveReqVO.getContent());
|
segment.setContentLength(saveReqVO.getContent().length());
|
segment.setTokens(estimateTokens(saveReqVO.getContent()));
|
segment.setStatus(saveReqVO.getStatus() != null ? saveReqVO.getStatus()
|
: CommonStatusEnum.ENABLE.getStatus());
|
segment.setRetrievalCount(0);
|
segment.setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY);
|
segmentMapper.insert(segment);
|
|
// 向量化
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
Map<String, Object> metadata = new HashMap<>();
|
metadata.put(METADATA_KNOWLEDGE_ID, knowledge.getId());
|
metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId());
|
metadata.put(METADATA_SEGMENT_ID, segment.getId());
|
Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata);
|
try {
|
vectorStore.add(Collections.singletonList(doc));
|
segment.setVectorId(segment.getId().toString());
|
segmentMapper.updateById(segment);
|
} catch (Exception e) {
|
log.error("手动创建分段向量化失败,segmentId={}", segment.getId(), e);
|
segment.setStatus(CommonStatusEnum.DISABLE.getStatus());
|
segmentMapper.updateById(segment);
|
}
|
return segment.getId();
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void updateSegment(AiKnowledgeSegmentSaveReqVO saveReqVO) {
|
AiKnowledgeSegmentDO segment = validateSegmentExists(saveReqVO.getId());
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(segment.getKnowledgeId());
|
// 删除旧向量
|
if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) {
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
try {
|
vectorStore.delete(Collections.singletonList(segment.getVectorId()));
|
} catch (Exception e) {
|
log.warn("删除旧向量失败: {}", segment.getVectorId(), e);
|
}
|
}
|
segment.setContent(saveReqVO.getContent());
|
segment.setContentLength(saveReqVO.getContent().length());
|
segment.setTokens(estimateTokens(saveReqVO.getContent()));
|
if (saveReqVO.getStatus() != null) segment.setStatus(saveReqVO.getStatus());
|
segmentMapper.updateById(segment);
|
|
// 重新向量化
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
Map<String, Object> metadata = new HashMap<>();
|
metadata.put(METADATA_KNOWLEDGE_ID, knowledge.getId());
|
metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId());
|
metadata.put(METADATA_SEGMENT_ID, segment.getId());
|
Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata);
|
try {
|
vectorStore.add(Collections.singletonList(doc));
|
segment.setVectorId(segment.getId().toString());
|
segmentMapper.updateById(segment);
|
} catch (Exception e) {
|
log.error("更新分段向量化失败,segmentId={}", segment.getId(), e);
|
segment.setStatus(CommonStatusEnum.DISABLE.getStatus());
|
segmentMapper.updateById(segment);
|
}
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void deleteSegment(Long id) {
|
AiKnowledgeSegmentDO segment = validateSegmentExists(id);
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(segment.getKnowledgeId());
|
if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) {
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
try {
|
vectorStore.delete(Collections.singletonList(segment.getVectorId()));
|
} catch (Exception e) {
|
log.warn("删除向量失败: {}", segment.getVectorId(), e);
|
}
|
}
|
segmentMapper.deleteById(id);
|
}
|
|
@Override
|
public void updateSegmentStatus(AiKnowledgeSegmentUpdateStatusReqVO updateStatusReqVO) {
|
AiKnowledgeSegmentDO segment = validateSegmentExists(updateStatusReqVO.getId());
|
segment.setStatus(updateStatusReqVO.getStatus());
|
segmentMapper.updateById(segment);
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void deleteSegmentsByDocumentId(Long documentId) {
|
List<AiKnowledgeSegmentDO> segments = segmentMapper.selectListByDocumentId(documentId);
|
if (CollUtil.isEmpty(segments)) return;
|
|
AiKnowledgeDO knowledge = knowledgeService.getKnowledge(segments.get(0).getKnowledgeId());
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
|
for (AiKnowledgeSegmentDO segment : segments) {
|
if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) {
|
try {
|
vectorStore.delete(Collections.singletonList(segment.getVectorId()));
|
} catch (Exception e) {
|
log.warn("删除向量[{}]失败: {}", segment.getVectorId(), e.getMessage());
|
}
|
}
|
}
|
segmentMapper.delete(new com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper<AiKnowledgeSegmentDO>()
|
.eq(AiKnowledgeSegmentDO::getDocumentId, documentId));
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void saveSegments(List<AiKnowledgeSegmentDO> segments, Long knowledgeId, Long embeddingModelId) {
|
if (CollUtil.isEmpty(segments)) return;
|
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(knowledgeId);
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(embeddingModelId, buildMetadataFields());
|
|
// 批量插入分段记录
|
for (AiKnowledgeSegmentDO segment : segments) {
|
segmentMapper.insert(segment);
|
}
|
|
// 转换为 Spring AI Document 并写入向量库
|
List<Document> documents = new ArrayList<>();
|
for (AiKnowledgeSegmentDO segment : segments) {
|
Map<String, Object> metadata = new HashMap<>();
|
metadata.put(METADATA_KNOWLEDGE_ID, knowledgeId);
|
metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId());
|
metadata.put(METADATA_SEGMENT_ID, segment.getId());
|
Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata);
|
documents.add(doc);
|
}
|
|
try {
|
vectorStore.add(documents);
|
// 更新 vectorId(Milvus 返回的 ID 就是传入的 docId)
|
for (AiKnowledgeSegmentDO segment : segments) {
|
segment.setVectorId(segment.getId().toString());
|
segmentMapper.updateById(segment);
|
}
|
log.info("批量写入向量库成功,知识库: {}, 分段数: {}", knowledgeId, segments.size());
|
} catch (Exception e) {
|
log.error("写入向量库失败,知识库: {}", knowledgeId, e);
|
for (AiKnowledgeSegmentDO segment : segments) {
|
segment.setStatus(CommonStatusEnum.DISABLE.getStatus());
|
segmentMapper.updateById(segment);
|
}
|
throw new RuntimeException("向量库写入失败: " + e.getMessage(), e);
|
}
|
}
|
|
@Override
|
public List<AiKnowledgeSegmentSearchRespBO> searchSegments(AiKnowledgeSegmentSearchReqBO searchReqBO) {
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(searchReqBO.getKnowledgeId());
|
VectorStore vectorStore = modelService.getOrCreateVectorStore(
|
knowledge.getEmbeddingModelId(), buildMetadataFields());
|
|
int topK = searchReqBO.getTopK() != null ? searchReqBO.getTopK() : knowledge.getTopK();
|
double similarityThreshold = searchReqBO.getSimilarityThreshold() != null
|
? searchReqBO.getSimilarityThreshold() : knowledge.getSimilarityThreshold();
|
|
Filter.Expression filterExpression = new FilterExpressionBuilder()
|
.eq(METADATA_KNOWLEDGE_ID, searchReqBO.getKnowledgeId())
|
.build();
|
|
SearchRequest request = SearchRequest.builder()
|
.query(searchReqBO.getContent())
|
.topK(topK)
|
.similarityThreshold(similarityThreshold)
|
.filterExpression(filterExpression)
|
.build();
|
|
log.info("向量检索: query={}, topK={}, threshold={}, filter=kid={}",
|
searchReqBO.getContent(), topK, similarityThreshold, searchReqBO.getKnowledgeId());
|
List<Document> results = vectorStore.similaritySearch(request);
|
log.info("向量检索结果数: {}", results != null ? results.size() : 0);
|
if (CollUtil.isEmpty(results)) return Collections.emptyList();
|
|
// 更新检索次数
|
List<String> vectorIds = results.stream()
|
.map(Document::getId)
|
.filter(StrUtil::isNotEmpty)
|
.collect(Collectors.toList());
|
if (CollUtil.isNotEmpty(vectorIds)) {
|
List<AiKnowledgeSegmentDO> hitSegments = segmentMapper.selectListByVectorIds(vectorIds);
|
if (CollUtil.isNotEmpty(hitSegments)) {
|
List<Long> hitIds = hitSegments.stream().map(AiKnowledgeSegmentDO::getId).collect(Collectors.toList());
|
segmentMapper.updateRetrievalCountIncrByIds(hitIds);
|
Set<Long> docIds = hitSegments.stream().map(AiKnowledgeSegmentDO::getDocumentId).collect(Collectors.toSet());
|
documentMapper.updateRetrievalCountIncr(docIds);
|
}
|
}
|
|
// 收集文档 ID 并批量查询文档名称
|
Set<Long> resultDocIds = results.stream()
|
.map(doc -> convertToLong(doc.getMetadata() != null
|
? doc.getMetadata().get(METADATA_DOCUMENT_ID) : null))
|
.filter(id -> id != null)
|
.collect(Collectors.toSet());
|
final Map<Long, String> docNameMap;
|
if (CollUtil.isNotEmpty(resultDocIds)) {
|
docNameMap = documentMapper.selectBatchIds(resultDocIds).stream()
|
.collect(Collectors.toMap(
|
cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO::getId,
|
cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO::getName,
|
(a, b) -> a));
|
} else {
|
docNameMap = Collections.emptyMap();
|
}
|
|
return results.stream()
|
.map(doc -> {
|
AiKnowledgeSegmentSearchRespBO bo = new AiKnowledgeSegmentSearchRespBO();
|
bo.setContent(doc.getText());
|
bo.setScore(doc.getScore() != null ? doc.getScore() : 0.0);
|
if (doc.getMetadata() != null) {
|
bo.setId(convertToLong(doc.getMetadata().get(METADATA_SEGMENT_ID)));
|
bo.setDocumentId(convertToLong(doc.getMetadata().get(METADATA_DOCUMENT_ID)));
|
bo.setDocumentName(docNameMap.getOrDefault(bo.getDocumentId(), "未知文档"));
|
bo.setKnowledgeId(convertToLong(doc.getMetadata().get(METADATA_KNOWLEDGE_ID)));
|
log.debug("检索结果元数据: id={}, docId={}, kid={}, metadataKeys={}",
|
bo.getId(), bo.getDocumentId(), bo.getKnowledgeId(),
|
doc.getMetadata().keySet());
|
}
|
if (bo.getContent() != null) {
|
bo.setContentLength(bo.getContent().length());
|
bo.setTokens(estimateTokens(bo.getContent()));
|
}
|
return bo;
|
})
|
.collect(Collectors.toList());
|
}
|
|
private Map<String, Class<?>> buildMetadataFields() {
|
return MapUtil.<String, Class<?>>builder()
|
.put(METADATA_KNOWLEDGE_ID, Long.class)
|
.put(METADATA_DOCUMENT_ID, Long.class)
|
.put(METADATA_SEGMENT_ID, Long.class)
|
.build();
|
}
|
|
/**
|
* 将 Milvus 返回的元数据值转为 Long。
|
* Gson 反序列化 JSON 数字默认为 Double,需要兼容多种类型。
|
*/
|
private Long convertToLong(Object value) {
|
if (value == null) return null;
|
if (value instanceof Long v) return v;
|
if (value instanceof Integer v) return v.longValue();
|
if (value instanceof Double v) return v.longValue();
|
if (value instanceof Float v) return v.longValue();
|
if (value instanceof Number v) return v.longValue();
|
if (value instanceof String v) {
|
try { return Long.valueOf(v); } catch (NumberFormatException e) { return null; }
|
}
|
log.warn("无法转换元数据值类型: {} = {}", value.getClass().getName(), value);
|
return null;
|
}
|
|
private Integer estimateTokens(String text) {
|
if (StrUtil.isEmpty(text)) return 0;
|
int chineseChars = 0, englishWords = 0;
|
for (char c : text.toCharArray()) {
|
if (c >= 0x4E00 && c <= 0x9FA5) chineseChars++;
|
}
|
for (String word : text.split("\\s+")) {
|
if (word.matches(".*[a-zA-Z].*")) englishWords++;
|
}
|
return chineseChars + (int) (englishWords * 1.3);
|
}
|
|
}
|