import functools import json import logging import os import chromadb from llama_index.core import QueryBundle, VectorStoreIndex from llama_index.core.retrievers import BaseRetriever, VectorIndexRetriever from llama_index.core.schema import NodeWithScore, TextNode from llama_index.core.vector_stores import FilterOperator, MetadataFilter, MetadataFilters from llama_index.embeddings.huggingface import HuggingFaceEmbedding from llama_index.postprocessor.sbert_rerank import SentenceTransformerRerank from llama_index.retrievers.bm25 import BM25Retriever from llama_index.vector_stores.chroma import ChromaVectorStore from rag.config import ( BOOKS, CANDIDATE_TOP_K, CHROMA_COLLECTION, CHROMA_DIR, CHROMA_DISTANCE, CHUNKS_PATH, EMBED_MODEL, HEADING_SEPARATOR, HF_INDEX_REPO_ENV, QUERY_INSTRUCTION, RERANK_MODEL, ) from rag.types import ChunkMetadata RETRIEVAL_MODES = ("vector", "bm25", "hybrid", "rerank") logger = logging.getLogger(__name__) @functools.cache def make_embed_model() -> HuggingFaceEmbedding: return HuggingFaceEmbedding( model_name=EMBED_MODEL, query_instruction=QUERY_INSTRUCTION, ) def flatten_metadata(metadata: ChunkMetadata) -> dict[str, str | bool]: code_tags = metadata["code_tags"] return { "book": metadata["book"], "book_title": metadata["book_title"], "chapter": metadata["chapter"], "part": metadata["part"], "heading_path": HEADING_SEPARATOR.join(metadata["heading_path"]), "path": metadata["path"], "url": metadata["url"], "has_code": metadata["has_code"], "code_tags": " | ".join(code_tags), } def download_index_if_missing() -> None: if CHROMA_DIR.exists() and CHUNKS_PATH.exists(): return repo_id = os.environ.get(HF_INDEX_REPO_ENV, "").strip() if not repo_id: raise RuntimeError( f"No index at {CHROMA_DIR} and {HF_INDEX_REPO_ENV} is unset. " "Build one with `uv run python -m ingest.build_index`, or set " f"{HF_INDEX_REPO_ENV}." ) from huggingface_hub import snapshot_download # We download the whole data directory, not just the collection, because BM25 # builds its term index from `chunks.jsonl` rather than from the vector # database. logger.info("downloading index from %s", repo_id) CHROMA_DIR.parent.mkdir(parents=True, exist_ok=True) snapshot_download(repo_id=repo_id, repo_type="dataset", local_dir=str(CHROMA_DIR.parent)) def open_collection(fresh: bool = False): client = chromadb.PersistentClient(path=str(CHROMA_DIR)) if fresh: try: client.delete_collection(CHROMA_COLLECTION) except Exception: pass return client.get_or_create_collection( CHROMA_COLLECTION, metadata={"hnsw:space": CHROMA_DISTANCE} ) def make_vector_store(fresh: bool = False) -> ChromaVectorStore: return ChromaVectorStore(chroma_collection=open_collection(fresh=fresh)) @functools.cache def load_nodes() -> tuple[TextNode, ...]: with CHUNKS_PATH.open(encoding="utf-8") as handle: chunks = [json.loads(line) for line in handle] return tuple( TextNode(id_=chunk["id"], text=chunk["text"], metadata=flatten_metadata(chunk["metadata"])) for chunk in chunks ) def book_filter(book: str | None) -> MetadataFilters | None: if not book or book not in BOOKS: return None return MetadataFilters( filters=[MetadataFilter(key="book", value=book, operator=FilterOperator.EQ)] ) class RerankingRetriever(BaseRetriever): def __init__(self, retriever: BaseRetriever, reranker: SentenceTransformerRerank): self._retriever = retriever self._reranker = reranker super().__init__() def _retrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]: candidates = self._retriever.retrieve(query_bundle) return self._reranker.postprocess_nodes(candidates, query_bundle=query_bundle) # The constant from the original reciprocal-rank-fusion paper, and the same value # LlamaIndex uses. RRF_K = 60.0 class HybridRetriever(BaseRetriever): """Vector and BM25 candidates merged by reciprocal rank fusion. We write the fusion out ourselves rather than using `QueryFusionRetriever`. Its fusion is the same arithmetic, but it also generates query variations with an LLM, which would spend the user's key inside retrieval and may duplicate the query the agent already wrote. """ def __init__(self, retrievers: list[BaseRetriever], top_k: int): self._retrievers = retrievers self._top_k = top_k super().__init__() def _retrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]: scores: dict[str, float] = {} found: dict[str, NodeWithScore] = {} for retriever in self._retrievers: ranked = sorted( retriever.retrieve(query_bundle), key=lambda n: n.score or 0.0, reverse=True ) for rank, node in enumerate(ranked): key = node.node.hash found[key] = node scores[key] = scores.get(key, 0.0) + 1.0 / (rank + RRF_K) best = sorted(scores, key=lambda key: scores[key], reverse=True)[: self._top_k] return [NodeWithScore(node=found[key].node, score=scores[key]) for key in best] def _fused(final_top_k: int, book: str | None) -> HybridRetriever: embed_model = make_embed_model() filters = book_filter(book) index = VectorStoreIndex.from_vector_store( vector_store=make_vector_store(), embed_model=embed_model ) return HybridRetriever( retrievers=[ VectorIndexRetriever( index=index, similarity_top_k=CANDIDATE_TOP_K, embed_model=embed_model, filters=filters, ), BM25Retriever.from_defaults( nodes=list(load_nodes()), similarity_top_k=CANDIDATE_TOP_K, filters=filters, ), ], top_k=final_top_k, ) @functools.cache def load_retriever( mode: str, top_k: int, book: str | None = None, rerank_model: str | None = None ) -> BaseRetriever: if mode == "vector": embed_model = make_embed_model() index = VectorStoreIndex.from_vector_store( vector_store=make_vector_store(), embed_model=embed_model ) return VectorIndexRetriever( index=index, similarity_top_k=top_k, embed_model=embed_model, filters=book_filter(book), ) if mode == "bm25": return BM25Retriever.from_defaults( nodes=list(load_nodes()), similarity_top_k=top_k, filters=book_filter(book) ) if mode == "hybrid": return _fused(top_k, book) if mode == "rerank": return RerankingRetriever( retriever=_fused(CANDIDATE_TOP_K, book), reranker=SentenceTransformerRerank(model=rerank_model or RERANK_MODEL, top_n=top_k), ) raise ValueError(f"Unknown retrieval mode {mode!r}. Expected one of: {RETRIEVAL_MODES}")