"""RAG engine: file text -> chunks -> HF-API embeddings -> LanceDB. Designed for HF Spaces Free CPU Basic: * No local embedding/LLM weights are downloaded. * LanceDB lives on ephemeral disk locally, but its files are mirrored into the linked HF Dataset so cold starts can restore the vector cache instead of rebuilding from rubric files every time. """ from __future__ import annotations import logging from pathlib import Path from threading import Lock from typing import Iterable, List, Optional from llama_index.core import Document, StorageContext, VectorStoreIndex, Settings from llama_index.core.node_parser import MarkdownNodeParser from llama_index.core.schema import MetadataMode, NodeWithScore, TextNode from llama_index.vector_stores.lancedb import LanceDBVectorStore from config import EMBED_MODEL, HF_TOKEN, LANCEDB_PATH from document_parser import parse_file_text from hf_embedding import SyncHuggingFaceInferenceEmbedding logger = logging.getLogger(__name__) TABLE_NAME = "rubrics" FILENAME_KEY = "source_filename" class RagEngine: """Thread-safe singleton wrapper around the LanceDB-backed index.""" def __init__(self) -> None: self._lock = Lock() self._index: Optional[VectorStoreIndex] = None self._vector_store: Optional[LanceDBVectorStore] = None if not HF_TOKEN: logger.warning( "HF_TOKEN is not set — embedding calls will fail. " "Set it in .env before uploading or querying rubrics." ) Settings.embed_model = None # type: ignore[assignment] else: Settings.embed_model = SyncHuggingFaceInferenceEmbedding( model_name=EMBED_MODEL, token=HF_TOKEN, ) Settings.llm = None Settings.node_parser = MarkdownNodeParser() Path(LANCEDB_PATH).mkdir(parents=True, exist_ok=True) self._vector_store = LanceDBVectorStore( uri=LANCEDB_PATH, table_name=TABLE_NAME, mode="overwrite" if not self._table_exists() else "append", ) def _table_exists(self) -> bool: import lancedb try: db = lancedb.connect(LANCEDB_PATH) return TABLE_NAME in db.table_names() except Exception: # noqa: BLE001 return False def _ensure_index(self) -> VectorStoreIndex: if self._index is None: storage = StorageContext.from_defaults(vector_store=self._vector_store) if self._table_exists(): self._index = VectorStoreIndex.from_vector_store( vector_store=self._vector_store, storage_context=storage, ) else: self._index = VectorStoreIndex.from_documents( [], storage_context=storage ) return self._index @staticmethod def _file_to_documents(file_path: Path, filename: str) -> List[Document]: try: text = parse_file_text(file_path, filename) except Exception as exc: # noqa: BLE001 logger.warning("Failed to parse %s: %s", filename, exc) return [] if not text.strip(): return [] return [ Document( text=text, metadata={FILENAME_KEY: filename}, excluded_llm_metadata_keys=[FILENAME_KEY], excluded_embed_metadata_keys=[FILENAME_KEY], ) ] def index_file(self, file_path: Path, filename: str) -> int: with self._lock: self.delete_by_filename(filename, _locked=True) docs = self._file_to_documents(file_path, filename) if not docs: return 0 index = self._ensure_index() nodes = Settings.node_parser.get_nodes_from_documents(docs) for n in nodes: n.metadata[FILENAME_KEY] = filename index.insert_nodes(nodes) return len(nodes) def index_many(self, items: Iterable[tuple[Path, str]]) -> int: total = 0 for path, name in items: total += self.index_file(path, name) return total def delete_by_filename(self, filename: str, *, _locked: bool = False) -> int: def _do() -> int: import lancedb if not self._table_exists(): return 0 db = lancedb.connect(LANCEDB_PATH) tbl = db.open_table(TABLE_NAME) try: before = tbl.count_rows() tbl.delete(f"metadata.{FILENAME_KEY} = '{filename}'") removed = before - tbl.count_rows() except Exception as exc: # noqa: BLE001 logger.warning("delete_by_filename failed: %s", exc) removed = 0 self._index = None return removed if _locked: return _do() with self._lock: return _do() def retrieve(self, query: str, top_k: int = 4) -> List[NodeWithScore]: with self._lock: if not self._table_exists(): return [] index = self._ensure_index() retriever = index.as_retriever(similarity_top_k=top_k) return retriever.retrieve(query) @staticmethod def format_context(nodes: List[NodeWithScore]) -> str: if not nodes: return "" chunks = [] for i, n in enumerate(nodes, 1): src = n.node.metadata.get(FILENAME_KEY, "unknown") text = n.node.get_content(metadata_mode=MetadataMode.NONE).strip() chunks.append(f"[{i}] (source: {src})\n{text}") return "\n\n---\n\n".join(chunks)