Files
MemRelay/backend/tests/test_source_extractors.py

192 lines
6.8 KiB
Python

from __future__ import annotations
import gzip
import tarfile
import zipfile
from pathlib import Path
from types import SimpleNamespace
import pytest
from docx import Document
from odf import text
from odf.opendocument import OpenDocumentText
from openpyxl import Workbook
from PIL import Image
from pptx import Presentation
from pptx.util import Inches
from pypdf import PdfWriter
from memrelay import source_extractors
from memrelay.source_extractors import ExtractionError, ExtractionLimits, extract_file
def extracted_text(path: Path, cache_dir: Path) -> tuple[str, str, list[str]]:
result = extract_file(path, cache_dir)
return "\n".join(chunk.text for chunk in result.chunks), result.extractor, result.warnings
def test_text_csv_tsv_cache_and_single_file_compression(tmp_path: Path) -> None:
cache = tmp_path / "cache"
plain = tmp_path / "notes.md"
plain.write_text("第一行\nsecond line", encoding="utf-8")
content, extractor, _ = extracted_text(plain, cache)
assert extractor == "text"
assert "第一行" in content
assert len(list(cache.glob("*.json"))) == 1
csv_path = tmp_path / "records.csv"
csv_path.write_text('name,value\n"hello,world",2\n', encoding="utf-8")
content, extractor, _ = extracted_text(csv_path, cache)
assert extractor == "csv"
assert '"hello,world",2' in content
tsv_path = tmp_path / "records.tsv"
tsv_path.write_text("name\tvalue\nhello\t2\n", encoding="utf-8")
content, extractor, _ = extracted_text(tsv_path, cache)
assert extractor == "tsv"
assert "hello\t2" in content
compressed = tmp_path / "manual.txt.gz"
with gzip.open(compressed, "wb") as output:
output.write("普通 gzip 内容".encode())
content, extractor, _ = extracted_text(compressed, cache)
assert extractor == "archive"
assert "普通 gzip 内容" in content
def test_office_open_document_pdf_and_image_extractors(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
cache = tmp_path / "cache"
docx_path = tmp_path / "document.docx"
document = Document()
document.add_paragraph("DOCX paragraph")
table = document.add_table(rows=1, cols=2)
table.cell(0, 0).text = "left"
table.cell(0, 1).text = "right"
document.save(docx_path)
content, extractor, _ = extracted_text(docx_path, cache)
assert extractor == "docx"
assert "DOCX paragraph" in content and "left\tright" in content
xlsx_path = tmp_path / "workbook.xlsx"
workbook = Workbook()
worksheet = workbook.active
worksheet.title = "Data"
worksheet.append(["name", "value"])
worksheet.append(["alpha", 42])
workbook.save(xlsx_path)
content, extractor, _ = extracted_text(xlsx_path, cache)
assert extractor == "xlsx"
assert "alpha\t42" in content
pptx_path = tmp_path / "slides.pptx"
presentation = Presentation()
slide = presentation.slides.add_slide(presentation.slide_layouts[6])
text_box = slide.shapes.add_textbox(Inches(1), Inches(1), Inches(5), Inches(1))
text_box.text = "Slide body"
slide.notes_slide.notes_text_frame.text = "Speaker note"
presentation.save(pptx_path)
content, extractor, _ = extracted_text(pptx_path, cache)
assert extractor == "pptx"
assert "Slide body" in content and "Speaker note" in content
odt_path = tmp_path / "notes.odt"
odt = OpenDocumentText()
odt.text.addElement(text.P(text="OpenDocument body"))
odt.save(str(odt_path))
content, extractor, _ = extracted_text(odt_path, cache)
assert extractor == "odf"
assert "OpenDocument body" in content
pdf_path = tmp_path / "blank.pdf"
writer = PdfWriter()
writer.add_blank_page(width=72, height=72)
with pdf_path.open("wb") as output:
writer.write(output)
monkeypatch.setattr(source_extractors.shutil, "which", lambda _command: None)
content, extractor, warnings = extracted_text(pdf_path, cache)
assert extractor == "pdf"
assert content == ""
assert warnings == ["扫描 PDF 需要 pdftoppm 与 tesseract"]
image_path = tmp_path / "scan.png"
Image.new("RGB", (32, 16), "white").save(image_path)
monkeypatch.setattr(source_extractors, "_ocr", lambda _path: "OCR result")
content, extractor, _ = extracted_text(image_path, cache)
assert extractor == "image-ocr"
assert content == "OCR result"
def test_xls_adapter_uses_sheet_rows(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
class Sheet:
name = "Legacy"
nrows = 2
ncols = 2
@staticmethod
def cell_value(row: int, col: int) -> str:
return (("name", "value"), ("old", "7"))[row][col]
workbook = SimpleNamespace(
sheets=lambda: [Sheet()],
release_resources=lambda: None,
)
monkeypatch.setattr(source_extractors, "open_workbook", lambda *_args, **_kwargs: workbook)
path = tmp_path / "legacy.xls"
path.write_bytes(b"legacy workbook")
content, extractor, _ = extracted_text(path, tmp_path / "cache")
assert extractor == "xls"
assert "old\t7" in content
def test_zip_tar_safety_limits_and_nested_depth(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
cache = tmp_path / "cache"
zip_path = tmp_path / "bundle.zip"
with zipfile.ZipFile(zip_path, "w") as archive:
archive.writestr("docs/readme.txt", "safe zip text")
archive.writestr("../outside.txt", "must not extract")
content, extractor, warnings = extracted_text(zip_path, cache)
assert extractor == "archive"
assert "safe zip text" in content
assert "must not extract" not in content
assert any("不安全路径" in warning for warning in warnings)
assert not (tmp_path / "outside.txt").exists()
source = tmp_path / "inside.txt"
source.write_text("safe tar text", encoding="utf-8")
tar_path = tmp_path / "bundle.tar"
with tarfile.open(tar_path, "w") as archive:
archive.add(source, arcname="inside.txt")
content, extractor, _ = extracted_text(tar_path, cache)
assert extractor == "archive"
assert "safe tar text" in content
with pytest.raises(ExtractionError, match="嵌套"):
extract_file(zip_path, cache, use_cache=False, _archive_depth=3)
monkeypatch.setattr(source_extractors, "BINARY_LIMIT", 4)
with pytest.raises(ExtractionError, match="大小|限制"):
extract_file(tar_path, cache, use_cache=False)
def test_configurable_text_limit_is_applied_and_isolated_in_cache(tmp_path: Path) -> None:
path = tmp_path / "configurable.txt"
cache = tmp_path / "cache"
path.write_text("x" * 2048, encoding="utf-8")
permissive = extract_file(
path,
cache,
limits=ExtractionLimits(text_bytes=4096),
)
assert permissive.chunks
with pytest.raises(ExtractionError, match="限制"):
extract_file(
path,
cache,
limits=ExtractionLimits(text_bytes=1024),
)