fix(aihr): harden shared document parsing

This commit is contained in:
2026-07-12 02:45:56 +08:00
parent 36ff3c79f8
commit 806f90a5de
6 changed files with 315 additions and 36 deletions
@@ -1,9 +1,47 @@
package org.dromara.aihr.knowledge.parse; package org.dromara.aihr.knowledge.parse;
import java.io.IOException;
import java.io.InputStream;
/** /**
* Stateless byte-document parser shared by knowledge ingestion flows. * Stateless byte-document parser shared by knowledge ingestion flows.
*/ */
public interface KnowledgeDocumentParser { public interface KnowledgeDocumentParser {
ParsedDocument parse(String fileName, String contentType, byte[] bytes); ParsedDocument parse(String fileName, String contentType, byte[] bytes);
default ParsedDocument parse(String fileName, String contentType, InputStream input) {
if (input == null) {
throw new ParseException(Failure.EMPTY, "document content is empty");
}
try {
return parse(fileName, contentType, input.readAllBytes());
} catch (IOException e) {
throw new ParseException(Failure.INVALID, "document reading failed", e);
}
}
enum Failure {
EMPTY,
TOO_LARGE,
INVALID
}
final class ParseException extends IllegalArgumentException {
private final Failure failure;
public ParseException(Failure failure, String message) {
super(message);
this.failure = failure;
}
public ParseException(Failure failure, String message, Throwable cause) {
super(message, cause);
this.failure = failure;
}
public Failure failure() {
return failure;
}
}
} }
@@ -20,15 +20,16 @@ public record ParsedDocument(String text, String mimeType, Map<String, String> m
return List.of(); return List.of();
} }
int[] codePoints = text.codePoints().toArray();
List<String> chunks = new ArrayList<>(); List<String> chunks = new ArrayList<>();
int step = blockSize - overlap; int step = blockSize - overlap;
for (int start = 0; start < text.length(); start += step) { for (int start = 0; start < codePoints.length; start += step) {
int end = Math.min(text.length(), start + blockSize); int end = Math.min(codePoints.length, start + blockSize);
String chunk = text.substring(start, end).trim(); String chunk = new String(codePoints, start, end - start).trim();
if (!chunk.isEmpty()) { if (!chunk.isEmpty()) {
chunks.add(chunk); chunks.add(chunk);
} }
if (end == text.length()) { if (end == codePoints.length) {
break; break;
} }
} }
@@ -1,14 +1,23 @@
package org.dromara.aihr.knowledge.parse; package org.dromara.aihr.knowledge.parse;
import org.apache.tika.exception.WriteLimitReachedException; import org.apache.tika.exception.WriteLimitReachedException;
import org.apache.tika.extractor.EmbeddedDocumentExtractor;
import org.apache.tika.io.BoundedInputStream;
import org.apache.tika.io.TemporaryResources;
import org.apache.tika.io.TikaInputStream;
import org.apache.tika.metadata.Metadata; import org.apache.tika.metadata.Metadata;
import org.apache.tika.metadata.TikaCoreProperties; import org.apache.tika.metadata.TikaCoreProperties;
import org.apache.tika.mime.MediaType;
import org.apache.tika.parser.AutoDetectParser; import org.apache.tika.parser.AutoDetectParser;
import org.apache.tika.parser.ParseContext; import org.apache.tika.parser.ParseContext;
import org.apache.tika.parser.Parser;
import org.apache.tika.sax.BodyContentHandler; import org.apache.tika.sax.BodyContentHandler;
import org.xml.sax.ContentHandler;
import org.springframework.stereotype.Component; import org.springframework.stereotype.Component;
import java.io.ByteArrayInputStream; import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.InputStream;
import java.util.LinkedHashMap; import java.util.LinkedHashMap;
import java.util.Locale; import java.util.Locale;
import java.util.Map; import java.util.Map;
@@ -17,6 +26,7 @@ import java.util.Map;
public class TikaKnowledgeDocumentParser implements KnowledgeDocumentParser { public class TikaKnowledgeDocumentParser implements KnowledgeDocumentParser {
static final int DEFAULT_MAX_EXPANDED_CHARS = 2_000_000; static final int DEFAULT_MAX_EXPANDED_CHARS = 2_000_000;
static final long MAX_INPUT_BYTES = 100L * 1024 * 1024;
private final int maxExpandedChars; private final int maxExpandedChars;
@@ -34,43 +44,70 @@ public class TikaKnowledgeDocumentParser implements KnowledgeDocumentParser {
@Override @Override
public ParsedDocument parse(String fileName, String contentType, byte[] bytes) { public ParsedDocument parse(String fileName, String contentType, byte[] bytes) {
if (bytes == null || bytes.length == 0) { if (bytes == null || bytes.length == 0) {
throw new IllegalArgumentException("document content is empty"); throw new ParseException(Failure.EMPTY, "document content is empty");
}
return parse(fileName, contentType, new ByteArrayInputStream(bytes));
}
@Override
public ParsedDocument parse(String fileName, String contentType, InputStream input) {
if (input == null) {
throw new ParseException(Failure.EMPTY, "document content is empty");
} }
Metadata metadata = new Metadata(); Metadata metadata = new Metadata();
if (fileName != null && !fileName.isBlank()) { if (fileName != null && !fileName.isBlank()) {
metadata.set(TikaCoreProperties.RESOURCE_NAME_KEY, fileName.trim()); metadata.set(TikaCoreProperties.RESOURCE_NAME_KEY, fileName.trim());
} }
if (contentType != null && !contentType.isBlank()) { AutoDetectParser parser = new AutoDetectParser();
metadata.set(Metadata.CONTENT_TYPE, contentType.trim());
}
BodyContentHandler handler = new BodyContentHandler(maxExpandedChars + 1); BodyContentHandler handler = new BodyContentHandler(maxExpandedChars + 1);
try (ByteArrayInputStream input = new ByteArrayInputStream(bytes)) { BoundedInputStream bounded = new BoundedInputStream(MAX_INPUT_BYTES + 1, input);
new AutoDetectParser().parse(input, handler, metadata, new ParseContext()); try (TemporaryResources temporaryResources = new TemporaryResources();
TikaInputStream tikaInput = TikaInputStream.get(bounded, temporaryResources, metadata)) {
tikaInput.mark(Integer.MAX_VALUE);
MediaType detected = parser.getDetector().detect(tikaInput, metadata);
tikaInput.reset();
String mimeType = resolvedMimeType(detected, contentType);
metadata.set(Metadata.CONTENT_TYPE, mimeType);
ParseContext context = new ParseContext();
context.set(Parser.class, parser);
context.set(EmbeddedDocumentExtractor.class, NO_EMBEDDED_DOCUMENTS);
parser.parse(tikaInput, handler, metadata, context);
rejectOversizedInput(bounded);
return parsedDocument(handler, metadata, mimeType);
} catch (Exception e) { } catch (Exception e) {
if (WriteLimitReachedException.isWriteLimitReached(e)) { if (e instanceof ParseException parseException) {
throw new IllegalArgumentException("document expanded text exceeds limit", e); throw parseException;
} }
throw new IllegalArgumentException("document parsing failed", e); if (bounded.hasHitBound() || bounded.getPos() > MAX_INPUT_BYTES) {
throw new ParseException(Failure.TOO_LARGE, "document input exceeds limit", e);
}
if (WriteLimitReachedException.isWriteLimitReached(e)) {
throw new ParseException(Failure.TOO_LARGE, "document expanded text exceeds limit", e);
}
throw new ParseException(Failure.INVALID, "document parsing failed", e);
} }
String text = handler.toString().trim();
if (text.isEmpty()) {
throw new IllegalArgumentException("document contains no text");
}
if (text.length() > maxExpandedChars) {
throw new IllegalArgumentException("document expanded text exceeds limit");
}
return new ParsedDocument(text, resolveMimeType(contentType, metadata), metadataMap(metadata));
} }
private static String resolveMimeType(String suppliedContentType, Metadata metadata) { private ParsedDocument parsedDocument(BodyContentHandler handler, Metadata metadata, String mimeType) {
String candidate = suppliedContentType; String text = handler.toString().trim();
if (candidate == null || candidate.isBlank()) { if (text.isEmpty()) {
candidate = metadata.get(Metadata.CONTENT_TYPE); throw new ParseException(Failure.EMPTY, "document contains no text");
} }
if (text.length() > maxExpandedChars) {
throw new ParseException(Failure.TOO_LARGE, "document expanded text exceeds limit");
}
return new ParsedDocument(text, mimeType, metadataMap(metadata));
}
private static String resolvedMimeType(MediaType detected, String suppliedContentType) {
String detectedMime = detected == null ? "" : detected.getBaseType().toString();
if (!detectedMime.isBlank() && !MediaType.OCTET_STREAM.toString().equals(detectedMime)) {
return detectedMime;
}
String candidate = suppliedContentType;
if (candidate == null || candidate.isBlank()) { if (candidate == null || candidate.isBlank()) {
return "application/octet-stream"; return "application/octet-stream";
} }
@@ -79,6 +116,25 @@ public class TikaKnowledgeDocumentParser implements KnowledgeDocumentParser {
return mimeType.isEmpty() ? "application/octet-stream" : mimeType.toLowerCase(Locale.ROOT); return mimeType.isEmpty() ? "application/octet-stream" : mimeType.toLowerCase(Locale.ROOT);
} }
private static void rejectOversizedInput(BoundedInputStream bounded) {
if (bounded.hasHitBound() || bounded.getPos() > MAX_INPUT_BYTES) {
throw new ParseException(Failure.TOO_LARGE, "document input exceeds limit");
}
}
private static final EmbeddedDocumentExtractor NO_EMBEDDED_DOCUMENTS = new EmbeddedDocumentExtractor() {
@Override
public boolean shouldParseEmbedded(Metadata metadata) {
return false;
}
@Override
public void parseEmbedded(InputStream stream, ContentHandler handler, Metadata metadata, boolean outputHtml)
throws IOException {
// Embedded payloads are deliberately excluded to bound recursive expansion.
}
};
private static Map<String, String> metadataMap(Metadata metadata) { private static Map<String, String> metadataMap(Metadata metadata) {
Map<String, String> values = new LinkedHashMap<>(); Map<String, String> values = new LinkedHashMap<>();
for (String name : metadata.names()) { for (String name : metadata.names()) {
@@ -2924,7 +2924,7 @@ public class AihrSopSeedService {
.send(builder.build(), HttpResponse.BodyHandlers.ofString()); .send(builder.build(), HttpResponse.BodyHandlers.ofString());
} }
private String readContent(MultipartFile file, String fileName) { String readContent(MultipartFile file, String fileName) {
if (AihrVideoService.videoFile(fileName)) { if (AihrVideoService.videoFile(fileName)) {
// 视频加工耗时数分钟,只允许走异步队列(暂存路径版本),同步接口直接拒绝 // 视频加工耗时数分钟,只允许走异步队列(暂存路径版本),同步接口直接拒绝
throw new ServiceException("视频请使用资料处理中心的批量导入(异步队列)上传"); throw new ServiceException("视频请使用资料处理中心的批量导入(异步队列)上传");
@@ -2946,14 +2946,16 @@ public class AihrSopSeedService {
} }
return ""; return "";
} }
try { try (InputStream input = file.getInputStream()) {
return parseDocument(fileName, file.getContentType(), file.getBytes()); return parseDocument(fileName, file.getContentType(), input);
} catch (ServiceException e) {
throw e;
} catch (Exception e) { } catch (Exception e) {
throw new ServiceException("文件解析失败"); throw new ServiceException("文件解析失败");
} }
} }
private String readContent(Path file, String fileName) { String readContent(Path file, String fileName) {
if (AihrVideoService.videoFile(fileName)) { if (AihrVideoService.videoFile(fileName)) {
return videoService.extractText(file, fileName, (bytes, mimeType) -> { return videoService.extractText(file, fileName, (bytes, mimeType) -> {
Optional<ChatRuntime> runtime = visionRuntime(); Optional<ChatRuntime> runtime = visionRuntime();
@@ -2985,15 +2987,24 @@ public class AihrSopSeedService {
} }
return ""; return "";
} }
try { try (InputStream input = Files.newInputStream(file)) {
return parseDocument(fileName, Files.probeContentType(file), Files.readAllBytes(file)); return parseDocument(fileName, Files.probeContentType(file), input);
} catch (ServiceException e) {
throw e;
} catch (Exception e) { } catch (Exception e) {
throw new ServiceException("文件解析失败"); throw new ServiceException("文件解析失败");
} }
} }
private String parseDocument(String fileName, String contentType, byte[] bytes) { private String parseDocument(String fileName, String contentType, InputStream input) {
return normalizeExtractedText(knowledgeDocumentParser.parse(fileName, contentType, bytes).text()); try {
return normalizeExtractedText(knowledgeDocumentParser.parse(fileName, contentType, input).text());
} catch (KnowledgeDocumentParser.ParseException e) {
if (e.failure() == KnowledgeDocumentParser.Failure.EMPTY) {
throw new ServiceException("文件内容不能为空");
}
throw new ServiceException("文件解析失败");
}
} }
private static FileFingerprint fileFingerprint(MultipartFile file) { private static FileFingerprint fileFingerprint(MultipartFile file) {
@@ -1,8 +1,15 @@
package org.dromara.aihr.knowledge.parse; package org.dromara.aihr.knowledge.parse;
import org.apache.pdfbox.pdmodel.PDDocument;
import org.apache.pdfbox.pdmodel.PDPage;
import org.apache.pdfbox.pdmodel.PDPageContentStream;
import org.apache.pdfbox.pdmodel.font.PDType1Font;
import org.apache.pdfbox.pdmodel.font.Standard14Fonts;
import org.junit.jupiter.api.Tag; import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertEquals;
@@ -82,4 +89,41 @@ class TikaKnowledgeDocumentParserTest {
assertEquals("1234567890", accepted.text()); assertEquals("1234567890", accepted.text());
assertTrue(rejected.getMessage().contains("exceeds")); assertTrue(rejected.getMessage().contains("exceeds"));
} }
@Test
void detectsActualPdfMimeWhenDeclaredTypeConflicts() throws IOException {
KnowledgeDocumentParser parser = new TikaKnowledgeDocumentParser();
byte[] pdf = pdfBytes("Fee guide");
ParsedDocument declaredText = parser.parse("fee-guide.pdf", "text/plain", pdf);
ParsedDocument declaredBinary = parser.parse("fee-guide.pdf", "application/octet-stream", pdf);
assertEquals("application/pdf", declaredText.mimeType());
assertEquals("application/pdf", declaredBinary.mimeType());
assertTrue(declaredText.text().contains("Fee guide"));
}
@Test
void chunksOnUnicodeCodePointBoundaries() {
ParsedDocument document = new ParsedDocument("A😀BC😀D", "text/plain", java.util.Map.of());
assertEquals(java.util.List.of("A😀B", "BC😀", "😀D"), document.chunks(3, 1));
assertTrue(document.chunks(3, 1).stream().noneMatch(chunk -> chunk.contains("�")));
}
private static byte[] pdfBytes(String text) throws IOException {
try (PDDocument document = new PDDocument(); ByteArrayOutputStream output = new ByteArrayOutputStream()) {
PDPage page = new PDPage();
document.addPage(page);
try (PDPageContentStream content = new PDPageContentStream(document, page)) {
content.beginText();
content.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12);
content.newLineAtOffset(72, 720);
content.showText(text);
content.endText();
}
document.save(output);
return output.toByteArray();
}
}
} }
@@ -1,13 +1,25 @@
package org.dromara.aihr.service; package org.dromara.aihr.service;
import org.dromara.aihr.domain.AihrSopDto; import org.dromara.aihr.domain.AihrSopDto;
import org.dromara.aihr.knowledge.parse.KnowledgeDocumentParser;
import org.dromara.aihr.knowledge.parse.ParsedDocument;
import org.dromara.common.core.exception.ServiceException;
import org.junit.jupiter.api.Tag; import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
import org.springframework.mock.web.MockMultipartFile;
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.List; import java.util.List;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.assertTrue;
public class AihrSopSeedServiceTest { public class AihrSopSeedServiceTest {
@@ -80,7 +92,124 @@ public class AihrSopSeedServiceTest {
assertFalse(cleaned.contains("う")); assertFalse(cleaned.contains("う"));
} }
@Test
@Tag("dev")
public void genericMultipartAndPathDocumentsUseStreamingParser() throws IOException {
RecordingParser parser = new RecordingParser();
AihrSopSeedService service = service(parser);
MockMultipartFile multipart = new MockMultipartFile(
"file", "guide.pdf", "application/pdf", "multipart".getBytes(StandardCharsets.UTF_8)
);
Path staged = Files.createTempFile("aihr-parser-", ".pdf");
Files.writeString(staged, "staged", StandardCharsets.UTF_8);
try {
assertEquals("parsed multipart", service.readContent(multipart, "guide.pdf"));
assertEquals("parsed staged", service.readContent(staged, "guide.pdf"));
} finally {
Files.deleteIfExists(staged);
}
assertEquals(2, parser.streamCalls);
assertEquals(0, parser.byteArrayCalls);
}
@Test
@Tag("dev")
public void markdownKeepsExplicitUtf8PathWithoutCallingParser() {
RecordingParser parser = new RecordingParser();
AihrSopSeedService service = service(parser);
MockMultipartFile markdown = new MockMultipartFile(
"file",
"note.md",
"text/markdown",
"---\ntitle: 测试\n---\n收费沟通先说明费用构成".getBytes(StandardCharsets.UTF_8)
);
assertEquals("收费沟通先说明费用构成", service.readContent(markdown, "note.md"));
assertEquals(0, parser.streamCalls);
assertEquals(0, parser.byteArrayCalls);
}
@Test
@Tag("dev")
public void emptyParserResultMapsToActionableServiceError() {
RecordingParser parser = new RecordingParser();
parser.failure = new KnowledgeDocumentParser.ParseException(
KnowledgeDocumentParser.Failure.EMPTY,
"document contains no text"
);
AihrSopSeedService service = service(parser);
MockMultipartFile multipart = new MockMultipartFile(
"file", "empty.pdf", "application/pdf", "not-empty-input".getBytes(StandardCharsets.UTF_8)
);
ServiceException error = assertThrows(
ServiceException.class,
() -> service.readContent(multipart, "empty.pdf")
);
assertEquals("文件内容不能为空", error.getMessage());
assertNull(error.getCause());
}
@Test
@Tag("dev")
public void parserFailureMapsWithoutLeakingInternalCause() {
RecordingParser parser = new RecordingParser();
parser.failure = new KnowledgeDocumentParser.ParseException(
KnowledgeDocumentParser.Failure.INVALID,
"internal parser detail",
new IllegalStateException("sensitive stack detail")
);
AihrSopSeedService service = service(parser);
MockMultipartFile multipart = new MockMultipartFile(
"file", "broken.pdf", "application/pdf", "broken".getBytes(StandardCharsets.UTF_8)
);
ServiceException error = assertThrows(
ServiceException.class,
() -> service.readContent(multipart, "broken.pdf")
);
assertEquals("文件解析失败", error.getMessage());
assertNull(error.getCause());
}
private static AihrSopSeedService service(KnowledgeDocumentParser parser) {
return new AihrSopSeedService(null, null, null, "", null, null, parser);
}
private static AihrSopSeedService.KnowledgeHit hit(Long fragmentId, String title) { private static AihrSopSeedService.KnowledgeHit hit(Long fragmentId, String title) {
return new AihrSopSeedService.KnowledgeHit(fragmentId, title, "sop", "", "doc-" + fragmentId, "片段内容", 1, 1.0); return new AihrSopSeedService.KnowledgeHit(fragmentId, title, "sop", "", "doc-" + fragmentId, "片段内容", 1, 1.0);
} }
private static final class RecordingParser implements KnowledgeDocumentParser {
private int streamCalls;
private int byteArrayCalls;
private ParseException failure;
@Override
public ParsedDocument parse(String fileName, String contentType, byte[] bytes) {
byteArrayCalls++;
throw new AssertionError("enterprise paths must not buffer the entire document");
}
@Override
public ParsedDocument parse(String fileName, String contentType, InputStream input) {
streamCalls++;
if (failure != null) {
throw failure;
}
try {
return new ParsedDocument(
"parsed " + new String(input.readAllBytes(), StandardCharsets.UTF_8),
contentType,
Map.of()
);
} catch (IOException e) {
throw new ParseException(Failure.INVALID, "test read failed", e);
}
}
}
} }