diff --git a/application/api/user/routes.py b/application/api/user/routes.py index 9464bef9..e946075e 100644 --- a/application/api/user/routes.py +++ b/application/api/user/routes.py @@ -21,7 +21,12 @@ from .models import models_ns from .prompts import prompts_ns from .schedules import schedules_ns from .sharing import sharing_ns -from .sources import sources_chunks_ns, sources_ns, sources_upload_ns +from .sources import ( + sources_chunks_ns, + sources_ns, + sources_search_ns, + sources_upload_ns, +) from .teams import teams_ns from .tools import tools_mcp_ns, tools_ns from .workflows import workflows_ns @@ -63,9 +68,10 @@ api.add_namespace(schedules_ns) # Sharing api.add_namespace(sharing_ns) -# Sources (main, chunks, upload) +# Sources (main, chunks, retrieval test, upload) api.add_namespace(sources_ns) api.add_namespace(sources_chunks_ns) +api.add_namespace(sources_search_ns) api.add_namespace(sources_upload_ns) # Teams (CRUD, membership, resource-sharing grants) diff --git a/application/api/user/sources/__init__.py b/application/api/user/sources/__init__.py index 07b380b7..d45fc136 100644 --- a/application/api/user/sources/__init__.py +++ b/application/api/user/sources/__init__.py @@ -1,7 +1,13 @@ """Sources module.""" from .chunks import sources_chunks_ns +from .retrieval_test import sources_search_ns from .routes import sources_ns from .upload import sources_upload_ns -__all__ = ["sources_ns", "sources_chunks_ns", "sources_upload_ns"] +__all__ = [ + "sources_ns", + "sources_chunks_ns", + "sources_search_ns", + "sources_upload_ns", +] diff --git a/application/api/user/sources/retrieval_test.py b/application/api/user/sources/retrieval_test.py new file mode 100644 index 00000000..de052adb --- /dev/null +++ b/application/api/user/sources/retrieval_test.py @@ -0,0 +1,254 @@ +"""Retrieval test — preview the chunks a query actually retrieves from a source. + +The chunk browser's ``search`` (``/api/get_chunks``) is a substring filter over +stored text; it says nothing about what a RAG query would retrieve. This +endpoint runs the *production* retrieval path against a single source so a user +can see the real ranked chunks for a query, and try retrieval settings without +saving them. + +It deliberately builds a ``Dispatcher`` — the same object the answer pipeline +uses — so retriever selection, per-source ``chunks`` / ``score_threshold``, the +prescreen stage and the token budget all behave exactly as they do at answer +time, rather than drifting from them. +""" + +import logging +import math +import time + +from flask import jsonify, make_response, request +from flask_restx import fields, Namespace, Resource +from pydantic import ValidationError + +from application.api import api +from application.api.user.sources.routes import _resolve_readable_source +from application.core.model_utils import get_default_model_id +from application.retriever.dispatcher import Dispatcher +from application.retriever.retriever_creator import RetrieverCreator +from application.storage.db.session import db_readonly +from application.storage.db.source_config import RetrievalConfig, SourceConfig +from application.utils import num_tokens_from_string + +logger = logging.getLogger(__name__) + +sources_search_ns = Namespace( + "sources", description="Source retrieval testing", path="/api" +) + +# A retrieval query is a search query, not a document. +MAX_QUERY_LENGTH = 2000 + +# The answer pipeline's default; keeps the preview's budgeting identical to it. +DOC_TOKEN_LIMIT = 50000 + +# Cost-attribution tag for any LLM call the preview makes (prescreen only — +# with no chat history the rephrase side-call is skipped entirely). +USAGE_SOURCE = "retrieval_test" + +# Ceiling on the prescreen LLM calls a single test may trigger. +MAX_PRESCREEN_BATCHES = 20 + + +@sources_search_ns.route("/sources//search") +class SourceSearch(Resource): + search_model = api.model( + "SourceSearchModel", + { + "query": fields.String( + required=True, description="The query to retrieve for" + ), + "retrieval": fields.Raw( + required=False, + description="Ad-hoc retrieval config to test with. Omit to use " + "the source's saved config. Never persisted.", + ), + }, + ) + + @api.expect(search_model) + @api.doc( + description="Run the real retrieval pipeline against one source and " + "return the ranked chunks it produces, optionally under an ad-hoc " + "retrieval config. Read-only: nothing is saved." + ) + def post(self, source_id): + decoded_token = request.decoded_token + if not decoded_token: + return make_response(jsonify({"success": False}), 401) + user = decoded_token.get("sub") + + body = request.get_json(silent=True) + if not isinstance(body, dict): + return make_response( + jsonify({"success": False, "message": "Invalid request body"}), 400 + ) + query = (body.get("query") or "").strip() + if not query: + return make_response( + jsonify({"success": False, "message": "Query is required"}), 400 + ) + if len(query) > MAX_QUERY_LENGTH: + return make_response( + jsonify( + { + "success": False, + "message": f"Query must be at most {MAX_QUERY_LENGTH} characters", + } + ), + 400, + ) + + try: + # Read access = owner or any team grant (viewer included), matching + # the other source read endpoints (wiki pages, graph). + with db_readonly() as conn: + doc = _resolve_readable_source(conn, source_id, user) + except Exception as e: + # An unresolvable id yields None (→ 404); reaching here means the + # lookup itself failed, which is ours, not the caller's. + logger.error(f"Error resolving source: {e}", exc_info=True) + return make_response( + jsonify({"success": False, "message": "Could not resolve source"}), 500 + ) + if not doc: + return make_response( + jsonify( + {"success": False, "message": "Source not found or access denied"} + ), + 404, + ) + resolved_id = str(doc["id"]) + + # A supplied config is validated exactly as strictly as a saved one (D7 + # strict-on-write), with the same static message so validation internals + # aren't echoed back. Absent → the source's saved config. + saved = SourceConfig.parse(doc.get("config")).retrieval + if body.get("retrieval") is not None: + try: + retrieval = RetrievalConfig.model_validate(body["retrieval"]) + except ValidationError: + return make_response( + jsonify( + { + "success": False, + "message": "Invalid retrieval config: one or more " + "fields failed validation.", + } + ), + 400, + ) + else: + retrieval = saved + + # "Ad-hoc" means the caller is testing something the source is NOT + # already configured to do. The client always sends the config it has on + # screen — including an untouched one — so compare by value rather than + # trusting the body's presence, or a source whose saved config is + # expensive could never be tested at all (see the ceiling below). + ad_hoc = retrieval != saved + + # RetrievalConfig.retriever is a free string, so an unknown key would + # only blow up inside RetrieverCreator — a caller's typo must be a 400, + # not a 500. + if retrieval.retriever not in RetrieverCreator.retrievers: + return make_response( + jsonify( + { + "success": False, + "message": f"Unknown retriever '{retrieval.retriever}'.", + } + ), + 400, + ) + + # Prescreen screens candidate_k candidates in batches of batch_size, one + # synchronous LLM call each, so candidate_k=500 / batch_size=1 is 500 + # provider calls in one request; RetrievalConfig bounds each field but + # not their ratio. Only an ad-hoc config is capped: a config already + # saved on the source runs at answer time anyway, so refusing to test it + # would defeat the endpoint. + if ad_hoc: + ps = retrieval.prescreen_config() + if ps is not None: + batches = math.ceil(ps.candidate_k / ps.batch_size) + if batches > MAX_PRESCREEN_BATCHES: + return make_response( + jsonify( + { + "success": False, + "message": ( + f"This prescreen config needs {batches} LLM " + f"calls to test (limit " + f"{MAX_PRESCREEN_BATCHES}). Raise batch_size " + "or lower candidate_k." + ), + } + ), + 400, + ) + + # A test has no chat behind it, so no model was requested. Resolve the + # same instance default the answer path falls back to: with a bogus id + # the prescreen stage would call the provider, fail, and silently keep + # every candidate. On an instance with no models configured this is + # None — leave the retriever's own default in place rather than + # forwarding it. + dispatcher_kwargs = {} + default_model_id = get_default_model_id() + if default_model_id: + dispatcher_kwargs["model_id"] = default_model_id + + try: + started = time.monotonic() + # Dispatcher directly, NOT build_dispatcher: the + # PER_SOURCE_RETRIEVAL_ENABLED kill-switch would fall back to a + # stock classic retriever and ignore the config under test. + retriever = Dispatcher( + source={"active_docs": [resolved_id], "question": query}, + chat_history=[], # no history ⇒ no rephrase side-call + chunks=retrieval.chunks, + doc_token_limit=DOC_TOKEN_LIMIT, + decoded_token=decoded_token, + sources=[{"id": resolved_id, "retrieval": retrieval}], + include_scores=True, + usage_source=USAGE_SOURCE, + **dispatcher_kwargs, + ) + docs = retriever.search(query) or [] + latency_ms = int((time.monotonic() - started) * 1000) + except Exception as e: + logger.error(f"Retrieval test failed: {e}", exc_info=True) + return make_response( + jsonify({"success": False, "message": "Retrieval failed"}), 500 + ) + + chunks = [ + { + "rank": idx, + "text": d.get("text", ""), + "title": d.get("title"), + "filename": d.get("filename"), + "source": d.get("source"), + "tokens": num_tokens_from_string(d.get("text", "")), + # None for retrievers/stores that produce no comparable score + # (graphrag's PPR ranking, stores without a score seam). + "score": d.get("score"), + "score_kind": d.get("score_kind"), + } + for idx, d in enumerate(docs, start=1) + ] + + return make_response( + jsonify( + { + "success": True, + "query": query, + "retriever": retrieval.retriever, + "retrieval": retrieval.model_dump(), + "total": len(chunks), + "latency_ms": latency_ms, + "chunks": chunks, + } + ), + 200, + ) diff --git a/application/retriever/classic_rag.py b/application/retriever/classic_rag.py index f8cdc560..dbf395c6 100644 --- a/application/retriever/classic_rag.py +++ b/application/retriever/classic_rag.py @@ -9,6 +9,10 @@ from application.vectorstore.vector_creator import VectorCreator class ClassicRAG(BaseRetriever): + # The group's real top-k, set by the Dispatcher when it inflates ``chunks`` + # for a prescreen fetch. None → ``chunks`` is already the top-k. + base_chunks = None + def __init__( self, source, @@ -25,7 +29,9 @@ class ClassicRAG(BaseRetriever): model_user_id=None, defer_rephrase=False, request_id=None, + include_scores=False, ): + self.include_scores = include_scores self.original_question = source.get("question", "") self.chat_history = chat_history if chat_history is not None else [] self.prompt = prompt @@ -147,8 +153,10 @@ class ClassicRAG(BaseRetriever): def _fetch_candidates(self, docsearch, question, src_k, score_threshold): """Fetch candidate hits for one vector store (vector search). - Subclasses override this to change candidate sourcing (e.g. RRF fusion) - while inheriting the surrounding per-source resolution and budgeting. + Returns plain hits, or ``(hit, score)`` pairs when ``include_scores`` is + set. Subclasses override this to change candidate sourcing (e.g. RRF + fusion) while inheriting the surrounding per-source resolution and + budgeting. """ # ``score_threshold`` is honoured by pgvector/mongodb and safely ignored # by stores whose ``search`` swallows kwargs. The candidate count is @@ -157,8 +165,14 @@ class ClassicRAG(BaseRetriever): search_kwargs = {"k": k} if score_threshold is not None: search_kwargs["score_threshold"] = score_threshold + if self.include_scores: + return docsearch.search_with_scores(question, **search_kwargs) return docsearch.search(question, **search_kwargs) + def _score_kind(self, docsearch): + """Label for the scores ``_fetch_candidates`` attaches (None if unscored).""" + return getattr(docsearch, "score_kind", None) + def _get_data(self): if self.chunks == 0 or not self.vectorstores: logging.info( @@ -168,7 +182,13 @@ class ClassicRAG(BaseRetriever): return [] all_docs = [] - chunks_per_source = max(1, self.chunks // len(self.vectorstores)) + # The Dispatcher inflates ``chunks`` to a prescreen source's candidate_k + # so the fetch is large enough for the screening stage. That inflated + # number must not become the top-k of the *other* sources in the group, + # so the fallback splits the group's real top-k (``base_chunks``) when + # the Dispatcher supplied one. + base_chunks = self.base_chunks if self.base_chunks is not None else self.chunks + chunks_per_source = max(1, base_chunks // len(self.vectorstores)) token_budget = max(int(self.doc_token_limit * 0.9), 100) cumulative_tokens = 0 @@ -211,11 +231,25 @@ class ClassicRAG(BaseRetriever): docs_temp = self._fetch_candidates( docsearch, question, src_k, score_threshold ) + score_kind = ( + self._score_kind(docsearch) if self.include_scores else None + ) + + # ``_fetch_candidates`` over-fetches (k >= 20) so a prescreen + # stage has candidates to filter; trim back to src_k so + # ``chunks`` is the final top-k it claims to be. With + # prescreen on, src_k is already raised to candidate_k above, + # so the stage still sees its full candidate set. + kept = 0 for doc in docs_temp: - if cumulative_tokens >= token_budget: + if kept >= src_k or cumulative_tokens >= token_budget: break + score = None + if isinstance(doc, tuple): + doc, score = doc + if hasattr(doc, "page_content") and hasattr(doc, "metadata"): page_content = doc.page_content metadata = doc.metadata @@ -231,8 +265,13 @@ class ClassicRAG(BaseRetriever): doc_tokens = num_tokens_from_string(doc_text_with_header) if cumulative_tokens + doc_tokens < token_budget: - all_docs.append({"text": page_content, **labels}) + entry = {"text": page_content, **labels} + if self.include_scores: + entry["score"] = score + entry["score_kind"] = score_kind + all_docs.append(entry) cumulative_tokens += doc_tokens + kept += 1 if cumulative_tokens >= token_budget: break diff --git a/application/retriever/dispatcher.py b/application/retriever/dispatcher.py index 38ebd325..21d419a5 100644 --- a/application/retriever/dispatcher.py +++ b/application/retriever/dispatcher.py @@ -68,6 +68,8 @@ class Dispatcher(BaseRetriever): request_id=None, sources: Optional[List[Dict[str, Any]]] = None, stages: Optional[List[Stage]] = None, + include_scores: bool = False, + usage_source: str = "rag_prescreen", ): """Build the dispatcher. @@ -83,8 +85,15 @@ class Dispatcher(BaseRetriever): falls back to the single classic group over ``source``. stages: Optional post-retrieval stages applied to each group's candidates before final budgeting. Default: none (pass-through). + include_scores: Ask each group's retriever to attach the raw store + score to every doc it returns. Off for the answer pipeline; on + for the retrieval-test endpoint. + usage_source: Cost-attribution tag for the prescreen stage's LLM + calls. """ + self._usage_source = usage_source self._ctor_kwargs = dict( + include_scores=include_scores, chat_history=chat_history, prompt=prompt, doc_token_limit=doc_token_limit, @@ -229,9 +238,12 @@ class Dispatcher(BaseRetriever): kwargs["defer_rephrase"] = True retriever = RetrieverCreator.create_retriever(retriever_key, **kwargs) # Hand the per-source retrieval configs to the classic retriever so it - # can honour per-source chunks/score_threshold/rephrase in its loop. + # can honour per-source chunks/score_threshold/rephrase in its loop, plus + # the un-inflated top-k so a prescreen source's candidate_k doesn't + # become the top-k of the sources beside it in the group. if group["retrievals"]: setattr(retriever, "per_source_retrieval", group["retrievals"]) + setattr(retriever, "base_chunks", self.chunks) return retriever def _group_stages(self, group: Dict[str, Any]) -> List[Stage]: @@ -251,6 +263,7 @@ class Dispatcher(BaseRetriever): agent_id=self._ctor_kwargs.get("agent_id"), model_user_id=self._ctor_kwargs.get("model_user_id"), request_id=self._ctor_kwargs.get("request_id"), + usage_source=self._usage_source, ) return list(self.stages) + prescreen diff --git a/application/retriever/graph_rag.py b/application/retriever/graph_rag.py index 9e41a678..240d0f19 100644 --- a/application/retriever/graph_rag.py +++ b/application/retriever/graph_rag.py @@ -39,6 +39,10 @@ def _idf(doc_freq: Any) -> float: class GraphRAGRetriever(BaseRetriever): """Per-source PPR retriever; falls back to ClassicRAG when a source has no graph.""" + # Set by the Dispatcher (see ClassicRAG.base_chunks); forwarded to the inner + # ClassicRAG on the fallback path. + base_chunks = None + def __init__( self, source, @@ -55,7 +59,11 @@ class GraphRAGRetriever(BaseRetriever): model_user_id=None, defer_rephrase=False, request_id=None, + include_scores=False, ): + # Graph docs are ranked by PPR, which yields no per-chunk similarity, so + # they stay unscored; the flag only matters to the classic fallback, + # which retrieves the sources that have no graph. self._classic = ClassicRAG( source=source, chat_history=chat_history, @@ -71,6 +79,7 @@ class GraphRAGRetriever(BaseRetriever): model_user_id=model_user_id, defer_rephrase=defer_rephrase, request_id=request_id, + include_scores=include_scores, ) self.original_question = self._classic.original_question self.chunks = self._classic.chunks @@ -129,6 +138,25 @@ class GraphRAGRetriever(BaseRetriever): candidates = max(self.chunks * 2, self.chunks + 5) return ranked[: max(1, candidates)] + def _source_top_k(self, source_id) -> int: + """How many chunks this source may contribute — its own top-k. + + Mirrors ClassicRAG's resolution: a per-source override wins (raised to + candidate_k when it prescreens, so the stage has candidates to filter); + otherwise the group's real top-k is split across the sources. Without + this, a prescreen source elsewhere in the group inflates ``chunks`` and + this source would return that inflated count. + """ + cfg = self.per_source_retrieval.get(source_id) + if cfg is not None: + top_k = max(1, int(cfg.chunks)) + ps = cfg.prescreen_config() if hasattr(cfg, "prescreen_config") else None + if ps is not None: + top_k = max(top_k, int(ps.candidate_k)) + return top_k + base = self.base_chunks if self.base_chunks is not None else self.chunks + return max(1, base // max(1, len(self.vectorstores))) + def _graph_docs_for_source(self, store, source_id) -> List[Dict[str, Any]]: """Local PPR retrieval for one source (caller guarantees it has a graph).""" question = self._classic._get_rephrased_question() @@ -160,8 +188,9 @@ class GraphRAGRetriever(BaseRetriever): docs: List[Dict[str, Any]] = [] token_budget = max(int(self.doc_token_limit * 0.9), 100) cumulative_tokens = 0 + source_top_k = self._source_top_k(source_id) for chunk_id in chunk_ids: - if len(docs) >= self.chunks: + if len(docs) >= source_top_k: break chunk = chunk_data.get(chunk_id) text = chunk.get("text") if chunk else None @@ -179,15 +208,20 @@ class GraphRAGRetriever(BaseRetriever): """Reuse the composed ClassicRAG to retrieve one source's chunks.""" original = self._classic.vectorstores original_overrides = self._classic.per_source_retrieval + original_base = self._classic.base_chunks try: self._classic.vectorstores = [source_id] self._classic.per_source_retrieval = { k: v for k, v in self.per_source_retrieval.items() if k == source_id } + # The Dispatcher sets these on *this* object; the inner retriever is + # the one that reads them. + self._classic.base_chunks = self.base_chunks return self._classic._get_data() finally: self._classic.vectorstores = original self._classic.per_source_retrieval = original_overrides + self._classic.base_chunks = original_base def _get_data(self) -> List[Dict[str, Any]]: if not self.vectorstores: diff --git a/application/retriever/hybrid_rag.py b/application/retriever/hybrid_rag.py index b789ddc5..cf739253 100644 --- a/application/retriever/hybrid_rag.py +++ b/application/retriever/hybrid_rag.py @@ -25,13 +25,14 @@ def _doc_key(doc): return (source, content) -def reciprocal_rank_fusion(vector_hits, keyword_hits, k=RRF_K): - """Fuse two ranked hit lists into one by Reciprocal Rank Fusion. +def fuse_with_scores(vector_hits, keyword_hits, k=RRF_K): + """Fuse two ranked hit lists by RRF, keeping each hit's fused score. Each list contributes ``1 / (k + rank)`` per document (rank 0-based); - documents are returned ordered by summed score, highest first. A document - present in only one list is ranked solely on that list's contribution, so - an empty ``keyword_hits`` yields exactly the vector ordering. + documents are returned as ``(doc, fused_score)`` ordered by score, highest + first. A document present in only one list is ranked solely on that list's + contribution, so an empty ``keyword_hits`` yields exactly the vector + ordering. """ scores = {} docs = {} @@ -42,12 +43,18 @@ def reciprocal_rank_fusion(vector_hits, keyword_hits, k=RRF_K): if key not in docs: docs[key] = doc ordered = sorted(docs.keys(), key=lambda key: scores[key], reverse=True) - return [docs[key] for key in ordered] + return [(docs[key], scores[key]) for key in ordered] class HybridRetriever(ClassicRAG): """ClassicRAG variant that fuses vector + keyword search with RRF.""" + def _score_kind(self, docsearch): + """RRF scores rank hits against each other, not against a similarity + cutoff — they are not comparable to the store's cosine scores, so they + carry their own label.""" + return "rrf" + def _fetch_candidates(self, docsearch, question, src_k, score_threshold): """Return RRF-fused vector+keyword hits for one vector store. @@ -59,4 +66,7 @@ class HybridRetriever(ClassicRAG): candidate_k = min(max(src_k * 2, 20), 500) vector_hits = docsearch.search(question, k=candidate_k) keyword_hits = docsearch.keyword_search(question, k=candidate_k) - return reciprocal_rank_fusion(vector_hits, keyword_hits) + fused = fuse_with_scores(vector_hits, keyword_hits) + if self.include_scores: + return fused + return [doc for doc, _ in fused] diff --git a/application/retriever/stages/prescreen.py b/application/retriever/stages/prescreen.py index 4050784a..4c1e61bf 100644 --- a/application/retriever/stages/prescreen.py +++ b/application/retriever/stages/prescreen.py @@ -56,6 +56,7 @@ class PreScreenStage: agent_id: Optional[str] = None, model_user_id: Optional[str] = None, request_id: Optional[str] = None, + usage_source: str = "rag_prescreen", ): """Build the stage. @@ -68,6 +69,7 @@ class PreScreenStage: decoded_token: Caller identity for BYOM resolution. agent_id: Agent context for BYOM resolution. model_user_id: BYOM-resolution scope for shared-agent dispatch. + usage_source: Cost-attribution tag for the screening calls. """ self.config = config self.llm_name = llm_name @@ -78,6 +80,7 @@ class PreScreenStage: self.agent_id = agent_id self.model_user_id = model_user_id self.request_id = request_id + self.usage_source = usage_source def _resolve_model(self) -> Optional[str]: """Use the configured model, else fall back to the request model.""" @@ -96,7 +99,7 @@ class PreScreenStage: ) # Tag rows so the screening calls land as a distinct cost source, and # stamp the originating request so the rows correlate to it. - llm._token_usage_source = "rag_prescreen" + llm._token_usage_source = self.usage_source llm._request_id = self.request_id return llm @@ -188,6 +191,7 @@ def build_prescreen_stages( agent_id: Optional[str] = None, model_user_id: Optional[str] = None, request_id: Optional[str] = None, + usage_source: str = "rag_prescreen", ) -> List[Stage]: """Build prescreen stages from a group's per-source retrieval configs. @@ -222,6 +226,7 @@ def build_prescreen_stages( agent_id=agent_id, model_user_id=model_user_id, request_id=request_id, + usage_source=usage_source, ) ) return stages diff --git a/application/vectorstore/base.py b/application/vectorstore/base.py index 142f75f8..b0f73836 100644 --- a/application/vectorstore/base.py +++ b/application/vectorstore/base.py @@ -218,6 +218,26 @@ class BaseVectorStore(ABC): """ return [] + # What ``search_with_scores`` reports, so a caller can label the number. + # ``cosine_similarity`` is higher-is-better in [0, 1]; ``l2_distance`` is + # lower-is-better and unbounded. None = this store reports no score. + score_kind = None + + def search_with_scores(self, question, k=2, *args, **kwargs): + """Search, pairing each hit with its raw relevance score. + + Default pairs every hit from :meth:`search` with ``None`` so stores that + surface no score still satisfy the contract. Stores that already compute + one override this and set :attr:`score_kind`. + + Returns: + A list of ``(Document, score | None)`` in the same rank order + :meth:`search` would return. + """ + return [ + (doc, None) for doc in self.search(question, k, *args, **kwargs) or [] + ] + @abstractmethod def add_texts(self, texts, metadatas=None, *args, **kwargs): """Add texts with their embeddings to the vectorstore""" diff --git a/application/vectorstore/faiss.py b/application/vectorstore/faiss.py index c40a09c9..795669e3 100644 --- a/application/vectorstore/faiss.py +++ b/application/vectorstore/faiss.py @@ -83,12 +83,26 @@ class FaissStore(BaseVectorStore): self.assert_embedding_dimensions(self.embeddings) + # LangChain's FAISS wrapper ranks by L2 distance (lower is better), not by + # a cosine similarity — so the number here is NOT comparable to the + # ``score_threshold`` the other stores honour, and must not be shown as one. + score_kind = "l2_distance" + def search(self, *args, **kwargs): # FAISS has no relevance-threshold knob; drop it so the per-source # score_threshold is safely ignored rather than crashing the forward. kwargs.pop("score_threshold", None) return self.docsearch.similarity_search(*args, **kwargs) + def search_with_scores(self, *args, **kwargs): + """Same search as :meth:`search`, pairing each hit with its L2 distance.""" + kwargs.pop("score_threshold", None) + results = self.docsearch.similarity_search_with_score(*args, **kwargs) + # The Documents come straight from the live in-memory docstore — return + # them untouched (the score rides alongside, never in their metadata) so + # the index can't be polluted and later persisted back to storage. + return [(doc, float(score)) for doc, score in results] + def add_texts(self, *args, **kwargs): return self.docsearch.add_texts(*args, **kwargs) diff --git a/application/vectorstore/mongodb.py b/application/vectorstore/mongodb.py index 2faa13dd..64e8591a 100644 --- a/application/vectorstore/mongodb.py +++ b/application/vectorstore/mongodb.py @@ -52,6 +52,8 @@ class MongoDBVectorStore(BaseVectorStore): def _collection(self): return self._database[self._collection_name] + score_kind = "cosine_similarity" + def search(self, question, k=2, *args, score_threshold=None, **kwargs): """Search via Atlas ``$vectorSearch``. @@ -61,6 +63,19 @@ class MongoDBVectorStore(BaseVectorStore): score_threshold: Optional ``vectorSearchScore`` floor in ``[0, 1]``; results scoring below it are dropped. """ + return [ + doc + for doc, _ in self.search_with_scores( + question, k, *args, score_threshold=score_threshold, **kwargs + ) + ] + + def search_with_scores(self, question, k=2, *args, score_threshold=None, **kwargs): + """Same search as :meth:`search`, pairing each hit with its score. + + The score is Atlas' ``vectorSearchScore`` — the same quantity + ``score_threshold`` is compared against. + """ query_vector = self._embedding.embed_query(question) pipeline = [ @@ -73,10 +88,10 @@ class MongoDBVectorStore(BaseVectorStore): "index": self._index_name, "filter": {"source_id": {"$eq": self._source_id}}, } - } + }, + {"$addFields": {"_score": {"$meta": "vectorSearchScore"}}}, ] if score_threshold is not None: - pipeline.append({"$addFields": {"_score": {"$meta": "vectorSearchScore"}}}) pipeline.append({"$match": {"_score": {"$gte": float(score_threshold)}}}) cursor = self._collection.aggregate(pipeline) @@ -87,9 +102,11 @@ class MongoDBVectorStore(BaseVectorStore): doc.pop("_id") doc.pop(self._text_key) doc.pop(self._embedding_key) - doc.pop("_score", None) + score = doc.pop("_score", None) metadata = doc - results.append(Document(text, metadata)) + results.append( + (Document(text, metadata), None if score is None else float(score)) + ) return results def _insert_texts(self, texts, metadatas): diff --git a/application/vectorstore/pgvector.py b/application/vectorstore/pgvector.py index 5fd8b117..c4b445d1 100644 --- a/application/vectorstore/pgvector.py +++ b/application/vectorstore/pgvector.py @@ -122,6 +122,8 @@ class PGVectorStore(BaseVectorStore): finally: cursor.close() + score_kind = "cosine_similarity" + def search( self, question: str, @@ -139,6 +141,27 @@ class PGVectorStore(BaseVectorStore): Cosine distance = ``1 - similarity``; rows with similarity below the threshold (distance above ``1 - threshold``) are dropped. """ + return [ + doc + for doc, _ in self.search_with_scores( + question, k, *args, score_threshold=score_threshold, **kwargs + ) + ] + + def search_with_scores( + self, + question: str, + k: int = 2, + *args, + score_threshold: float = None, + **kwargs, + ) -> List[tuple]: + """Same search as :meth:`search`, pairing each hit with its similarity. + + The score is the cosine similarity (``1 - cosine_distance``) — the exact + quantity ``score_threshold`` is compared against, so a caller can read a + result's score and pick a threshold from it directly. + """ query_vector = self._embedding.embed_query(question) conn = self._get_connection() @@ -167,10 +190,13 @@ class PGVectorStore(BaseVectorStore): if max_distance is not None and distance is not None and distance > max_distance: continue metadata = metadata or {} - documents.append(Document(page_content=text, metadata=metadata)) + score = None if distance is None else 1.0 - float(distance) + documents.append( + (Document(page_content=text, metadata=metadata), score) + ) return documents - + except Exception as e: logging.error(f"Error searching documents: {e}", exc_info=True) return [] diff --git a/frontend/src/api/endpoints.ts b/frontend/src/api/endpoints.ts index 8f23f896..edc89b97 100644 --- a/frontend/src/api/endpoints.ts +++ b/frontend/src/api/endpoints.ts @@ -56,6 +56,7 @@ const endpoints = { SYNC_SOURCE: '/api/sync_source', REINGEST_SOURCE: '/api/sources/reingest', SOURCE_CONFIG: (id: string) => `/api/sources/${id}/config`, + SOURCE_SEARCH: (id: string) => `/api/sources/${id}/search`, CREATE_WIKI: '/api/sources/wiki', CONVERT_TO_WIKI: (id: string) => `/api/sources/${id}/wiki/convert`, ENABLE_GRAPHRAG: (id: string) => `/api/sources/${id}/graphrag/enable`, diff --git a/frontend/src/api/services/userService.ts b/frontend/src/api/services/userService.ts index bd65491c..563743b5 100644 --- a/frontend/src/api/services/userService.ts +++ b/frontend/src/api/services/userService.ts @@ -102,6 +102,12 @@ const userService = { token: string | null, ): Promise => apiClient.patch(endpoints.USER.SOURCE_CONFIG(sourceId), config, token), + testSourceRetrieval: ( + sourceId: string, + data: { query: string; retrieval?: any }, + token: string | null, + ): Promise => + apiClient.post(endpoints.USER.SOURCE_SEARCH(sourceId), data, token), createWiki: ( data: { name: string; initial_content?: string }, token: string | null, diff --git a/frontend/src/components/Chunks.tsx b/frontend/src/components/Chunks.tsx index 7b0dfad3..e01ead13 100644 --- a/frontend/src/components/Chunks.tsx +++ b/frontend/src/components/Chunks.tsx @@ -108,6 +108,8 @@ interface ChunksProps { displayPath?: string; onFileSearch?: (query: string) => SearchResult[]; onFileSelect?: (path: string) => void; + /** Extra header control, rendered left of the chunk actions. */ + headerAction?: React.ReactNode; } const Chunks: React.FC = ({ @@ -118,6 +120,7 @@ const Chunks: React.FC = ({ displayPath, onFileSearch, onFileSelect, + headerAction, }) => { const [fileSearchQuery, setFileSearchQuery] = useState(''); const [fileSearchResults, setFileSearchResults] = useState( @@ -349,6 +352,7 @@ const Chunks: React.FC = ({
+ {headerAction} {editingChunk ? ( !isEditing ? ( <> diff --git a/frontend/src/components/ConnectorTree.tsx b/frontend/src/components/ConnectorTree.tsx index db8ff6db..99bb30ca 100644 --- a/frontend/src/components/ConnectorTree.tsx +++ b/frontend/src/components/ConnectorTree.tsx @@ -16,12 +16,15 @@ interface ConnectorTreeProps { docId: string; sourceName: string; onBackToDocuments: () => void; + /** Extra header control, rendered left of the Sync button. */ + headerAction?: React.ReactNode; } const ConnectorTree: React.FC = ({ docId, sourceName, onBackToDocuments, + headerAction, }) => { const { t } = useTranslation(); const token = useSelector(selectToken); @@ -101,33 +104,36 @@ const ConnectorTree: React.FC = ({ }; const topRightAction = ( - + : t('settings.sources.sync')} + + ); const extraContent = ( diff --git a/frontend/src/components/FileTree.tsx b/frontend/src/components/FileTree.tsx index e9e6fbf1..34586151 100644 --- a/frontend/src/components/FileTree.tsx +++ b/frontend/src/components/FileTree.tsx @@ -26,12 +26,15 @@ interface FileTreeProps { docId: string; sourceName: string; onBackToDocuments: () => void; + /** Extra header control, rendered left of "Add file". */ + headerAction?: React.ReactNode; } const FileTree: React.FC = ({ docId, sourceName, onBackToDocuments, + headerAction, }) => { const { t } = useTranslation(); const token = useSelector(selectToken); @@ -231,15 +234,22 @@ const FileTree: React.FC = ({ : t('settings.sources.deletingTitle') : null; - const topRightAction = !isProcessing ? ( - - ) : null; + // headerAction stays visible while an upload/delete is in flight — only the + // Add file button is suppressed then. + const topRightAction = ( + <> + {headerAction} + {!isProcessing ? ( + + ) : null} + + ); const extraContent = ( void; + /** Extra header control, right-aligned in the title row. */ + headerAction?: React.ReactNode; } const GRAPH_LIMIT = 100; @@ -35,6 +37,7 @@ const GraphView: React.FC = ({ docId, sourceName, onBackToDocuments, + headerAction, }) => { const { t } = useTranslation(); const token = useSelector(selectToken); @@ -169,6 +172,7 @@ const GraphView: React.FC = ({ {sourceName} + {headerAction ?
{headerAction}
: null}
diff --git a/frontend/src/components/WikiViewer.tsx b/frontend/src/components/WikiViewer.tsx index e5e7ca6d..6ef4181d 100644 --- a/frontend/src/components/WikiViewer.tsx +++ b/frontend/src/components/WikiViewer.tsx @@ -23,6 +23,8 @@ interface WikiViewerProps { sourceName: string; canEdit?: boolean; onBackToDocuments: () => void; + /** Extra header control, right-aligned in the title row. */ + headerAction?: React.ReactNode; } const markdownComponents = { @@ -66,6 +68,7 @@ const WikiViewer: React.FC = ({ sourceName, canEdit = false, onBackToDocuments, + headerAction, }) => { const { t } = useTranslation(); const token = useSelector(selectToken); @@ -220,6 +223,7 @@ const WikiViewer: React.FC = ({ {sourceName} + {headerAction ?
{headerAction}
: null}
diff --git a/frontend/src/components/tree/TreeBrowser.tsx b/frontend/src/components/tree/TreeBrowser.tsx index 4bb87ef5..64ce9e3f 100644 --- a/frontend/src/components/tree/TreeBrowser.tsx +++ b/frontend/src/components/tree/TreeBrowser.tsx @@ -633,7 +633,7 @@ const TreeBrowser: React.FC = ({ /> {searchQuery && ( -
+
{searchResults.length === 0 ? (
@@ -744,7 +744,7 @@ const TreeBrowser: React.FC = ({
) : ( -
+
{renderPathNavigation()}
diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index d474b801..3150465a 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -189,6 +189,31 @@ "close": "Close" } }, + "testRetrieval": { + "action": "Test retrieval", + "title": "Test retrieval", + "subtitle": "See the chunks a query actually retrieves from \"{{name}}\", using the real retrieval pipeline.", + "subtitleGeneric": "See the chunks a query actually retrieves, using the real retrieval pipeline.", + "queryPlaceholder": "Ask something this source should answer...", + "run": "Run", + "notSavedHint": "These settings are only used for this test — they are not saved to the source. Change them in Source settings to make them permanent.", + "resultSummary": "{{total}} chunks retrieved · {{retriever}}", + "latency": "{{ms}} ms", + "empty": "No chunks retrieved for this query.", + "emptyWithThreshold": "No chunks scored at or above the score threshold ({{threshold}}). Lower it to see near misses.", + "showMore": "Show full chunk", + "showLess": "Show less", + "noScore": "no score", + "noScoreHint": "This retriever ranks chunks without producing a comparable per-chunk score.", + "scoreKinds": { + "cosine_similarity": "similarity", + "l2_distance": "distance", + "rrf": "RRF" + }, + "errors": { + "failed": "Retrieval test failed. Please try again." + } + }, "editConfig": "Source settings", "shareWithTeam": "Share with team", "deleteWarning": "Are you sure you want to delete \"{{name}}\"?", diff --git a/frontend/src/settings/Sources.tsx b/frontend/src/settings/Sources.tsx index e9892e24..c6847247 100644 --- a/frontend/src/settings/Sources.tsx +++ b/frontend/src/settings/Sources.tsx @@ -59,6 +59,7 @@ import ConvertToWikiModal from './ConvertToWikiModal'; import EnableGraphRAGModal from './EnableGraphRAGModal'; import { clearGraphBuild, selectGraphBuilds } from './graphBuildSlice'; import SourceConfigModal from './SourceConfigModal'; +import TestRetrievalModal from './TestRetrievalModal'; type SourceMenuOption = { icon: string | LucideIcon; @@ -123,6 +124,9 @@ export default function Sources({ ); const [configModalState, setConfigModalState] = useState('INACTIVE'); + const [documentToTest, setDocumentToTest] = useState(null); + const [testRetrievalState, setTestRetrievalState] = + useState('INACTIVE'); const [documentToConvert, setDocumentToConvert] = useState(null); const [convertModalState, setConvertModalState] = useState('INACTIVE'); @@ -430,6 +434,20 @@ export default function Sources({ }); } + if (document.id) { + actions.push({ + icon: SearchIcon, + label: t('settings.sources.testRetrieval.action'), + onClick: () => { + setDocumentToTest(document); + setTestRetrievalState('ACTIVE'); + }, + iconWidth: 16, + iconHeight: 16, + variant: 'default', + }); + } + if ( document.id && !isWiki && @@ -530,6 +548,21 @@ export default function Sources({ }; }, [graphRAGAvailable, token]); + // Rendered inside the open source view's own header row. + const testRetrievalAction = documentToView ? ( + + ) : null; + return documentToView ? (
{documentToView.config?.kind === 'wiki' || @@ -542,12 +575,14 @@ export default function Sources({ documentToView.team_access === 'editor' } onBackToDocuments={() => setDocumentToView(undefined)} + headerAction={testRetrievalAction} /> ) : documentToView.config?.kind === 'graphrag' ? ( setDocumentToView(undefined)} + headerAction={testRetrievalAction} /> ) : documentToView.isNested ? ( documentToView.type === 'connector:file' ? ( @@ -555,12 +590,14 @@ export default function Sources({ docId={documentToView.id || ''} sourceName={documentToView.name} onBackToDocuments={() => setDocumentToView(undefined)} + headerAction={testRetrievalAction} /> ) : ( setDocumentToView(undefined)} + headerAction={testRetrievalAction} /> ) ) : ( @@ -568,8 +605,17 @@ export default function Sources({ documentId={documentToView.id || ''} documentName={documentToView.name} handleGoBack={() => setDocumentToView(undefined)} + headerAction={testRetrievalAction} /> )} +
) : (
@@ -924,6 +970,20 @@ export default function Sources({ }} /> + { + setTestRetrievalState(state); + if (state === 'INACTIVE') { + setDocumentToTest(null); + } + }} + document={documentToTest} + hybridAvailable={hybridAvailable} + graphRAGAvailable={graphRAGAvailable} + availableModels={availableModels} + /> + { diff --git a/frontend/src/settings/TestRetrievalModal.test.tsx b/frontend/src/settings/TestRetrievalModal.test.tsx new file mode 100644 index 00000000..59383a61 --- /dev/null +++ b/frontend/src/settings/TestRetrievalModal.test.tsx @@ -0,0 +1,70 @@ +import i18n from 'i18next'; +import { renderToStaticMarkup } from 'react-dom/server'; +import { I18nextProvider, initReactI18next } from 'react-i18next'; +import { beforeAll, describe, expect, it } from 'vitest'; + +import en from '../locale/en.json'; +import { ScoreBadge, type RetrievedChunk } from './TestRetrievalModal'; + +const testI18n = i18n.createInstance(); + +beforeAll(async () => { + await testI18n.use(initReactI18next).init({ + lng: 'en', + fallbackLng: 'en', + resources: { en: { translation: en } }, + }); +}); + +const chunk = (over: Partial): RetrievedChunk => ({ + rank: 1, + text: 't', + title: 't', + filename: 'f.md', + source: 'f.md', + tokens: 10, + score: null, + score_kind: null, + ...over, +}); + +const render = (c: RetrievedChunk): string => + renderToStaticMarkup( + + + , + ); + +describe('ScoreBadge', () => { + // The three score kinds are NOT interchangeable — a cosine similarity is + // higher-is-better in [0,1] and is what score_threshold compares against, an + // L2 distance is lower-is-better, and an RRF score only ranks hits against + // each other. Each must be labelled as itself. + it('labels a cosine similarity as a similarity', () => { + const html = render( + chunk({ score: 0.8241, score_kind: 'cosine_similarity' }), + ); + expect(html).toContain('similarity'); + expect(html).toContain('0.824'); + expect(html).not.toContain('distance'); + }); + + it('labels a FAISS L2 score as a distance, not a similarity', () => { + const html = render(chunk({ score: 1.6515, score_kind: 'l2_distance' })); + expect(html).toContain('distance'); + expect(html).toContain('1.651'); + expect(html).not.toContain('similarity'); + }); + + it('labels a fused hybrid score as RRF', () => { + const html = render(chunk({ score: 0.0167, score_kind: 'rrf' })); + expect(html).toContain('RRF'); + expect(html).toContain('0.017'); + }); + + it('says "no score" rather than inventing one', () => { + const html = render(chunk({ score: null, score_kind: null })); + expect(html).toContain('no score'); + expect(html).not.toContain('0.000'); + }); +}); diff --git a/frontend/src/settings/TestRetrievalModal.tsx b/frontend/src/settings/TestRetrievalModal.tsx new file mode 100644 index 00000000..6cf5daf1 --- /dev/null +++ b/frontend/src/settings/TestRetrievalModal.tsx @@ -0,0 +1,315 @@ +import { useEffect, useState } from 'react'; +import { useTranslation } from 'react-i18next'; +import { useSelector } from 'react-redux'; + +import userService from '../api/services/userService'; +import Spinner from '../components/Spinner'; +import { Button } from '../components/ui/button'; +import { Input } from '../components/ui/input'; +import { Modal } from '../components/ui/modal'; +import { ActiveState, Doc } from '../models/misc'; +import type { Model } from '../models/types'; +import { selectToken } from '../preferences/preferenceSlice'; + +import RetrievalOptions, { + configToOptions, + isPrescreenConfigValid, + optionsToConfig, + type RetrievalOptionsValue, +} from './components/RetrievalOptions'; + +/** One retrieved chunk as returned by POST /api/sources//search. */ +export type RetrievedChunk = { + rank: number; + text: string; + title: string | null; + filename: string | null; + source: string | null; + tokens: number; + // null when the retriever/store produces no comparable score (graphrag's PPR + // ranking, stores without a score seam). + score: number | null; + score_kind: 'cosine_similarity' | 'l2_distance' | 'rrf' | null; +}; + +type RetrievalResult = { + query: string; + retriever: string; + // The config the run actually used, echoed back by the backend. + retrieval: { score_threshold: number | null }; + total: number; + latency_ms: number; + chunks: RetrievedChunk[]; +}; + +interface TestRetrievalModalProps { + modalState: ActiveState; + setModalState: (state: ActiveState) => void; + document: Doc | null; + hybridAvailable?: boolean; + graphRAGAvailable?: boolean; + availableModels?: Model[]; +} + +/** + * Renders a chunk's score with the label its kind earns. The kinds are NOT + * interchangeable: cosine similarity is higher-is-better in [0,1] and is what + * `score_threshold` is compared against; an L2 distance (FAISS) is + * lower-is-better and unbounded; an RRF score only ranks hits against each + * other. Showing a bare number would invite the user to read one as another. + */ +export function ScoreBadge({ chunk }: { chunk: RetrievedChunk }) { + const { t } = useTranslation(); + const tr = (key: string) => t(`settings.sources.testRetrieval.${key}`); + + if (chunk.score === null || chunk.score_kind === null) { + return ( + + {tr('noScore')} + + ); + } + + const label = + chunk.score_kind === 'cosine_similarity' + ? tr('scoreKinds.cosine_similarity') + : chunk.score_kind === 'l2_distance' + ? tr('scoreKinds.l2_distance') + : tr('scoreKinds.rrf'); + + return ( + + {label} + {chunk.score.toFixed(3)} + + ); +} + +export default function TestRetrievalModal({ + modalState, + setModalState, + document, + hybridAvailable = false, + graphRAGAvailable = false, + availableModels = [], +}: TestRetrievalModalProps) { + const { t } = useTranslation(); + const token = useSelector(selectToken); + const tr = (key: string, opts?: Record) => + t(`settings.sources.testRetrieval.${key}`, opts ?? {}); + + const [query, setQuery] = useState(''); + const [options, setOptions] = useState(() => + configToOptions(document?.config), + ); + const [running, setRunning] = useState(false); + const [result, setResult] = useState(null); + const [error, setError] = useState(null); + const [expanded, setExpanded] = useState>(new Set()); + + useEffect(() => { + if (modalState === 'ACTIVE') { + setOptions(configToOptions(document?.config)); + setQuery(''); + setResult(null); + setError(null); + setRunning(false); + setExpanded(new Set()); + } + }, [modalState, document]); + + const closeModal = () => setModalState('INACTIVE'); + + const prescreenValid = isPrescreenConfigValid(options); + const canRun = !!document?.id && !!query.trim() && !running && prescreenValid; + + const handleRun = async () => { + if (!canRun || !document?.id) return; + setRunning(true); + setError(null); + try { + const response = await userService.testSourceRetrieval( + document.id, + { + query: query.trim(), + // Send the form's retrieval block as an ad-hoc override — the backend + // never persists it, so the source's saved config is untouched. + retrieval: optionsToConfig(options).retrieval, + }, + token, + ); + const data = await response.json().catch(() => ({})); + if (!response.ok || !data?.success) { + setError(data?.message || tr('errors.failed')); + setResult(null); + return; + } + setResult(data as RetrievalResult); + setExpanded(new Set()); + } catch { + setError(tr('errors.failed')); + setResult(null); + } finally { + setRunning(false); + } + }; + + const toggleExpanded = (rank: number) => { + setExpanded((prev) => { + const next = new Set(prev); + if (next.has(rank)) next.delete(rank); + else next.add(rank); + return next; + }); + }; + + // A threshold that filtered everything out is the single most likely reason + // for an empty result, so say so instead of a bare "no chunks". Read the + // threshold off the completed run, not the live form — editing the field + // after a run must not rewrite that run's explanation. + const ranThreshold = result?.retrieval?.score_threshold ?? null; + const emptyMessage = ranThreshold + ? tr('emptyWithThreshold', { threshold: ranThreshold }) + : tr('empty'); + + return ( + !o && closeModal()} + hideTitle + title={tr('title')} + size="lg" + mobileVariant="sheet" + // Same width ramp and padding as PromptsModal so the two large modals + // read as one family. + className="bg-card dark:bg-card w-[95vw] max-w-[650px] rounded-2xl px-4 py-4 sm:px-6 sm:py-6 md:max-w-[860px] md:px-8 md:py-6 lg:max-w-[980px]" + contentClassName="max-h-[70vh]" + > +
+

+ {tr('title')} +

+

+ {document?.name + ? tr('subtitle', { name: document.name }) + : tr('subtitleGeneric')} +

+ +
+
+ setQuery(e.target.value)} + onKeyDown={(e) => { + if (e.key === 'Enter') handleRun(); + }} + /> + +
+ + + +

{tr('notSavedHint')}

+ + {!prescreenValid && ( +
+ {t('settings.sources.configModal.prescreenInvalidHint')} +
+ )} + + {error && ( +
+ {error} +
+ )} + + {result && ( +
+
+ + {tr('resultSummary', { + total: result.total, + retriever: result.retriever, + })} + + {tr('latency', { ms: result.latency_ms })} +
+ + {result.chunks.length === 0 ? ( +
+ {emptyMessage} +
+ ) : ( + result.chunks.map((chunk) => { + const isOpen = expanded.has(chunk.rank); + return ( +
+
+
+ + #{chunk.rank} + + + {chunk.filename || chunk.title || chunk.source} + +
+
+ + + {chunk.tokens} {t('settings.sources.tokensUnit')} + +
+
+ {/* Chunks routinely start and end with blank lines; left + in, the collapsed clamp spends its 3 lines on nothing + and the preview looks empty. */} +

+ {chunk.text.trim()} +

+ +
+ ); + }) + )} +
+ )} +
+
+
+ ); +} diff --git a/frontend/src/settings/components/RetrievalOptions.tsx b/frontend/src/settings/components/RetrievalOptions.tsx index f4afb3a7..2d05dd14 100644 --- a/frontend/src/settings/components/RetrievalOptions.tsx +++ b/frontend/src/settings/components/RetrievalOptions.tsx @@ -323,6 +323,12 @@ type RetrievalOptionsProps = { graphRAGAvailable?: boolean; // Models for the graph extraction-model picker, same shape as the agent form. availableModels?: Model[]; + // Shows only the knobs that change what a query retrieves, hiding the + // ingest-time groups (chunking, graph extraction) and `exposure` (which picks + // *when* a source is searched at answer time, not what comes back). Used by + // the retrieval tester, where those knobs cannot affect the result and + // showing them would imply they do. + queryOnly?: boolean; }; /** @@ -338,6 +344,7 @@ export default function RetrievalOptions({ hybridAvailable = false, graphRAGAvailable = false, availableModels = [], + queryOnly = false, }: RetrievalOptionsProps) { const { t } = useTranslation(); const [open, setOpen] = useState(false); @@ -486,50 +493,56 @@ export default function RetrievalOptions({ )} - - - setRetrieval({ rephrase_query: checked }) - } - /> - - - - + setRetrieval({ exposure: v as RetrievalExposure }) + } > - - - - - {tr('retrieval.exposures.prefetch')} - - - {tr('retrieval.exposures.agentic_tool')} - - - - + + + + + + {tr('retrieval.exposures.prefetch')} + + + {tr('retrieval.exposures.agentic_tool')} + + + + + )} {/* Graph extraction group (graphrag only; re-ingest required to apply) */} - {isGraphRAG && ( + {isGraphRAG && !queryOnly && (
@@ -676,7 +689,7 @@ export default function RetrievalOptions({ )} {/* Chunking group (re-ingest required) */} -
+
diff --git a/tests/api/user/sources/test_retrieval_test.py b/tests/api/user/sources/test_retrieval_test.py new file mode 100644 index 00000000..21b2989d --- /dev/null +++ b/tests/api/user/sources/test_retrieval_test.py @@ -0,0 +1,484 @@ +"""Tests for application/api/user/sources/retrieval_test.py.""" + +import json +import uuid +from contextlib import contextmanager +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask + + +@pytest.fixture +def app(): + return Flask(__name__) + + +@contextmanager +def _patch_db(conn): + @contextmanager + def _yield(): + yield conn + + with patch("application.api.user.sources.retrieval_test.db_readonly", _yield): + yield + + +def _grant_team_access(pg_conn, owner, member, source_id, access_level): + from application.storage.db.repositories.team_members import ( + TeamMembersRepository, + ) + from application.storage.db.repositories.team_resource_grants import ( + TeamResourceGrantsRepository, + ) + from application.storage.db.repositories.teams import TeamsRepository + + team = TeamsRepository(pg_conn).create( + "Acme", f"acme-{uuid.uuid4().hex[:8]}", owner + ) + TeamMembersRepository(pg_conn).add_member(team["id"], member, role="team_member") + TeamResourceGrantsRepository(pg_conn).grant( + team["id"], + "source", + source_id, + owner_id=owner, + granted_by=owner, + access_level=access_level, + ) + + +def _seed_source(pg_conn, user="u", name="src", config=None): + from application.storage.db.repositories.sources import SourcesRepository + + repo = SourcesRepository(pg_conn) + src = repo.create(name, user_id=user) + if config is not None: + repo.update(str(src["id"]), user, {"config": config}) + src = repo.get_any(str(src["id"]), user) + return src + + +def _post(app, source_id, body, user="u"): + from application.api.user.sources.retrieval_test import SourceSearch + + with app.test_request_context( + f"/api/sources/{source_id}/search", + method="POST", + data=json.dumps(body), + content_type="application/json", + ): + from flask import request + + request.decoded_token = {"sub": user} + return SourceSearch().post(source_id) + + +class TestSourceSearchGuards: + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.retrieval_test import SourceSearch + + with app.test_request_context( + "/api/sources/abc/search", + method="POST", + data=json.dumps({"query": "hi"}), + content_type="application/json", + ): + from flask import request + + request.decoded_token = None + response = SourceSearch().post("abc") + assert response.status_code == 401 + + def test_returns_400_without_query(self, app, pg_conn): + src = _seed_source(pg_conn, user="u-noq") + with _patch_db(pg_conn): + response = _post(app, str(src["id"]), {"query": " "}, user="u-noq") + assert response.status_code == 400 + + def test_returns_400_for_overlong_query(self, app, pg_conn): + from application.api.user.sources.retrieval_test import MAX_QUERY_LENGTH + + src = _seed_source(pg_conn, user="u-long") + with _patch_db(pg_conn): + response = _post( + app, + str(src["id"]), + {"query": "x" * (MAX_QUERY_LENGTH + 1)}, + user="u-long", + ) + assert response.status_code == 400 + + def test_returns_404_when_source_missing(self, app, pg_conn): + with _patch_db(pg_conn): + response = _post( + app, + "00000000-0000-0000-0000-000000000000", + {"query": "hi"}, + user="u", + ) + assert response.status_code == 404 + + def test_returns_400_for_invalid_retrieval_config(self, app, pg_conn): + src = _seed_source(pg_conn, user="u-bad") + with _patch_db(pg_conn): + response = _post( + app, + str(src["id"]), + # chunks must be >= 1 (RetrievalConfig._bounded_chunks) + {"query": "hi", "retrieval": {"chunks": 0}}, + user="u-bad", + ) + assert response.status_code == 400 + + def test_rejects_unknown_retrieval_field(self, app, pg_conn): + src = _seed_source(pg_conn, user="u-extra") + with _patch_db(pg_conn): + response = _post( + app, + str(src["id"]), + # RetrievalConfig forbids extras — a typo must not silently no-op + {"query": "hi", "retrieval": {"chunkz": 5}}, + user="u-extra", + ) + assert response.status_code == 400 + + +class TestSourceSearchAccess: + """Read access = owner or any team grant, matching the wiki/graph reads.""" + + def test_stranger_gets_404(self, app, pg_conn): + src = _seed_source(pg_conn, user="u-owner-x") + + fake = MagicMock() + fake.search.return_value = [{"text": "secret", "filename": "f.md"}] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post( + app, str(src["id"]), {"query": "q"}, user="u-stranger-x" + ) + + assert response.status_code == 404 + # Never even touch the vector store for a source we can't read. + dispatcher.assert_not_called() + + def test_team_viewer_can_run_a_test(self, app, pg_conn): + owner = "alice-retrieval" + viewer = "bob-retrieval-viewer" + src = _seed_source(pg_conn, user=owner) + _grant_team_access(pg_conn, owner, viewer, str(src["id"]), "viewer") + + fake = MagicMock() + fake.search.return_value = [{"text": "shared chunk", "filename": "f.md"}] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ): + response = _post(app, str(src["id"]), {"query": "q"}, user=viewer) + + assert response.status_code == 200 + assert response.json["total"] == 1 + + +class TestSourceSearchRetrieval: + def test_returns_ranked_chunks_with_scores(self, app, pg_conn): + user = "u-search" + src = _seed_source(pg_conn, user=user) + + fake = MagicMock() + fake.search.return_value = [ + { + "text": "chunk one", + "title": "T1", + "filename": "f.md", + "source": "f.md", + "score": 0.82, + "score_kind": "cosine_similarity", + }, + { + "text": "chunk two", + "title": "T2", + "filename": "f.md", + "source": "f.md", + "score": 0.71, + "score_kind": "cosine_similarity", + }, + ] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ): + response = _post(app, str(src["id"]), {"query": "what runs"}, user=user) + + assert response.status_code == 200 + data = response.json + assert data["total"] == 2 + assert [c["rank"] for c in data["chunks"]] == [1, 2] + assert data["chunks"][0]["score"] == 0.82 + assert data["chunks"][0]["score_kind"] == "cosine_similarity" + assert data["chunks"][0]["tokens"] > 0 + fake.search.assert_called_once_with("what runs") + + def test_unscored_retriever_yields_null_scores(self, app, pg_conn): + """graphrag ranks by PPR and attaches no score — it must not be faked.""" + user = "u-noscore" + src = _seed_source(pg_conn, user=user) + + fake = MagicMock() + fake.search.return_value = [{"text": "graph chunk", "filename": "g.md"}] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ): + response = _post(app, str(src["id"]), {"query": "q"}, user=user) + + assert response.status_code == 200 + assert response.json["chunks"][0]["score"] is None + assert response.json["chunks"][0]["score_kind"] is None + + def test_ad_hoc_config_is_passed_to_dispatcher_and_not_persisted( + self, app, pg_conn + ): + from application.storage.db.repositories.sources import SourcesRepository + + user = "u-adhoc" + src = _seed_source(pg_conn, user=user) + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post( + app, + str(src["id"]), + {"query": "q", "retrieval": {"chunks": 7, "score_threshold": 0.5}}, + user=user, + ) + + assert response.status_code == 200 + kwargs = dispatcher.call_args.kwargs + assert kwargs["chunks"] == 7 + assert kwargs["include_scores"] is True + assert kwargs["usage_source"] == "retrieval_test" + retrieval = kwargs["sources"][0]["retrieval"] + assert retrieval.chunks == 7 + assert retrieval.score_threshold == 0.5 + # Echoed back so the UI can show what actually ran. + assert response.json["retrieval"]["chunks"] == 7 + + # The source's stored config must be untouched by a test run. + stored = SourcesRepository(pg_conn).get_any(str(src["id"]), user) + assert (stored.get("config") or {}).get("retrieval") is None + + def test_falls_back_to_saved_config(self, app, pg_conn): + user = "u-saved" + src = _seed_source( + pg_conn, + user=user, + config={"retrieval": {"chunks": 9, "retriever": "classic"}}, + ) + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post(app, str(src["id"]), {"query": "q"}, user=user) + + assert response.status_code == 200 + assert dispatcher.call_args.kwargs["chunks"] == 9 + assert response.json["retrieval"]["chunks"] == 9 + + def test_returns_500_when_retrieval_raises(self, app, pg_conn): + user = "u-boom" + src = _seed_source(pg_conn, user=user) + + fake = MagicMock() + fake.search.side_effect = RuntimeError("vector store down") + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ): + response = _post(app, str(src["id"]), {"query": "q"}, user=user) + + assert response.status_code == 500 + + +class TestPrescreenCostCeiling: + """One request must not be able to fan out into hundreds of LLM calls.""" + + def test_rejects_a_prescreen_config_that_needs_too_many_llm_calls( + self, app, pg_conn + ): + src = _seed_source(pg_conn, user="u-costly") + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher" + ) as dispatcher: + response = _post( + app, + str(src["id"]), + { + "query": "q", + # 500 candidates screened one at a time = 500 provider calls. + "retrieval": { + "chunks": 1, + "prescreen": { + "candidate_k": 500, + "batch_size": 1, + "max_keep": 1, + }, + }, + }, + user="u-costly", + ) + + assert response.status_code == 400 + dispatcher.assert_not_called() + + def test_allows_a_prescreen_config_within_the_ceiling(self, app, pg_conn): + src = _seed_source(pg_conn, user="u-ok") + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post( + app, + str(src["id"]), + { + "query": "q", + # 40 candidates / batches of 10 = 4 calls. + "retrieval": { + "chunks": 2, + "prescreen": { + "candidate_k": 40, + "batch_size": 10, + "max_keep": 8, + }, + }, + }, + user="u-ok", + ) + + assert response.status_code == 200 + dispatcher.assert_called_once() + + def test_saved_config_is_testable_even_above_the_ceiling(self, app, pg_conn): + """The client echoes the on-screen config back, so an untouched saved + config arrives in the body. It must still be judged 'saved', or a source + configured beyond the ceiling could never be tested — which is the whole + point of the endpoint.""" + expensive = { + "chunks": 2, + "prescreen": {"candidate_k": 300, "batch_size": 10, "max_keep": 8}, + } + src = _seed_source( + pg_conn, user="u-saved-costly", config={"retrieval": expensive} + ) + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post( + app, + str(src["id"]), + # 30 batches — over the ad-hoc ceiling, but it IS the saved config. + {"query": "q", "retrieval": expensive}, + user="u-saved-costly", + ) + + assert response.status_code == 200 + dispatcher.assert_called_once() + + def test_editing_the_saved_config_upward_is_still_capped(self, app, pg_conn): + saved = { + "chunks": 2, + "prescreen": {"candidate_k": 300, "batch_size": 10, "max_keep": 8}, + } + src = _seed_source(pg_conn, user="u-edit-costly", config={"retrieval": saved}) + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.Dispatcher" + ) as dispatcher: + response = _post( + app, + str(src["id"]), + { + "query": "q", + # Same source, but the caller cranked batch_size down to 1 + # → 300 LLM calls. That is ad-hoc, and capped. + "retrieval": { + "chunks": 2, + "prescreen": { + "candidate_k": 300, + "batch_size": 1, + "max_keep": 8, + }, + }, + }, + user="u-edit-costly", + ) + + assert response.status_code == 400 + dispatcher.assert_not_called() + + +class TestModelResolution: + def test_passes_the_instance_default_model(self, app, pg_conn): + """Without a real model id the prescreen stage calls the provider with a + placeholder, fails, and silently keeps every candidate.""" + src = _seed_source(pg_conn, user="u-model") + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.get_default_model_id", + return_value="gpt-4o", + ), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post(app, str(src["id"]), {"query": "q"}, user="u-model") + + assert response.status_code == 200 + assert dispatcher.call_args.kwargs["model_id"] == "gpt-4o" + + def test_no_configured_models_leaves_the_retriever_default(self, app, pg_conn): + """get_default_model_id() is None on an instance with no models; passing + that through would override the retriever's own default with None.""" + src = _seed_source(pg_conn, user="u-nomodel") + + fake = MagicMock() + fake.search.return_value = [] + + with _patch_db(pg_conn), patch( + "application.api.user.sources.retrieval_test.get_default_model_id", + return_value=None, + ), patch( + "application.api.user.sources.retrieval_test.Dispatcher", + return_value=fake, + ) as dispatcher: + response = _post(app, str(src["id"]), {"query": "q"}, user="u-nomodel") + + assert response.status_code == 200 + assert "model_id" not in dispatcher.call_args.kwargs diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 65649879..294eb524 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -416,3 +416,36 @@ class TestGetChunkTexts: store, cursor = self._store_with_mock_conn() assert store.get_chunk_texts("sid", []) == {} cursor.execute.assert_not_called() + + +class TestGraphRAGTopK: + """A prescreen source elsewhere in the group inflates ``chunks``; a graph + source must still contribute only its own top-k.""" + + @patch("application.retriever.graph_rag.num_tokens_from_string", return_value=10) + @patch("application.retriever.graph_rag.GraphStore") + @patch("application.retriever.graph_rag.graphrag_available", return_value=True) + def test_inflated_chunks_do_not_raise_a_graph_source_top_k( + self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed + ): + nodes = [{"id": f"n{i}", "doc_freq": 1} for i in range(1, 5)] + edges = [ + {"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}, + {"src_node_id": "n2", "dst_node_id": "n3", "weight": 1.0}, + {"src_node_id": "n3", "dst_node_id": "n4", "weight": 1.0}, + ] + node_chunks = {f"n{i}": [f"c{i}"] for i in range(1, 5)} + chunk_texts = {f"c{i}": f"text {i}" for i in range(1, 5)} + seed_rows = [{"id": "n1", "distance": 0.0}] + mock_store_cls.return_value = _store_with_graph( + nodes, edges, node_chunks, chunk_texts, seed_rows + ) + + # What the Dispatcher does when another source in the group prescreens + # at candidate_k=40: chunks inflated to 40, base_chunks left at the real 2. + rag = _make_retriever(chunks=40) + rag.base_chunks = 2 + + docs = rag._get_data() + + assert len(docs) == 2 diff --git a/tests/retriever/test_hybrid.py b/tests/retriever/test_hybrid.py index 7f5efa5f..7db1c719 100644 --- a/tests/retriever/test_hybrid.py +++ b/tests/retriever/test_hybrid.py @@ -4,7 +4,7 @@ from unittest.mock import MagicMock, Mock, patch import pytest -from application.retriever.hybrid_rag import HybridRetriever, reciprocal_rank_fusion +from application.retriever.hybrid_rag import fuse_with_scores, HybridRetriever from application.retriever.retriever_creator import RetrieverCreator @@ -53,27 +53,27 @@ class TestReciprocalRankFusion: vector_hits = [only_vec, shared] keyword_hits = [shared, only_kw] - fused = reciprocal_rank_fusion(vector_hits, keyword_hits) + fused = [doc for doc, _ in fuse_with_scores(vector_hits, keyword_hits)] assert fused[0].page_content == "shared" assert {d.page_content for d in fused} == {"shared", "vec_only", "kw_only"} def test_empty_keyword_is_vector_only_order(self): vector_hits = [_make_doc("a", source="a"), _make_doc("b", source="b")] - fused = reciprocal_rank_fusion(vector_hits, []) + fused = [doc for doc, _ in fuse_with_scores(vector_hits, [])] assert [d.page_content for d in fused] == ["a", "b"] def test_dedupes_same_doc(self): d_vec = _make_doc("same", source="same") d_kw = _make_doc("same", source="same") - fused = reciprocal_rank_fusion([d_vec], [d_kw]) + fused = [doc for doc, _ in fuse_with_scores([d_vec], [d_kw])] assert len(fused) == 1 def test_higher_keyword_rank_can_promote(self): # Vector top is "v0"; keyword strongly favours "kw" (rank 0 vs v0's rank 1). v0 = _make_doc("v0", source="v0") kw = _make_doc("kw", source="kw") - fused = reciprocal_rank_fusion([v0, kw], [kw]) + fused = [doc for doc, _ in fuse_with_scores([v0, kw], [kw])] assert fused[0].page_content == "kw" diff --git a/tests/retriever/test_include_scores.py b/tests/retriever/test_include_scores.py new file mode 100644 index 00000000..65fea66f --- /dev/null +++ b/tests/retriever/test_include_scores.py @@ -0,0 +1,206 @@ +"""Score passthrough (``include_scores``) across the retrievers. + +The retrieval-test endpoint needs each chunk's raw store score; the answer +pipeline does not. These tests pin both halves of that: opted in, the score and +its kind ride on every doc; opted out (the default everywhere else), the doc +dicts are exactly what they were before the flag existed. +""" + +from unittest.mock import Mock, patch + +import pytest + +from application.retriever.classic_rag import ClassicRAG +from application.retriever.hybrid_rag import HybridRetriever + + +@pytest.fixture +def _patch_llm_creator(mock_llm, monkeypatch): + monkeypatch.setattr( + "application.retriever.classic_rag.LLMCreator.create_llm", + Mock(return_value=mock_llm), + ) + return mock_llm + + +def _make_doc(page_content, source="s", title="t"): + doc = Mock() + doc.page_content = page_content + doc.metadata = {"title": title, "source": source} + return doc + + +def _make_store(score_kind="cosine_similarity"): + store = Mock() + store.score_kind = score_kind + store.search.return_value = [_make_doc("hit one"), _make_doc("hit two")] + store.search_with_scores.return_value = [ + (_make_doc("hit one"), 0.82), + (_make_doc("hit two"), 0.71), + ] + store.keyword_search.return_value = [] + return store + + +def _retrieve(retriever_cls, store, **overrides): + kwargs = dict( + source={"question": "q", "active_docs": ["vs1"]}, + chat_history=None, + prompt="", + chunks=2, + doc_token_limit=50000, + model_id="test-model", + llm_name="openai", + api_key="fake", + decoded_token={"sub": "user1"}, + ) + kwargs.update(overrides) + retriever = retriever_cls(**kwargs) + with patch( + "application.retriever.classic_rag.VectorCreator.create_vectorstore", + return_value=store, + ): + return retriever.search("q") + + +@pytest.mark.unit +class TestClassicRAGScores: + def test_off_by_default_leaves_docs_untouched(self, _patch_llm_creator): + """The answer pipeline must see exactly the doc dict it saw before.""" + store = _make_store() + docs = _retrieve(ClassicRAG, store) + + assert docs + assert set(docs[0]) == {"text", "title", "source", "filename"} + store.search.assert_called_once() + store.search_with_scores.assert_not_called() + + def test_include_scores_attaches_score_and_kind(self, _patch_llm_creator): + store = _make_store() + docs = _retrieve(ClassicRAG, store, include_scores=True) + + assert [d["score"] for d in docs] == [0.82, 0.71] + assert {d["score_kind"] for d in docs} == {"cosine_similarity"} + store.search_with_scores.assert_called_once() + store.search.assert_not_called() + + def test_unscored_store_yields_null_scores(self, _patch_llm_creator): + """A store with no score seam reports None rather than a fabricated 0.""" + store = _make_store(score_kind=None) + store.search_with_scores.return_value = [ + (_make_doc("hit one"), None), + (_make_doc("hit two"), None), + ] + docs = _retrieve(ClassicRAG, store, include_scores=True) + + assert [d["score"] for d in docs] == [None, None] + assert [d["score_kind"] for d in docs] == [None, None] + + +@pytest.mark.unit +class TestTopK: + """``chunks`` is the final top-k, not just a floor on the fetch size.""" + + def test_returns_at_most_chunks_docs(self, _patch_llm_creator): + store = Mock() + store.score_kind = None + store.search.return_value = [_make_doc(f"hit {i}") for i in range(20)] + store.keyword_search.return_value = [] + + docs = _retrieve(ClassicRAG, store, chunks=2) + + assert len(docs) == 2 + assert [d["text"] for d in docs] == ["hit 0", "hit 1"] + # The over-fetch itself is intact — only the tail is dropped. + assert store.search.call_args.kwargs["k"] == 20 + + def test_token_budget_still_caps_below_top_k(self, _patch_llm_creator): + """The budget remains the harder of the two limits.""" + store = Mock() + store.score_kind = None + store.search.return_value = [_make_doc("word " * 500) for _ in range(10)] + store.keyword_search.return_value = [] + + docs = _retrieve(ClassicRAG, store, chunks=10, doc_token_limit=600) + + assert 0 < len(docs) < 10 + + def test_hybrid_respects_top_k_too(self, _patch_llm_creator): + store = Mock() + store.score_kind = None + store.search.return_value = [_make_doc(f"hit {i}") for i in range(20)] + store.keyword_search.return_value = [] + + docs = _retrieve(HybridRetriever, store, chunks=3) + + assert len(docs) == 3 + + +@pytest.mark.unit +class TestHybridScores: + def test_reports_rrf_not_the_store_kind(self, _patch_llm_creator): + """RRF fuses two rankings — the fused number is not the store's cosine + score, so it must not be labelled as one.""" + store = _make_store() + store.keyword_search.return_value = [_make_doc("hit two")] + + docs = _retrieve(HybridRetriever, store, include_scores=True) + + assert {d["score_kind"] for d in docs} == {"rrf"} + # A doc found by both searches outranks one found by vector search alone. + assert docs[0]["text"] == "hit two" + assert docs[0]["score"] > docs[1]["score"] + + def test_off_by_default_leaves_docs_untouched(self, _patch_llm_creator): + store = _make_store() + docs = _retrieve(HybridRetriever, store) + + assert docs + assert set(docs[0]) == {"text", "title", "source", "filename"} + + +@pytest.mark.unit +class TestCandidateKDoesNotLeakAcrossSources: + """A prescreen source's inflated fetch must not become its neighbour's top-k. + + The Dispatcher raises the group's ``chunks`` to the prescreen candidate_k so + the fetch is big enough to screen. A source in that same group with no + override still has to fall back to the group's *real* top-k. + """ + + def test_default_source_uses_base_chunks_not_the_inflated_fetch( + self, _patch_llm_creator + ): + store = Mock() + store.score_kind = None + store.search.return_value = [_make_doc(f"hit {i}") for i in range(40)] + store.keyword_search.return_value = [] + + # What the Dispatcher does for a group whose other source prescreens at + # candidate_k=40: chunks inflated to 40, base_chunks kept at the real 2. + retriever = ClassicRAG( + source={"question": "q", "active_docs": ["vs1", "vs2"]}, + chunks=40, + doc_token_limit=50000, + decoded_token={"sub": "u"}, + ) + retriever.base_chunks = 2 + + with patch( + "application.retriever.classic_rag.VectorCreator.create_vectorstore", + return_value=store, + ): + docs = retriever.search("q") + + # 2 sources sharing a top-k of 2 → 1 chunk each, not 20 each. + assert len(docs) == 2 + + def test_absent_base_chunks_keeps_chunks_as_the_top_k(self, _patch_llm_creator): + store = Mock() + store.score_kind = None + store.search.return_value = [_make_doc(f"hit {i}") for i in range(40)] + store.keyword_search.return_value = [] + + docs = _retrieve(ClassicRAG, store, chunks=4) + + assert len(docs) == 4 diff --git a/tests/vectorstore/test_base.py b/tests/vectorstore/test_base.py index 7e0346aa..12d3ef79 100644 --- a/tests/vectorstore/test_base.py +++ b/tests/vectorstore/test_base.py @@ -384,3 +384,34 @@ class TestBaseVectorStore: result = store._get_embeddings("some_custom_embedding") assert result is mock_emb mock_get_instance.assert_called_with("some_custom_embedding") + + +@pytest.mark.unit +class TestSearchWithScoresDefault: + def test_pairs_hits_with_none(self): + """A store that reports no score still satisfies the contract, so the + retriever never has to special-case it.""" + from application.vectorstore.base import BaseVectorStore + + class _Store(BaseVectorStore): + def search(self, question, k=2, *args, **kwargs): + return ["a", "b"] + + def add_texts(self, texts, metadatas=None, *args, **kwargs): + return [] + + store = _Store() + assert store.score_kind is None + assert store.search_with_scores("q", k=2) == [("a", None), ("b", None)] + + def test_handles_store_returning_none(self): + from application.vectorstore.base import BaseVectorStore + + class _Store(BaseVectorStore): + def search(self, question, k=2, *args, **kwargs): + return None + + def add_texts(self, texts, metadatas=None, *args, **kwargs): + return [] + + assert _Store().search_with_scores("q") == [] diff --git a/tests/vectorstore/test_faiss.py b/tests/vectorstore/test_faiss.py index 7646f329..b8840de9 100644 --- a/tests/vectorstore/test_faiss.py +++ b/tests/vectorstore/test_faiss.py @@ -562,3 +562,67 @@ class TestFaissStoreAssertEmbeddingDimensionsMatch: # Should not raise since embedding name is not the huggingface one store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) assert store is not None + + +@pytest.mark.unit +class TestFaissSearchWithScores: + @staticmethod + def _build(mock_faiss, mock_get_emb, mock_settings, mock_storage_creator, ds): + mock_settings.EMBEDDINGS_NAME = "test_model" + mock_get_emb.return_value = Mock(dimension=3) + mock_faiss.from_documents.return_value = ds + mock_storage_creator.get_storage.return_value = Mock() + + from application.vectorstore.faiss import FaissStore + + return FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_reports_l2_distance_and_drops_threshold( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + doc = Mock(page_content="text1", metadata={"source": "a"}) + ds = Mock(index=Mock(d=3)) + ds.similarity_search_with_score.return_value = [(doc, 0.42)] + store = self._build( + mock_faiss, mock_get_emb, mock_settings, mock_storage_creator, ds + ) + + results = store.search_with_scores("query", k=5, score_threshold=0.9) + + # LangChain's FAISS ranks by L2 distance, NOT cosine similarity — the + # label must say so, and score_threshold must be dropped (FAISS has no + # such knob and would crash on it). + assert store.score_kind == "l2_distance" + ds.similarity_search_with_score.assert_called_once_with("query", k=5) + assert results == [(doc, 0.42)] + + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_does_not_mutate_docstore_metadata( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + # similarity_search_with_score hands back the LIVE docstore Documents; + # writing a score into their metadata would pollute the in-memory index + # and could be persisted back to storage by a later add_chunk. + doc = Mock(page_content="text1", metadata={"source": "a"}) + ds = Mock(index=Mock(d=3)) + ds.similarity_search_with_score.return_value = [(doc, 0.42)] + store = self._build( + mock_faiss, mock_get_emb, mock_settings, mock_storage_creator, ds + ) + + store.search_with_scores("query", k=1) + + assert doc.metadata == {"source": "a"} diff --git a/tests/vectorstore/test_mongodb.py b/tests/vectorstore/test_mongodb.py index a3030edf..20291253 100644 --- a/tests/vectorstore/test_mongodb.py +++ b/tests/vectorstore/test_mongodb.py @@ -287,3 +287,40 @@ class TestMongoDBVectorStoreDeleteChunk: result = store.delete_chunk("bad_id") assert result is False + + +@pytest.mark.unit +class TestMongoDBSearchWithScores: + def test_reports_vector_search_score(self): + store, mock_collection, _ = _make_mongodb_store() + mock_collection.aggregate.return_value = iter( + [ + { + "_id": "id1", + "text": "hello", + "embedding": [0.1], + "source": "a", + "_score": 0.83, + } + ] + ) + + results = store.search_with_scores("query", k=1) + + assert store.score_kind == "cosine_similarity" + assert results[0][0].page_content == "hello" + assert results[0][1] == pytest.approx(0.83) + # The score must not leak into metadata — it rides alongside the doc. + assert "_score" not in results[0][0].metadata + + def test_score_is_added_even_without_a_threshold(self): + """The $addFields stage is unconditional, so an unfiltered search still + carries a score (the whole point of the retrieval tester).""" + store, mock_collection, _ = _make_mongodb_store() + mock_collection.aggregate.return_value = iter([]) + + store.search_with_scores("query", k=1) + + pipeline = mock_collection.aggregate.call_args[0][0] + assert any("$addFields" in stage for stage in pipeline) + assert not any("$match" in stage for stage in pipeline) diff --git a/tests/vectorstore/test_pgvector.py b/tests/vectorstore/test_pgvector.py index 78db6de2..0611653c 100644 --- a/tests/vectorstore/test_pgvector.py +++ b/tests/vectorstore/test_pgvector.py @@ -403,3 +403,38 @@ class TestPGVectorStoreConnection: store.__del__() mock_conn.close.assert_called_once() + + +@pytest.mark.unit +class TestPGVectorSearchWithScores: + def test_reports_cosine_similarity(self): + store, _, mock_cursor, _ = _make_store() + mock_cursor.fetchall.return_value = [ + ("close", {"source": "a.txt"}, 0.10), + ("far", {"source": "b.txt"}, 0.40), + ] + + results = store.search_with_scores("query", k=2) + + assert store.score_kind == "cosine_similarity" + assert [doc.page_content for doc, _ in results] == ["close", "far"] + # similarity = 1 - cosine distance, the quantity score_threshold uses. + assert results[0][1] == pytest.approx(0.90) + assert results[1][1] == pytest.approx(0.60) + + def test_honours_score_threshold(self): + store, _, mock_cursor, _ = _make_store() + mock_cursor.fetchall.return_value = [ + ("close", {}, 0.10), # sim 0.90 → kept + ("far", {}, 0.40), # sim 0.60 → dropped + ] + + results = store.search_with_scores("query", k=5, score_threshold=0.85) + + assert [doc.page_content for doc, _ in results] == ["close"] + + def test_returns_empty_on_error(self): + store, _, mock_cursor, _ = _make_store() + mock_cursor.execute.side_effect = Exception("connection lost") + + assert store.search_with_scores("query") == []