Chunks preview

This commit is contained in:
Pavel committed 2026-07-13 23:22:17 +03:00
1 parent 73c3dfb5c4
commit 0447eb9b8d
33 files changed
+1961 -109

No files matched your search

+8 -2
View File
@@ -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)
+7 -1
View File
@@ -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",
]
@@ -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/<string:source_id>/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,
)
+44 -5
View File
@@ -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
+14 -1
View File
@@ -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
+35 -1
View File
@@ -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:
+17 -7
View File
@@ -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]
+6 -1
View File
@@ -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
+20
View File
@@ -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"""
+14
View File
@@ -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)
+21 -4
View File
@@ -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):
+28 -2
View File
@@ -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 []
+1
View File
@@ -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`,
+6
View File
@@ -102,6 +102,12 @@ const userService = {
token: string | null,
): Promise<Response> =>
apiClient.patch(endpoints.USER.SOURCE_CONFIG(sourceId), config, token),
testSourceRetrieval: (
sourceId: string,
data: { query: string; retrieval?: any },
token: string | null,
): Promise<Response> =>
apiClient.post(endpoints.USER.SOURCE_SEARCH(sourceId), data, token),
createWiki: (
data: { name: string; initial_content?: string },
token: string | null,
+4
View File
@@ -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<ChunksProps> = ({
@@ -118,6 +120,7 @@ const Chunks: React.FC<ChunksProps> = ({
displayPath,
onFileSearch,
onFileSelect,
headerAction,
}) => {
const [fileSearchQuery, setFileSearchQuery] = useState('');
const [fileSearchResults, setFileSearchResults] = useState<SearchResult[]>(
@@ -349,6 +352,7 @@ const Chunks: React.FC<ChunksProps> = ({
</div>
<div className="mt-2 flex w-full flex-row flex-nowrap items-center justify-end gap-2 overflow-x-auto sm:mt-0 sm:w-auto">
{headerAction}
{editingChunk ? (
!isEditing ? (
<>
+31 -25
View File
@@ -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<ConnectorTreeProps> = ({
docId,
sourceName,
onBackToDocuments,
headerAction,
}) => {
const { t } = useTranslation();
const token = useSelector(selectToken);
@@ -101,33 +104,36 @@ const ConnectorTree: React.FC<ConnectorTreeProps> = ({
};
const topRightAction = (
<button
onClick={() => setSyncConfirmationModal('ACTIVE')}
disabled={isSyncing}
className={`flex h-[38px] min-w-[108px] items-center justify-center rounded-full px-4 text-sm font-medium whitespace-nowrap transition-colors ${
isSyncing
? 'dark:bg-muted dark:text-muted-foreground cursor-not-allowed bg-gray-300 text-gray-600'
: 'bg-primary hover:bg-primary/90 text-white'
}`}
title={
isSyncing
? `${t('settings.sources.syncing')} ${syncProgress}%`
<>
{headerAction}
<button
onClick={() => setSyncConfirmationModal('ACTIVE')}
disabled={isSyncing}
className={`flex h-[38px] min-w-[108px] items-center justify-center rounded-full px-4 text-sm font-medium whitespace-nowrap transition-colors ${
isSyncing
? 'dark:bg-muted dark:text-muted-foreground cursor-not-allowed bg-gray-300 text-gray-600'
: 'bg-primary hover:bg-primary/90 text-white'
}`}
title={
isSyncing
? `${t('settings.sources.syncing')} ${syncProgress}%`
: syncDone
? 'Done'
: t('settings.sources.sync')
}
>
<img
src={syncDone ? CheckmarkIcon : SyncIcon}
alt={t('settings.sources.sync')}
className={`mr-2 h-4 w-4 brightness-0 invert filter ${isSyncing ? 'animate-spin' : ''}`}
/>
{isSyncing
? `${syncProgress}%`
: syncDone
? 'Done'
: t('settings.sources.sync')
}
>
<img
src={syncDone ? CheckmarkIcon : SyncIcon}
alt={t('settings.sources.sync')}
className={`mr-2 h-4 w-4 brightness-0 invert filter ${isSyncing ? 'animate-spin' : ''}`}
/>
{isSyncing
? `${syncProgress}%`
: syncDone
? 'Done'
: t('settings.sources.sync')}
</button>
: t('settings.sources.sync')}
</button>
</>
);
const extraContent = (
+19 -9
View File
@@ -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<FileTreeProps> = ({
docId,
sourceName,
onBackToDocuments,
headerAction,
}) => {
const { t } = useTranslation();
const token = useSelector(selectToken);
@@ -231,15 +234,22 @@ const FileTree: React.FC<FileTreeProps> = ({
: t('settings.sources.deletingTitle')
: null;
const topRightAction = !isProcessing ? (
<button
onClick={handleAddFile}
className="bg-primary hover:bg-primary/90 flex h-[38px] min-w-[108px] items-center justify-center rounded-full px-4 text-sm font-medium whitespace-nowrap text-white"
title={t('settings.sources.addFile')}
>
{t('settings.sources.addFile')}
</button>
) : null;
// headerAction stays visible while an upload/delete is in flight — only the
// Add file button is suppressed then.
const topRightAction = (
<>
{headerAction}
{!isProcessing ? (
<button
onClick={handleAddFile}
className="bg-primary hover:bg-primary/90 flex h-[38px] min-w-[108px] items-center justify-center rounded-full px-4 text-sm font-medium whitespace-nowrap text-white"
title={t('settings.sources.addFile')}
>
{t('settings.sources.addFile')}
</button>
) : null}
</>
);
const extraContent = (
<ConfirmationModal
+4
View File
@@ -26,6 +26,8 @@ interface GraphViewProps {
docId: string;
sourceName: string;
onBackToDocuments: () => 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<GraphViewProps> = ({
docId,
sourceName,
onBackToDocuments,
headerAction,
}) => {
const { t } = useTranslation();
const token = useSelector(selectToken);
@@ -169,6 +172,7 @@ const GraphView: React.FC<GraphViewProps> = ({
<span className="text-primary font-semibold wrap-break-word">
{sourceName}
</span>
{headerAction ? <div className="ml-auto">{headerAction}</div> : null}
</div>
<div className="bg-muted/60 text-muted-foreground dark:bg-accent/40 mb-4 flex items-start gap-2 rounded-xl px-4 py-3 text-xs">
+4
View File
@@ -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<WikiViewerProps> = ({
sourceName,
canEdit = false,
onBackToDocuments,
headerAction,
}) => {
const { t } = useTranslation();
const token = useSelector(selectToken);
@@ -220,6 +223,7 @@ const WikiViewer: React.FC<WikiViewerProps> = ({
<span className="text-primary font-semibold wrap-break-word">
{sourceName}
</span>
{headerAction ? <div className="ml-auto">{headerAction}</div> : null}
</div>
<div className="bg-muted/60 text-muted-foreground dark:bg-accent/40 mb-4 flex items-start gap-2 rounded-xl px-4 py-3 text-xs">
+2 -2
View File
@@ -633,7 +633,7 @@ const TreeBrowser: React.FC<TreeBrowserProps> = ({
/>
{searchQuery && (
<div className="border-border bg-card dark:border-border dark:bg-card absolute top-full right-0 left-0 z-10 max-h-[calc(100vh-200px)] w-full overflow-hidden rounded-b-xl border border-t-0 shadow-lg transition-all duration-200">
<div className="border-border bg-card dark:border-border dark:bg-card absolute top-full right-0 left-0 z-20 max-h-[calc(100vh-200px)] w-full overflow-hidden rounded-b-xl border border-t-0 shadow-lg transition-all duration-200">
<div className="max-h-[calc(100vh-200px)] overflow-x-hidden overflow-y-auto overscroll-contain">
{searchResults.length === 0 ? (
<div className="text-muted-foreground py-2 text-center text-sm">
@@ -744,7 +744,7 @@ const TreeBrowser: React.FC<TreeBrowserProps> = ({
</div>
</div>
) : (
<div className="flex w-full max-w-full flex-col overflow-hidden">
<div className="flex w-full max-w-full flex-col overflow-x-clip">
<div className="mb-2">{renderPathNavigation()}</div>
<div className="w-full">
+25
View File
@@ -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}}\"?",
+60
View File
@@ -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<ActiveState>('INACTIVE');
const [documentToTest, setDocumentToTest] = useState<Doc | null>(null);
const [testRetrievalState, setTestRetrievalState] =
useState<ActiveState>('INACTIVE');
const [documentToConvert, setDocumentToConvert] = useState<Doc | null>(null);
const [convertModalState, setConvertModalState] =
useState<ActiveState>('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 ? (
<Button
type="button"
variant="outline"
className="h-[38px] rounded-full px-4 text-sm font-medium whitespace-nowrap"
onClick={() => {
setDocumentToTest(documentToView);
setTestRetrievalState('ACTIVE');
}}
>
{t('settings.sources.testRetrieval.action')}
</Button>
) : null;
return documentToView ? (
<div className="mt-8 flex flex-col">
{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' ? (
<GraphView
docId={documentToView.id || ''}
sourceName={documentToView.name}
onBackToDocuments={() => 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}
/>
) : (
<FileTree
docId={documentToView.id || ''}
sourceName={documentToView.name}
onBackToDocuments={() => setDocumentToView(undefined)}
headerAction={testRetrievalAction}
/>
)
) : (
@@ -568,8 +605,17 @@ export default function Sources({
documentId={documentToView.id || ''}
documentName={documentToView.name}
handleGoBack={() => setDocumentToView(undefined)}
headerAction={testRetrievalAction}
/>
)}
<TestRetrievalModal
modalState={testRetrievalState}
setModalState={setTestRetrievalState}
document={documentToTest}
hybridAvailable={hybridAvailable}
graphRAGAvailable={graphRAGAvailable}
availableModels={availableModels}
/>
</div>
) : (
<div className="mt-8 flex w-full max-w-full flex-col">
@@ -924,6 +970,20 @@ export default function Sources({
}}
/>
<TestRetrievalModal
modalState={testRetrievalState}
setModalState={(state) => {
setTestRetrievalState(state);
if (state === 'INACTIVE') {
setDocumentToTest(null);
}
}}
document={documentToTest}
hybridAvailable={hybridAvailable}
graphRAGAvailable={graphRAGAvailable}
availableModels={availableModels}
/>
<ConvertToWikiModal
modalState={convertModalState}
setModalState={(state) => {
@@ -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>): 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(
<I18nextProvider i18n={testI18n}>
<ScoreBadge chunk={c} />
</I18nextProvider>,
);
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');
});
});
@@ -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/<id>/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 (
<span className="text-muted-foreground text-xs" title={tr('noScoreHint')}>
{tr('noScore')}
</span>
);
}
const label =
chunk.score_kind === 'cosine_similarity'
? tr('scoreKinds.cosine_similarity')
: chunk.score_kind === 'l2_distance'
? tr('scoreKinds.l2_distance')
: tr('scoreKinds.rrf');
return (
<span className="text-muted-foreground font-mono text-xs">
<span className="mr-1 font-sans">{label}</span>
{chunk.score.toFixed(3)}
</span>
);
}
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<string, unknown>) =>
t(`settings.sources.testRetrieval.${key}`, opts ?? {});
const [query, setQuery] = useState('');
const [options, setOptions] = useState<RetrievalOptionsValue>(() =>
configToOptions(document?.config),
);
const [running, setRunning] = useState(false);
const [result, setResult] = useState<RetrievalResult | null>(null);
const [error, setError] = useState<string | null>(null);
const [expanded, setExpanded] = useState<Set<number>>(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 (
<Modal
open={modalState === 'ACTIVE'}
onOpenChange={(o) => !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]"
>
<div className="flex flex-col">
<p className="mb-1 text-xl font-semibold text-[#2B2B2B] dark:text-white">
{tr('title')}
</p>
<p className="dark:text-muted-foreground mb-6 text-sm text-[#6B6B6B]">
{document?.name
? tr('subtitle', { name: document.name })
: tr('subtitleGeneric')}
</p>
<div className="flex flex-col gap-4">
<div className="flex flex-row items-center gap-2">
<Input
type="text"
value={query}
autoFocus
placeholder={tr('queryPlaceholder')}
className="h-[42px] flex-1 rounded-3xl px-4"
onChange={(e) => setQuery(e.target.value)}
onKeyDown={(e) => {
if (e.key === 'Enter') handleRun();
}}
/>
<Button
type="button"
disabled={!canRun}
onClick={handleRun}
className="h-[42px] min-w-[96px] shrink-0 rounded-3xl px-6 text-sm font-medium"
>
{running ? <Spinner size="small" /> : tr('run')}
</Button>
</div>
<RetrievalOptions
value={options}
onChange={setOptions}
queryOnly
hybridAvailable={hybridAvailable}
graphRAGAvailable={graphRAGAvailable}
availableModels={availableModels}
/>
<p className="text-muted-foreground text-xs">{tr('notSavedHint')}</p>
{!prescreenValid && (
<div className="rounded-xl bg-amber-50 p-3 text-xs text-amber-800 dark:bg-amber-900/30 dark:text-amber-200">
{t('settings.sources.configModal.prescreenInvalidHint')}
</div>
)}
{error && (
<div className="rounded-xl bg-red-50 p-3 text-sm text-red-700 dark:bg-red-900/40 dark:text-red-300">
{error}
</div>
)}
{result && (
<div className="flex flex-col gap-3">
<div className="text-muted-foreground flex flex-row items-center justify-between text-xs">
<span>
{tr('resultSummary', {
total: result.total,
retriever: result.retriever,
})}
</span>
<span>{tr('latency', { ms: result.latency_ms })}</span>
</div>
{result.chunks.length === 0 ? (
<div className="border-border text-muted-foreground rounded-xl border border-dashed p-6 text-center text-sm">
{emptyMessage}
</div>
) : (
result.chunks.map((chunk) => {
const isOpen = expanded.has(chunk.rank);
return (
<div
key={chunk.rank}
className="border-border bg-muted/40 rounded-xl border p-4"
>
<div className="mb-2 flex flex-row items-center justify-between gap-2">
<div className="flex min-w-0 flex-row items-center gap-2">
<span className="bg-muted text-muted-foreground shrink-0 rounded-md px-2 py-0.5 font-mono text-xs">
#{chunk.rank}
</span>
<span
className="text-foreground truncate text-sm font-medium"
title={chunk.source ?? undefined}
>
{chunk.filename || chunk.title || chunk.source}
</span>
</div>
<div className="flex shrink-0 flex-row items-center gap-3">
<ScoreBadge chunk={chunk} />
<span className="text-muted-foreground text-xs">
{chunk.tokens} {t('settings.sources.tokensUnit')}
</span>
</div>
</div>
{/* Chunks routinely start and end with blank lines; left
in, the collapsed clamp spends its 3 lines on nothing
and the preview looks empty. */}
<p
className={`text-muted-foreground text-sm whitespace-pre-wrap ${
isOpen ? '' : 'line-clamp-3'
}`}
>
{chunk.text.trim()}
</p>
<Button
type="button"
variant="link"
onClick={() => toggleExpanded(chunk.rank)}
className="text-muted-foreground h-auto px-0 py-1 text-xs"
>
{isOpen ? tr('showLess') : tr('showMore')}
</Button>
</div>
);
})
)}
</div>
)}
</div>
</div>
</Modal>
);
}
@@ -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({
</SettingRow>
)}
<SettingRow
label={tr('retrieval.rephraseQuery')}
htmlFor="retrieval-rephrase"
>
<Switch
id="retrieval-rephrase"
checked={value.retrieval.rephrase_query}
disabled={disabled}
onCheckedChange={(checked) =>
setRetrieval({ rephrase_query: checked })
}
/>
</SettingRow>
<SettingRow
label={tr('retrieval.exposure')}
htmlFor="retrieval-exposure"
description={tr('retrieval.exposureHint')}
alignStart
>
<Select
value={value.retrieval.exposure}
disabled={disabled}
onValueChange={(v) =>
setRetrieval({ exposure: v as RetrievalExposure })
}
{/* Rephrasing only fires when there is chat history to rephrase
against, and a test has none — so the switch would do nothing. */}
{!queryOnly && (
<SettingRow
label={tr('retrieval.rephraseQuery')}
htmlFor="retrieval-rephrase"
>
<SelectTrigger
id="retrieval-exposure"
className="w-52 rounded-md"
size="lg"
<Switch
id="retrieval-rephrase"
checked={value.retrieval.rephrase_query}
disabled={disabled}
onCheckedChange={(checked) =>
setRetrieval({ rephrase_query: checked })
}
/>
</SettingRow>
)}
{!queryOnly && (
<SettingRow
label={tr('retrieval.exposure')}
htmlFor="retrieval-exposure"
description={tr('retrieval.exposureHint')}
alignStart
>
<Select
value={value.retrieval.exposure}
disabled={disabled}
onValueChange={(v) =>
setRetrieval({ exposure: v as RetrievalExposure })
}
>
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="prefetch">
{tr('retrieval.exposures.prefetch')}
</SelectItem>
<SelectItem value="agentic_tool">
{tr('retrieval.exposures.agentic_tool')}
</SelectItem>
</SelectContent>
</Select>
</SettingRow>
<SelectTrigger
id="retrieval-exposure"
className="w-52 rounded-md"
size="lg"
>
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="prefetch">
{tr('retrieval.exposures.prefetch')}
</SelectItem>
<SelectItem value="agentic_tool">
{tr('retrieval.exposures.agentic_tool')}
</SelectItem>
</SelectContent>
</Select>
</SettingRow>
)}
<SettingRow
label={tr('prescreen.enable')}
@@ -599,7 +612,7 @@ export default function RetrievalOptions({
</div>
{/* Graph extraction group (graphrag only; re-ingest required to apply) */}
{isGraphRAG && (
{isGraphRAG && !queryOnly && (
<div className="flex flex-col gap-3">
<GroupHeader title={tr('graph.title')} tag={tr('graph.tag')} />
@@ -676,7 +689,7 @@ export default function RetrievalOptions({
)}
{/* Chunking group (re-ingest required) */}
<div className="flex flex-col gap-3">
<div className={cn('flex flex-col gap-3', queryOnly && 'hidden')}>
<GroupHeader title={tr('chunking.title')} tag={tr('chunking.tag')} />
<div className="divide-border/50 divide-y">
@@ -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
+33
View File
@@ -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
+5 -5
View File
@@ -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"
+206
View File
@@ -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
+31
View File
@@ -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") == []
+64
View File
@@ -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"}
+37
View File
@@ -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)
+35
View File
@@ -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") == []