289 lines
10 KiB
Python
289 lines
10 KiB
Python
"""FastAPI endpoints for document chunking + demo frontend."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import mimetypes
|
|
import re
|
|
import shutil
|
|
import tempfile
|
|
import time
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
|
|
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from fastapi.responses import FileResponse, JSONResponse
|
|
from fastapi.staticfiles import StaticFiles
|
|
from pydantic import BaseModel, Field
|
|
|
|
from rag_cut.models import ChunkResult, SplitConfig, SplitMode
|
|
from rag_cut.pipeline import chunk_document
|
|
from rag_cut.retrieval import recall_chunks
|
|
|
|
ROOT = Path(__file__).resolve().parent.parent.parent
|
|
FRONTEND_DIR = ROOT / "frontend"
|
|
STORAGE_DIR = ROOT / "storage"
|
|
RESULTS_DIR = STORAGE_DIR / "results"
|
|
UPLOADS_DIR = STORAGE_DIR / "uploads"
|
|
DOC_ID_RE = re.compile(r"^[a-f0-9]{12}$")
|
|
WORD_EXTS = {".doc", ".docx"}
|
|
|
|
app = FastAPI(title="RAG-cut Chunking API", version="0.1.0")
|
|
|
|
app.add_middleware(
|
|
CORSMiddleware,
|
|
allow_origins=["*"],
|
|
allow_credentials=False,
|
|
allow_methods=["*"],
|
|
allow_headers=["*"],
|
|
)
|
|
|
|
(STORAGE_DIR / "assets").mkdir(parents=True, exist_ok=True)
|
|
RESULTS_DIR.mkdir(parents=True, exist_ok=True)
|
|
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
|
|
app.mount("/assets", StaticFiles(directory=STORAGE_DIR / "assets"), name="assets")
|
|
|
|
|
|
class RecallRequest(BaseModel):
|
|
doc_id: str
|
|
query: str = Field(min_length=1, max_length=500)
|
|
top_k: int = Field(default=5, ge=1, le=20)
|
|
|
|
|
|
def _result_path(doc_id: str) -> Path:
|
|
if not DOC_ID_RE.fullmatch(doc_id):
|
|
raise HTTPException(status_code=400, detail="Invalid doc_id")
|
|
return RESULTS_DIR / f"{doc_id}.json"
|
|
|
|
|
|
def _save_result(result: ChunkResult) -> None:
|
|
_result_path(result.doc_id).write_text(result.model_dump_json(indent=2), encoding="utf-8")
|
|
|
|
|
|
def _find_original(doc_id: str) -> Path | None:
|
|
uploads_dir = UPLOADS_DIR / doc_id
|
|
if not uploads_dir.is_dir():
|
|
return None
|
|
files = sorted(p for p in uploads_dir.iterdir() if p.is_file())
|
|
return files[0] if files else None
|
|
|
|
|
|
def _find_preview(doc_id: str) -> Path | None:
|
|
"""Return a browser-previewable source, preferring converted Word PDFs."""
|
|
original = _find_original(doc_id)
|
|
if original is None:
|
|
return None
|
|
if original.suffix.lower() == ".pdf":
|
|
return original
|
|
if original.suffix.lower() not in WORD_EXTS:
|
|
return None
|
|
|
|
conversion_dir = STORAGE_DIR / "assets" / doc_id / "_conversion"
|
|
exact = conversion_dir / f"{original.stem}.pdf"
|
|
if exact.is_file():
|
|
return exact
|
|
if conversion_dir.is_dir():
|
|
return next(iter(sorted(conversion_dir.glob("*.pdf"))), None)
|
|
return None
|
|
|
|
|
|
def _load_result(doc_id: str) -> ChunkResult:
|
|
result_path = _result_path(doc_id)
|
|
if not result_path.exists():
|
|
raise HTTPException(status_code=404, detail="Chunk result not found")
|
|
try:
|
|
return ChunkResult.model_validate_json(result_path.read_text(encoding="utf-8"))
|
|
except Exception as exc:
|
|
raise HTTPException(status_code=500, detail=f"Invalid result file: {exc}") from exc
|
|
|
|
|
|
def _result_summary(path: Path) -> dict | None:
|
|
"""Build list summary without loading the full chunks array."""
|
|
try:
|
|
with path.open("r", encoding="utf-8") as fh:
|
|
head = fh.read(32768)
|
|
idx = head.find('\n "chunks"')
|
|
if idx == -1:
|
|
idx = head.find('"chunks"')
|
|
if idx != -1:
|
|
data = json.loads(head[:idx].rstrip().rstrip(",") + "\n}")
|
|
else:
|
|
data = json.loads(path.read_text(encoding="utf-8"))
|
|
except (OSError, json.JSONDecodeError):
|
|
return None
|
|
doc_id = data.get("doc_id") or path.stem
|
|
if not DOC_ID_RE.fullmatch(str(doc_id)):
|
|
return None
|
|
mtime = path.stat().st_mtime
|
|
return {
|
|
"doc_id": doc_id,
|
|
"filename": data.get("filename") or path.name,
|
|
"split_mode": data.get("split_mode"),
|
|
"chunk_count": data.get("chunk_count", 0),
|
|
"block_count": data.get("block_count", 0),
|
|
"has_original": _find_original(doc_id) is not None,
|
|
"saved_at": datetime.fromtimestamp(mtime, tz=timezone.utc).isoformat(),
|
|
"saved_ts": mtime,
|
|
}
|
|
|
|
|
|
@app.get("/health")
|
|
def health():
|
|
return {"status": "ok"}
|
|
|
|
|
|
@app.get("/api/results")
|
|
def list_results():
|
|
"""List persisted chunk results (newest first)."""
|
|
items: list[dict] = []
|
|
for path in RESULTS_DIR.glob("*.json"):
|
|
summary = _result_summary(path)
|
|
if summary:
|
|
items.append(summary)
|
|
items.sort(key=lambda item: item["saved_ts"], reverse=True)
|
|
for item in items:
|
|
item.pop("saved_ts", None)
|
|
return {"count": len(items), "results": items}
|
|
|
|
|
|
@app.get("/api/results/{doc_id}")
|
|
def get_result(doc_id: str):
|
|
"""Load a full persisted chunk result by doc_id."""
|
|
chunk_result = _load_result(doc_id)
|
|
payload = json.loads(chunk_result.model_dump_json())
|
|
payload["has_original"] = _find_original(doc_id) is not None
|
|
return JSONResponse(content=payload)
|
|
|
|
|
|
@app.get("/api/results/{doc_id}/original")
|
|
def get_original(doc_id: str):
|
|
"""Serve the original uploaded file for a historical result (if still on disk)."""
|
|
if not DOC_ID_RE.fullmatch(doc_id):
|
|
raise HTTPException(status_code=400, detail="Invalid doc_id")
|
|
original = _find_original(doc_id)
|
|
if original is None:
|
|
raise HTTPException(status_code=404, detail="Original file not found")
|
|
media_type, _ = mimetypes.guess_type(original.name)
|
|
return FileResponse(
|
|
path=original,
|
|
media_type=media_type or "application/octet-stream",
|
|
filename=original.name,
|
|
)
|
|
|
|
|
|
@app.get("/api/results/{doc_id}/preview")
|
|
def get_preview(doc_id: str):
|
|
"""Serve an inline PDF preview for PDF and converted Word documents."""
|
|
if not DOC_ID_RE.fullmatch(doc_id):
|
|
raise HTTPException(status_code=400, detail="Invalid doc_id")
|
|
preview = _find_preview(doc_id)
|
|
if preview is None:
|
|
raise HTTPException(status_code=404, detail="PDF preview not found")
|
|
return FileResponse(path=preview, media_type="application/pdf")
|
|
|
|
|
|
def _delete_tree(path: Path) -> bool:
|
|
if not path.exists():
|
|
return False
|
|
if path.is_dir():
|
|
shutil.rmtree(path)
|
|
else:
|
|
path.unlink()
|
|
return True
|
|
|
|
|
|
@app.delete("/api/results/{doc_id}")
|
|
def delete_result(doc_id: str):
|
|
"""Delete a persisted chunk result and its uploads/assets."""
|
|
result_path = _result_path(doc_id)
|
|
if not result_path.exists():
|
|
raise HTTPException(status_code=404, detail="Chunk result not found")
|
|
|
|
removed = {
|
|
"result": _delete_tree(result_path),
|
|
"uploads": _delete_tree(UPLOADS_DIR / doc_id),
|
|
"assets": _delete_tree(STORAGE_DIR / "assets" / doc_id),
|
|
}
|
|
return {"doc_id": doc_id, "deleted": True, "removed": removed}
|
|
|
|
|
|
@app.post("/api/chunk")
|
|
async def chunk_file(
|
|
file: UploadFile = File(...),
|
|
mode: str | None = Form(None),
|
|
delimiter: str | None = Form(None),
|
|
parent_delimiter: str | None = Form(None),
|
|
child_delimiter: str | None = Form(None),
|
|
max_chunk_size: int | None = Form(None),
|
|
child_max_size: int | None = Form(None),
|
|
overlap: int | None = Form(None),
|
|
header_row_start: int | None = Form(None),
|
|
header_row_end: int | None = Form(None),
|
|
start_row: int | None = Form(None),
|
|
rows_per_chunk: int | None = Form(None),
|
|
):
|
|
config = None
|
|
if mode is not None:
|
|
try:
|
|
split_mode = SplitMode(mode)
|
|
except ValueError as exc:
|
|
raise HTTPException(status_code=400, detail=f"Invalid mode: {mode}") from exc
|
|
|
|
config = SplitConfig(
|
|
mode=split_mode,
|
|
delimiter=delimiter,
|
|
parent_delimiter=parent_delimiter,
|
|
child_delimiter=child_delimiter,
|
|
max_chunk_size=1500 if max_chunk_size is None else max_chunk_size,
|
|
child_max_size=512 if child_max_size is None else child_max_size,
|
|
overlap=150 if overlap is None else overlap,
|
|
header_row_start=1 if header_row_start is None else header_row_start,
|
|
header_row_end=1 if header_row_end is None else header_row_end,
|
|
start_row=2 if start_row is None else start_row,
|
|
rows_per_chunk=1 if rows_per_chunk is None else rows_per_chunk,
|
|
)
|
|
|
|
filename = Path(file.filename or "upload").name
|
|
payload = await file.read()
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
tmp_path = Path(tmp) / filename
|
|
tmp_path.write_bytes(payload)
|
|
try:
|
|
# Heavy parsing (MinerU / LibreOffice / OCR) must not block the event loop,
|
|
# otherwise /health and other uploads hang while one document is processing.
|
|
result = await asyncio.to_thread(chunk_document, tmp_path, config)
|
|
_save_result(result)
|
|
return JSONResponse(content=json.loads(result.model_dump_json()))
|
|
except ValueError as exc:
|
|
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
|
except FileNotFoundError as exc:
|
|
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
|
except RuntimeError as exc:
|
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
|
|
|
|
|
@app.post("/api/recall")
|
|
def recall(request: RecallRequest):
|
|
result_path = _result_path(request.doc_id)
|
|
if not result_path.exists():
|
|
raise HTTPException(status_code=404, detail="Chunk result not found; upload and split the document again")
|
|
|
|
started = time.perf_counter()
|
|
chunk_result = ChunkResult.model_validate_json(result_path.read_text(encoding="utf-8"))
|
|
results, candidate_count = recall_chunks(request.query, chunk_result.chunks, request.top_k)
|
|
return {
|
|
"doc_id": request.doc_id,
|
|
"query": request.query,
|
|
"top_k": request.top_k,
|
|
"candidate_count": candidate_count,
|
|
"returned_count": len(results),
|
|
"elapsed_ms": round((time.perf_counter() - started) * 1000, 3),
|
|
"results": results,
|
|
}
|
|
|
|
|
|
if FRONTEND_DIR.exists():
|
|
app.mount("/", StaticFiles(directory=FRONTEND_DIR, html=True), name="frontend")
|