senlinyy's picture
feat: update for feedback chatbot
5ce9fab
Raw
History Blame Contribute Delete
5.78 kB
"""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)