@@ -0,0 +1,244 @@
|
||||
"""API integration tests for chunk persistence and recall."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from api.main import app
|
||||
from rag_cut.models import Chunk, ChunkResult, SplitMode
|
||||
|
||||
|
||||
def result(filename: str = "manual.txt") -> ChunkResult:
|
||||
content = "交易密码可在账户安全页面修改"
|
||||
return ChunkResult(
|
||||
filename=filename,
|
||||
doc_id="abcdef123456",
|
||||
split_mode=SplitMode.DEFAULT,
|
||||
block_count=2,
|
||||
chunk_count=2,
|
||||
chunks=[
|
||||
Chunk(index=0, content="", char_count=18, block_types=["image"], meta={"retrieval": False}),
|
||||
Chunk(index=1, content=content, char_count=len(content), block_types=["paragraph"], meta={"retrieval": True}),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class ApiTest(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.client = TestClient(app)
|
||||
|
||||
def test_chunk_preserves_upload_filename_and_persists_result(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp, patch("api.main.RESULTS_DIR", Path(tmp)):
|
||||
with patch("api.main.chunk_document", side_effect=lambda path, config=None: result(path.name)):
|
||||
response = self.client.post("/api/chunk", files={"file": ("用户手册.txt", b"body", "text/plain")})
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json()["filename"], "用户手册.txt")
|
||||
self.assertTrue((Path(tmp) / "abcdef123456.json").exists())
|
||||
|
||||
def test_recall_endpoint_excludes_preview_only_chunk(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp, patch("api.main.RESULTS_DIR", Path(tmp)):
|
||||
(Path(tmp) / "abcdef123456.json").write_text(result().model_dump_json(), encoding="utf-8")
|
||||
response = self.client.post(
|
||||
"/api/recall",
|
||||
json={"doc_id": "abcdef123456", "query": "交易密码", "top_k": 5},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
payload = response.json()
|
||||
self.assertEqual(payload["candidate_count"], 1)
|
||||
self.assertEqual(payload["results"][0]["chunk_index"], 1)
|
||||
|
||||
def test_recall_rejects_invalid_doc_id(self) -> None:
|
||||
response = self.client.post("/api/recall", json={"doc_id": "../secret", "query": "test"})
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_list_and_get_historical_results(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
results_dir = Path(tmp) / "results"
|
||||
uploads_dir = Path(tmp) / "uploads"
|
||||
results_dir.mkdir()
|
||||
uploads_dir.mkdir()
|
||||
(results_dir / "abcdef123456.json").write_text(result().model_dump_json(), encoding="utf-8")
|
||||
upload_doc = uploads_dir / "abcdef123456"
|
||||
upload_doc.mkdir()
|
||||
(upload_doc / "manual.txt").write_text("hello original", encoding="utf-8")
|
||||
|
||||
with (
|
||||
patch("api.main.RESULTS_DIR", results_dir),
|
||||
patch("api.main.UPLOADS_DIR", uploads_dir),
|
||||
):
|
||||
listed = self.client.get("/api/results")
|
||||
self.assertEqual(listed.status_code, 200)
|
||||
body = listed.json()
|
||||
self.assertEqual(body["count"], 1)
|
||||
self.assertEqual(body["results"][0]["doc_id"], "abcdef123456")
|
||||
self.assertEqual(body["results"][0]["filename"], "manual.txt")
|
||||
self.assertTrue(body["results"][0]["has_original"])
|
||||
|
||||
detail = self.client.get("/api/results/abcdef123456")
|
||||
self.assertEqual(detail.status_code, 200)
|
||||
self.assertEqual(detail.json()["chunk_count"], 2)
|
||||
self.assertTrue(detail.json()["has_original"])
|
||||
|
||||
original = self.client.get("/api/results/abcdef123456/original")
|
||||
self.assertEqual(original.status_code, 200)
|
||||
self.assertEqual(original.content, b"hello original")
|
||||
|
||||
def test_get_result_not_found(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp, patch("api.main.RESULTS_DIR", Path(tmp)):
|
||||
response = self.client.get("/api/results/abcdef123456")
|
||||
self.assertEqual(response.status_code, 404)
|
||||
|
||||
def test_word_preview_serves_converted_pdf(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = Path(tmp)
|
||||
uploads_dir = root / "uploads"
|
||||
doc_dir = uploads_dir / "abcdef123456"
|
||||
conversion_dir = root / "assets" / "abcdef123456" / "_conversion"
|
||||
doc_dir.mkdir(parents=True)
|
||||
conversion_dir.mkdir(parents=True)
|
||||
(doc_dir / "manual.docx").write_bytes(b"word")
|
||||
preview = conversion_dir / "manual.pdf"
|
||||
preview.write_bytes(b"%PDF-1.7 preview")
|
||||
|
||||
with (
|
||||
patch("api.main.UPLOADS_DIR", uploads_dir),
|
||||
patch("api.main.STORAGE_DIR", root),
|
||||
):
|
||||
response = self.client.get("/api/results/abcdef123456/preview")
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.headers["content-type"], "application/pdf")
|
||||
self.assertEqual(response.content, b"%PDF-1.7 preview")
|
||||
|
||||
def test_word_preview_returns_not_found_without_conversion(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = Path(tmp)
|
||||
uploads_dir = root / "uploads"
|
||||
doc_dir = uploads_dir / "abcdef123456"
|
||||
doc_dir.mkdir(parents=True)
|
||||
(doc_dir / "manual.doc").write_bytes(b"word")
|
||||
|
||||
with (
|
||||
patch("api.main.UPLOADS_DIR", uploads_dir),
|
||||
patch("api.main.STORAGE_DIR", root),
|
||||
):
|
||||
response = self.client.get("/api/results/abcdef123456/preview")
|
||||
|
||||
self.assertEqual(response.status_code, 404)
|
||||
|
||||
def test_delete_historical_result_removes_files(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = Path(tmp)
|
||||
results_dir = root / "results"
|
||||
uploads_dir = root / "uploads"
|
||||
assets_dir = root / "assets" / "abcdef123456"
|
||||
results_dir.mkdir()
|
||||
uploads_dir.mkdir()
|
||||
assets_dir.mkdir(parents=True)
|
||||
(results_dir / "abcdef123456.json").write_text(result().model_dump_json(), encoding="utf-8")
|
||||
upload_doc = uploads_dir / "abcdef123456"
|
||||
upload_doc.mkdir()
|
||||
(upload_doc / "manual.txt").write_text("hello", encoding="utf-8")
|
||||
(assets_dir / "img.png").write_bytes(b"png")
|
||||
|
||||
with (
|
||||
patch("api.main.RESULTS_DIR", results_dir),
|
||||
patch("api.main.UPLOADS_DIR", uploads_dir),
|
||||
patch("api.main.STORAGE_DIR", root),
|
||||
):
|
||||
deleted = self.client.delete("/api/results/abcdef123456")
|
||||
self.assertEqual(deleted.status_code, 200)
|
||||
body = deleted.json()
|
||||
self.assertTrue(body["deleted"])
|
||||
self.assertTrue(body["removed"]["result"])
|
||||
self.assertTrue(body["removed"]["uploads"])
|
||||
self.assertTrue(body["removed"]["assets"])
|
||||
|
||||
self.assertFalse((results_dir / "abcdef123456.json").exists())
|
||||
self.assertFalse(upload_doc.exists())
|
||||
self.assertFalse(assets_dir.exists())
|
||||
|
||||
listed = self.client.get("/api/results")
|
||||
self.assertEqual(listed.json()["count"], 0)
|
||||
|
||||
missing = self.client.delete("/api/results/abcdef123456")
|
||||
self.assertEqual(missing.status_code, 404)
|
||||
|
||||
def test_delete_rejects_invalid_doc_id(self) -> None:
|
||||
response = self.client.delete("/api/results/not-a-valid")
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_get_original_rejects_invalid_doc_id(self) -> None:
|
||||
response = self.client.get("/api/results/../secret/original")
|
||||
self.assertIn(response.status_code, (400, 404))
|
||||
|
||||
def test_chunk_accepts_manual_mode_and_preserves_zero_overlap(self) -> None:
|
||||
captured: dict = {}
|
||||
|
||||
def fake_chunk(path, config=None):
|
||||
captured["config"] = config
|
||||
return result(path.name)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp, patch("api.main.RESULTS_DIR", Path(tmp)):
|
||||
with patch("api.main.chunk_document", side_effect=fake_chunk):
|
||||
response = self.client.post(
|
||||
"/api/chunk",
|
||||
files={"file": ("sheet.csv", b"a,b\n1,2\n", "text/csv")},
|
||||
data={
|
||||
"mode": "by_row",
|
||||
"overlap": "0",
|
||||
"max_chunk_size": "2400",
|
||||
"header_row_start": "1",
|
||||
"header_row_end": "1",
|
||||
"start_row": "2",
|
||||
"rows_per_chunk": "5",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
config = captured["config"]
|
||||
self.assertIsNotNone(config)
|
||||
self.assertEqual(config.mode, SplitMode.BY_ROW)
|
||||
self.assertEqual(config.overlap, 0)
|
||||
self.assertEqual(config.rows_per_chunk, 5)
|
||||
|
||||
def test_chunk_accepts_parent_child_mode_fields(self) -> None:
|
||||
captured: dict = {}
|
||||
|
||||
def fake_chunk(path, config=None):
|
||||
captured["config"] = config
|
||||
return result(path.name)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp, patch("api.main.RESULTS_DIR", Path(tmp)):
|
||||
with patch("api.main.chunk_document", side_effect=fake_chunk):
|
||||
response = self.client.post(
|
||||
"/api/chunk",
|
||||
files={"file": ("guide.md", b"a##b###c", "text/markdown")},
|
||||
data={
|
||||
"mode": "parent_child",
|
||||
"parent_delimiter": "##",
|
||||
"child_delimiter": "###",
|
||||
"max_chunk_size": "2000",
|
||||
"child_max_size": "512",
|
||||
"overlap": "0",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
config = captured["config"]
|
||||
self.assertEqual(config.mode, SplitMode.PARENT_CHILD)
|
||||
self.assertEqual(config.parent_delimiter, "##")
|
||||
self.assertEqual(config.child_delimiter, "###")
|
||||
self.assertEqual(config.child_max_size, 512)
|
||||
self.assertEqual(config.overlap, 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,290 @@
|
||||
"""Tests for generic heading/layout multimodal splitting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from rag_cut.models import Block, BlockType, SplitConfig, SplitMode
|
||||
from rag_cut.pipeline import chunk_document
|
||||
from rag_cut.splitters.pdf_semantic import CHUNK_STRATEGY, split_pdf_semantic
|
||||
|
||||
|
||||
def h(text: str, page: int = 1, level: int = 0, y: int = 100) -> Block:
|
||||
return Block(
|
||||
type=BlockType.HEADING if level else BlockType.PARAGRAPH,
|
||||
text=text,
|
||||
level=level,
|
||||
meta={"page": page, "bbox": [72, y, 520, y + 20], "page_height": 800},
|
||||
)
|
||||
|
||||
|
||||
def p(text: str, page: int = 1, y: int = 130) -> Block:
|
||||
return Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text=text,
|
||||
meta={"page": page, "bbox": [72, y, 520, y + 20], "page_height": 800},
|
||||
)
|
||||
|
||||
|
||||
def img(name: str, page: int = 1, y: int = 180) -> Block:
|
||||
return Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id=name,
|
||||
image_path=f"assets/test/{name}.png",
|
||||
ocr_text="Login Submit",
|
||||
meta={"page": page, "bbox": [100, y, 500, y + 180], "page_height": 800},
|
||||
)
|
||||
|
||||
|
||||
def table(page: int = 1, y: int = 220) -> Block:
|
||||
return Block(
|
||||
type=BlockType.TABLE,
|
||||
markdown="| Field | Description |\n| --- | --- |\n| status | Account status |",
|
||||
image_id="table1",
|
||||
image_path="assets/test/table1.png",
|
||||
meta={
|
||||
"page": page,
|
||||
"pages": [page],
|
||||
"bbox": [72, y, 520, y + 120],
|
||||
"bboxes": [{"page": page, "bbox": [72, y, 520, y + 120]}],
|
||||
"row_count": 2,
|
||||
"col_count": 2,
|
||||
"cross_page": False,
|
||||
"crop_path": "assets/test/table1.png",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class HeadingLayoutMultimodalTest(unittest.TestCase):
|
||||
def test_numbered_headings_create_same_level_boundaries(self) -> None:
|
||||
chunks = split_pdf_semantic(
|
||||
[
|
||||
h("1 Overview"),
|
||||
p("Intro text."),
|
||||
h("1.1 Account Status", y=180),
|
||||
p("Account status details.", y=210),
|
||||
h("1.2 Trade Detail", y=260),
|
||||
p("Trade detail body.", y=290),
|
||||
],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
headings = [c.meta["heading"] for c in chunks]
|
||||
self.assertEqual(headings, ["1 Overview", "1.1 Account Status", "1.2 Trade Detail"])
|
||||
self.assertTrue(all(c.meta["chunk_strategy"] == CHUNK_STRATEGY for c in chunks))
|
||||
|
||||
def test_letter_subheadings_bind_nearest_images(self) -> None:
|
||||
chunks = split_pdf_semantic(
|
||||
[
|
||||
h("1 Getting Started"),
|
||||
h("A. Log in", y=140),
|
||||
p("Open the login page.", y=170),
|
||||
img("login", y=210),
|
||||
h("B) Authentication", y=430),
|
||||
p("Approve the authentication request.", y=460),
|
||||
img("auth", y=500),
|
||||
],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
self.assertEqual([c.meta["heading"] for c in chunks], ["A. Log in", "B) Authentication"])
|
||||
self.assertEqual(chunks[0].meta["heading_path"], ["1 Getting Started", "A. Log in"])
|
||||
self.assertEqual(chunks[0].meta["images"][0]["image_id"], "login")
|
||||
self.assertEqual(chunks[1].meta["images"][0]["image_id"], "auth")
|
||||
|
||||
def test_table_stays_atomic_with_markdown_and_position(self) -> None:
|
||||
chunks = split_pdf_semantic(
|
||||
[h("2 Settings"), p("The fields are listed below."), table(y=180)],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
self.assertEqual(len(chunks), 1)
|
||||
self.assertIn("| Field | Description |", chunks[0].content)
|
||||
self.assertEqual(chunks[0].meta["tables"][0]["row_count"], 2)
|
||||
self.assertEqual(chunks[0].meta["tables"][0]["bbox"], [72, 180, 520, 300])
|
||||
|
||||
def test_unlabelled_image_only_chunk_is_preview_only(self) -> None:
|
||||
cover_logo = img("cover-logo")
|
||||
cover_logo.ocr_text = ""
|
||||
|
||||
chunks = split_pdf_semantic(
|
||||
[cover_logo, h("1 Overview", page=2), p("Useful body.", page=2)],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
self.assertEqual(len(chunks), 2)
|
||||
self.assertFalse(chunks[0].meta["retrieval"])
|
||||
self.assertTrue(chunks[1].meta["retrieval"])
|
||||
self.assertIn("cover-logo", chunks[0].content)
|
||||
|
||||
def test_image_only_chunk_with_ocr_remains_retrievable(self) -> None:
|
||||
chunks = split_pdf_semantic(
|
||||
[img("workflow-screenshot")],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
self.assertEqual(len(chunks), 1)
|
||||
self.assertTrue(chunks[0].meta["retrieval"])
|
||||
|
||||
def test_noise_blocks_are_skipped(self) -> None:
|
||||
blocks = []
|
||||
for page in range(1, 5):
|
||||
blocks.append(p("Product Manual", page=page, y=20))
|
||||
blocks.append(p(str(page), page=page, y=760))
|
||||
blocks.extend(
|
||||
[
|
||||
p("Contents"),
|
||||
p("1 Intro ........ 3"),
|
||||
p("2 Setup ........ 8"),
|
||||
h("1 Intro", page=2),
|
||||
p("Useful body.", page=2),
|
||||
]
|
||||
)
|
||||
|
||||
chunks = split_pdf_semantic(blocks, SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000))
|
||||
|
||||
content = "\n".join(c.content for c in chunks)
|
||||
self.assertIn("Useful body.", content)
|
||||
self.assertNotIn("Product Manual", content)
|
||||
self.assertNotIn("Contents", content)
|
||||
self.assertNotIn("........", content)
|
||||
self.assertNotIn("2 Setup", content)
|
||||
|
||||
def test_cross_page_same_heading_is_one_chunk(self) -> None:
|
||||
chunks = split_pdf_semantic(
|
||||
[h("3 Client"), p("Page one text.", page=1), p("Page two continuation.", page=2)],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
)
|
||||
|
||||
self.assertEqual(len(chunks), 1)
|
||||
self.assertEqual(chunks[0].meta["pages"], [1, 2])
|
||||
self.assertIn("Page two continuation.", chunks[0].content)
|
||||
|
||||
def test_instructional_steps_stay_in_parent_section(self) -> None:
|
||||
"""Step 1/2/3 lines must not become headings that get flushed away."""
|
||||
chunks = split_pdf_semantic(
|
||||
[
|
||||
h("5.7 CRS SUBMISSION TO IRD", level=2),
|
||||
h("5.7.1 Submission at 1st time", level=3, y=140),
|
||||
p(
|
||||
"For CRS submission at first time by G3SB, prior consent has to be obtained "
|
||||
"from IRD by submitting test data file to the AEOI Portal for validation.",
|
||||
y=170,
|
||||
),
|
||||
img("workflow", y=210),
|
||||
p("Step 1 Users prepare the XML as per sections 5.1– 5.5.", y=420),
|
||||
p('Step 2 Set parameter "Export in test data format?" value to "Y".', y=450),
|
||||
p("Step 3 Check XML file tag DocTypeIndic was using OECD11.", y=480),
|
||||
p(
|
||||
"Step 4 Go to Registration and login page of AEOI Portal: "
|
||||
"https://aeoi1.ird.gov.hk/portal/landing/",
|
||||
y=510,
|
||||
),
|
||||
p("Download the encryption tools to encrypt the XML file.", y=540),
|
||||
],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=4000),
|
||||
)
|
||||
|
||||
content = "\n".join(c.content for c in chunks)
|
||||
self.assertIn("Step 1 Users prepare the XML", content)
|
||||
self.assertIn("Export in test data format", content)
|
||||
self.assertIn("OECD11", content)
|
||||
self.assertIn("Step 4 Go to Registration", content)
|
||||
self.assertTrue(any(c.meta["heading"] == "5.7.1 Submission at 1st time" for c in chunks))
|
||||
|
||||
def test_oversized_section_splits_by_length_without_parent(self) -> None:
|
||||
blocks = [h("4 Reports")] + [p(f"Long paragraph {i} " + "x" * 80, y=130 + i) for i in range(8)]
|
||||
chunks = split_pdf_semantic(blocks, SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=260))
|
||||
|
||||
self.assertGreater(len(chunks), 1)
|
||||
self.assertTrue(all(c.meta.get("retrieval", True) for c in chunks))
|
||||
self.assertFalse(any(c.meta.get("is_section_parent") for c in chunks))
|
||||
self.assertFalse(any(c.meta.get("is_sub_chunk") for c in chunks))
|
||||
self.assertFalse(any(c.meta.get("parent_chunk_id") is not None for c in chunks))
|
||||
|
||||
def test_catalog_titles_fold_into_section_not_dropped(self) -> None:
|
||||
"""TOC-style '1.xxx' lines between catalog headings must not vanish on flush."""
|
||||
chunks = split_pdf_semantic(
|
||||
[
|
||||
h("14.W. Whitney George", level=1, y=100),
|
||||
p("Famous fund manager bio.", y=130),
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="形態指標",
|
||||
level=1,
|
||||
meta={
|
||||
"page": 1,
|
||||
"bbox": [72, 200, 200, 220],
|
||||
"page_height": 800,
|
||||
"font_size": 16,
|
||||
"body_font_size": 12,
|
||||
},
|
||||
),
|
||||
p("共包含以下11個形態指標說明:", y=240),
|
||||
h("1.頭肩頂形態", level=1, y=260),
|
||||
h("2.頭肩底形態", level=1, y=280),
|
||||
h("1.頭肩頂形態", level=1, page=2, y=120),
|
||||
p("整體介紹:頂部反轉形態說明。", page=2, y=150),
|
||||
],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=4000),
|
||||
)
|
||||
|
||||
morph = next(c for c in chunks if c.meta.get("heading") == "形態指標")
|
||||
self.assertEqual(morph.meta.get("heading_path"), ["形態指標"])
|
||||
self.assertIn("共包含以下11個形態指標說明:", morph.content)
|
||||
self.assertIn("1.頭肩頂形態", morph.content)
|
||||
self.assertIn("2.頭肩底形態", morph.content)
|
||||
|
||||
detail = next(
|
||||
c for c in chunks if c.meta.get("heading") == "1.頭肩頂形態" and "整體介紹" in c.content
|
||||
)
|
||||
self.assertIn("整體介紹", detail.content)
|
||||
|
||||
def test_empty_explicit_section_does_not_pollute_previous_chunk(self) -> None:
|
||||
empty_heading = h("2 REFERENCES", level=1, page=3, y=100)
|
||||
empty_heading.meta["source_heading"] = True
|
||||
chunks = split_pdf_semantic(
|
||||
[
|
||||
h("1 GLOSSARY", level=1, page=2, y=100),
|
||||
table(page=2, y=140),
|
||||
empty_heading,
|
||||
h("3 DOCUMENT HISTORY", level=1, page=3, y=180),
|
||||
p("History body.", page=3, y=220),
|
||||
],
|
||||
SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=4000),
|
||||
)
|
||||
|
||||
glossary = next(c for c in chunks if c.meta.get("heading") == "1 GLOSSARY")
|
||||
self.assertNotIn("2 REFERENCES", glossary.content)
|
||||
self.assertFalse(any("2 REFERENCES" in c.content for c in chunks))
|
||||
|
||||
|
||||
class DummyParser:
|
||||
def parse(self, path: Path, assets_dir: Path) -> list[Block]:
|
||||
return [h("1 Overview"), p("Pipeline body.")]
|
||||
|
||||
|
||||
class PipelineSemanticIntegrationTest(unittest.TestCase):
|
||||
def test_pdf_default_mode_uses_heading_layout_multimodal(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
tmp_root = Path(tmp)
|
||||
source = tmp_root / "sample.pdf"
|
||||
source.write_bytes(b"%PDF-1.4\n")
|
||||
|
||||
with patch("rag_cut.pipeline.get_parser", return_value=DummyParser()):
|
||||
result = chunk_document(
|
||||
source,
|
||||
config=SplitConfig(mode=SplitMode.DEFAULT, max_chunk_size=2000),
|
||||
storage_root=tmp_root / "storage",
|
||||
)
|
||||
|
||||
self.assertEqual(result.chunk_count, 1)
|
||||
self.assertEqual(result.split_config["chunk_strategy"], CHUNK_STRATEGY)
|
||||
self.assertEqual(result.chunks[0].meta["chunk_strategy"], CHUNK_STRATEGY)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Tests for PDF text merge and heading classification."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from rag_cut.layout_meta import enrich_layout_metadata
|
||||
from rag_cut.models import Block, BlockType
|
||||
from rag_cut.parsers.pdf.layout import TextLine, _merge_lines_to_paragraphs
|
||||
from rag_cut.parsers.pdf.text_extract import MAX_HEADING_CHARS, _font_heading_level, _looks_like_ui_label
|
||||
from rag_cut.parsers.pdf.noise_filter import is_margin_noise_block
|
||||
from rag_cut.splitters.heading_splitter import is_heading_block, normalize_heading_block, numbered_heading_level
|
||||
|
||||
|
||||
class ParagraphMergeTest(unittest.TestCase):
|
||||
def test_does_not_merge_title_and_body_with_different_font_sizes(self) -> None:
|
||||
lines = [
|
||||
TextLine(x0=90, y0=90, x1=400, y1=110, text="Master Collection/Pattern Scanning", font_size=18),
|
||||
TextLine(x0=90, y0=115, x1=480, y1=132, text="及所有大師/形態/策略指標的位置", font_size=17.5),
|
||||
TextLine(
|
||||
x0=90,
|
||||
y0=145,
|
||||
x1=500,
|
||||
y1=200,
|
||||
text="頁面左側【篩選】欄下方即是【Master Collection】文件夾。點擊此檔夾即可出現下拉頁面。",
|
||||
font_size=11,
|
||||
),
|
||||
]
|
||||
merged = _merge_lines_to_paragraphs(lines, page_width=600)
|
||||
texts = [ln.text for ln in merged]
|
||||
self.assertEqual(len(merged), 2)
|
||||
self.assertIn("Master Collection", texts[0])
|
||||
self.assertTrue(texts[1].startswith("頁面左側"))
|
||||
self.assertNotIn("頁面左側", texts[0])
|
||||
|
||||
|
||||
class FontHeadingTest(unittest.TestCase):
|
||||
def test_short_large_text_is_heading(self) -> None:
|
||||
level = _font_heading_level(18, 11, "Master Collection 位置說明")
|
||||
self.assertIsNotNone(level)
|
||||
|
||||
def test_title_plus_body_blob_is_not_heading(self) -> None:
|
||||
blob = (
|
||||
"Master Collection/Pattern Scanning/Strategy Scanning 及所有大師/形態/策略指標的位置 "
|
||||
"頁面左側【篩選】欄下方即是【Master Collection】文件夾。點擊此檔夾即可出現下拉頁面,"
|
||||
"其中包含了十五個大師的選股策略。如圖所示:"
|
||||
)
|
||||
self.assertGreater(len(blob), MAX_HEADING_CHARS)
|
||||
self.assertIsNone(_font_heading_level(18, 11, blob))
|
||||
|
||||
def test_compact_numbered_chinese_title_is_heading(self) -> None:
|
||||
# FAQ style: "10.上升三角形態" with no space after the dot.
|
||||
level = _font_heading_level(16, 12, "10.上升三角形態")
|
||||
self.assertEqual(level, 1)
|
||||
self.assertFalse(_looks_like_ui_label("10.上升三角形態"))
|
||||
|
||||
def test_cjk_section_banner_is_level1_not_ui_label(self) -> None:
|
||||
self.assertFalse(_looks_like_ui_label("形態指標"))
|
||||
self.assertEqual(_font_heading_level(16, 12, "形態指標"), 1)
|
||||
|
||||
|
||||
|
||||
class HeadingDemoteTest(unittest.TestCase):
|
||||
def test_normalize_demotes_overlong_heading(self) -> None:
|
||||
blob = "T" * (MAX_HEADING_CHARS + 20)
|
||||
block = Block(type=BlockType.HEADING, text=blob, level=1, meta={"page": 1})
|
||||
normalized = normalize_heading_block(block)
|
||||
self.assertEqual(normalized.type, BlockType.PARAGRAPH)
|
||||
self.assertFalse(is_heading_block(normalized))
|
||||
|
||||
def test_enrich_demotes_overlong_heading_block(self) -> None:
|
||||
blob = (
|
||||
"Master Collection/Pattern Scanning/Strategy Scanning Master Collection/"
|
||||
"Pattern Scanning/Strategy Scanning 及所有大師 /形態/策略指標的位置 "
|
||||
"頁面左側【篩選】欄下方即是【Master Collection】文件夾。點擊此檔夾即可出 "
|
||||
"現下拉頁面,其中包含了十五個大師的選股策略。"
|
||||
)
|
||||
blocks = [
|
||||
Block(type=BlockType.HEADING, text=blob, level=1, meta={"page": 1, "bbox": [90, 90, 500, 240]}),
|
||||
Block(type=BlockType.IMAGE, image_id="a.png", image_path="a.png", meta={"page": 1, "bbox": [90, 250, 500, 480]}),
|
||||
]
|
||||
enriched = enrich_layout_metadata(blocks)
|
||||
self.assertEqual(enriched[0].type, BlockType.PARAGRAPH)
|
||||
self.assertNotEqual(enriched[1].meta.get("bound_heading"), blob)
|
||||
|
||||
def test_top_of_page_numbered_title_not_margin_noise(self) -> None:
|
||||
block = Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="10.上升三角形態",
|
||||
meta={"page": 50, "bbox": [90, 74, 209, 90], "page_height": 842},
|
||||
)
|
||||
self.assertFalse(is_margin_noise_block(block, 842))
|
||||
self.assertEqual(numbered_heading_level("10.上升三角形態"), 1)
|
||||
self.assertTrue(is_heading_block(normalize_heading_block(block)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,166 @@
|
||||
"""Tests for MinerU content_list → Block mapping."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import subprocess
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from rag_cut.models import BlockType
|
||||
from rag_cut.parsers.mineru_adapter import (
|
||||
_blocks_from_content_list,
|
||||
_find_content_list,
|
||||
_run_mineru,
|
||||
)
|
||||
|
||||
|
||||
class MineruAdapterProcessTest(unittest.TestCase):
|
||||
def test_default_timeout_allows_large_manual_to_finish(self) -> None:
|
||||
with patch.dict("os.environ", {}, clear=True):
|
||||
from rag_cut.parsers.mineru_adapter import _mineru_timeout_sec
|
||||
|
||||
self.assertEqual(_mineru_timeout_sec(), 540.0)
|
||||
|
||||
def test_configured_api_url_is_passed_to_mineru(self) -> None:
|
||||
process = Mock()
|
||||
process.communicate.return_value = ("", "")
|
||||
process.returncode = 0
|
||||
|
||||
with (
|
||||
tempfile.TemporaryDirectory() as tmp,
|
||||
patch("rag_cut.parsers.mineru_adapter._mineru_command", return_value="mineru"),
|
||||
patch.dict("os.environ", {"RAG_CUT_MINERU_API_URL": "http://127.0.0.1:30000"}),
|
||||
patch("rag_cut.parsers.mineru_adapter.subprocess.Popen", return_value=process) as popen,
|
||||
):
|
||||
ok = _run_mineru(Path(tmp) / "manual.pdf", Path(tmp) / "output")
|
||||
|
||||
self.assertTrue(ok)
|
||||
args = popen.call_args.args[0]
|
||||
self.assertEqual(args[-2:], ["--api-url", "http://127.0.0.1:30000"])
|
||||
|
||||
def test_timeout_terminates_windows_process_tree(self) -> None:
|
||||
process = Mock()
|
||||
process.pid = 4321
|
||||
process.communicate.side_effect = subprocess.TimeoutExpired("mineru", 30)
|
||||
process.poll.return_value = None
|
||||
|
||||
with (
|
||||
tempfile.TemporaryDirectory() as tmp,
|
||||
patch("rag_cut.parsers.mineru_adapter._mineru_command", return_value="mineru"),
|
||||
patch("rag_cut.parsers.mineru_adapter._mineru_timeout_sec", return_value=30),
|
||||
patch("rag_cut.parsers.mineru_adapter.os.name", "nt"),
|
||||
patch("rag_cut.parsers.mineru_adapter.subprocess.Popen", return_value=process),
|
||||
patch("rag_cut.parsers.mineru_adapter.subprocess.run") as run,
|
||||
):
|
||||
ok = _run_mineru(Path(tmp) / "manual.pdf", Path(tmp) / "output")
|
||||
|
||||
self.assertFalse(ok)
|
||||
run.assert_called_once_with(
|
||||
["taskkill", "/PID", "4321", "/T", "/F"],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
|
||||
class MineruAdapterMappingTest(unittest.TestCase):
|
||||
def test_chart_becomes_image_block(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
assets = Path(tmp)
|
||||
img = assets / "chart.jpg"
|
||||
img.write_bytes(b"fake")
|
||||
items = [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "第二十節 TradingView",
|
||||
"text_level": 2,
|
||||
"page_idx": 30,
|
||||
"bbox": [127, 86, 363, 104],
|
||||
},
|
||||
{
|
||||
"type": "chart",
|
||||
"img_path": "chart.jpg",
|
||||
"page_idx": 30,
|
||||
"bbox": [129, 171, 878, 428],
|
||||
"image_caption": [],
|
||||
},
|
||||
]
|
||||
with patch("rag_cut.parsers.mineru_adapter._mineru_ocr_enabled", return_value=False):
|
||||
blocks = _blocks_from_content_list(items, assets / "content_list.json", assets)
|
||||
|
||||
types = [b.type for b in blocks]
|
||||
self.assertEqual(types, [BlockType.HEADING, BlockType.IMAGE])
|
||||
self.assertEqual(blocks[1].image_id, "chart.jpg")
|
||||
self.assertEqual(blocks[1].meta.get("mineru_type"), "chart")
|
||||
self.assertEqual(blocks[1].meta.get("page"), 31)
|
||||
|
||||
def test_table_keeps_markdown_and_screenshot(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
assets = Path(tmp)
|
||||
img = assets / "table.png"
|
||||
img.write_bytes(b"png")
|
||||
items = [
|
||||
{
|
||||
"type": "table",
|
||||
"img_path": "table.png",
|
||||
"page_idx": 1,
|
||||
"bbox": [100, 100, 400, 300],
|
||||
"table_body": "<table><tr><td>A</td><td>B</td></tr></table>",
|
||||
}
|
||||
]
|
||||
with patch("rag_cut.parsers.mineru_adapter._mineru_ocr_enabled", return_value=False):
|
||||
blocks = _blocks_from_content_list(items, assets / "content_list.json", assets)
|
||||
|
||||
self.assertEqual(len(blocks), 1)
|
||||
self.assertEqual(blocks[0].type, BlockType.TABLE)
|
||||
self.assertIn("<table>", blocks[0].markdown or "")
|
||||
self.assertEqual(blocks[0].image_id, "table.png")
|
||||
self.assertTrue(blocks[0].image_path)
|
||||
|
||||
def test_page_footnote_kept_as_paragraph(self) -> None:
|
||||
items = [
|
||||
{
|
||||
"type": "page_footnote",
|
||||
"text": "<sup>1</sup> See Market Master.",
|
||||
"page_idx": 8,
|
||||
"bbox": [85, 889, 912, 916],
|
||||
}
|
||||
]
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
blocks = _blocks_from_content_list(items, Path(tmp) / "content_list.json", Path(tmp))
|
||||
self.assertEqual(len(blocks), 1)
|
||||
self.assertEqual(blocks[0].type, BlockType.PARAGRAPH)
|
||||
self.assertTrue(blocks[0].meta.get("is_footnote"))
|
||||
|
||||
def test_list_items_become_paragraph(self) -> None:
|
||||
items = [
|
||||
{
|
||||
"type": "list",
|
||||
"list_items": ["第一點說明", {"text": "第二點說明"}],
|
||||
"page_idx": 2,
|
||||
"bbox": [100, 200, 400, 260],
|
||||
}
|
||||
]
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
blocks = _blocks_from_content_list(items, Path(tmp) / "content_list.json", Path(tmp))
|
||||
self.assertEqual(len(blocks), 1)
|
||||
self.assertEqual(blocks[0].type, BlockType.PARAGRAPH)
|
||||
self.assertIn("第一點說明", blocks[0].text or "")
|
||||
self.assertIn("第二點說明", blocks[0].text or "")
|
||||
|
||||
def test_find_content_list_prefers_non_v2(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = Path(tmp) / "doc" / "auto"
|
||||
root.mkdir(parents=True)
|
||||
v1 = root / "doc_content_list.json"
|
||||
v2 = root / "doc_content_list_v2.json"
|
||||
v1.write_text("[]", encoding="utf-8")
|
||||
v2.write_text("[]", encoding="utf-8")
|
||||
found = _find_content_list(Path(tmp))
|
||||
self.assertEqual(found, v1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Tests for parent/child delimiter splitting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from rag_cut.models import Block, BlockType, SplitConfig, SplitMode
|
||||
from rag_cut.pipeline import chunk_document
|
||||
from rag_cut.splitters.parent_child import split_by_parent_child
|
||||
|
||||
|
||||
def p(text: str) -> Block:
|
||||
return Block(type=BlockType.PARAGRAPH, text=text)
|
||||
|
||||
|
||||
class ParentChildSplitTest(unittest.TestCase):
|
||||
def test_parent_and_child_chunks_are_linked(self) -> None:
|
||||
blocks = [p("Intro A##Detail A1###Detail A2##Intro B###Detail B1")]
|
||||
groups = split_by_parent_child(
|
||||
blocks,
|
||||
SplitConfig(
|
||||
mode=SplitMode.PARENT_CHILD,
|
||||
parent_delimiter="##",
|
||||
child_delimiter="###",
|
||||
max_chunk_size=2000,
|
||||
child_max_size=500,
|
||||
overlap=0,
|
||||
),
|
||||
)
|
||||
|
||||
parents = [g for g in groups if g.meta.get("is_section_parent")]
|
||||
children = [g for g in groups if g.meta.get("is_sub_chunk")]
|
||||
self.assertGreaterEqual(len(parents), 2)
|
||||
self.assertGreaterEqual(len(children), 2)
|
||||
self.assertTrue(all(g.meta.get("retrieval") is False for g in parents))
|
||||
self.assertTrue(all(g.meta.get("retrieval") is True for g in children))
|
||||
self.assertTrue(all(g.meta.get("parent_section_id") for g in children))
|
||||
|
||||
def test_requires_parent_delimiter(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
split_by_parent_child(
|
||||
[p("hello")],
|
||||
SplitConfig(mode=SplitMode.PARENT_CHILD, parent_delimiter=None),
|
||||
)
|
||||
|
||||
def test_pipeline_persists_parent_chunk_ids(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
source = Path(tmp) / "pc.md"
|
||||
source.write_text(
|
||||
"Alpha parent##A1 child###A2 child##Beta parent###B1 child",
|
||||
encoding="utf-8",
|
||||
)
|
||||
result = chunk_document(
|
||||
source,
|
||||
config=SplitConfig(
|
||||
mode=SplitMode.PARENT_CHILD,
|
||||
parent_delimiter="##",
|
||||
child_delimiter="###",
|
||||
max_chunk_size=2000,
|
||||
child_max_size=400,
|
||||
overlap=0,
|
||||
),
|
||||
storage_root=Path(tmp) / "storage",
|
||||
)
|
||||
|
||||
self.assertEqual(result.split_mode, SplitMode.PARENT_CHILD)
|
||||
self.assertEqual(result.split_config.get("chunk_strategy"), "parent_child_delimiter")
|
||||
parents = [c for c in result.chunks if c.meta.get("is_section_parent")]
|
||||
children = [c for c in result.chunks if c.meta.get("is_sub_chunk")]
|
||||
self.assertTrue(parents)
|
||||
self.assertTrue(children)
|
||||
for child in children:
|
||||
self.assertIn("parent_chunk_id", child.meta)
|
||||
self.assertIsInstance(child.meta["parent_chunk_id"], int)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Tests for PDF margin noise filtering."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from rag_cut.layout_meta import enrich_layout_metadata, sort_blocks_reading_order
|
||||
from rag_cut.models import Block, BlockType
|
||||
from rag_cut.parsers.mineru_adapter import _blocks_from_content_list
|
||||
from rag_cut.parsers.pdf.noise_filter import (
|
||||
detect_running_header_texts,
|
||||
filter_noise_blocks,
|
||||
filter_toc_blocks,
|
||||
is_margin_noise_block,
|
||||
is_tiny_image_block,
|
||||
is_toc_entry_line,
|
||||
is_toc_title_text,
|
||||
)
|
||||
|
||||
|
||||
class MineruNoiseFilterTest(unittest.TestCase):
|
||||
def test_skips_header_and_page_number_types(self) -> None:
|
||||
items = [
|
||||
{"type": "text", "text": "目錄", "text_level": 2, "page_idx": 1},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "1. 登入 ...3\n2. 股票報價. ....4",
|
||||
"page_idx": 1,
|
||||
"bbox": [114, 190, 883, 517],
|
||||
},
|
||||
{"type": "header", "text": "H5 i-Trade 用戶使用手冊", "page_idx": 1},
|
||||
{"type": "header", "text": "NEN WAFE:SOLUTIONS", "page_idx": 1},
|
||||
{"type": "page_number", "text": "2", "page_idx": 1},
|
||||
]
|
||||
blocks = _blocks_from_content_list(items, Path("content_list.json"), Path("assets"))
|
||||
# MinerU adapter may still emit 目录 text; pipeline noise filter removes it.
|
||||
texts = [block.text for block in blocks]
|
||||
self.assertIn("目錄", texts)
|
||||
self.assertEqual(texts.count("H5 i-Trade 用戶使用手冊"), 0)
|
||||
self.assertEqual(texts.count("2"), 0)
|
||||
filtered = filter_noise_blocks(blocks)
|
||||
filtered_texts = [block.text for block in filtered]
|
||||
self.assertNotIn("目錄", filtered_texts)
|
||||
self.assertFalse(any("登入" in (t or "") for t in filtered_texts))
|
||||
|
||||
|
||||
class NoiseFilterHeuristicTest(unittest.TestCase):
|
||||
def test_margin_page_number_is_noise(self) -> None:
|
||||
block = Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="14",
|
||||
meta={"page": 8, "bbox": [867, 942, 882, 954], "page_height": 1000},
|
||||
)
|
||||
self.assertTrue(is_margin_noise_block(block, 1000))
|
||||
|
||||
def test_body_paragraph_is_kept(self) -> None:
|
||||
block = Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="開啟AFEH5 i-Trade網站,用戶無需登入即可查閱各項延遲15分鐘資訊。",
|
||||
meta={"page": 3, "bbox": [116, 200, 880, 260], "page_height": 1000},
|
||||
)
|
||||
self.assertFalse(is_margin_noise_block(block, 1000))
|
||||
|
||||
def test_running_header_detection(self) -> None:
|
||||
blocks = []
|
||||
for page in range(1, 11):
|
||||
blocks.append(
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="H5 i-Trade 用戶使用手冊",
|
||||
meta={"page": page, "bbox": [651, 79, 882, 97], "page_height": 1000},
|
||||
)
|
||||
)
|
||||
blocks.append(
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text=f"{page}. 章節",
|
||||
level=2,
|
||||
meta={"page": page, "bbox": [116, 122, 191, 146], "page_height": 1000},
|
||||
)
|
||||
)
|
||||
running = detect_running_header_texts(blocks)
|
||||
self.assertIn("H5 i-Trade 用戶使用手冊", running)
|
||||
self.assertNotIn("1. 章節", running)
|
||||
|
||||
def test_filter_noise_blocks_removes_margin_and_running_text(self) -> None:
|
||||
blocks = [
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="目錄",
|
||||
level=2,
|
||||
meta={"page": 2, "bbox": [119, 123, 176, 146], "page_height": 1000},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="1. 登入 ...3",
|
||||
meta={"page": 2, "bbox": [114, 190, 883, 517], "page_height": 1000},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="2",
|
||||
meta={"page": 2, "bbox": [868, 942, 880, 954], "page_height": 1000},
|
||||
),
|
||||
]
|
||||
for page in range(1, 6):
|
||||
blocks.append(
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="H5 i-Trade 用戶使用手冊",
|
||||
meta={"page": page, "bbox": [651, 79, 882, 97], "page_height": 1000},
|
||||
)
|
||||
)
|
||||
|
||||
filtered = filter_noise_blocks(blocks)
|
||||
texts = [block.text for block in filtered]
|
||||
self.assertNotIn("目錄", texts)
|
||||
self.assertNotIn("1. 登入 ...3", texts)
|
||||
self.assertNotIn("2", texts)
|
||||
self.assertNotIn("H5 i-Trade 用戶使用手冊", texts)
|
||||
|
||||
def test_filter_toc_drops_contents_keeps_in_section_catalog(self) -> None:
|
||||
blocks = [
|
||||
Block(type=BlockType.HEADING, text="Contents", level=1, meta={"page": 1, "bbox": [72, 80, 200, 100]}),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="1 Intro ........ 3",
|
||||
meta={"page": 1, "bbox": [72, 120, 400, 140]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="2 Setup ........ 8",
|
||||
meta={"page": 1, "bbox": [72, 150, 400, 170]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="1 Intro",
|
||||
level=1,
|
||||
meta={"page": 2, "bbox": [72, 100, 200, 120]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="Useful body paragraph describing the product workflow in detail.",
|
||||
meta={"page": 2, "bbox": [72, 140, 500, 200]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="形態指標",
|
||||
level=1,
|
||||
meta={"page": 45, "bbox": [72, 200, 200, 220]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="共包含以下11個形態指標說明:",
|
||||
meta={"page": 45, "bbox": [72, 240, 400, 260]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="1.頭肩頂形態",
|
||||
level=1,
|
||||
meta={"page": 45, "bbox": [72, 270, 220, 290]},
|
||||
),
|
||||
]
|
||||
self.assertTrue(is_toc_title_text("目錄"))
|
||||
self.assertTrue(is_toc_entry_line("1 Intro ........ 3"))
|
||||
self.assertFalse(is_toc_entry_line("1.頭肩頂形態"))
|
||||
|
||||
filtered = filter_toc_blocks(blocks)
|
||||
texts = [b.text for b in filtered]
|
||||
self.assertNotIn("Contents", texts)
|
||||
self.assertNotIn("1 Intro ........ 3", texts)
|
||||
self.assertIn("1 Intro", texts)
|
||||
self.assertIn("Useful body paragraph describing the product workflow in detail.", texts)
|
||||
self.assertIn("形態指標", texts)
|
||||
self.assertIn("1.頭肩頂形態", texts)
|
||||
|
||||
def test_filter_toc_drops_mineru_compacted_contents_block(self) -> None:
|
||||
contents = Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text=(
|
||||
"CONTENTS..\n"
|
||||
"1 GLOSSARY ..2\n"
|
||||
"5.1 AEOI ID SETTING.. 4"
|
||||
),
|
||||
meta={"page": 2, "mineru_type": "text"},
|
||||
)
|
||||
|
||||
self.assertTrue(is_toc_title_text("CONTENTS.."))
|
||||
self.assertTrue(is_toc_entry_line("1 GLOSSARY ..2"))
|
||||
self.assertTrue(is_toc_entry_line("5.1 AEOI ID SETTING.. 4"))
|
||||
self.assertEqual(filter_toc_blocks([contents]), [])
|
||||
|
||||
def test_filter_toc_keeps_real_section_below_directory_on_same_page(self) -> None:
|
||||
blocks = [
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="CONTENTS",
|
||||
level=2,
|
||||
meta={"page": 2, "bbox": [115, 129, 284, 152]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="CONTENTS..\n1 GLOSSARY ..2\n2 REFERENCES 3\n3 DOCUMENT HISTORY.. .3",
|
||||
meta={"page": 2, "bbox": [117, 167, 884, 718]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="1 GLOSSARY",
|
||||
level=2,
|
||||
meta={"page": 2, "bbox": [117, 771, 321, 793]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.TABLE,
|
||||
markdown="| Abbreviation | Description |\n| --- | --- |",
|
||||
meta={"page": 2, "bbox": [127, 809, 885, 891]},
|
||||
),
|
||||
]
|
||||
|
||||
filtered = filter_toc_blocks(blocks)
|
||||
self.assertEqual([block.type for block in filtered], [BlockType.HEADING, BlockType.TABLE])
|
||||
self.assertEqual(filtered[0].text, "1 GLOSSARY")
|
||||
|
||||
def test_tiny_logo_image_is_filtered_wide_flowchart_kept(self) -> None:
|
||||
logo = Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="logo",
|
||||
image_path="assets/logo.png",
|
||||
meta={"page": 14, "bbox": [121, 420, 156, 445], "page_height": 1000},
|
||||
)
|
||||
flowchart = Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="flow",
|
||||
image_path="assets/flow.png",
|
||||
meta={"page": 13, "bbox": [119, 614, 949, 689], "page_height": 1000},
|
||||
)
|
||||
self.assertTrue(is_tiny_image_block(logo))
|
||||
self.assertFalse(is_tiny_image_block(flowchart))
|
||||
|
||||
filtered = filter_noise_blocks([logo, flowchart])
|
||||
ids = [block.image_id for block in filtered]
|
||||
self.assertEqual(ids, ["flow"])
|
||||
|
||||
def test_nested_image_fragment_dropped(self) -> None:
|
||||
outer = Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="screenshot",
|
||||
image_path="assets/shot.png",
|
||||
meta={"page": 14, "bbox": [100, 400, 500, 700], "page_height": 1000},
|
||||
)
|
||||
nested = Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="icon",
|
||||
image_path="assets/icon.png",
|
||||
meta={"page": 14, "bbox": [200, 450, 280, 530], "page_height": 1000},
|
||||
)
|
||||
filtered = filter_noise_blocks([outer, nested])
|
||||
ids = [block.image_id for block in filtered]
|
||||
self.assertEqual(ids, ["screenshot"])
|
||||
|
||||
|
||||
class LayoutOrderAndBindTest(unittest.TestCase):
|
||||
def test_sort_puts_page_top_before_lower_fragment(self) -> None:
|
||||
misordered = [
|
||||
Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="logo",
|
||||
image_path="assets/logo.png",
|
||||
meta={"page": 14, "bbox": [121, 420, 156, 445]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="5.8 APPENDIX",
|
||||
meta={"page": 14, "bbox": [114, 85, 324, 107]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="5.9 FAQ",
|
||||
meta={"page": 14, "bbox": [114, 334, 231, 355]},
|
||||
),
|
||||
]
|
||||
ordered = sort_blocks_reading_order(misordered)
|
||||
texts = [(b.image_id or b.text) for b in ordered]
|
||||
self.assertEqual(texts, ["5.8 APPENDIX", "5.9 FAQ", "logo"])
|
||||
|
||||
def test_enrich_binds_image_to_spatial_heading(self) -> None:
|
||||
blocks = [
|
||||
Block(
|
||||
type=BlockType.HEADING,
|
||||
text="5.7.2Subsequent submissions",
|
||||
level=3,
|
||||
meta={"page": 13, "bbox": [115, 468, 465, 488]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.IMAGE,
|
||||
image_id="wrong_early",
|
||||
image_path="assets/x.png",
|
||||
meta={"page": 14, "bbox": [341, 513, 431, 648]},
|
||||
),
|
||||
Block(
|
||||
type=BlockType.PARAGRAPH,
|
||||
text="5.9 FAQ",
|
||||
meta={"page": 14, "bbox": [114, 334, 231, 355]},
|
||||
),
|
||||
]
|
||||
enriched = enrich_layout_metadata(blocks)
|
||||
faq_img = next(b for b in enriched if b.type == BlockType.IMAGE)
|
||||
self.assertEqual(faq_img.meta.get("bound_heading"), "5.9 FAQ")
|
||||
self.assertLess(
|
||||
next(b for b in enriched if b.text == "5.9 FAQ").meta["order_index"],
|
||||
faq_img.meta["order_index"],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Tests for recall ranking and retrieval eligibility."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from rag_cut.models import Chunk
|
||||
from rag_cut.retrieval import recall_chunks
|
||||
|
||||
|
||||
def chunk(index: int, content: str, *, retrieval: bool = True, heading: str = "") -> Chunk:
|
||||
return Chunk(
|
||||
index=index,
|
||||
content=content,
|
||||
char_count=len(content),
|
||||
block_types=["paragraph"],
|
||||
meta={"retrieval": retrieval, "heading": heading},
|
||||
)
|
||||
|
||||
|
||||
class RecallChunksTest(unittest.TestCase):
|
||||
def test_relevant_chunk_ranks_first(self) -> None:
|
||||
results, candidate_count = recall_chunks(
|
||||
"股票交易密码",
|
||||
[
|
||||
chunk(0, "登录后可以修改股票交易密码", heading="账户安全"),
|
||||
chunk(1, "年度报告及公司治理", heading="公司资料"),
|
||||
],
|
||||
top_k=2,
|
||||
)
|
||||
|
||||
self.assertEqual(candidate_count, 2)
|
||||
self.assertEqual(results[0]["chunk_index"], 0)
|
||||
self.assertGreater(results[0]["score"], 0)
|
||||
|
||||
def test_preview_only_chunks_are_not_candidates(self) -> None:
|
||||
results, candidate_count = recall_chunks(
|
||||
"logo",
|
||||
[chunk(0, "", retrieval=False), chunk(1, "Useful body")],
|
||||
)
|
||||
|
||||
self.assertEqual(candidate_count, 1)
|
||||
self.assertEqual(results, [])
|
||||
|
||||
def test_top_k_is_respected(self) -> None:
|
||||
results, _ = recall_chunks("account", [chunk(i, f"account details {i}") for i in range(5)], top_k=2)
|
||||
self.assertEqual(len(results), 2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Tests for spreadsheet layout detection and row splitting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from rag_cut.parsers.pdf.tables import detect_spreadsheet_layout, is_qa_style_table
|
||||
from rag_cut.pipeline import chunk_document
|
||||
from rag_cut.splitters.by_row import split_by_row
|
||||
from rag_cut.models import Block, BlockType, SplitConfig, SplitMode
|
||||
|
||||
|
||||
BACKMAN_HEADER = [
|
||||
[
|
||||
"Session会话(必填):用于标识1个对话",
|
||||
"query 用户输入(必填):消息内容",
|
||||
"用户ID(必填)",
|
||||
"使用的大语言模型(必填)",
|
||||
"要求AI回复的语言(必填)",
|
||||
"reference_output 标准答案(可选)",
|
||||
],
|
||||
["session", "query", "userid", "model", "lang", "reference_output"],
|
||||
["1", "账户余额是多少?", "SUPPORT", "gpt-4o-mini", "zh", "SELECT 1"],
|
||||
["2", "今日有哪些账户透支?", "SUPPORT", "gpt-4o-mini", "zh", "SELECT 2"],
|
||||
]
|
||||
|
||||
|
||||
class SpreadsheetLayoutTest(unittest.TestCase):
|
||||
def test_detect_template_description_and_header_rows(self) -> None:
|
||||
layout = detect_spreadsheet_layout(BACKMAN_HEADER)
|
||||
self.assertEqual(layout["preamble_rows"], 1)
|
||||
self.assertEqual(layout["header_row_start"], 2)
|
||||
self.assertEqual(layout["header_row_end"], 2)
|
||||
self.assertEqual(layout["data_start_row"], 3)
|
||||
self.assertTrue(is_qa_style_table(BACKMAN_HEADER, layout))
|
||||
|
||||
def test_split_by_row_uses_real_header_and_one_row_per_chunk(self) -> None:
|
||||
layout = detect_spreadsheet_layout(BACKMAN_HEADER)
|
||||
block = Block(
|
||||
type=BlockType.TABLE,
|
||||
markdown="",
|
||||
meta={"rows": BACKMAN_HEADER, "table_title": "testset", **layout},
|
||||
)
|
||||
groups = split_by_row(
|
||||
[block],
|
||||
SplitConfig(
|
||||
mode=SplitMode.BY_ROW,
|
||||
header_row_start=layout["header_row_start"],
|
||||
header_row_end=layout["header_row_end"],
|
||||
start_row=layout["data_start_row"],
|
||||
rows_per_chunk=1,
|
||||
),
|
||||
)
|
||||
self.assertEqual(len(groups), 2)
|
||||
first_md = groups[0][0].markdown
|
||||
self.assertIn("| session | query | userid | model | lang | reference_output |", first_md)
|
||||
self.assertIn("账户余额是多少?", first_md)
|
||||
self.assertNotIn("Session会话(必填)", first_md)
|
||||
|
||||
|
||||
class BackmanFixtureTest(unittest.TestCase):
|
||||
def test_backman_upload_chunks_one_row_each(self) -> None:
|
||||
from pathlib import Path
|
||||
|
||||
uploads = list((Path(__file__).resolve().parents[2] / "storage" / "uploads").rglob("BackmanAI*.xlsx"))
|
||||
if not uploads:
|
||||
self.skipTest("BackmanAI fixture not uploaded")
|
||||
result = chunk_document(uploads[0])
|
||||
self.assertEqual(result.split_config["rows_per_chunk"], 1)
|
||||
self.assertEqual(result.split_config["header_row_start"], 2)
|
||||
self.assertEqual(result.split_config["start_row"], 3)
|
||||
self.assertGreaterEqual(result.chunk_count, 40)
|
||||
first = result.chunks[0].content
|
||||
self.assertIn("| session | query |", first)
|
||||
self.assertNotIn("Session会话(必填)", first)
|
||||
self.assertLess(result.chunks[0].char_count, 1200)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user