245 lines
10 KiB
Python
245 lines
10 KiB
Python
"""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()
|