first commit

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
陈辅元
2026-07-16 11:12:17 +08:00
co-authored by Cursor
commit 4003624b8c
80 changed files with 9990 additions and 0 deletions
+244
View File
@@ -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="![logo](logo.png)", 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()