package cn.iocoder.yudao.module.ai.service.knowledge;
|
|
import cn.hutool.core.collection.CollUtil;
|
import cn.hutool.core.io.FileUtil;
|
import cn.hutool.core.io.IoUtil;
|
import cn.hutool.core.util.ArrayUtil;
|
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.document.*;
|
import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentProcessRespVO;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDO;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeSegmentDO;
|
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiModelDO;
|
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.model.AiModelService;
|
import jakarta.annotation.Resource;
|
import lombok.extern.slf4j.Slf4j;
|
import org.springframework.ai.document.Document;
|
import org.springframework.ai.transformer.splitter.TokenTextSplitter;
|
import org.springframework.stereotype.Service;
|
import org.springframework.transaction.annotation.Transactional;
|
|
import org.apache.tika.parser.AutoDetectParser;
|
import org.apache.tika.parser.ParseContext;
|
import org.apache.tika.parser.Parser;
|
import org.apache.tika.extractor.EmbeddedDocumentExtractor;
|
import org.apache.tika.sax.BodyContentHandler;
|
import org.xml.sax.ContentHandler;
|
import java.io.ByteArrayOutputStream;
|
import java.io.InputStream;
|
import java.net.HttpURLConnection;
|
import java.net.URI;
|
import java.net.URL;
|
import java.net.URLEncoder;
|
import java.nio.charset.StandardCharsets;
|
import java.util.ArrayList;
|
import java.util.Collections;
|
import java.util.List;
|
|
import static cn.iocoder.yudao.framework.common.exception.util.ServiceExceptionUtil.exception;
|
import static cn.iocoder.yudao.module.ai.enums.ErrorCodeConstants.*;
|
|
@Slf4j
|
@Service
|
public class AiKnowledgeDocumentServiceImpl implements AiKnowledgeDocumentService {
|
|
@Resource
|
private AiKnowledgeDocumentMapper documentMapper;
|
@Resource
|
private AiKnowledgeSegmentMapper segmentMapper;
|
@Resource
|
private AiKnowledgeService knowledgeService;
|
@Resource
|
private AiKnowledgeSegmentService segmentService;
|
@Resource
|
private AiModelService modelService;
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public List<Long> createDocuments(AiKnowledgeDocumentCreateReqVO createReqVO) {
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(createReqVO.getKnowledgeId());
|
List<AiKnowledgeDocumentCreateReqVO.DocumentItem> list = createReqVO.getList();
|
if (CollUtil.isEmpty(list)) {
|
throw new IllegalArgumentException("文档列表不能为空");
|
}
|
int defaultSegmentMaxTokens = createReqVO.getSegmentMaxTokens() != null
|
? createReqVO.getSegmentMaxTokens() : 800;
|
|
List<Long> ids = new ArrayList<>();
|
for (AiKnowledgeDocumentCreateReqVO.DocumentItem item : list) {
|
AiKnowledgeDocumentCreateReqVO.UrlInfo urlInfo = item.getUrl();
|
String fileUrl = urlInfo != null ? urlInfo.getUrl() : null;
|
AiKnowledgeDocumentDO document = new AiKnowledgeDocumentDO();
|
document.setKnowledgeId(knowledge.getId());
|
document.setName(item.getName());
|
document.setUrl(fileUrl);
|
document.setStatus(CommonStatusEnum.ENABLE.getStatus());
|
document.setSegmentMaxTokens(defaultSegmentMaxTokens);
|
documentMapper.insert(document);
|
ids.add(document.getId());
|
|
if (fileUrl != null) {
|
try {
|
loadDocumentContent(document);
|
documentMapper.updateById(document);
|
// 自动执行分段和向量化
|
processDocumentSegmentsInternal(document, knowledge);
|
} catch (Exception e) {
|
log.error("文档[{}]处理失败: {}", document.getId(), e.getMessage());
|
document.setStatus(CommonStatusEnum.DISABLE.getStatus());
|
documentMapper.updateById(document);
|
}
|
}
|
}
|
return ids;
|
}
|
|
private void loadDocumentContent(AiKnowledgeDocumentDO document) {
|
try {
|
byte[] fileBytes = downloadFile(document.getUrl());
|
String content = extractText(fileBytes, document.getName());
|
if (StrUtil.isEmpty(content)) {
|
content = new String(fileBytes, java.nio.charset.Charset.forName("UTF-8"));
|
}
|
document.setContent(content);
|
document.setContentLength(content.length());
|
document.setTokens(estimateTokens(content));
|
} catch (Exception e) {
|
throw new RuntimeException("文档内容加载失败: " + e.getMessage(), e);
|
}
|
}
|
|
private byte[] downloadFile(String fileUrl) {
|
try {
|
// 手动编码 URL 路径中的非 ASCII 字符,避免 URI 构造函数报错
|
String encodedUrl = encodeUrlPath(fileUrl);
|
URL url = URL.of(new URI(encodedUrl), null);
|
HttpURLConnection conn = (HttpURLConnection) url.openConnection();
|
conn.setConnectTimeout(30000);
|
conn.setReadTimeout(60000);
|
conn.setRequestMethod("GET");
|
try (InputStream is = conn.getInputStream();
|
ByteArrayOutputStream bos = new ByteArrayOutputStream()) {
|
IoUtil.copy(is, bos);
|
return bos.toByteArray();
|
}
|
} catch (Exception e) {
|
throw new RuntimeException("文件下载失败: " + e.getMessage(), e);
|
}
|
}
|
|
private static String encodeUrlPath(String urlString) {
|
// 分离 scheme://authority 和 path?query#fragment
|
int schemeEnd = urlString.indexOf("://");
|
if (schemeEnd < 0) return urlString;
|
int pathStart = urlString.indexOf('/', schemeEnd + 3);
|
if (pathStart < 0) return urlString; // 无 path
|
|
String base = urlString.substring(0, pathStart);
|
String pathAndRest = urlString.substring(pathStart);
|
|
// 分离 path 和 query + fragment
|
int queryStart = pathAndRest.indexOf('?');
|
int fragStart = pathAndRest.indexOf('#');
|
String rawPath, query, fragment;
|
if (queryStart >= 0) {
|
rawPath = pathAndRest.substring(0, queryStart);
|
if (fragStart >= 0 && fragStart > queryStart) {
|
query = pathAndRest.substring(queryStart, fragStart);
|
fragment = pathAndRest.substring(fragStart);
|
} else {
|
query = pathAndRest.substring(queryStart);
|
fragment = "";
|
}
|
} else if (fragStart >= 0) {
|
rawPath = pathAndRest.substring(0, fragStart);
|
query = "";
|
fragment = pathAndRest.substring(fragStart);
|
} else {
|
rawPath = pathAndRest;
|
query = "";
|
fragment = "";
|
}
|
|
// 对路径每个段做 URL 编码
|
StringBuilder encodedPath = new StringBuilder();
|
for (String segment : rawPath.split("/")) {
|
if (!segment.isEmpty()) {
|
encodedPath.append("/").append(URLEncoder.encode(segment, StandardCharsets.UTF_8)
|
.replace("+", "%20"));
|
}
|
}
|
if (rawPath.endsWith("/")) encodedPath.append("/");
|
|
return base + encodedPath + query + fragment;
|
}
|
|
private String extractText(byte[] fileBytes, String fileName) {
|
try {
|
String ext = FileUtil.extName(fileName).toLowerCase();
|
if (ArrayUtil.contains(new String[]{"txt", "md", "json", "xml", "csv", "yaml", "yml"}, ext)) {
|
return new String(fileBytes, StandardCharsets.UTF_8);
|
}
|
// 使用 AutoDetectParser + EmbeddedDocumentExtractor 跳过嵌入图片,避免提取到二进制乱码
|
Parser parser = new AutoDetectParser();
|
ParseContext context = new ParseContext();
|
context.set(EmbeddedDocumentExtractor.class, new EmbeddedDocumentExtractor() {
|
@Override
|
public boolean shouldParseEmbedded(org.apache.tika.metadata.Metadata metadata) {
|
return false;
|
}
|
@Override
|
public void parseEmbedded(InputStream inputStream, ContentHandler contentHandler,
|
org.apache.tika.metadata.Metadata metadata, boolean outputHtml) {
|
}
|
});
|
BodyContentHandler handler = new BodyContentHandler(-1);
|
try (InputStream is = new java.io.ByteArrayInputStream(fileBytes)) {
|
parser.parse(is, handler, new org.apache.tika.metadata.Metadata(), context);
|
}
|
return handler.toString().trim();
|
} catch (Exception e) {
|
log.warn("Tika 解析文档失败,尝试 UTF-8 文本读取: {}", e.getMessage());
|
return "";
|
}
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void updateDocument(AiKnowledgeDocumentUpdateReqVO updateReqVO) {
|
AiKnowledgeDocumentDO document = validateDocumentExists(updateReqVO.getId());
|
boolean urlChanged = StrUtil.isNotEmpty(updateReqVO.getUrl())
|
&& !updateReqVO.getUrl().equals(document.getUrl());
|
if (StrUtil.isNotEmpty(updateReqVO.getName())) document.setName(updateReqVO.getName());
|
if (urlChanged) document.setUrl(updateReqVO.getUrl());
|
documentMapper.updateById(document);
|
// URL 变更时重新提取内容、重新分段和向量化
|
if (urlChanged) {
|
try {
|
loadDocumentContent(document);
|
documentMapper.updateById(document);
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(document.getKnowledgeId());
|
// 删除旧分段和向量
|
segmentService.deleteSegmentsByDocumentId(document.getId());
|
// 重新分段 + 向量化
|
processDocumentSegmentsInternal(document, knowledge);
|
} catch (Exception e) {
|
log.error("文档[{}] URL 更新后处理失败: {}", document.getId(), e.getMessage());
|
document.setStatus(CommonStatusEnum.DISABLE.getStatus());
|
documentMapper.updateById(document);
|
}
|
}
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void deleteDocument(Long id) {
|
AiKnowledgeDocumentDO document = validateDocumentExists(id);
|
segmentService.deleteSegmentsByDocumentId(id);
|
documentMapper.deleteById(id);
|
}
|
|
@Override
|
public AiKnowledgeDocumentDO getDocument(Long id) {
|
return documentMapper.selectById(id);
|
}
|
|
@Override
|
public AiKnowledgeDocumentDO validateDocumentExists(Long id) {
|
AiKnowledgeDocumentDO document = documentMapper.selectById(id);
|
if (document == null) throw exception(KNOWLEDGE_DOCUMENT_NOT_EXISTS);
|
return document;
|
}
|
|
@Override
|
public PageResult<AiKnowledgeDocumentDO> getDocumentPage(AiKnowledgeDocumentPageReqVO pageReqVO) {
|
return documentMapper.selectPage(pageReqVO);
|
}
|
|
@Override
|
@Transactional(rollbackFor = Exception.class)
|
public void updateDocumentStatus(AiKnowledgeDocumentUpdateStatusReqVO updateStatusReqVO) {
|
AiKnowledgeDocumentDO document = validateDocumentExists(updateStatusReqVO.getId());
|
document.setStatus(updateStatusReqVO.getStatus());
|
documentMapper.updateById(document);
|
}
|
|
@Override
|
public List<AiKnowledgeSegmentProcessRespVO> getDocumentProcessingProgress(List<Long> documentIds) {
|
if (CollUtil.isEmpty(documentIds)) return Collections.emptyList();
|
return segmentMapper.selectProcessList(documentIds);
|
}
|
|
@Override
|
public void processDocumentSegments(Long documentId) {
|
AiKnowledgeDocumentDO document = validateDocumentExists(documentId);
|
if (StrUtil.isEmpty(document.getContent())) {
|
throw exception(KNOWLEDGE_DOCUMENT_FILE_EMPTY);
|
}
|
AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(document.getKnowledgeId());
|
processDocumentSegmentsInternal(document, knowledge);
|
}
|
|
private void processDocumentSegmentsInternal(AiKnowledgeDocumentDO document, AiKnowledgeDO knowledge) {
|
if (StrUtil.isEmpty(document.getContent())) return;
|
|
AiModelDO embeddingModel = modelService.validateModel(knowledge.getEmbeddingModelId());
|
|
// 删除旧分段
|
segmentService.deleteSegmentsByDocumentId(document.getId());
|
|
// 文本切片
|
int segmentMaxTokens = document.getSegmentMaxTokens() != null ? document.getSegmentMaxTokens() : 800;
|
List<String> segmentTexts = splitContent(document.getContent(), segmentMaxTokens);
|
if (CollUtil.isEmpty(segmentTexts)) return;
|
|
// 向量化并存储
|
List<AiKnowledgeSegmentDO> segments = new ArrayList<>();
|
for (String content : segmentTexts) {
|
if (StrUtil.isEmpty(content.trim())) continue;
|
segments.add(buildSegmentDO(knowledge.getId(), document.getId(), content));
|
}
|
segmentService.saveSegments(segments, knowledge.getId(), embeddingModel.getId());
|
}
|
|
private List<String> splitContent(String content, int maxTokens) {
|
TokenTextSplitter splitter = TokenTextSplitter.builder()
|
.withChunkSize(maxTokens)
|
.withMinChunkSizeChars(50)
|
.withMinChunkLengthToEmbed(10)
|
.withMaxNumChunks(1000)
|
.withKeepSeparator(true)
|
.build();
|
List<Document> docs = splitter.apply(Collections.singletonList(new Document(content)));
|
List<String> result = new ArrayList<>();
|
for (Document doc : docs) {
|
if (StrUtil.isNotEmpty(doc.getText())) {
|
result.add(doc.getText().trim());
|
}
|
}
|
return result;
|
}
|
|
@Override
|
public List<String> previewSplit(String url, String name, Integer segmentMaxTokens) {
|
byte[] fileBytes = downloadFile(url);
|
String text = extractText(fileBytes, name);
|
if (StrUtil.isEmpty(text)) {
|
text = new String(fileBytes, java.nio.charset.Charset.forName("UTF-8"));
|
}
|
int maxTokens = segmentMaxTokens != null && segmentMaxTokens > 0 ? segmentMaxTokens : 800;
|
return splitContent(text, maxTokens);
|
}
|
|
private AiKnowledgeSegmentDO buildSegmentDO(Long knowledgeId, Long documentId, String content) {
|
AiKnowledgeSegmentDO segment = new AiKnowledgeSegmentDO();
|
segment.setKnowledgeId(knowledgeId);
|
segment.setDocumentId(documentId);
|
segment.setContent(content);
|
segment.setContentLength(content.length());
|
segment.setTokens(estimateTokens(content));
|
segment.setStatus(CommonStatusEnum.ENABLE.getStatus());
|
segment.setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY);
|
segment.setRetrievalCount(0);
|
return segment;
|
}
|
|
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);
|
}
|
|
}
|