"""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()