benroshan commited on
Commit
09d72bf
·
1 Parent(s): 0155d2c

feat: add filter_docs to HybridRetriever and get_retriever_filtered helper

Browse files
server/retriever.py CHANGED
@@ -57,6 +57,7 @@ class HybridRetriever(BaseRetriever):
57
  workspace_id: str = "default"
58
  use_hyde: bool = False
59
  use_multi_query: bool = False
 
60
 
61
  class Config:
62
  arbitrary_types_allowed = True
@@ -121,7 +122,10 @@ class HybridRetriever(BaseRetriever):
121
  return query
122
 
123
  def _dense_retrieve(self, query: str, k: int) -> list[dict]:
124
- results = self.vectorstore.similarity_search_with_relevance_scores(query, k=k)
 
 
 
125
  output = []
126
  for doc, score in results:
127
  output.append({
@@ -173,7 +177,10 @@ class HybridRetriever(BaseRetriever):
173
  for q in queries:
174
  dense_query = self._hyde_expand(q) if self.use_hyde else q
175
  d_results = self._dense_retrieve(dense_query, k=self.retrieve_k)
176
- s_results = get_index(self.workspace_id).search(q, k=self.retrieve_k)
 
 
 
177
 
178
  for rank, doc in enumerate(d_results):
179
  key = doc["content"][:120]
@@ -226,6 +233,23 @@ def get_retriever(workspace_id: str = DEFAULT_WORKSPACE) -> HybridRetriever:
226
  return _retriever_cache[workspace_id]
227
 
228
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
229
  def retrieve_with_scores(query: str, k: int = 5) -> list[dict]:
230
  """Compatibility shim for precision eval — returns top-k chunks as dicts."""
231
  retriever = get_retriever()
 
57
  workspace_id: str = "default"
58
  use_hyde: bool = False
59
  use_multi_query: bool = False
60
+ filter_docs: list[str] | None = None
61
 
62
  class Config:
63
  arbitrary_types_allowed = True
 
122
  return query
123
 
124
  def _dense_retrieve(self, query: str, k: int) -> list[dict]:
125
+ filter_arg = {"source": {"$in": self.filter_docs}} if self.filter_docs else None
126
+ results = self.vectorstore.similarity_search_with_relevance_scores(
127
+ query, k=k, filter=filter_arg
128
+ )
129
  output = []
130
  for doc, score in results:
131
  output.append({
 
177
  for q in queries:
178
  dense_query = self._hyde_expand(q) if self.use_hyde else q
179
  d_results = self._dense_retrieve(dense_query, k=self.retrieve_k)
180
+ filter_sources = set(self.filter_docs) if self.filter_docs else None
181
+ s_results = get_index(self.workspace_id).search(
182
+ q, k=self.retrieve_k, filter_sources=filter_sources
183
+ )
184
 
185
  for rank, doc in enumerate(d_results):
186
  key = doc["content"][:120]
 
233
  return _retriever_cache[workspace_id]
234
 
235
 
236
+ def get_retriever_filtered(workspace_id: str, filter_docs: list[str]) -> HybridRetriever:
237
+ """One-off retriever with doc filter applied. Reuses cached vectorstore; does NOT cache itself."""
238
+ config = load_config()
239
+ retrieval_cfg = config.get("retrieval", {})
240
+ return HybridRetriever(
241
+ vectorstore=get_vectorstore(workspace_id),
242
+ dense_weight=retrieval_cfg.get("dense_weight", 0.7),
243
+ sparse_weight=retrieval_cfg.get("sparse_weight", 0.3),
244
+ retrieve_k=retrieval_cfg.get("retrieve_k", 10),
245
+ rerank_k=retrieval_cfg.get("rerank_k", 5),
246
+ workspace_id=workspace_id,
247
+ use_hyde=retrieval_cfg.get("hyde_enabled", False),
248
+ use_multi_query=retrieval_cfg.get("multi_query_enabled", False),
249
+ filter_docs=filter_docs,
250
+ )
251
+
252
+
253
  def retrieve_with_scores(query: str, k: int = 5) -> list[dict]:
254
  """Compatibility shim for precision eval — returns top-k chunks as dicts."""
255
  retriever = get_retriever()
tests/test_retriever_filter.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from unittest.mock import MagicMock, patch
2
+ from langchain_core.documents import Document
3
+ from server.retriever import HybridRetriever, get_retriever_filtered
4
+
5
+
6
+ def _make_retriever(filter_docs=None):
7
+ mock_vs = MagicMock()
8
+ mock_vs.similarity_search_with_relevance_scores.return_value = [
9
+ (Document(page_content="UPI grew in 2024", metadata={"source": "rbi.pdf", "page": 1, "chunk_index": 0}), 0.85),
10
+ (Document(page_content="NPCI payments record", metadata={"source": "npci.pdf", "page": 1, "chunk_index": 0}), 0.72),
11
+ ]
12
+ return HybridRetriever(
13
+ vectorstore=mock_vs,
14
+ dense_weight=0.7,
15
+ sparse_weight=0.3,
16
+ retrieve_k=10,
17
+ rerank_k=5,
18
+ workspace_id="test",
19
+ use_hyde=False,
20
+ use_multi_query=False,
21
+ filter_docs=filter_docs,
22
+ )
23
+
24
+
25
+ def test_filter_docs_field_defaults_to_none():
26
+ retriever = _make_retriever()
27
+ assert retriever.filter_docs is None
28
+
29
+
30
+ def test_filter_docs_field_set():
31
+ retriever = _make_retriever(filter_docs=["rbi.pdf"])
32
+ assert retriever.filter_docs == ["rbi.pdf"]
33
+
34
+
35
+ def test_dense_retrieve_passes_filter_to_chroma():
36
+ retriever = _make_retriever(filter_docs=["rbi.pdf"])
37
+ retriever._dense_retrieve("UPI", k=5)
38
+ call_kwargs = retriever.vectorstore.similarity_search_with_relevance_scores.call_args
39
+ assert call_kwargs.kwargs.get("filter") == {"source": {"$in": ["rbi.pdf"]}}
40
+
41
+
42
+ def test_dense_retrieve_no_filter_when_none():
43
+ retriever = _make_retriever(filter_docs=None)
44
+ retriever._dense_retrieve("UPI", k=5)
45
+ call_kwargs = retriever.vectorstore.similarity_search_with_relevance_scores.call_args
46
+ assert call_kwargs.kwargs.get("filter") is None
47
+
48
+
49
+ def test_get_retriever_filtered_returns_retriever_with_filter():
50
+ mock_vs = MagicMock()
51
+ with patch("server.retriever.get_vectorstore", return_value=mock_vs), \
52
+ patch("server.retriever.load_config", return_value={
53
+ "retrieval": {"dense_weight": 0.7, "sparse_weight": 0.3,
54
+ "retrieve_k": 10, "rerank_k": 5,
55
+ "hyde_enabled": False, "multi_query_enabled": False}
56
+ }):
57
+ r = get_retriever_filtered("default", ["rbi.pdf"])
58
+ assert r.filter_docs == ["rbi.pdf"]
59
+ assert r.vectorstore is mock_vs
60
+
61
+
62
+ def test_get_retriever_filtered_does_not_modify_singleton_cache():
63
+ from server.retriever import _retriever_cache
64
+ initial_keys = set(_retriever_cache.keys())
65
+ mock_vs = MagicMock()
66
+ with patch("server.retriever.get_vectorstore", return_value=mock_vs), \
67
+ patch("server.retriever.load_config", return_value={
68
+ "retrieval": {"dense_weight": 0.7, "sparse_weight": 0.3,
69
+ "retrieve_k": 10, "rerank_k": 5,
70
+ "hyde_enabled": False, "multi_query_enabled": False}
71
+ }):
72
+ get_retriever_filtered("workspace-filtered", ["doc.pdf"])
73
+ assert "workspace-filtered" not in _retriever_cache
74
+ assert set(_retriever_cache.keys()) == initial_keys