@@ -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()
|
||||
Reference in New Issue
Block a user