omarsol Claude Fable 5 commited on
Commit
72f3c77
·
1 Parent(s): 0f36b6a

fix(graphrag): query-relevant community reports, rerank floor parity

Browse files

_community_context accepted the query but never used it: every search
injected the same global top-k community reports sorted by static rank,
contradicting the module's "relevant community reports" contract
(unfinished wiring; the scaffold commit documents the data model but
never connected it). Select reports by reranking the top-30
static-ranked candidates against the query with the existing Cohere
client (one bounded rerank call per query); scores stay pinned at 0.0
so reports remain context-only.

Also applies the classical retriever's RERANK_SCORE_FLOOR in
_apply_token_budget (exempting the score-0.0 community reports) per the
module's fairness contract, and skips oversized results instead of
breaking, matching the classical fix.

Note: a future re-run of the graphrag eval arm will not be directly
comparable to the recorded F28 numbers, which measured the
static-report behavior.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

Files changed (2) hide show
  1. app/graph_rag.py +44 -9
  2. tests/test_graph_rag.py +161 -0
app/graph_rag.py CHANGED
@@ -40,6 +40,7 @@ from typing import Any
40
  from .chroma_rag import (
41
  DEFAULT_CONTEXT_TOKEN_BUDGET,
42
  DEFAULT_RERANK_MODEL,
 
43
  SearchResult,
44
  get_token_encoding,
45
  rerank_results,
@@ -63,6 +64,10 @@ GRAPHRAG_COMMUNITY_SOURCE = "graphrag_community"
63
  DEFAULT_ENTITY_TOP_K = 20
64
  DEFAULT_COMMUNITY_TOP_K = 3
65
  DEFAULT_RERANK_TOP_K = 5
 
 
 
 
66
 
67
 
68
  class GraphRAGIndexNotBuilt(RuntimeError):
@@ -211,8 +216,8 @@ class GraphRAGRetriever:
211
  """Assemble GraphRAG context as reranked SearchResult chunks.
212
 
213
  Steps: embed query -> nearest entities (LanceDB) -> their text units
214
- (mapped to real source/url via document_id) + top community reports ->
215
- Cohere rerank -> token budget. Returns [] on any failure so the agent
216
  degrades to KB browsing, matching the classical backend's contract.
217
  """
218
  try:
@@ -287,19 +292,27 @@ class GraphRAGRetriever:
287
  )
288
 
289
  def _community_context(self, query: str) -> list[SearchResult]:
290
- """Top community reports as context-only chunks (no source -> no recall
291
- credit, by design)."""
 
 
 
 
 
 
 
 
292
  if self._reports is None or self._community_top_k <= 0:
293
  return []
294
  reports = self._reports
295
  if "rank" in reports.columns:
296
  reports = reports.sort_values("rank", ascending=False)
297
- results: list[SearchResult] = []
298
- for _, row in reports.head(self._community_top_k).iterrows():
299
  content = str(row.get("full_content") or row.get("summary") or "")
300
  if not content:
301
  continue
302
- results.append(
303
  SearchResult(
304
  chunk_id=f"community:{row.get('community', '')}",
305
  doc_id="",
@@ -315,7 +328,18 @@ class GraphRAGRetriever:
315
  retrieval_method="graphrag_community",
316
  )
317
  )
318
- return results
 
 
 
 
 
 
 
 
 
 
 
319
 
320
  def _apply_token_budget(
321
  self, results: list[SearchResult], token_budget: int | None
@@ -324,11 +348,22 @@ class GraphRAGRetriever:
324
  filtered: list[SearchResult] = []
325
  total = 0
326
  for result in results:
 
 
 
 
 
 
 
 
 
327
  result_tokens = len(
328
  self._encoding.encode(result.content, disallowed_special=())
329
  )
330
  if total + result_tokens > budget:
331
- break
 
 
332
  total += result_tokens
333
  filtered.append(result)
334
  return filtered
 
40
  from .chroma_rag import (
41
  DEFAULT_CONTEXT_TOKEN_BUDGET,
42
  DEFAULT_RERANK_MODEL,
43
+ RERANK_SCORE_FLOOR,
44
  SearchResult,
45
  get_token_encoding,
46
  rerank_results,
 
64
  DEFAULT_ENTITY_TOP_K = 20
65
  DEFAULT_COMMUNITY_TOP_K = 3
66
  DEFAULT_RERANK_TOP_K = 5
67
+ # Candidate pool for query-relevant community-report selection: only the top
68
+ # reports by the index's static rank are reranked against the query, so the
69
+ # per-query cost stays bounded (one rerank call over at most this many docs).
70
+ COMMUNITY_RERANK_CANDIDATES = 30
71
 
72
 
73
  class GraphRAGIndexNotBuilt(RuntimeError):
 
216
  """Assemble GraphRAG context as reranked SearchResult chunks.
217
 
218
  Steps: embed query -> nearest entities (LanceDB) -> their text units
219
+ (mapped to real source/url via document_id) + query-relevant community
220
+ reports -> Cohere rerank -> token budget. Returns [] on any failure so the agent
221
  degrades to KB browsing, matching the classical backend's contract.
222
  """
223
  try:
 
292
  )
293
 
294
  def _community_context(self, query: str) -> list[SearchResult]:
295
+ """Query-relevant community reports as context-only chunks (synthetic
296
+ source, no url -> no recall/citation credit, by design).
297
+
298
+ Candidates are the ``COMMUNITY_RERANK_CANDIDATES`` highest-ranked
299
+ reports (the index's static rank); the ``community_top_k`` most relevant
300
+ to the query are selected with the same Cohere rerank used for text
301
+ units. Scores are pinned back to 0.0 so the reports stay context-only
302
+ chunks (and stay exempt from the rerank-score floor in
303
+ ``_apply_token_budget``).
304
+ """
305
  if self._reports is None or self._community_top_k <= 0:
306
  return []
307
  reports = self._reports
308
  if "rank" in reports.columns:
309
  reports = reports.sort_values("rank", ascending=False)
310
+ candidates: list[SearchResult] = []
311
+ for _, row in reports.head(COMMUNITY_RERANK_CANDIDATES).iterrows():
312
  content = str(row.get("full_content") or row.get("summary") or "")
313
  if not content:
314
  continue
315
+ candidates.append(
316
  SearchResult(
317
  chunk_id=f"community:{row.get('community', '')}",
318
  doc_id="",
 
328
  retrieval_method="graphrag_community",
329
  )
330
  )
331
+ if len(candidates) <= self._community_top_k:
332
+ return candidates
333
+ selected = rerank_results(
334
+ self._cohere,
335
+ query,
336
+ candidates,
337
+ model=self._rerank_model,
338
+ top_n=self._community_top_k,
339
+ )
340
+ for result in selected:
341
+ result.score = 0.0
342
+ return selected
343
 
344
  def _apply_token_budget(
345
  self, results: list[SearchResult], token_budget: int | None
 
348
  filtered: list[SearchResult] = []
349
  total = 0
350
  for result in results:
351
+ # Same low-relevance floor as the classical retriever, so both eval
352
+ # arms drop weak reranked hits (the module's fairness contract).
353
+ # Community reports are context-only chunks pinned at score 0.0 and
354
+ # are exempt, or the floor would drop them all.
355
+ if (
356
+ result.source != GRAPHRAG_COMMUNITY_SOURCE
357
+ and result.score < RERANK_SCORE_FLOOR
358
+ ):
359
+ continue
360
  result_tokens = len(
361
  self._encoding.encode(result.content, disallowed_special=())
362
  )
363
  if total + result_tokens > budget:
364
+ # Skip an oversized result instead of cutting off the whole
365
+ # list; keep filling the budget with later results that fit.
366
+ continue
367
  total += result_tokens
368
  filtered.append(result)
369
  return filtered
tests/test_graph_rag.py CHANGED
@@ -14,7 +14,9 @@ pytest.importorskip("graphrag", reason="install the optional 'graphrag' extra to
14
  import pandas as pd
15
 
16
  from app.chat_types import ChatRequest
 
17
  from app.graph_rag import (
 
18
  GRAPHRAG_COMMUNITY_SOURCE,
19
  GraphRAGIndexNotBuilt,
20
  GraphRAGRetriever,
@@ -22,6 +24,40 @@ from app.graph_rag import (
22
  )
23
 
24
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
25
  class _FakeSearch:
26
  def __init__(self, df: pd.DataFrame) -> None:
27
  self._df = df
@@ -120,6 +156,8 @@ class GraphRagMappingTestCase(unittest.TestCase):
120
  self.assertEqual({res.source for res in only_a}, {"src_a"})
121
 
122
  def test_community_reports_are_context_only(self) -> None:
 
 
123
  r = _bare_retriever()
124
  r._community_top_k = 2
125
  r._reports = pd.DataFrame(
@@ -137,9 +175,132 @@ class GraphRagMappingTestCase(unittest.TestCase):
137
  # no url -> never resolves as a cited source card.
138
  self.assertEqual(res.source, GRAPHRAG_COMMUNITY_SOURCE)
139
  self.assertEqual(res.url, "")
 
140
  # Sorted by rank desc: highest-rank community first.
141
  self.assertEqual(results[0].title, "C1")
142
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
143
 
144
  if __name__ == "__main__":
145
  unittest.main()
 
14
  import pandas as pd
15
 
16
  from app.chat_types import ChatRequest
17
+ from app.chroma_rag import SearchResult, get_token_encoding
18
  from app.graph_rag import (
19
+ COMMUNITY_RERANK_CANDIDATES,
20
  GRAPHRAG_COMMUNITY_SOURCE,
21
  GraphRAGIndexNotBuilt,
22
  GraphRAGRetriever,
 
24
  )
25
 
26
 
27
+ class _FakeRerankItem:
28
+ def __init__(self, index: int, score: float) -> None:
29
+ self.index = index
30
+ self.relevance_score = score
31
+
32
+
33
+ class _FakeRerankResponse:
34
+ def __init__(self, items: list[_FakeRerankItem]) -> None:
35
+ self.results = items
36
+
37
+
38
+ class _FakeCohere:
39
+ """Captures rerank calls; ranks documents containing the query text first."""
40
+
41
+ def __init__(self) -> None:
42
+ self.calls: list[dict] = []
43
+
44
+ def rerank(self, *, model, query, documents, top_n): # type: ignore[no-untyped-def]
45
+ self.calls.append(
46
+ {"model": model, "query": query, "documents": list(documents)}
47
+ )
48
+ order = sorted(
49
+ range(len(documents)),
50
+ key=lambda i: (query.lower() in documents[i].lower(), -i),
51
+ reverse=True,
52
+ )
53
+ return _FakeRerankResponse(
54
+ [
55
+ _FakeRerankItem(doc_index, 1.0 - rank * 0.1)
56
+ for rank, doc_index in enumerate(order[:top_n])
57
+ ]
58
+ )
59
+
60
+
61
  class _FakeSearch:
62
  def __init__(self, df: pd.DataFrame) -> None:
63
  self._df = df
 
156
  self.assertEqual({res.source for res in only_a}, {"src_a"})
157
 
158
  def test_community_reports_are_context_only(self) -> None:
159
+ # With candidates <= community_top_k every report is returned without a
160
+ # rerank call (no _cohere set on the bare retriever proves that).
161
  r = _bare_retriever()
162
  r._community_top_k = 2
163
  r._reports = pd.DataFrame(
 
175
  # no url -> never resolves as a cited source card.
176
  self.assertEqual(res.source, GRAPHRAG_COMMUNITY_SOURCE)
177
  self.assertEqual(res.url, "")
178
+ self.assertEqual(res.score, 0.0)
179
  # Sorted by rank desc: highest-rank community first.
180
  self.assertEqual(results[0].title, "C1")
181
 
182
+ def test_community_reports_selected_by_query_relevance(self) -> None:
183
+ # The query, not the static community rank, decides which reports come
184
+ # back: the top-ranked candidates are reranked against the query.
185
+ r = _bare_retriever()
186
+ r._community_top_k = 1
187
+ r._rerank_model = "fake-rerank"
188
+ fake = _FakeCohere()
189
+ r._cohere = fake
190
+ r._reports = pd.DataFrame(
191
+ {
192
+ "community": [1, 2, 3, 4],
193
+ "title": ["C1", "C2", "C3", "C4"],
194
+ "full_content": [
195
+ "agents and tools",
196
+ "prompt engineering tips",
197
+ "vector databases overview",
198
+ "fine-tuning walkthrough",
199
+ ],
200
+ "rank": [9.0, 8.0, 7.0, 6.0],
201
+ }
202
+ )
203
+
204
+ results = r._community_context("vector databases")
205
+
206
+ # Not C1 (highest static rank): the query-relevant report wins.
207
+ self.assertEqual([res.title for res in results], ["C3"])
208
+ # Every candidate reached the reranker in one call.
209
+ self.assertEqual(len(fake.calls), 1)
210
+ self.assertEqual(len(fake.calls[0]["documents"]), 4)
211
+ # Still context-only: synthetic source, no url, score pinned to 0.0.
212
+ for res in results:
213
+ self.assertEqual(res.source, GRAPHRAG_COMMUNITY_SOURCE)
214
+ self.assertEqual(res.url, "")
215
+ self.assertEqual(res.score, 0.0)
216
+ self.assertEqual(res.retrieval_method, "graphrag_community")
217
+
218
+ def test_community_rerank_candidate_pool_is_bounded(self) -> None:
219
+ # Cost stays bounded: only the top COMMUNITY_RERANK_CANDIDATES reports
220
+ # by static rank are sent to the reranker.
221
+ r = _bare_retriever()
222
+ r._community_top_k = 1
223
+ r._rerank_model = "fake-rerank"
224
+ fake = _FakeCohere()
225
+ r._cohere = fake
226
+ total = COMMUNITY_RERANK_CANDIDATES + 5
227
+ r._reports = pd.DataFrame(
228
+ {
229
+ "community": list(range(total)),
230
+ "title": [f"C{i}" for i in range(total)],
231
+ "full_content": [f"report {i}" for i in range(total)],
232
+ "rank": [float(total - i) for i in range(total)],
233
+ }
234
+ )
235
+
236
+ r._community_context("q")
237
+
238
+ documents = fake.calls[0]["documents"]
239
+ self.assertEqual(len(documents), COMMUNITY_RERANK_CANDIDATES)
240
+ self.assertIn("report 0", documents) # highest static rank kept
241
+ self.assertNotIn(f"report {total - 1}", documents) # lowest dropped
242
+
243
+
244
+ class GraphRagTokenBudgetTestCase(unittest.TestCase):
245
+ def _retriever(self) -> GraphRAGRetriever:
246
+ r = _bare_retriever()
247
+ r._encoding = get_token_encoding(None)
248
+ r._token_budget = 100_000
249
+ return r
250
+
251
+ def _result(
252
+ self,
253
+ chunk_id: str,
254
+ score: float,
255
+ *,
256
+ content: str = "x",
257
+ source: str = "src_a",
258
+ method: str = "graphrag",
259
+ ) -> SearchResult:
260
+ return SearchResult(
261
+ chunk_id=chunk_id,
262
+ doc_id=chunk_id,
263
+ title=chunk_id,
264
+ url="",
265
+ source=source,
266
+ retrieve_doc=False,
267
+ tokens=10,
268
+ score=score,
269
+ content=content,
270
+ chunk_content=content,
271
+ heading_path="",
272
+ retrieval_method=method,
273
+ )
274
+
275
+ def test_budget_skips_oversized_result_and_fills_with_smaller(self) -> None:
276
+ # Same regression as the classical retriever: an oversized rank-1
277
+ # result under a small per-request budget must not empty the list.
278
+ r = self._retriever()
279
+ oversized = self._result("big", 0.9, content="word " * 500)
280
+ small = self._result("small", 0.8, content="word " * 15)
281
+
282
+ kept = r._apply_token_budget([oversized, small], token_budget=50)
283
+
284
+ self.assertEqual([res.chunk_id for res in kept], ["small"])
285
+
286
+ def test_low_score_floor_exempts_community_reports(self) -> None:
287
+ # Fairness with the classical arm: weak reranked text units are dropped
288
+ # by the same score floor, but community reports (context-only chunks
289
+ # pinned at score 0.0) are exempt.
290
+ r = self._retriever()
291
+ strong = self._result("strong", 0.5)
292
+ weak = self._result("weak", 0.05)
293
+ community = self._result(
294
+ "community:1",
295
+ 0.0,
296
+ source=GRAPHRAG_COMMUNITY_SOURCE,
297
+ method="graphrag_community",
298
+ )
299
+
300
+ kept = r._apply_token_budget([strong, weak, community], token_budget=None)
301
+
302
+ self.assertEqual([res.chunk_id for res in kept], ["strong", "community:1"])
303
+
304
 
305
  if __name__ == "__main__":
306
  unittest.main()