Backend: - Deterministic SQLite connection cleanup. The RAG code used bare `with get_connection() as conn:`, which commits but never closes, leaning on GC to release handles (the rest of studio_db closes explicitly). Add a closing_connection() context manager that commits/rolls back like sqlite3's own manager and always closes, and route all 30 RAG call sites through it. - filter_by_min_score no longer drops BM25-only and figure-ref hits. min_score is a cosine floor, so it now gates only hits that carry a dense_score; lexical and figure-ref hits (dense_score is None) pass through instead of being silently discarded when the floor is raised. - Fix two tests that could not pass against the production code: the RRF fusion test asserted the wrong winner (c edges out b: 0.032266 vs 0.032258), and two tool-handler scope tests stubbed retrieve_hybrid without accepting the embedder_model kwarg the handler now passes (TypeError was swallowed, leaving captured["scope"] unset). Frontend: - Removing an in-flight upload chip now routes through the teardown thunk already registered for the aggregate-progress toast (abort, unsubscribe, release the index slot, delete the backend doc with the correct kb/thread scope key it closed over) and clears the toast entry. Deleting directly leaked the concurrency slot and hardcoded the thread scope, mis-targeting KB-scoped docs. Applied in both the composer hook and the compare-view composer; drop the now-vestigial chip-scope-key tracking and unused activeThreadId selectors. Add index-progress-store.remove(id).
1751 lines
56 KiB
Python
1751 lines
56 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""RAG API: KB CRUD, document upload (KB + per-thread), ingestion SSE, search."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import json
|
|
import os
|
|
import queue as queue_module
|
|
import time
|
|
from datetime import datetime, timedelta, timezone
|
|
from pathlib import Path
|
|
from typing import Any, Literal, Optional
|
|
from urllib.parse import quote
|
|
from uuid import uuid4
|
|
|
|
import jwt
|
|
from fastapi import (
|
|
APIRouter,
|
|
Depends,
|
|
Header,
|
|
HTTPException,
|
|
Query,
|
|
Request,
|
|
UploadFile,
|
|
)
|
|
from fastapi.responses import FileResponse, Response, StreamingResponse
|
|
from pydantic import BaseModel, Field
|
|
|
|
from auth.authentication import get_current_subject, get_current_subject_sse
|
|
from auth.storage import get_jwt_secret
|
|
|
|
|
|
async def _sse_auth(
|
|
token: str | None = Query(None),
|
|
authorization: str | None = Header(None),
|
|
) -> str:
|
|
return await get_current_subject_sse(token, authorization)
|
|
|
|
|
|
from core.rag import embeddings, ingestion, reranker, retrieval, vector_store
|
|
from core.rag.authorization import document_for_subject_or_404
|
|
from core.rag.locators import backfill_document_locators
|
|
from core.rag.vector_store import kb_scope, thread_scope
|
|
from loggers import get_logger
|
|
from storage.studio_db import (
|
|
closing_connection,
|
|
list_chat_settings,
|
|
upsert_chat_settings_merge,
|
|
)
|
|
from utils.paths.storage_roots import ensure_dir, rag_uploads_root, resolve_under_root
|
|
from utils.rag.config import (
|
|
RAG_MAX_UPLOAD_MB,
|
|
RAG_RERANK_CANDIDATE_K,
|
|
RAG_UPLOAD_EXTS,
|
|
)
|
|
|
|
router = APIRouter()
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
# --- Pydantic schemas ---
|
|
|
|
ChunkingStrategy = Literal["standard", "late"]
|
|
KBMode = Literal["text", "multimodal"]
|
|
|
|
|
|
class CreateKBRequest(BaseModel):
|
|
name: str = Field(min_length = 1, max_length = 200)
|
|
description: str | None = None
|
|
embedding_model: str | None = None
|
|
chunking_strategy: ChunkingStrategy = "standard"
|
|
mode: KBMode = "text"
|
|
|
|
|
|
class KBResponse(BaseModel):
|
|
id: str
|
|
name: str
|
|
description: str | None
|
|
embedding_model: str
|
|
chunking_strategy: ChunkingStrategy
|
|
mode: KBMode
|
|
created_at: int
|
|
|
|
|
|
class KBListResponse(BaseModel):
|
|
knowledge_bases: list[KBResponse]
|
|
|
|
|
|
class DocumentResponse(BaseModel):
|
|
id: str
|
|
kb_id: str | None
|
|
thread_id: str | None
|
|
filename: str
|
|
content_type: str | None
|
|
status: str
|
|
num_chunks: int
|
|
byte_size: int
|
|
error: str | None
|
|
created_at: int
|
|
|
|
|
|
class DocumentListResponse(BaseModel):
|
|
documents: list[DocumentResponse]
|
|
|
|
|
|
class ThreadIndexSummary(BaseModel):
|
|
thread_id: str
|
|
title: str | None
|
|
num_documents: int
|
|
num_chunks: int
|
|
|
|
|
|
class ThreadIndexListResponse(BaseModel):
|
|
threads: list[ThreadIndexSummary]
|
|
|
|
|
|
class UploadResponse(BaseModel):
|
|
document_id: str
|
|
job_id: str
|
|
filename: str
|
|
# Identical content hash already indexed in this scope; no job started, job_id "".
|
|
already_indexed: bool = False
|
|
|
|
|
|
class SearchRequest(BaseModel):
|
|
query: str = Field(min_length = 1, max_length = 4000)
|
|
kb_id: str | None = None
|
|
thread_id: str | None = None
|
|
top_k: int = Field(default = 10, ge = 1, le = 100)
|
|
mode: Literal["bm25", "dense", "hybrid"] = "hybrid"
|
|
document_ids: list[str] | None = None
|
|
enable_rerank: bool = False
|
|
reranker_model: str | None = None
|
|
min_score: float = Field(default = 0.0, ge = 0.0, le = 1.0)
|
|
|
|
|
|
class SearchHit(BaseModel):
|
|
chunk_id: str
|
|
document_id: str
|
|
chunk_index: int
|
|
text: str
|
|
score: float
|
|
page_number: int | None = None
|
|
filename: str | None = None
|
|
kind: str = "text"
|
|
image_url: str | None = None
|
|
source_page_index: int | None = None
|
|
page_char_start: int | None = None
|
|
page_char_end: int | None = None
|
|
line_start: int | None = None
|
|
line_end: int | None = None
|
|
|
|
|
|
class SearchResponse(BaseModel):
|
|
hits: list[SearchHit]
|
|
|
|
|
|
# --- Helpers ---
|
|
|
|
|
|
def _sanitize_filename(filename: str) -> str:
|
|
name = Path(filename).name.strip().replace("\x00", "")
|
|
return name or "document"
|
|
|
|
|
|
def _now_ms() -> int:
|
|
return int(time.time())
|
|
|
|
|
|
from core.rag.scope import resolve_scope_embedder as _resolve_scope_embedder # noqa: E402
|
|
|
|
|
|
def _row_to_kb(row: Any) -> KBResponse:
|
|
keys = row.keys() if hasattr(row, "keys") else ()
|
|
chunking_strategy = (
|
|
row["chunking_strategy"] if "chunking_strategy" in keys else "standard"
|
|
)
|
|
mode = row["mode"] if "mode" in keys else "text"
|
|
return KBResponse(
|
|
id = row["id"],
|
|
name = row["name"],
|
|
description = row["description"],
|
|
embedding_model = row["embedding_model"],
|
|
chunking_strategy = chunking_strategy,
|
|
mode = mode,
|
|
created_at = row["created_at"],
|
|
)
|
|
|
|
|
|
def _validate_mode_combo(mode: KBMode, chunking_strategy: ChunkingStrategy) -> None:
|
|
"""Reject (multimodal, late) — no embedder supports both at once."""
|
|
if mode == "multimodal" and chunking_strategy == "late":
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = (
|
|
"Late chunking is not supported in multimodal mode — "
|
|
"the multimodal embedder does not expose per-token "
|
|
"embeddings. Pick 'standard' chunking or 'text' mode."
|
|
),
|
|
)
|
|
|
|
|
|
def _row_to_document(row: Any) -> DocumentResponse:
|
|
return DocumentResponse(
|
|
id = row["id"],
|
|
kb_id = row["kb_id"],
|
|
thread_id = row["thread_id"],
|
|
filename = row["filename"],
|
|
content_type = row["content_type"],
|
|
status = row["status"],
|
|
num_chunks = row["num_chunks"],
|
|
byte_size = row["byte_size"],
|
|
error = row["error"],
|
|
created_at = row["created_at"],
|
|
)
|
|
|
|
|
|
def _kb_or_404(kb_id: str) -> Any:
|
|
with closing_connection() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM rag_knowledge_bases WHERE id = ?",
|
|
(kb_id,),
|
|
).fetchone()
|
|
if not row:
|
|
raise HTTPException(status_code = 404, detail = "Knowledge base not found")
|
|
return row
|
|
|
|
|
|
def _thread_or_404(thread_id: str) -> None:
|
|
with closing_connection() as conn:
|
|
row = conn.execute(
|
|
"SELECT id FROM chat_threads WHERE id = ?",
|
|
(thread_id,),
|
|
).fetchone()
|
|
if not row:
|
|
raise HTTPException(status_code = 404, detail = "Thread not found")
|
|
|
|
|
|
def _document_or_404(document_id: str) -> Any:
|
|
with closing_connection() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM rag_documents WHERE id = ?",
|
|
(document_id,),
|
|
).fetchone()
|
|
if not row:
|
|
raise HTTPException(status_code = 404, detail = "Document not found")
|
|
return row
|
|
|
|
|
|
async def _save_upload(file: UploadFile) -> tuple[Path, str, int, str]:
|
|
import anyio
|
|
|
|
filename = _sanitize_filename(file.filename or "document")
|
|
ext = Path(filename).suffix.lower()
|
|
if ext not in RAG_UPLOAD_EXTS:
|
|
allowed = ", ".join(sorted(RAG_UPLOAD_EXTS))
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = f"Unsupported file type: {ext}. Allowed: {allowed}",
|
|
)
|
|
upload_dir = ensure_dir(rag_uploads_root())
|
|
stored_name = f"{uuid4().hex}_{Path(filename).stem}{ext}"
|
|
stored_path = upload_dir / stored_name
|
|
max_bytes = RAG_MAX_UPLOAD_MB * 1024 * 1024
|
|
written = 0
|
|
# Hash bytes while streaming to dedup identical re-uploads within a scope.
|
|
hasher = hashlib.sha256()
|
|
# anyio worker thread keeps the event loop free. Outer try/except cleans up
|
|
# partial files after the async-with closes the fd (Windows refuses unlink on an open fd).
|
|
try:
|
|
async with await anyio.open_file(stored_path, "wb") as f:
|
|
while True:
|
|
chunk = await file.read(1024 * 1024)
|
|
if not chunk:
|
|
break
|
|
written += len(chunk)
|
|
if written > max_bytes:
|
|
raise HTTPException(
|
|
status_code = 413,
|
|
detail = f"File exceeds {RAG_MAX_UPLOAD_MB} MB limit",
|
|
)
|
|
hasher.update(chunk)
|
|
await f.write(chunk)
|
|
except HTTPException:
|
|
stored_path.unlink(missing_ok = True)
|
|
raise
|
|
if written == 0:
|
|
stored_path.unlink(missing_ok = True)
|
|
raise HTTPException(status_code = 400, detail = "Empty upload payload")
|
|
return stored_path, filename, written, hasher.hexdigest()
|
|
|
|
|
|
def _start_ingestion(
|
|
*,
|
|
filename: str,
|
|
stored_path: Path,
|
|
byte_size: int,
|
|
content_type: str | None,
|
|
kb_id: str | None,
|
|
thread_id: str | None,
|
|
embedding_model: str,
|
|
chunking_strategy: str = "standard",
|
|
mode: str = "text",
|
|
caption_images: bool = True,
|
|
content_hash: str | None = None,
|
|
) -> UploadResponse:
|
|
document_id = str(uuid4())
|
|
with closing_connection() as conn:
|
|
# Dedup: skip re-ingestion if the same content hash is already indexed
|
|
# in this scope. Only 'completed' counts — failed/in-flight may retry.
|
|
# Scope is the target kb_id or thread_id (a file in two KBs indexes in each).
|
|
if content_hash:
|
|
if kb_id is not None:
|
|
existing = conn.execute(
|
|
"SELECT id, filename FROM rag_documents "
|
|
"WHERE kb_id = ? AND content_hash = ? AND status = 'completed' "
|
|
"LIMIT 1",
|
|
(kb_id, content_hash),
|
|
).fetchone()
|
|
else:
|
|
existing = conn.execute(
|
|
"SELECT id, filename FROM rag_documents "
|
|
"WHERE thread_id = ? AND content_hash = ? AND status = 'completed' "
|
|
"LIMIT 1",
|
|
(thread_id, content_hash),
|
|
).fetchone()
|
|
if existing is not None:
|
|
# Drop the redundant upload; the already-indexed copy is source of truth.
|
|
_unlink_if_under_uploads(stored_path)
|
|
return UploadResponse(
|
|
document_id = existing["id"],
|
|
job_id = "",
|
|
filename = existing["filename"],
|
|
already_indexed = True,
|
|
)
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO rag_documents
|
|
(id, kb_id, thread_id, filename, content_type, stored_path,
|
|
status, num_chunks, byte_size, content_hash, created_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, 'pending', 0, ?, ?, ?)
|
|
""",
|
|
(
|
|
document_id,
|
|
kb_id,
|
|
thread_id,
|
|
filename,
|
|
content_type,
|
|
str(stored_path),
|
|
byte_size,
|
|
content_hash,
|
|
_now_ms(),
|
|
),
|
|
)
|
|
conn.commit()
|
|
job_id = ingestion.enqueue_ingestion(
|
|
document_id = document_id,
|
|
stored_path = stored_path,
|
|
kb_id = kb_id,
|
|
thread_id = thread_id,
|
|
embedding_model = embedding_model,
|
|
chunking_strategy = chunking_strategy,
|
|
mode = mode,
|
|
enable_captions = caption_images,
|
|
)
|
|
return UploadResponse(document_id = document_id, job_id = job_id, filename = filename)
|
|
|
|
|
|
def _unlink_if_under_uploads(path: Path) -> None:
|
|
try:
|
|
real = Path(os.path.realpath(path))
|
|
root = Path(os.path.realpath(rag_uploads_root()))
|
|
real.relative_to(root)
|
|
except (OSError, ValueError):
|
|
return
|
|
real.unlink(missing_ok = True)
|
|
|
|
|
|
# --- Knowledge bases ---
|
|
|
|
|
|
@router.post("/knowledge-bases", response_model = KBResponse)
|
|
def create_knowledge_base(
|
|
payload: CreateKBRequest,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> KBResponse:
|
|
from utils.rag.config import resolve_embedder
|
|
|
|
_validate_mode_combo(payload.mode, payload.chunking_strategy)
|
|
|
|
kb_id = str(uuid4())
|
|
# No override: resolve from (mode, strategy) matrix.
|
|
embedding_model = payload.embedding_model or resolve_embedder(
|
|
payload.mode, payload.chunking_strategy
|
|
)
|
|
created_at = _now_ms()
|
|
with closing_connection() as conn:
|
|
try:
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO rag_knowledge_bases
|
|
(id, name, description, owner_user_id, embedding_model,
|
|
chunking_strategy, mode, created_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
kb_id,
|
|
payload.name,
|
|
payload.description,
|
|
current_subject,
|
|
embedding_model,
|
|
payload.chunking_strategy,
|
|
payload.mode,
|
|
created_at,
|
|
),
|
|
)
|
|
conn.commit()
|
|
except Exception as exc:
|
|
raise HTTPException(
|
|
status_code = 409,
|
|
detail = f"Could not create KB: {exc}",
|
|
) from exc
|
|
return KBResponse(
|
|
id = kb_id,
|
|
name = payload.name,
|
|
description = payload.description,
|
|
embedding_model = embedding_model,
|
|
chunking_strategy = payload.chunking_strategy,
|
|
mode = payload.mode,
|
|
created_at = created_at,
|
|
)
|
|
|
|
|
|
@router.get("/knowledge-bases", response_model = KBListResponse)
|
|
def list_knowledge_bases(
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> KBListResponse:
|
|
with closing_connection() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM rag_knowledge_bases ORDER BY created_at DESC"
|
|
).fetchall()
|
|
return KBListResponse(knowledge_bases = [_row_to_kb(r) for r in rows])
|
|
|
|
|
|
class RagDefaults(BaseModel):
|
|
chunking_strategy: ChunkingStrategy = "standard"
|
|
mode: KBMode = "text"
|
|
embedding_model: str | None = None
|
|
|
|
|
|
class UpdateRagDefaultsRequest(BaseModel):
|
|
"""Patch shape — only fields present overwrite stored values."""
|
|
|
|
chunking_strategy: ChunkingStrategy | None = None
|
|
mode: KBMode | None = None
|
|
embedding_model: str | None = None
|
|
|
|
|
|
_DEFAULTS_KEY = "rag.defaults"
|
|
|
|
|
|
def _load_rag_defaults() -> RagDefaults:
|
|
settings = list_chat_settings()
|
|
raw = settings.get(_DEFAULTS_KEY) or {}
|
|
if not isinstance(raw, dict):
|
|
raw = {}
|
|
return RagDefaults(
|
|
chunking_strategy = raw.get("chunking_strategy") or "standard",
|
|
mode = raw.get("mode") or "text",
|
|
embedding_model = raw.get("embedding_model"),
|
|
)
|
|
|
|
|
|
@router.get("/defaults", response_model = RagDefaults)
|
|
def get_rag_defaults(
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> RagDefaults:
|
|
return _load_rag_defaults()
|
|
|
|
|
|
@router.post("/warmup")
|
|
def warmup_rag_embedder(
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
"""Preload the configured default embedder so the first retrieval is warm.
|
|
|
|
Called from the frontend when the user enables the RAG pill — moves
|
|
the cold-load latency out of the first chat-completion path, where a
|
|
multi-second load can race the llama-server prefill timeout.
|
|
"""
|
|
from utils.rag.config import resolve_embedder
|
|
|
|
defaults = _load_rag_defaults()
|
|
model_name = defaults.embedding_model or resolve_embedder(
|
|
defaults.mode,
|
|
defaults.chunking_strategy,
|
|
)
|
|
try:
|
|
embeddings.get_embedder(model_name)
|
|
except Exception as exc: # noqa: BLE001
|
|
# Log details server-side; return a generic message so paths/stack stay hidden.
|
|
logger.warning("RAG warmup failed for %s: %s", model_name, exc)
|
|
return {"ok": False, "model": model_name, "error": "Failed to load embedder"}
|
|
return {"ok": True, "model": model_name}
|
|
|
|
|
|
@router.post("/reranker/precache")
|
|
def precache_rag_reranker(
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
"""Download the reranker weights (~1.1 GB) into the HF cache.
|
|
|
|
Called from the frontend the moment the user flips the "Use
|
|
reranker" switch ON so the cost lands on the explicit toggle
|
|
instead of the first chat turn — where a multi-minute download
|
|
looks like a hung tool call.
|
|
"""
|
|
from core.rag.reranker import precache_reranker
|
|
from utils.rag.config import RAG_RERANKER_MODEL
|
|
|
|
try:
|
|
precache_reranker()
|
|
except Exception as exc: # noqa: BLE001
|
|
# Log details server-side; client gets a generic message so paths/stack stay hidden.
|
|
logger.warning(
|
|
"RAG reranker precache failed",
|
|
model = RAG_RERANKER_MODEL,
|
|
error = str(exc),
|
|
)
|
|
return {
|
|
"ok": False,
|
|
"model": RAG_RERANKER_MODEL,
|
|
"error": "Failed to download reranker",
|
|
}
|
|
return {"ok": True, "model": RAG_RERANKER_MODEL}
|
|
|
|
|
|
@router.put("/defaults", response_model = RagDefaults)
|
|
def set_rag_defaults(
|
|
payload: UpdateRagDefaultsRequest,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> RagDefaults:
|
|
current = _load_rag_defaults()
|
|
new_strategy = payload.chunking_strategy or current.chunking_strategy
|
|
new_mode = payload.mode or current.mode
|
|
# PATCH-style: empty string clears, null/missing keeps current.
|
|
if payload.embedding_model is None:
|
|
new_embedder = current.embedding_model
|
|
elif payload.embedding_model.strip() == "":
|
|
new_embedder = None
|
|
else:
|
|
new_embedder = payload.embedding_model.strip()
|
|
_validate_mode_combo(new_mode, new_strategy)
|
|
|
|
upsert_chat_settings_merge(
|
|
{
|
|
_DEFAULTS_KEY: {
|
|
"chunking_strategy": new_strategy,
|
|
"mode": new_mode,
|
|
"embedding_model": new_embedder,
|
|
}
|
|
}
|
|
)
|
|
return RagDefaults(
|
|
chunking_strategy = new_strategy,
|
|
mode = new_mode,
|
|
embedding_model = new_embedder,
|
|
)
|
|
|
|
|
|
class ThreadRagSettings(BaseModel):
|
|
chunking_strategy: ChunkingStrategy = "standard"
|
|
mode: KBMode = "text"
|
|
embedding_model: str | None = None
|
|
|
|
|
|
class UpdateThreadRagSettingsRequest(BaseModel):
|
|
chunking_strategy: ChunkingStrategy | None = None
|
|
mode: KBMode | None = None
|
|
embedding_model: str | None = None
|
|
# Reingest-only (not persisted); omit or None keeps captioning on.
|
|
caption_images: bool | None = None
|
|
|
|
|
|
def _thread_settings_key(thread_id: str) -> str:
|
|
return f"thread:{thread_id}:rag"
|
|
|
|
|
|
def _load_thread_settings(thread_id: str) -> ThreadRagSettings:
|
|
"""Per-thread RAG settings (chat_settings['thread:<id>:rag']) with defaults fallback."""
|
|
settings = list_chat_settings()
|
|
raw = settings.get(_thread_settings_key(thread_id)) or {}
|
|
if not isinstance(raw, dict):
|
|
raw = {}
|
|
fallback = _load_rag_defaults()
|
|
return ThreadRagSettings(
|
|
chunking_strategy = (raw.get("chunking_strategy") or fallback.chunking_strategy),
|
|
mode = raw.get("mode") or fallback.mode,
|
|
embedding_model = raw.get("embedding_model") or fallback.embedding_model,
|
|
)
|
|
|
|
|
|
@router.get(
|
|
"/threads/{thread_id}/settings",
|
|
response_model = ThreadRagSettings,
|
|
)
|
|
def get_thread_rag_settings(
|
|
thread_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> ThreadRagSettings:
|
|
return _load_thread_settings(thread_id)
|
|
|
|
|
|
@router.put(
|
|
"/threads/{thread_id}/settings",
|
|
response_model = ThreadRagSettings,
|
|
)
|
|
def set_thread_rag_settings(
|
|
thread_id: str,
|
|
payload: UpdateThreadRagSettingsRequest,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> ThreadRagSettings:
|
|
current = _load_thread_settings(thread_id)
|
|
new_strategy = payload.chunking_strategy or current.chunking_strategy
|
|
new_mode = payload.mode or current.mode
|
|
if payload.embedding_model is None:
|
|
new_embedder = current.embedding_model
|
|
elif payload.embedding_model.strip() == "":
|
|
new_embedder = None
|
|
else:
|
|
new_embedder = payload.embedding_model.strip()
|
|
_validate_mode_combo(new_mode, new_strategy)
|
|
|
|
upsert_chat_settings_merge(
|
|
{
|
|
_thread_settings_key(thread_id): {
|
|
"chunking_strategy": new_strategy,
|
|
"mode": new_mode,
|
|
"embedding_model": new_embedder,
|
|
}
|
|
}
|
|
)
|
|
return ThreadRagSettings(
|
|
chunking_strategy = new_strategy,
|
|
mode = new_mode,
|
|
embedding_model = new_embedder,
|
|
)
|
|
|
|
|
|
class ReingestKBRequest(BaseModel):
|
|
"""All fields optional — omitting one keeps the KB's current value."""
|
|
|
|
chunking_strategy: ChunkingStrategy | None = None
|
|
mode: KBMode | None = None
|
|
embedding_model: str | None = None
|
|
# Not persisted on the KB; omit or None keeps captioning on for the rebuild.
|
|
caption_images: bool | None = None
|
|
|
|
|
|
class ReingestResponse(BaseModel):
|
|
job_ids: list[str]
|
|
document_ids: list[str]
|
|
|
|
|
|
def _reingest_scope(
|
|
*,
|
|
kb_id: str | None,
|
|
thread_id: str | None,
|
|
chunking_strategy: str,
|
|
mode: str,
|
|
embedding_model: str,
|
|
caption_images: bool = True,
|
|
) -> ReingestResponse:
|
|
"""Wipe scope artifacts and re-enqueue every document; metadata untouched."""
|
|
scope = kb_scope(kb_id) if kb_id else thread_scope(thread_id) # type: ignore[arg-type]
|
|
with closing_connection() as conn:
|
|
if kb_id:
|
|
rows = conn.execute(
|
|
"SELECT id, stored_path FROM rag_documents WHERE kb_id = ?",
|
|
(kb_id,),
|
|
).fetchall()
|
|
else:
|
|
rows = conn.execute(
|
|
"SELECT id, stored_path FROM rag_documents WHERE thread_id = ?",
|
|
(thread_id,),
|
|
).fetchall()
|
|
# Drop rag_documents (chunks cascade); disk files reused below.
|
|
doc_ids = [r["id"] for r in rows]
|
|
if doc_ids:
|
|
placeholders = ",".join("?" for _ in doc_ids)
|
|
conn.execute(
|
|
f"DELETE FROM rag_documents WHERE id IN ({placeholders})",
|
|
doc_ids,
|
|
)
|
|
conn.commit()
|
|
ingestion.delete_scope_artifacts(scope)
|
|
|
|
job_ids: list[str] = []
|
|
new_doc_ids: list[str] = []
|
|
for row in rows:
|
|
stored_path = Path(row["stored_path"])
|
|
if not stored_path.is_file():
|
|
continue
|
|
filename = stored_path.name
|
|
# Strip the upload-time UUID prefix; keep the original filename.
|
|
if "_" in filename:
|
|
_uuid_prefix, _, original = filename.partition("_")
|
|
if original:
|
|
filename = original
|
|
upload = _start_ingestion(
|
|
filename = filename,
|
|
stored_path = stored_path,
|
|
byte_size = stored_path.stat().st_size,
|
|
content_type = None,
|
|
kb_id = kb_id,
|
|
thread_id = thread_id,
|
|
embedding_model = embedding_model,
|
|
chunking_strategy = chunking_strategy,
|
|
mode = mode,
|
|
caption_images = caption_images,
|
|
)
|
|
job_ids.append(upload.job_id)
|
|
new_doc_ids.append(upload.document_id)
|
|
return ReingestResponse(job_ids = job_ids, document_ids = new_doc_ids)
|
|
|
|
|
|
@router.post(
|
|
"/knowledge-bases/{kb_id}/reingest",
|
|
response_model = ReingestResponse,
|
|
)
|
|
def reingest_knowledge_base(
|
|
kb_id: str,
|
|
payload: ReingestKBRequest,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> ReingestResponse:
|
|
from utils.rag.config import resolve_embedder
|
|
|
|
kb_row = _kb_or_404(kb_id)
|
|
keys = kb_row.keys() if hasattr(kb_row, "keys") else ()
|
|
current_strategy = (
|
|
kb_row["chunking_strategy"] if "chunking_strategy" in keys else "standard"
|
|
)
|
|
current_mode = kb_row["mode"] if "mode" in keys else "text"
|
|
current_embedder = kb_row["embedding_model"]
|
|
|
|
new_strategy = payload.chunking_strategy or current_strategy
|
|
new_mode = payload.mode or current_mode
|
|
_validate_mode_combo(new_mode, new_strategy)
|
|
|
|
new_embedder = payload.embedding_model or (
|
|
current_embedder
|
|
if (new_strategy == current_strategy and new_mode == current_mode)
|
|
else resolve_embedder(new_mode, new_strategy)
|
|
)
|
|
|
|
with closing_connection() as conn:
|
|
conn.execute(
|
|
"""
|
|
UPDATE rag_knowledge_bases
|
|
SET chunking_strategy = ?, mode = ?, embedding_model = ?
|
|
WHERE id = ?
|
|
""",
|
|
(new_strategy, new_mode, new_embedder, kb_id),
|
|
)
|
|
conn.commit()
|
|
|
|
return _reingest_scope(
|
|
kb_id = kb_id,
|
|
thread_id = None,
|
|
chunking_strategy = new_strategy,
|
|
mode = new_mode,
|
|
embedding_model = new_embedder,
|
|
caption_images = payload.caption_images is not False,
|
|
)
|
|
|
|
|
|
@router.post(
|
|
"/threads/{thread_id}/reingest",
|
|
response_model = ReingestResponse,
|
|
)
|
|
def reingest_thread_documents(
|
|
thread_id: str,
|
|
payload: UpdateThreadRagSettingsRequest | None = None,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> ReingestResponse:
|
|
"""Rebuild a thread's RAG index; optional body updates settings before reingest."""
|
|
from utils.rag.config import resolve_embedder
|
|
|
|
if payload is None:
|
|
payload = UpdateThreadRagSettingsRequest()
|
|
if (
|
|
payload.chunking_strategy is not None
|
|
or payload.mode is not None
|
|
or payload.embedding_model is not None
|
|
):
|
|
settings = set_thread_rag_settings(
|
|
thread_id,
|
|
payload,
|
|
current_subject = current_subject,
|
|
)
|
|
else:
|
|
settings = _load_thread_settings(thread_id)
|
|
|
|
embedder = settings.embedding_model or resolve_embedder(
|
|
settings.mode,
|
|
settings.chunking_strategy,
|
|
)
|
|
return _reingest_scope(
|
|
kb_id = None,
|
|
thread_id = thread_id,
|
|
chunking_strategy = settings.chunking_strategy,
|
|
mode = settings.mode,
|
|
embedding_model = embedder,
|
|
caption_images = payload.caption_images is not False,
|
|
)
|
|
|
|
|
|
@router.delete("/knowledge-bases/{kb_id}")
|
|
def delete_knowledge_base(
|
|
kb_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
_kb_or_404(kb_id)
|
|
with closing_connection() as conn:
|
|
doc_rows = conn.execute(
|
|
"SELECT stored_path FROM rag_documents WHERE kb_id = ?",
|
|
(kb_id,),
|
|
).fetchall()
|
|
conn.execute("DELETE FROM rag_knowledge_bases WHERE id = ?", (kb_id,))
|
|
conn.commit()
|
|
for row in doc_rows:
|
|
_unlink_if_under_uploads(Path(row["stored_path"]))
|
|
ingestion.delete_scope_artifacts(kb_scope(kb_id))
|
|
return {"ok": True}
|
|
|
|
|
|
# --- Document upload (KB and per-thread) ---
|
|
|
|
|
|
@router.post("/knowledge-bases/{kb_id}/documents", response_model = UploadResponse)
|
|
async def upload_kb_document(
|
|
kb_id: str,
|
|
file: UploadFile,
|
|
caption_images: bool = True,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> UploadResponse:
|
|
kb_row = _kb_or_404(kb_id)
|
|
stored_path, filename, byte_size, content_hash = await _save_upload(file)
|
|
# Tolerate pre-Phase-3 rows missing chunking_strategy/mode.
|
|
kb_keys = kb_row.keys() if hasattr(kb_row, "keys") else ()
|
|
chunking_strategy = (
|
|
kb_row["chunking_strategy"] if "chunking_strategy" in kb_keys else "standard"
|
|
)
|
|
mode = kb_row["mode"] if "mode" in kb_keys else "text"
|
|
return _start_ingestion(
|
|
filename = filename,
|
|
stored_path = stored_path,
|
|
byte_size = byte_size,
|
|
content_type = file.content_type,
|
|
kb_id = kb_id,
|
|
thread_id = None,
|
|
embedding_model = kb_row["embedding_model"],
|
|
chunking_strategy = chunking_strategy,
|
|
mode = mode,
|
|
caption_images = caption_images,
|
|
content_hash = content_hash,
|
|
)
|
|
|
|
|
|
@router.post("/threads/{thread_id}/documents", response_model = UploadResponse)
|
|
async def upload_thread_document(
|
|
thread_id: str,
|
|
file: UploadFile,
|
|
caption_images: bool = True,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> UploadResponse:
|
|
from utils.rag.config import resolve_embedder
|
|
|
|
# No chat_threads check — fresh threads aren't persisted until first run.
|
|
stored_path, filename, byte_size, content_hash = await _save_upload(file)
|
|
settings = _load_thread_settings(thread_id)
|
|
embedder = settings.embedding_model or resolve_embedder(
|
|
settings.mode,
|
|
settings.chunking_strategy,
|
|
)
|
|
return _start_ingestion(
|
|
filename = filename,
|
|
stored_path = stored_path,
|
|
byte_size = byte_size,
|
|
content_type = file.content_type,
|
|
kb_id = None,
|
|
thread_id = thread_id,
|
|
embedding_model = embedder,
|
|
chunking_strategy = settings.chunking_strategy,
|
|
mode = settings.mode,
|
|
caption_images = caption_images,
|
|
content_hash = content_hash,
|
|
)
|
|
|
|
|
|
# --- Document list / delete ---
|
|
|
|
|
|
@router.get("/knowledge-bases/{kb_id}/documents", response_model = DocumentListResponse)
|
|
def list_kb_documents(
|
|
kb_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> DocumentListResponse:
|
|
_kb_or_404(kb_id)
|
|
with closing_connection() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM rag_documents WHERE kb_id = ? ORDER BY created_at DESC",
|
|
(kb_id,),
|
|
).fetchall()
|
|
return DocumentListResponse(documents = [_row_to_document(r) for r in rows])
|
|
|
|
|
|
@router.get("/threads/{thread_id}/documents", response_model = DocumentListResponse)
|
|
def list_thread_documents(
|
|
thread_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> DocumentListResponse:
|
|
with closing_connection() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM rag_documents WHERE thread_id = ? ORDER BY created_at DESC",
|
|
(thread_id,),
|
|
).fetchall()
|
|
return DocumentListResponse(documents = [_row_to_document(r) for r in rows])
|
|
|
|
|
|
@router.get("/images/{document_id}/{filename}")
|
|
def get_rag_image(
|
|
document_id: str,
|
|
filename: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> FileResponse:
|
|
"""Serve an extracted image; realpath-check against the uploads root."""
|
|
document_for_subject_or_404(document_id, current_subject)
|
|
if "/" in filename or "\\" in filename or filename.startswith("."):
|
|
raise HTTPException(status_code = 400, detail = "Invalid filename")
|
|
root = Path(os.path.realpath(rag_uploads_root() / "images"))
|
|
candidate = rag_uploads_root() / "images" / document_id / filename
|
|
try:
|
|
real = Path(os.path.realpath(candidate))
|
|
real.relative_to(root)
|
|
except (OSError, ValueError) as exc:
|
|
raise HTTPException(status_code = 404, detail = "Image not found") from exc
|
|
if not real.is_file():
|
|
raise HTTPException(status_code = 404, detail = "Image not found")
|
|
return FileResponse(str(real))
|
|
|
|
|
|
@router.delete("/documents/{document_id}")
|
|
def delete_document(
|
|
document_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
row = _document_or_404(document_id)
|
|
scope = kb_scope(row["kb_id"]) if row["kb_id"] else thread_scope(row["thread_id"])
|
|
with closing_connection() as conn:
|
|
conn.execute("DELETE FROM rag_documents WHERE id = ?", (document_id,))
|
|
conn.commit()
|
|
_unlink_if_under_uploads(Path(row["stored_path"]))
|
|
ingestion.delete_document_artifacts(document_id, scope)
|
|
return {"ok": True}
|
|
|
|
|
|
@router.get("/thread-indexes", response_model = ThreadIndexListResponse)
|
|
def list_thread_indexes(
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> ThreadIndexListResponse:
|
|
"""List threads with >=1 RAG doc. LEFT JOIN keeps unpersisted threads (null title)."""
|
|
with closing_connection() as conn:
|
|
rows = conn.execute(
|
|
"""
|
|
SELECT
|
|
d.thread_id AS thread_id,
|
|
t.title AS title,
|
|
COUNT(DISTINCT d.id) AS num_documents,
|
|
COALESCE(SUM(d.num_chunks), 0) AS num_chunks
|
|
FROM rag_documents d
|
|
LEFT JOIN chat_threads t ON t.id = d.thread_id
|
|
WHERE d.thread_id IS NOT NULL
|
|
GROUP BY d.thread_id, t.title
|
|
ORDER BY MAX(d.created_at) DESC
|
|
"""
|
|
).fetchall()
|
|
return ThreadIndexListResponse(
|
|
threads = [
|
|
ThreadIndexSummary(
|
|
thread_id = r["thread_id"],
|
|
title = r["title"],
|
|
num_documents = int(r["num_documents"]),
|
|
num_chunks = int(r["num_chunks"]),
|
|
)
|
|
for r in rows
|
|
]
|
|
)
|
|
|
|
|
|
@router.delete("/threads/{thread_id}/documents")
|
|
def clear_thread_documents(
|
|
thread_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
"""Drop all RAG artifacts for thread_id; chat thread itself untouched."""
|
|
ingestion.purge_thread_documents([thread_id])
|
|
return {"ok": True}
|
|
|
|
|
|
# --- Ingestion job SSE ---
|
|
|
|
|
|
@router.post("/jobs/{job_id}/cancel")
|
|
def cancel_job(
|
|
job_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> dict:
|
|
"""Stop an in-flight ingestion job. The caller deletes the document
|
|
afterwards to reset the index; this only halts the worker."""
|
|
cancelled = ingestion.cancel_ingestion(job_id)
|
|
return {"ok": True, "cancelled": cancelled}
|
|
|
|
|
|
@router.get("/jobs/{job_id}/events")
|
|
async def job_events(
|
|
job_id: str,
|
|
request: Request,
|
|
current_subject: str = Depends(_sse_auth),
|
|
) -> StreamingResponse:
|
|
state = ingestion.get_job_state(job_id)
|
|
if state is None:
|
|
with closing_connection() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM rag_ingestion_jobs WHERE id = ?",
|
|
(job_id,),
|
|
).fetchone()
|
|
if not row:
|
|
raise HTTPException(status_code = 404, detail = "Job not found")
|
|
return StreamingResponse(
|
|
_replay_terminal_state(row),
|
|
media_type = "text/event-stream",
|
|
)
|
|
|
|
consumer_queue = state.subscribe()
|
|
|
|
async def stream():
|
|
try:
|
|
initial = {
|
|
"type": "status",
|
|
"status": state.status,
|
|
"stage": state.stage,
|
|
"progress": state.progress,
|
|
}
|
|
yield f"data: {json.dumps(initial)}\n\n"
|
|
while True:
|
|
if await request.is_disconnected():
|
|
break
|
|
try:
|
|
event = await asyncio.get_event_loop().run_in_executor(
|
|
None,
|
|
consumer_queue.get,
|
|
True,
|
|
15.0,
|
|
)
|
|
except queue_module.Empty:
|
|
yield ": keep-alive\n\n"
|
|
if state.status in ("completed", "failed"):
|
|
break
|
|
continue
|
|
yield f"data: {json.dumps(event)}\n\n"
|
|
if event.get("type") in ("complete", "error"):
|
|
break
|
|
finally:
|
|
state.unsubscribe(consumer_queue)
|
|
|
|
return StreamingResponse(stream(), media_type = "text/event-stream")
|
|
|
|
|
|
async def _replay_terminal_state(row: Any):
|
|
payload = {
|
|
"type": "status",
|
|
"status": row["status"],
|
|
"stage": row["stage"],
|
|
"progress": row["progress"],
|
|
"error": row["error"],
|
|
}
|
|
yield f"data: {json.dumps(payload)}\n\n"
|
|
|
|
|
|
# --- Search ---
|
|
|
|
|
|
# --- Document preview (file + preview-target) ---
|
|
|
|
|
|
PreviewMediaKind = Literal["pdf", "text", "docx", "html", "image", "unknown"]
|
|
PreviewChunkKind = Literal["text", "image", "caption"]
|
|
|
|
|
|
class PreviewPdfRegion(BaseModel):
|
|
pageIndex: int
|
|
pageNumber: int | None = None
|
|
x: float
|
|
y: float
|
|
width: float
|
|
height: float
|
|
confidence: Literal["exact"]
|
|
source: str
|
|
|
|
|
|
class PreviewTargetResponse(BaseModel):
|
|
"""Per contracts.md §1.2 / §1.3 — single shape covering both
|
|
cited-chunk and document-row preview modes. The §1.3 metadata-only
|
|
mode returns ``None`` for ``chunkId``/``chunkIndex``/``targetPage``/
|
|
``snippet``/``kind``/``imageUrl`` (Q2: no first-chunk guessing).
|
|
"""
|
|
|
|
documentId: str
|
|
filename: str
|
|
contentType: str | None
|
|
mediaKind: PreviewMediaKind
|
|
byteSize: int
|
|
status: str
|
|
kbId: str | None
|
|
threadId: str | None
|
|
chunkId: str | None
|
|
chunkIndex: int | None
|
|
targetPage: int | None
|
|
snippet: str | None
|
|
kind: PreviewChunkKind | None
|
|
imageUrl: str | None
|
|
sourcePageIndex: int | None
|
|
pageCharStart: int | None
|
|
pageCharEnd: int | None
|
|
lineStart: int | None
|
|
lineEnd: int | None
|
|
pdfRegions: list[PreviewPdfRegion] = Field(default_factory = list)
|
|
|
|
|
|
class PreviewFileUrlResponse(BaseModel):
|
|
url: str
|
|
expiresAt: int
|
|
|
|
|
|
class LocatorBackfillResponse(BaseModel):
|
|
documentId: str
|
|
totalChunks: int
|
|
matched: int
|
|
alreadyLocated: int
|
|
ambiguous: int
|
|
missing: int
|
|
skipped: int
|
|
regionsMatched: int
|
|
pagesRefreshed: int
|
|
|
|
|
|
# Extension allowlist for inline rendering / disposition. Unlisted ext collapses
|
|
# to ("application/octet-stream", attachment, "unknown"). .html / .htm serve as
|
|
# text/plain attachment (decisions Q7 + Risk #3) so uploaded HTML can't execute in the app origin.
|
|
_PREVIEW_EXT_MAP: dict[str, tuple[str, str, PreviewMediaKind]] = {
|
|
".pdf": ("application/pdf", "inline", "pdf"),
|
|
".txt": ("text/plain; charset=utf-8", "inline", "text"),
|
|
".md": ("text/markdown; charset=utf-8", "inline", "text"),
|
|
".markdown": ("text/markdown; charset=utf-8", "inline", "text"),
|
|
".docx": (
|
|
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
|
"attachment",
|
|
"docx",
|
|
),
|
|
".html": ("text/plain; charset=utf-8", "attachment", "html"),
|
|
".htm": ("text/plain; charset=utf-8", "attachment", "html"),
|
|
".png": ("image/png", "inline", "image"),
|
|
".jpg": ("image/jpeg", "inline", "image"),
|
|
".jpeg": ("image/jpeg", "inline", "image"),
|
|
".gif": ("image/gif", "inline", "image"),
|
|
".webp": ("image/webp", "inline", "image"),
|
|
}
|
|
|
|
|
|
def _ascii_only(value: str) -> bool:
|
|
try:
|
|
value.encode("ascii")
|
|
except UnicodeEncodeError:
|
|
return False
|
|
return True
|
|
|
|
|
|
def _content_disposition_header(filename: str, disposition: str) -> str:
|
|
"""Build a Content-Disposition header. Non-ASCII filenames use RFC 5987
|
|
``filename*=UTF-8''…`` alongside an ASCII-only ``filename=`` fallback so
|
|
older clients still get something readable.
|
|
"""
|
|
from urllib.parse import quote as _urlquote
|
|
|
|
safe = filename.replace('"', "").replace("\r", "").replace("\n", "")
|
|
if _ascii_only(safe):
|
|
return f'{disposition}; filename="{safe}"'
|
|
ascii_fallback = safe.encode("ascii", "replace").decode("ascii")
|
|
encoded = _urlquote(safe, safe = "")
|
|
return f'{disposition}; filename="{ascii_fallback}"; ' f"filename*=UTF-8''{encoded}"
|
|
|
|
|
|
def _preview_file_metadata(filename: str) -> tuple[str, str, PreviewMediaKind]:
|
|
"""Map a stored filename to (content_type, disposition, mediaKind).
|
|
|
|
Unknown extensions always force ``application/octet-stream`` +
|
|
``attachment`` + ``unknown`` so the browser cannot sniff a sensitive
|
|
type and inline it (Risk #3).
|
|
"""
|
|
ext = Path(filename).suffix.lower()
|
|
return _PREVIEW_EXT_MAP.get(
|
|
ext, ("application/octet-stream", "attachment", "unknown")
|
|
)
|
|
|
|
|
|
_PREVIEW_FILE_AUDIENCE = "rag-preview-file"
|
|
_PREVIEW_FILE_TTL_SECONDS = 5 * 60
|
|
_JWT_ALGORITHM = "HS256"
|
|
|
|
|
|
def _parse_pdf_regions(value: str | None) -> list[PreviewPdfRegion]:
|
|
if not value:
|
|
return []
|
|
try:
|
|
raw = json.loads(value)
|
|
except (TypeError, json.JSONDecodeError):
|
|
return []
|
|
if not isinstance(raw, list):
|
|
return []
|
|
out: list[PreviewPdfRegion] = []
|
|
for item in raw:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
try:
|
|
region = PreviewPdfRegion(**item)
|
|
except Exception:
|
|
continue
|
|
if (
|
|
0 <= region.x <= 1
|
|
and 0 <= region.y <= 1
|
|
and region.width > 0
|
|
and region.height > 0
|
|
):
|
|
out.append(region)
|
|
return out
|
|
|
|
|
|
def _preview_file_token(
|
|
*,
|
|
subject: str,
|
|
document_id: str,
|
|
) -> tuple[str, int]:
|
|
secret = get_jwt_secret(subject)
|
|
if secret is None:
|
|
raise HTTPException(status_code = 401, detail = "Invalid or expired token")
|
|
expires = datetime.now(timezone.utc) + timedelta(seconds = _PREVIEW_FILE_TTL_SECONDS)
|
|
payload = {
|
|
"sub": subject,
|
|
"aud": _PREVIEW_FILE_AUDIENCE,
|
|
"document_id": document_id,
|
|
"exp": expires,
|
|
}
|
|
token = jwt.encode(payload, secret, algorithm = _JWT_ALGORITHM)
|
|
return token, int(expires.timestamp())
|
|
|
|
|
|
def _subject_from_preview_file_token(document_id: str, token: str) -> str:
|
|
try:
|
|
unverified = jwt.decode(
|
|
token,
|
|
options = {
|
|
"verify_signature": False,
|
|
"verify_exp": False,
|
|
"verify_aud": False,
|
|
},
|
|
)
|
|
except jwt.InvalidTokenError as exc:
|
|
raise HTTPException(status_code = 401, detail = "Invalid preview token") from exc
|
|
subject = unverified.get("sub")
|
|
if not isinstance(subject, str) or not subject:
|
|
raise HTTPException(status_code = 401, detail = "Invalid preview token")
|
|
secret = get_jwt_secret(subject)
|
|
if secret is None:
|
|
raise HTTPException(status_code = 401, detail = "Invalid preview token")
|
|
try:
|
|
payload = jwt.decode(
|
|
token,
|
|
secret,
|
|
algorithms = [_JWT_ALGORITHM],
|
|
audience = _PREVIEW_FILE_AUDIENCE,
|
|
)
|
|
except jwt.InvalidTokenError as exc:
|
|
raise HTTPException(status_code = 401, detail = "Invalid preview token") from exc
|
|
if payload.get("document_id") != document_id:
|
|
raise HTTPException(status_code = 401, detail = "Invalid preview token")
|
|
return subject
|
|
|
|
|
|
def _resolve_document_file_or_404(doc_row: Any, document_id: str) -> Path:
|
|
try:
|
|
resolved = resolve_under_root(
|
|
doc_row["stored_path"],
|
|
root = rag_uploads_root(),
|
|
)
|
|
except ValueError as exc:
|
|
# Escape (symlink / ``..`` / absolute outside root) collapses to
|
|
# "file not found" — the auth row exists, the bytes do not.
|
|
logger.warning(
|
|
"RAG preview: stored_path escaped uploads root for doc %s: %s",
|
|
document_id,
|
|
exc,
|
|
)
|
|
raise HTTPException(
|
|
status_code = 404,
|
|
detail = "Document file not found",
|
|
) from exc
|
|
|
|
if not resolved.is_file():
|
|
raise HTTPException(
|
|
status_code = 404,
|
|
detail = "Document file not found",
|
|
)
|
|
return resolved
|
|
|
|
|
|
def _parse_range_header(range_header: str | None, size: int) -> tuple[int, int] | None:
|
|
if not range_header:
|
|
return None
|
|
if not range_header.startswith("bytes="):
|
|
raise ValueError("unsupported range unit")
|
|
spec = range_header[len("bytes=") :].strip()
|
|
if "," in spec or "-" not in spec:
|
|
raise ValueError("multiple or malformed ranges are not supported")
|
|
start_s, end_s = spec.split("-", 1)
|
|
if not start_s and not end_s:
|
|
raise ValueError("empty range")
|
|
if not start_s:
|
|
suffix = int(end_s)
|
|
if suffix <= 0:
|
|
raise ValueError("invalid suffix range")
|
|
start = max(0, size - suffix)
|
|
end = size - 1
|
|
else:
|
|
start = int(start_s)
|
|
end = int(end_s) if end_s else size - 1
|
|
if start < 0 or end < start or start >= size:
|
|
raise ValueError("range outside file")
|
|
return start, min(end, size - 1)
|
|
|
|
|
|
def _iter_file_range(path: Path, start: int, end: int):
|
|
with path.open("rb") as fh:
|
|
fh.seek(start)
|
|
remaining = end - start + 1
|
|
while remaining > 0:
|
|
chunk = fh.read(min(64 * 1024, remaining))
|
|
if not chunk:
|
|
break
|
|
remaining -= len(chunk)
|
|
yield chunk
|
|
|
|
|
|
def _serve_document_file_row(
|
|
doc_row: Any,
|
|
document_id: str,
|
|
range_header: str | None,
|
|
) -> FileResponse | Response | StreamingResponse:
|
|
resolved = _resolve_document_file_or_404(doc_row, document_id)
|
|
content_type, disposition, _media_kind = _preview_file_metadata(doc_row["filename"])
|
|
safe_name = _sanitize_filename(doc_row["filename"])
|
|
headers = {
|
|
"Content-Disposition": _content_disposition_header(safe_name, disposition),
|
|
"X-Content-Type-Options": "nosniff",
|
|
"Cache-Control": "private, max-age=0, must-revalidate",
|
|
"Accept-Ranges": "bytes",
|
|
}
|
|
size = resolved.stat().st_size
|
|
|
|
try:
|
|
byte_range = _parse_range_header(range_header, size)
|
|
except (TypeError, ValueError):
|
|
range_headers = dict(headers)
|
|
range_headers["Content-Range"] = f"bytes */{size}"
|
|
return Response(status_code = 416, headers = range_headers)
|
|
|
|
if byte_range is not None:
|
|
start, end = byte_range
|
|
range_headers = dict(headers)
|
|
range_headers["Content-Range"] = f"bytes {start}-{end}/{size}"
|
|
range_headers["Content-Length"] = str(end - start + 1)
|
|
return StreamingResponse(
|
|
_iter_file_range(resolved, start, end),
|
|
status_code = 206,
|
|
media_type = content_type,
|
|
headers = range_headers,
|
|
)
|
|
|
|
# FileResponse handles ordinary downloads; still advertise Accept-Ranges so
|
|
# PDF.js can switch to explicit range requests via the signed URL path.
|
|
return FileResponse(
|
|
path = str(resolved),
|
|
media_type = content_type,
|
|
headers = headers,
|
|
)
|
|
|
|
|
|
@router.get(
|
|
"/documents/{document_id}/preview-target",
|
|
response_model = PreviewTargetResponse,
|
|
)
|
|
def get_document_preview_target(
|
|
document_id: str,
|
|
chunk_id: Optional[str] = Query(None),
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> PreviewTargetResponse:
|
|
"""Resolve preview metadata for a document, optionally focused on a chunk.
|
|
|
|
See contracts.md §1 for the response shape. ``chunk_id`` is the durable
|
|
``rag_chunks.id`` carried as ``backendChunkId`` on the frontend; when
|
|
supplied it MUST belong to ``document_id`` or we collapse to 404 (so a
|
|
probe cannot enumerate cross-document chunk ids).
|
|
"""
|
|
doc_row = document_for_subject_or_404(document_id, current_subject)
|
|
_ct, _disposition, media_kind = _preview_file_metadata(doc_row["filename"])
|
|
|
|
base = {
|
|
"documentId": doc_row["id"],
|
|
"filename": doc_row["filename"],
|
|
"contentType": doc_row["content_type"],
|
|
"mediaKind": media_kind,
|
|
"byteSize": int(doc_row["byte_size"]),
|
|
"status": doc_row["status"],
|
|
"kbId": doc_row["kb_id"],
|
|
"threadId": doc_row["thread_id"],
|
|
}
|
|
|
|
if not chunk_id:
|
|
# Q2: document-row preview is metadata-only — frontend MUST NOT fall back to "first chunk".
|
|
return PreviewTargetResponse(
|
|
**base,
|
|
chunkId = None,
|
|
chunkIndex = None,
|
|
targetPage = None,
|
|
snippet = None,
|
|
kind = None,
|
|
imageUrl = None,
|
|
sourcePageIndex = None,
|
|
pageCharStart = None,
|
|
pageCharEnd = None,
|
|
lineStart = None,
|
|
lineEnd = None,
|
|
pdfRegions = [],
|
|
)
|
|
|
|
# One connection enforces membership AND fetches the row in a single query.
|
|
# A separate membership check would open a second SQLite connection and a
|
|
# TOCTOU window — if the chunk is deleted between the two calls, the fetch
|
|
# returns None and the route 500s (D1.1). Cross-document collapses to the
|
|
# same 404 — never 400 (would leak doc existence).
|
|
with closing_connection() as conn:
|
|
chunk_row = conn.execute(
|
|
"""
|
|
SELECT id, chunk_index, page_number, text, kind, image_path,
|
|
source_page_index, page_char_start, page_char_end,
|
|
line_start, line_end, pdf_regions_json
|
|
FROM rag_chunks WHERE id = ? AND document_id = ?
|
|
""",
|
|
(chunk_id, document_id),
|
|
).fetchone()
|
|
|
|
if chunk_row is None:
|
|
raise HTTPException(
|
|
status_code = 404,
|
|
detail = "Document not found",
|
|
)
|
|
|
|
chunk_kind: PreviewChunkKind = chunk_row["kind"] or "text" # type: ignore[assignment]
|
|
image_url: str | None = None
|
|
if chunk_kind == "image" and chunk_row["image_path"]:
|
|
image_url = (
|
|
f"/api/rag/images/{doc_row['id']}/" f"{Path(chunk_row['image_path']).name}"
|
|
)
|
|
|
|
return PreviewTargetResponse(
|
|
**base,
|
|
chunkId = chunk_row["id"],
|
|
chunkIndex = int(chunk_row["chunk_index"]),
|
|
targetPage = (
|
|
int(chunk_row["page_number"])
|
|
if chunk_row["page_number"] is not None
|
|
else None
|
|
),
|
|
snippet = chunk_row["text"] or "",
|
|
kind = chunk_kind,
|
|
imageUrl = image_url,
|
|
sourcePageIndex = chunk_row["source_page_index"],
|
|
pageCharStart = chunk_row["page_char_start"],
|
|
pageCharEnd = chunk_row["page_char_end"],
|
|
lineStart = chunk_row["line_start"],
|
|
lineEnd = chunk_row["line_end"],
|
|
pdfRegions = _parse_pdf_regions(chunk_row["pdf_regions_json"]),
|
|
)
|
|
|
|
|
|
@router.post(
|
|
"/documents/{document_id}/locators/backfill",
|
|
response_model = LocatorBackfillResponse,
|
|
)
|
|
def backfill_document_locators_route(
|
|
document_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> LocatorBackfillResponse:
|
|
"""In-place locator backfill for existing citations.
|
|
|
|
This preserves ``document_id`` and ``chunk_id``. Chunks are updated only
|
|
when their text has one unambiguous match in the parsed document text;
|
|
duplicate or missing matches remain null.
|
|
"""
|
|
doc_row = document_for_subject_or_404(document_id, current_subject)
|
|
resolved = _resolve_document_file_or_404(doc_row, document_id)
|
|
try:
|
|
result = backfill_document_locators(document_id, resolved)
|
|
except Exception as exc:
|
|
logger.warning(
|
|
"RAG locator backfill failed",
|
|
document_id = document_id,
|
|
error = str(exc),
|
|
)
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = "Document locators could not be backfilled",
|
|
) from exc
|
|
return LocatorBackfillResponse(
|
|
documentId = result.document_id,
|
|
totalChunks = result.total_chunks,
|
|
matched = result.matched,
|
|
alreadyLocated = result.already_located,
|
|
ambiguous = result.ambiguous,
|
|
missing = result.missing,
|
|
skipped = result.skipped,
|
|
regionsMatched = result.regions_matched,
|
|
pagesRefreshed = result.pages_refreshed,
|
|
)
|
|
|
|
|
|
@router.get(
|
|
"/documents/{document_id}/file-url",
|
|
response_model = PreviewFileUrlResponse,
|
|
)
|
|
def get_document_file_url(
|
|
document_id: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> PreviewFileUrlResponse:
|
|
"""Mint a short-lived signed URL for PDF.js range requests.
|
|
|
|
The normal bearer-protected `/file` route stays available for blob
|
|
fallback; this route keeps bearer access tokens out of query strings.
|
|
"""
|
|
document_for_subject_or_404(document_id, current_subject)
|
|
token, expires_at = _preview_file_token(
|
|
subject = current_subject,
|
|
document_id = document_id,
|
|
)
|
|
url = (
|
|
f"/api/rag/documents/{quote(document_id, safe = '')}/file-signed"
|
|
f"?token={quote(token, safe = '')}"
|
|
)
|
|
return PreviewFileUrlResponse(url = url, expiresAt = expires_at)
|
|
|
|
|
|
@router.get("/documents/{document_id}/file-signed", response_model = None)
|
|
def get_signed_document_file(
|
|
document_id: str,
|
|
request: Request,
|
|
token: str = Query(...),
|
|
) -> FileResponse | Response | StreamingResponse:
|
|
subject = _subject_from_preview_file_token(document_id, token)
|
|
doc_row = document_for_subject_or_404(document_id, subject)
|
|
return _serve_document_file_row(
|
|
doc_row,
|
|
document_id,
|
|
request.headers.get("range"),
|
|
)
|
|
|
|
|
|
@router.get("/documents/{document_id}/file", response_model = None)
|
|
def get_document_file(
|
|
document_id: str,
|
|
request: Request,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> FileResponse | Response | StreamingResponse:
|
|
"""Serve the original uploaded bytes for ``document_id``.
|
|
|
|
Path resolution per contracts.md §2.1: the route NEVER accepts a
|
|
client-supplied filename. ``stored_path`` is DB-issued and we run it
|
|
through ``resolve_under_root`` which delegates to ``_assert_contained``
|
|
(realpath + symlink/junction-safe). Any escape (symlink, junction,
|
|
``..``, absolute outside root) collapses to the second 404.
|
|
|
|
Content-Type and disposition come from the extension allowlist
|
|
(``_preview_file_metadata``). HTML/DOCX/unknown serve as
|
|
``attachment`` with safe content-type so an uploaded ``.html`` can
|
|
never execute in the app origin (Risk #3 / decisions Q7).
|
|
"""
|
|
doc_row = document_for_subject_or_404(document_id, current_subject)
|
|
return _serve_document_file_row(
|
|
doc_row,
|
|
document_id,
|
|
request.headers.get("range"),
|
|
)
|
|
|
|
|
|
@router.post("/search", response_model = SearchResponse)
|
|
def search(
|
|
payload: SearchRequest,
|
|
current_subject: str = Depends(get_current_subject),
|
|
) -> SearchResponse:
|
|
if bool(payload.kb_id) == bool(payload.thread_id):
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = "exactly one of kb_id or thread_id must be supplied",
|
|
)
|
|
if payload.kb_id:
|
|
_kb_or_404(payload.kb_id)
|
|
scope = kb_scope(payload.kb_id)
|
|
else:
|
|
scope = thread_scope(payload.thread_id)
|
|
|
|
# Query must use the same embedder as the scope (dim must match).
|
|
scope_embedder = _resolve_scope_embedder(scope)
|
|
logger.info(
|
|
"RAG search: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r",
|
|
scope,
|
|
scope_embedder or "<default>",
|
|
payload.mode,
|
|
payload.top_k,
|
|
payload.min_score,
|
|
payload.enable_rerank,
|
|
payload.query[:120],
|
|
)
|
|
|
|
# Reranker needs a wider candidate pool than top_k.
|
|
candidate_k = (
|
|
max(payload.top_k, RAG_RERANK_CANDIDATE_K)
|
|
if payload.enable_rerank
|
|
else payload.top_k
|
|
)
|
|
|
|
if payload.mode == "bm25":
|
|
hits = retrieval.retrieve_bm25(scope, payload.query, candidate_k)
|
|
elif payload.mode == "dense":
|
|
hits = retrieval.retrieve_dense(
|
|
scope,
|
|
payload.query,
|
|
candidate_k,
|
|
document_ids = payload.document_ids,
|
|
embedder_model = scope_embedder,
|
|
)
|
|
else:
|
|
hits = retrieval.retrieve_hybrid(
|
|
scope,
|
|
payload.query,
|
|
k = candidate_k,
|
|
document_ids = payload.document_ids,
|
|
embedder_model = scope_embedder,
|
|
)
|
|
|
|
retrieved_count = len(hits)
|
|
if payload.min_score > 0.0:
|
|
hits = retrieval.filter_by_min_score(hits, payload.min_score)
|
|
logger.info(
|
|
"RAG search: retrieved=%d met_threshold=%d (min_score=%.3f)",
|
|
retrieved_count,
|
|
len(hits),
|
|
payload.min_score,
|
|
)
|
|
else:
|
|
logger.info("RAG search: retrieved=%d (no threshold)", retrieved_count)
|
|
|
|
chunk_ids = [h.chunk_id for h in hits]
|
|
chunk_lookup: dict[str, dict] = {}
|
|
if chunk_ids:
|
|
placeholders = ",".join("?" for _ in chunk_ids)
|
|
with closing_connection() as conn:
|
|
rows = conn.execute(
|
|
f"""
|
|
SELECT c.id AS chunk_id, c.document_id, c.chunk_index, c.text,
|
|
c.page_number, c.kind, c.image_path, c.linked_chunk_id,
|
|
c.source_page_index, c.page_char_start, c.page_char_end,
|
|
c.line_start, c.line_end,
|
|
d.filename
|
|
FROM rag_chunks c
|
|
JOIN rag_documents d ON d.id = c.document_id
|
|
WHERE c.id IN ({placeholders})
|
|
""",
|
|
chunk_ids,
|
|
).fetchall()
|
|
for r in rows:
|
|
chunk_lookup[r["chunk_id"]] = dict(r)
|
|
|
|
if payload.enable_rerank:
|
|
pairs = [
|
|
(hit, chunk_lookup[hit.chunk_id]["text"])
|
|
for hit in hits
|
|
if hit.chunk_id in chunk_lookup
|
|
]
|
|
hits = reranker.rerank(
|
|
payload.query,
|
|
pairs,
|
|
model_name = payload.reranker_model,
|
|
top_k = payload.top_k,
|
|
)
|
|
else:
|
|
hits = hits[: payload.top_k]
|
|
|
|
out: list[SearchHit] = []
|
|
for hit in hits:
|
|
meta = chunk_lookup.get(hit.chunk_id)
|
|
if not meta:
|
|
continue
|
|
kind = meta.get("kind", "text") or "text"
|
|
image_url: str | None = None
|
|
if kind == "image" and meta.get("image_path"):
|
|
image_url = (
|
|
f"/api/rag/images/{meta['document_id']}/{Path(meta['image_path']).name}"
|
|
)
|
|
out.append(
|
|
SearchHit(
|
|
chunk_id = hit.chunk_id,
|
|
document_id = meta["document_id"],
|
|
chunk_index = meta["chunk_index"],
|
|
text = meta["text"] or "",
|
|
score = hit.score,
|
|
page_number = meta.get("page_number"),
|
|
filename = meta.get("filename"),
|
|
kind = kind,
|
|
image_url = image_url,
|
|
source_page_index = meta.get("source_page_index"),
|
|
page_char_start = meta.get("page_char_start"),
|
|
page_char_end = meta.get("page_char_end"),
|
|
line_start = meta.get("line_start"),
|
|
line_end = meta.get("line_end"),
|
|
)
|
|
)
|
|
logger.info("RAG search: returning %d hits", len(out))
|
|
return SearchResponse(hits = out)
|