Cognitive-rag / backend / routes_kb.py
routes_kb.py
Raw
from __future__ import annotations

import shutil
import zipfile
from pathlib import Path
from typing import Any

from fastapi import APIRouter, File, HTTPException, Request, UploadFile
from fastapi.responses import StreamingResponse

from backend import pipeline
from backend.logging_audit import resolve_user


router = APIRouter(prefix="/kb", tags=["kb"])


def _safe_extract_zip(zip_path: Path, destination: Path) -> list[str]:
    extracted: list[str] = []
    with zipfile.ZipFile(zip_path, "r") as archive:
        for member in archive.infolist():
            target = destination / member.filename
            resolved = target.resolve()
            if not str(resolved).startswith(str(destination.resolve())):
                raise HTTPException(status_code=400, detail="Unsafe zip entry path detected.")
            if member.is_dir():
                resolved.mkdir(parents=True, exist_ok=True)
                continue
            resolved.parent.mkdir(parents=True, exist_ok=True)
            with archive.open(member) as src, resolved.open("wb") as dst:
                shutil.copyfileobj(src, dst)
            extracted.append(member.filename)
    return extracted


@router.post("/upload")
async def upload_kb(request: Request, files: list[UploadFile] = File(...)) -> dict[str, Any]:
    ip = request.client.host if request.client else "unknown"
    user_n = resolve_user(ip)
    job = pipeline.create_job(user_n=user_n, ip=ip)

    uploaded_names: list[str] = []
    for incoming in files:
        filename = incoming.filename or "uploaded_file"
        dest = job.uploads_dir / filename
        dest.parent.mkdir(parents=True, exist_ok=True)
        content = await incoming.read()
        with dest.open("wb") as handle:
            handle.write(content)
        if filename.lower().endswith(".zip"):
            uploaded_names.extend(_safe_extract_zip(dest, job.uploads_dir))
            dest.unlink(missing_ok=True)
        else:
            uploaded_names.append(filename)

    if not uploaded_names:
        raise HTTPException(status_code=400, detail="No files uploaded.")

    pipeline.record_uploads(job.job_id, uploaded_names)
    pipeline.start_job(job.job_id)
    return {"job_id": job.job_id, "status": "running"}


@router.get("/{job_id}/stream")
async def stream_kb(job_id: str):
    try:
        _ = pipeline.get_job(job_id)
    except KeyError as exc:
        raise HTTPException(status_code=404, detail=str(exc)) from exc

    async def _gen():
        async for item in pipeline.stream_events(job_id):
            yield item

    return StreamingResponse(
        _gen(),
        media_type="text/event-stream",
        headers={
            "Cache-Control": "no-cache",
            "Connection": "keep-alive",
            "X-Accel-Buffering": "no",
        },
    )


@router.get("/{job_id}/status")
async def kb_status(job_id: str) -> dict[str, Any]:
    try:
        return pipeline.job_status(job_id)
    except KeyError as exc:
        raise HTTPException(status_code=404, detail=str(exc)) from exc