unsloth/studio/backend/routes/rag.py
Daniel Han ab0828b976 Studio: fix RAG correctness bugs
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).
2026-05-31 09:56:23 +00:00

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)