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>
- app/graph_rag.py +44 -9
- tests/test_graph_rag.py +161 -0
|
@@ -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) +
|
| 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 |
-
"""
|
| 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 |
-
|
| 298 |
-
for _, row in reports.head(
|
| 299 |
content = str(row.get("full_content") or row.get("summary") or "")
|
| 300 |
if not content:
|
| 301 |
continue
|
| 302 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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
|
|
@@ -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()
|