feat: 导入 MemRelay 初始源码
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
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),
|
||||
)
|
||||
Reference in New Issue
Block a user