mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-03 09:12:55 +00:00
Chunks preview
This commit is contained in:
1 parent
73c3dfb5c4
commit
0447eb9b8d
33 files changed
+1961
-109
No files matched your search
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
@@ -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
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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`,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 ? (
|
||||
<>
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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}}\"?",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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") == []
|
||||
@@ -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"}
|
||||
@@ -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)
|
||||
@@ -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") == []
|
||||
Reference in new issue
Block a user