diff --git a/ai_engineering/rag_assistant/.env.example b/ai_engineering/rag_assistant/.env.example index 644edb64..f869f86a 100644 --- a/ai_engineering/rag_assistant/.env.example +++ b/ai_engineering/rag_assistant/.env.example @@ -3,3 +3,8 @@ OPENAI_API_KEY=sk-your-key-here OPENAI_MODEL=gpt-4o-mini EMBEDDING_MODEL=sentence-transformers/all-MiniLM-L6-v2 + +# Optional serving configuration. Leave RERANKER_MODEL unset to preserve +# embedding-only retrieval. +# RERANKER_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2 +RERANKER_CANDIDATES=20 diff --git a/ai_engineering/rag_assistant/README.md b/ai_engineering/rag_assistant/README.md index 7f977cc3..4ffd55dc 100644 --- a/ai_engineering/rag_assistant/README.md +++ b/ai_engineering/rag_assistant/README.md @@ -145,6 +145,59 @@ This is a 5-case smoke eval (`sample_docs/eval_cases.json`), not a large benchma - **Honest evaluation**: ships an actual `recall@k` + MRR eval harness with checked-in cases, instead of just claiming "high quality". - **Reproducible smoke pipeline**: `python -m smoke.run_smoke` exercises ingest + retrieve + eval end-to-end, no key required, and writes a markdown report a reviewer can read in 30 seconds. +## Optional Cross-Encoder Re-ranking + +Embedding retrieval remains the default. Add `--rerank` to `ask` or `eval` to +retrieve a larger candidate pool and re-score it with +`cross-encoder/ms-marco-MiniLM-L-6-v2`: + +```bash +python cli.py ask \ + --store ./index \ + --question "What is RAG?" \ + --rerank \ + --candidate-k 20 \ + -k 5 + +python cli.py eval \ + --store ./index \ + --cases ./sample_docs/eval_cases.json \ + --rerank \ + --candidate-k 20 \ + -k 3 +``` + +For the Flask service, set `RERANKER_MODEL` and optionally +`RERANKER_CANDIDATES`. The JSON response exposes both the final cross-encoder +`score` and the original `retrieval_score`. + +Run the baseline comparison with: + +```bash +python -m smoke.run_reranker_eval +``` + +The generated `reports/reranker_eval.md` records baseline and re-ranked +recall@k and MRR, their deltas, and per-case reciprocal-rank changes. The +current five-case suite is saturated, so the observed delta is provisional +until the expanded benchmark in issue #18 is available. + +### Measured Smoke Result + +A CPU run using `sentence-transformers/all-MiniLM-L6-v2` and +`cross-encoder/ms-marco-MiniLM-L-6-v2` produced: + +| Configuration | recall@3 | MRR | +|---|---:|---:| +| Embedding only | 1.000 | 0.900 | +| Cross-encoder re-ranked | 1.000 | 1.000 | +| Observed delta | +0.000 | +0.100 | + +The re-ranker moved the approximate-nearest-neighbor question from rank 2 to +rank 1. This records the observed result on the checked-in five-case smoke +suite; it does not establish a general quality lift. Issue #18 remains the +required follow-up for a discriminative benchmark. + ## Scope - The chunker is character-based and language-agnostic, which is good portability but slightly worse than tokenization-aware chunking for very long contexts. diff --git a/ai_engineering/rag_assistant/app.py b/ai_engineering/rag_assistant/app.py index 81bf0ed9..f7c1344c 100644 --- a/ai_engineering/rag_assistant/app.py +++ b/ai_engineering/rag_assistant/app.py @@ -5,6 +5,7 @@ Build with `build_app("./index")` so a saved store is loaded once at startup and reused across requests (no per-request rebuild). +Set RERANKER_MODEL to opt into cross-encoder re-ranking. """ from __future__ import annotations @@ -13,12 +14,23 @@ from flask import Flask, jsonify, request -from rag import RAGPipeline +from rag import CrossEncoderReranker, RAGPipeline def build_app(store_path: str) -> Flask: app = Flask(__name__) - pipeline = RAGPipeline.load(store_path) + reranker_model = os.environ.get("RERANKER_MODEL", "").strip() + reranker = ( + CrossEncoderReranker(model_name=reranker_model) + if reranker_model + else None + ) + candidate_k = int(os.environ.get("RERANKER_CANDIDATES", "20")) + pipeline = RAGPipeline.load( + store_path, + reranker=reranker, + candidate_k=candidate_k, + ) @app.route("/health", methods=["GET"]) def health(): @@ -39,6 +51,7 @@ def ask(): "retrieved": [ { "score": r.score, + "retrieval_score": r.retrieval_score, "source": r.chunk.source, "chunk_index": r.chunk.chunk_index, } diff --git a/ai_engineering/rag_assistant/cli.py b/ai_engineering/rag_assistant/cli.py index c1b2e0b8..4c757ac3 100644 --- a/ai_engineering/rag_assistant/cli.py +++ b/ai_engineering/rag_assistant/cli.py @@ -9,8 +9,10 @@ Usage: python cli.py ingest --docs-dir ./sample_docs --store ./index python cli.py ask --store ./index --question "What is RAG?" + python cli.py ask --store ./index --question "What is RAG?" --rerank python cli.py serve --store ./index --port 8080 - python cli.py eval --store ./index --cases ./sample_docs/eval_cases.json --k 3 + python cli.py eval --store ./index --cases ./sample_docs/eval_cases.json -k 3 + python cli.py eval --store ./index --cases ./sample_docs/eval_cases.json -k 3 --rerank """ from __future__ import annotations @@ -20,9 +22,18 @@ import os import sys from pathlib import Path -from typing import List - -from rag import Chunker, Document, RAGPipeline, eval_retrieval, EvalCase +from typing import List, Optional + +from rag import ( + Chunker, + CrossEncoderReranker, + DEFAULT_RERANKER_MODEL, + Document, + EvalCase, + RAGPipeline, + Reranker, + eval_retrieval, +) from rag.embedder import make_default_embedder from rag.vector_store import VectorStore @@ -48,6 +59,31 @@ def load_docs_from_dir(dir_path: str) -> List[Document]: return docs +def _build_reranker(args: argparse.Namespace) -> Optional[Reranker]: + if not getattr(args, "rerank", False): + return None + return CrossEncoderReranker(model_name=args.reranker_model) + + +def _add_reranker_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--rerank", + action="store_true", + help="re-rank a larger candidate pool with a cross-encoder", + ) + parser.add_argument( + "--reranker-model", + default=DEFAULT_RERANKER_MODEL, + help="sentence-transformers CrossEncoder model name", + ) + parser.add_argument( + "--candidate-k", + type=int, + default=20, + help="embedding candidates passed to the re-ranker", + ) + + def cmd_ingest(args: argparse.Namespace) -> None: docs = load_docs_from_dir(args.docs_dir) if not docs: @@ -60,7 +96,11 @@ def cmd_ingest(args: argparse.Namespace) -> None: def cmd_ask(args: argparse.Namespace) -> None: - pipeline = RAGPipeline.load(args.store) + pipeline = RAGPipeline.load( + args.store, + reranker=_build_reranker(args), + candidate_k=args.candidate_k, + ) result = pipeline.ask(args.question, k=args.k) print(f"\nAnswer:\n{result.answer}\n") if result.sources: @@ -82,7 +122,16 @@ def cmd_eval(args: argparse.Namespace) -> None: if embedder.dim != store.dim: sys.exit(f"embedder dim {embedder.dim} != store dim {store.dim}; reindex.") from rag.retriever import Retriever - result = eval_retrieval(Retriever(embedder, store), cases, k=args.k) + result = eval_retrieval( + Retriever( + embedder, + store, + reranker=_build_reranker(args), + candidate_k=args.candidate_k, + ), + cases, + k=args.k, + ) print(result) for row in result.per_case: print(f" - {row['question'][:60]:60s} recall={row['recall']:.0f} rr={row['reciprocal_rank']:.3f}") @@ -103,6 +152,7 @@ def main() -> None: p_ask.add_argument("--store", required=True) p_ask.add_argument("--question", required=True) p_ask.add_argument("-k", type=int, default=5) + _add_reranker_args(p_ask) p_ask.set_defaults(func=cmd_ask) p_serve = sub.add_parser("serve") @@ -114,6 +164,7 @@ def main() -> None: p_eval.add_argument("--store", required=True) p_eval.add_argument("--cases", required=True) p_eval.add_argument("-k", type=int, default=5) + _add_reranker_args(p_eval) p_eval.set_defaults(func=cmd_eval) args = parser.parse_args() diff --git a/ai_engineering/rag_assistant/rag/__init__.py b/ai_engineering/rag_assistant/rag/__init__.py index 706b4685..3dc1da9f 100644 --- a/ai_engineering/rag_assistant/rag/__init__.py +++ b/ai_engineering/rag_assistant/rag/__init__.py @@ -4,7 +4,8 @@ Chunker - split documents into overlapping chunks Embedder - convert chunks to dense vectors VectorStore - FAISS-backed nearest-neighbor index with metadata - Retriever - thin wrapper that wires Embedder + VectorStore together + Retriever - embedding retrieval with optional cross-encoder re-ranking + Reranker - protocol for re-ordering an initial candidate pool Generator - calls the OpenAI Chat Completions API with retrieved context RAGPipeline - end-to-end: ingest -> retrieve -> generate eval_retrieval - retrieval-quality metrics (recall@k, MRR) @@ -16,6 +17,11 @@ from .chunker import Chunker, Document, Chunk from .embedder import Embedder, HashEmbedder from .vector_store import VectorStore +from .reranker import ( + CrossEncoderReranker, + DEFAULT_RERANKER_MODEL, + Reranker, +) from .retriever import Retriever, RetrievedChunk from .generator import Generator from .pipeline import RAGPipeline @@ -25,6 +31,7 @@ "Chunker", "Document", "Chunk", "Embedder", "HashEmbedder", "VectorStore", + "Reranker", "CrossEncoderReranker", "DEFAULT_RERANKER_MODEL", "Retriever", "RetrievedChunk", "Generator", "RAGPipeline", diff --git a/ai_engineering/rag_assistant/rag/embedder.py b/ai_engineering/rag_assistant/rag/embedder.py index c056261b..3460958d 100644 --- a/ai_engineering/rag_assistant/rag/embedder.py +++ b/ai_engineering/rag_assistant/rag/embedder.py @@ -66,7 +66,14 @@ def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"): from sentence_transformers import SentenceTransformer # type: ignore self._model = SentenceTransformer(model_name) - self.dim = self._model.get_sentence_embedding_dimension() + dimension_getter = getattr( + self._model, + "get_embedding_dimension", + None, + ) + if dimension_getter is None: + dimension_getter = self._model.get_sentence_embedding_dimension + self.dim = dimension_getter() def embed(self, texts: Sequence[str]) -> np.ndarray: if not texts: diff --git a/ai_engineering/rag_assistant/rag/pipeline.py b/ai_engineering/rag_assistant/rag/pipeline.py index 6684fea4..983688b2 100644 --- a/ai_engineering/rag_assistant/rag/pipeline.py +++ b/ai_engineering/rag_assistant/rag/pipeline.py @@ -1,7 +1,8 @@ """End-to-end ingest + retrieve + generate pipeline. `RAGPipeline.ingest(docs)` runs the chunker + embedder + store. -`RAGPipeline.ask(question, k=5)` runs the retriever + generator. +`RAGPipeline.ask(question, k=5)` runs the retriever, optional re-ranker, +and generator. `from_env()` reads `OPENAI_API_KEY`, `OPENAI_MODEL`, `EMBEDDING_MODEL`, so the CLI and serving layers share configuration. @@ -18,6 +19,7 @@ from .chunker import Chunk, Chunker, Document from .embedder import Embedder, make_default_embedder from .generator import GenerationResult, Generator +from .reranker import Reranker from .retriever import RetrievedChunk, Retriever from .vector_store import VectorStore @@ -37,10 +39,17 @@ def __init__( store: VectorStore, generator: Optional[Generator] = None, chunker: Optional[Chunker] = None, + reranker: Optional[Reranker] = None, + candidate_k: int = 20, ): self.embedder = embedder self.store = store - self.retriever = Retriever(embedder, store) + self.retriever = Retriever( + embedder, + store, + reranker=reranker, + candidate_k=candidate_k, + ) self.generator = generator self.chunker = chunker or Chunker() @@ -95,6 +104,8 @@ def load( dir_path: str, embedder: Optional[Embedder] = None, generator: Optional[Generator] = None, + reranker: Optional[Reranker] = None, + candidate_k: int = 20, ) -> "RAGPipeline": store = VectorStore.load(dir_path) emb = embedder or make_default_embedder() @@ -103,4 +114,10 @@ def load( f"embedder dim {emb.dim} does not match stored index dim {store.dim}; " "did you change EMBEDDING_MODEL since indexing?" ) - return cls(embedder=emb, store=store, generator=generator or Generator()) + return cls( + embedder=emb, + store=store, + generator=generator or Generator(), + reranker=reranker, + candidate_k=candidate_k, + ) diff --git a/ai_engineering/rag_assistant/rag/reranker.py b/ai_engineering/rag_assistant/rag/reranker.py new file mode 100644 index 00000000..e142f745 --- /dev/null +++ b/ai_engineering/rag_assistant/rag/reranker.py @@ -0,0 +1,97 @@ +"""Optional cross-encoder re-ranking for retrieved chunks. + +The production implementation wraps sentence-transformers CrossEncoder, while +the constructor accepts an injected model so tests can remain deterministic +and network-free. +""" + +from __future__ import annotations + +from typing import List, Optional, Protocol, Sequence, TYPE_CHECKING + +if TYPE_CHECKING: + from .retriever import RetrievedChunk + + +DEFAULT_RERANKER_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2" + + +class CrossEncoderModel(Protocol): + def predict( + self, + sentences, + *, + batch_size: int, + show_progress_bar: bool, + ): ... + + +class Reranker(Protocol): + def rerank( + self, + query: str, + candidates: Sequence["RetrievedChunk"], + k: int, + ) -> List["RetrievedChunk"]: ... + + +class CrossEncoderReranker: + """Re-score query/chunk pairs with a sentence-transformers cross-encoder.""" + + def __init__( + self, + model_name: str = DEFAULT_RERANKER_MODEL, + *, + batch_size: int = 32, + model: Optional[CrossEncoderModel] = None, + ): + if batch_size < 1: + raise ValueError("batch_size must be >= 1") + self.model_name = model_name + self.batch_size = batch_size + if model is None: + from sentence_transformers import CrossEncoder # type: ignore + + model = CrossEncoder(model_name) + self._model = model + + def rerank( + self, + query: str, + candidates: Sequence["RetrievedChunk"], + k: int, + ) -> List["RetrievedChunk"]: + if k <= 0 or not query.strip() or not candidates: + return [] + + pairs = [(query, candidate.chunk.text) for candidate in candidates] + raw_scores = self._model.predict( + pairs, + batch_size=self.batch_size, + show_progress_bar=False, + ) + scores = [float(score) for score in raw_scores] + if len(scores) != len(candidates): + raise RuntimeError( + "cross-encoder returned " + f"{len(scores)} scores for {len(candidates)} candidates" + ) + + from .retriever import RetrievedChunk + + ranked = sorted( + enumerate(zip(candidates, scores)), + key=lambda item: (-item[1][1], item[0]), + ) + return [ + RetrievedChunk( + score=rerank_score, + chunk=candidate.chunk, + retrieval_score=( + candidate.retrieval_score + if candidate.retrieval_score is not None + else candidate.score + ), + ) + for _, (candidate, rerank_score) in ranked[:k] + ] diff --git a/ai_engineering/rag_assistant/rag/retriever.py b/ai_engineering/rag_assistant/rag/retriever.py index 048fc038..e962867b 100644 --- a/ai_engineering/rag_assistant/rag/retriever.py +++ b/ai_engineering/rag_assistant/rag/retriever.py @@ -2,16 +2,18 @@ Keeping retrieval as its own class (rather than a free function) means swapping in a different backend (Chroma, Pinecone, Weaviate) only -requires writing a new `Retriever` subclass. +requires writing a new `Retriever` subclass. An optional re-ranker can +re-score a larger candidate pool before the final top-k is returned. """ from __future__ import annotations from dataclasses import dataclass -from typing import List +from typing import List, Optional from .chunker import Chunk from .embedder import Embedder +from .reranker import Reranker from .vector_store import VectorStore @@ -19,20 +21,40 @@ class RetrievedChunk: score: float chunk: Chunk + retrieval_score: Optional[float] = None class Retriever: - def __init__(self, embedder: Embedder, store: VectorStore): + def __init__( + self, + embedder: Embedder, + store: VectorStore, + reranker: Optional[Reranker] = None, + candidate_k: int = 20, + ): if embedder.dim != store.dim: raise ValueError( f"embedder dim {embedder.dim} != store dim {store.dim}" ) + if candidate_k < 1: + raise ValueError("candidate_k must be >= 1") self.embedder = embedder self.store = store + self.reranker = reranker + self.candidate_k = candidate_k def retrieve(self, query: str, k: int = 5) -> List[RetrievedChunk]: - if not query.strip(): + if not query.strip() or k <= 0: return [] + + pool_k = max(k, self.candidate_k) if self.reranker is not None else k qvec = self.embedder.embed([query]) - scores, chunks = self.store.search(qvec, k=k) - return [RetrievedChunk(score=float(s), chunk=c) for s, c in zip(scores, chunks)] + scores, chunks = self.store.search(qvec, k=pool_k) + candidates = [ + RetrievedChunk(score=float(score), chunk=chunk) + for score, chunk in zip(scores, chunks) + ] + + if self.reranker is None: + return candidates[:k] + return self.reranker.rerank(query, candidates, k) diff --git a/ai_engineering/rag_assistant/smoke/run_reranker_eval.py b/ai_engineering/rag_assistant/smoke/run_reranker_eval.py new file mode 100644 index 00000000..394d5471 --- /dev/null +++ b/ai_engineering/rag_assistant/smoke/run_reranker_eval.py @@ -0,0 +1,119 @@ +"""Compare embedding-only retrieval against optional cross-encoder re-ranking. + +This script requires sentence-transformers model weights. It does not require +an OpenAI API key. The checked-in five-case suite is saturated, so the report +is a provisional measurement until the expanded benchmark in issue #18 lands. + +Outputs: + reports/reranker_eval.md +""" + +from __future__ import annotations + +import json +import os +from datetime import datetime, timezone + +from rag.chunker import Chunker, Document +from rag.embedder import make_default_embedder +from rag.eval import EvalCase, eval_retrieval +from rag.reranker import CrossEncoderReranker, DEFAULT_RERANKER_MODEL +from rag.retriever import Retriever +from rag.vector_store import VectorStore + + +SAMPLE_DIR = os.path.join(os.path.dirname(__file__), "..", "sample_docs") +REPORT_PATH = "reports/reranker_eval.md" +K = int(os.environ.get("RERANKER_EVAL_K", "3")) +CANDIDATE_K = int(os.environ.get("RERANKER_CANDIDATES", "20")) +MODEL_NAME = os.environ.get("RERANKER_MODEL", DEFAULT_RERANKER_MODEL) + + +def _load_docs() -> list[Document]: + docs = [] + for name in sorted(os.listdir(SAMPLE_DIR)): + if not name.endswith(".md"): + continue + with open(os.path.join(SAMPLE_DIR, name), "r", encoding="utf-8") as fh: + docs.append(Document(source=name, text=fh.read())) + return docs + + +def main() -> None: + os.makedirs("reports", exist_ok=True) + docs = _load_docs() + embedder = make_default_embedder() + + chunker = Chunker(chunk_size=400, chunk_overlap=80) + chunks = chunker.chunk_corpus(docs) + store = VectorStore(dim=embedder.dim) + store.add(chunks, embedder.embed([chunk.text for chunk in chunks])) + + with open( + os.path.join(SAMPLE_DIR, "eval_cases.json"), + "r", + encoding="utf-8", + ) as fh: + cases = [EvalCase(**case) for case in json.load(fh)] + + baseline = eval_retrieval( + Retriever(embedder, store), + cases, + k=K, + ) + reranked = eval_retrieval( + Retriever( + embedder, + store, + reranker=CrossEncoderReranker(model_name=MODEL_NAME), + candidate_k=CANDIDATE_K, + ), + cases, + k=K, + ) + + recall_delta = reranked.recall_at_k - baseline.recall_at_k + mrr_delta = reranked.mrr - baseline.mrr + rows = "\n".join( + ( + f"| {base['question']} | {base['reciprocal_rank']:.3f} | " + f"{rerank['reciprocal_rank']:.3f} | " + f"{rerank['reciprocal_rank'] - base['reciprocal_rank']:+.3f} |" + ) + for base, rerank in zip(baseline.per_case, reranked.per_case) + ) + + with open(REPORT_PATH, "w", encoding="utf-8") as fh: + fh.write( + "# RAG Assistant - Re-ranker Evaluation\n\n" + f"_Generated: {datetime.now(timezone.utc).replace(tzinfo=None).isoformat(timespec='seconds')}Z_\n\n" + f"- **Embedder**: `{type(embedder).__name__}`\n" + f"- **Re-ranker**: `{MODEL_NAME}`\n" + f"- **Cases**: {len(cases)}\n" + f"- **k**: {K}\n" + f"- **Candidate pool**: {CANDIDATE_K}\n\n" + "## Aggregate comparison\n\n" + "| Configuration | recall@k | MRR |\n" + "|---|---:|---:|\n" + f"| Embedding only | {baseline.recall_at_k:.3f} | {baseline.mrr:.3f} |\n" + f"| Cross-encoder re-ranked | {reranked.recall_at_k:.3f} | {reranked.mrr:.3f} |\n" + f"| Delta | {recall_delta:+.3f} | {mrr_delta:+.3f} |\n\n" + "## Per-case reciprocal-rank comparison\n\n" + "| Question | Baseline RR | Re-ranked RR | Delta |\n" + "|---|---:|---:|---:|\n" + f"{rows}\n\n" + "## Interpretation boundary\n\n" + "This is a five-case smoke evaluation over three small documents. " + "It records the observed delta but is not large enough to establish " + "a reliable quality lift. Re-run against the expanded benchmark from " + "issue #18 before making a comparative performance claim.\n" + ) + + print( + f"Wrote {REPORT_PATH}: " + f"recall delta={recall_delta:+.3f}, MRR delta={mrr_delta:+.3f}" + ) + + +if __name__ == "__main__": + main() diff --git a/ai_engineering/rag_assistant/tests/test_reranker.py b/ai_engineering/rag_assistant/tests/test_reranker.py new file mode 100644 index 00000000..39b2c9e7 --- /dev/null +++ b/ai_engineering/rag_assistant/tests/test_reranker.py @@ -0,0 +1,173 @@ +"""Network-free tests for optional cross-encoder re-ranking.""" + +import unittest + +from rag.chunker import Chunk, Chunker, Document +from rag.embedder import HashEmbedder +from rag.pipeline import RAGPipeline +from rag.reranker import CrossEncoderReranker +from rag.retriever import RetrievedChunk, Retriever +from rag.vector_store import VectorStore + + +def make_chunk(index: int, source: str, text: str) -> Chunk: + return Chunk( + chunk_id=f"doc#{index}", + doc_id="doc", + source=source, + chunk_index=index, + text=text, + ) + + +class FakeCrossEncoder: + def __init__(self, scores): + self.scores = scores + self.calls = [] + + def predict(self, pairs, *, batch_size, show_progress_bar): + self.calls.append( + { + "pairs": list(pairs), + "batch_size": batch_size, + "show_progress_bar": show_progress_bar, + } + ) + return self.scores[: len(pairs)] + + +class FakeEmbedder: + dim = 2 + + def embed(self, texts): + return [[1.0, 0.0] for _ in texts] + + +class FakeStore: + dim = 2 + + def __init__(self, chunks): + self.chunks = chunks + self.requested_k = None + + def search(self, query_vector, k): + self.requested_k = k + scores = [0.9, 0.8, 0.7, 0.6][:k] + return scores, self.chunks[:k] + + +class TestCrossEncoderReranker(unittest.TestCase): + def test_reorders_candidates_and_preserves_embedding_score(self): + candidates = [ + RetrievedChunk(0.9, make_chunk(0, "a.md", "alpha")), + RetrievedChunk(0.8, make_chunk(1, "b.md", "beta")), + RetrievedChunk(0.7, make_chunk(2, "c.md", "gamma")), + ] + model = FakeCrossEncoder([0.1, 0.95, 0.4]) + reranker = CrossEncoderReranker(model=model, batch_size=8) + + result = reranker.rerank("query", candidates, k=2) + + self.assertEqual([item.chunk.source for item in result], ["b.md", "c.md"]) + self.assertEqual([item.score for item in result], [0.95, 0.4]) + self.assertEqual([item.retrieval_score for item in result], [0.8, 0.7]) + self.assertEqual(model.calls[0]["batch_size"], 8) + self.assertFalse(model.calls[0]["show_progress_bar"]) + + def test_ties_preserve_original_candidate_order(self): + candidates = [ + RetrievedChunk(0.9, make_chunk(0, "a.md", "alpha")), + RetrievedChunk(0.8, make_chunk(1, "b.md", "beta")), + ] + reranker = CrossEncoderReranker( + model=FakeCrossEncoder([0.5, 0.5]) + ) + + result = reranker.rerank("query", candidates, k=2) + + self.assertEqual([item.chunk.source for item in result], ["a.md", "b.md"]) + + def test_empty_inputs_do_not_call_model(self): + model = FakeCrossEncoder([]) + reranker = CrossEncoderReranker(model=model) + + self.assertEqual(reranker.rerank("", [], k=3), []) + self.assertEqual(model.calls, []) + + +class TestRetrieverWithReranker(unittest.TestCase): + def setUp(self): + self.chunks = [ + make_chunk(0, "a.md", "alpha"), + make_chunk(1, "b.md", "beta"), + make_chunk(2, "c.md", "gamma"), + make_chunk(3, "d.md", "delta"), + ] + + def test_baseline_requests_only_k(self): + store = FakeStore(self.chunks) + retriever = Retriever(FakeEmbedder(), store) + + result = retriever.retrieve("query", k=2) + + self.assertEqual(store.requested_k, 2) + self.assertEqual([item.chunk.source for item in result], ["a.md", "b.md"]) + + def test_reranker_expands_candidate_pool(self): + store = FakeStore(self.chunks) + reranker = CrossEncoderReranker( + model=FakeCrossEncoder([0.1, 0.2, 0.95, 0.3]) + ) + retriever = Retriever( + FakeEmbedder(), + store, + reranker=reranker, + candidate_k=4, + ) + + result = retriever.retrieve("query", k=2) + + self.assertEqual(store.requested_k, 4) + self.assertEqual([item.chunk.source for item in result], ["c.md", "d.md"]) + + def test_non_positive_k_returns_empty(self): + store = FakeStore(self.chunks) + retriever = Retriever(FakeEmbedder(), store) + + self.assertEqual(retriever.retrieve("query", k=0), []) + self.assertIsNone(store.requested_k) + + +class TestPipelineRerankerIntegration(unittest.TestCase): + def test_pipeline_uses_injected_reranker_without_network(self): + embedder = HashEmbedder(dim=64) + store = VectorStore(dim=embedder.dim) + reranker = CrossEncoderReranker( + model=FakeCrossEncoder([0.1, 0.9, 0.4]) + ) + pipeline = RAGPipeline( + embedder=embedder, + store=store, + generator=None, + chunker=Chunker(chunk_size=80, chunk_overlap=10), + reranker=reranker, + candidate_k=3, + ) + pipeline.ingest( + [ + Document(source="a.md", text="alpha topic"), + Document(source="b.md", text="beta topic"), + Document(source="c.md", text="gamma topic"), + ] + ) + + result = pipeline.ask("topic", k=2) + + self.assertEqual(len(result.retrieved), 2) + self.assertTrue( + all(item.retrieval_score is not None for item in result.retrieved) + ) + + +if __name__ == "__main__": + unittest.main()