From 94a33c478122894ad3bf2e8d62ff8f50a34c264f Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:38:54 +0100 Subject: [PATCH 01/27] fix(graphrag): rank without scipy so graph retrieval stops falling back MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit networkx.pagerank delegates to a scipy implementation, and scipy is not a DocsGPT dependency — it only arrives transitively through the optional docling extra. In a default install every graph retrieval raised ModuleNotFoundError inside _ppr_scores, hit the per-source except, and degraded to ClassicRAG: the graph was built and paid for, then never used, with one ERROR line per source per query as the only signal. Rank with a local power iteration over the same row-normalized transition matrix: undirected edges normalized per endpoint, dangling nodes redistributed along the restart vector, and the restart vector normalized across the nodes the subgraph actually holds so seed mass cannot leak. Parity with networkx is asserted while scipy happens to be installed in the test env, and the retrieval path is exercised with the import blocked. --- docsgpt/retriever/graph_rag.py | 109 ++++++++++++++++++++++++-- tests/retriever/test_graph_rag.py | 122 ++++++++++++++++++++++++++++++ 2 files changed, 226 insertions(+), 5 deletions(-) diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index c7ac538e..5593a9b2 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -1,9 +1,13 @@ """GraphRAG local retriever — Personalized PageRank over a per-source graph. -Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx -Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR -mass -> shared token budget. No LLM call at query time beyond the (optional, -reused) rephrase. +Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> Personalized +PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR mass -> shared +token budget. No LLM call at query time beyond the (optional, reused) rephrase. + +``networkx`` supplies the graph structure, but the ranking is the local power +iteration in :func:`_personalized_pagerank`: ``nx.pagerank`` delegates to scipy, +which DocsGPT does not depend on, so calling it turned every graph retrieval +into a silent ClassicRAG fallback. Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the ``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the @@ -40,6 +44,99 @@ def _idf(doc_freq: Any) -> float: return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]: + """Normalized restart distribution over ``nodes``. + + Weights are clamped at zero (a cosine distance above 1 yields a negative + seed weight) and normalized across the nodes actually present in the graph, + so restart mass can never leak to a node the subgraph does not contain. An + absent or all-zero personalization collapses to a uniform restart. + """ + if personalization: + weights = { + node: max(float(personalization.get(node, 0.0) or 0.0), 0.0) + for node in nodes + } + total = sum(weights.values()) + if total > 0: + return {node: weight / total for node, weight in weights.items()} + uniform = 1.0 / len(nodes) + return {node: uniform for node in nodes} + + +def _personalized_pagerank( + graph: nx.Graph, + personalization: Dict[Any, float] | None = None, + *, + weight: str = "weight", + alpha: float = 0.85, + max_iter: int = 100, + tol: float = 1.0e-6, +) -> Dict[Any, float]: + """Personalized PageRank by power iteration — no scipy. + + ``networkx.pagerank`` delegates to a scipy implementation, and scipy is not + a DocsGPT dependency: in a default install the import raises and every graph + retrieval silently degrades to the ClassicRAG fallback. This is the same + algorithm over the same row-normalized transition matrix, so the ranking is + unchanged where scipy happens to be installed. + + Args: + graph: Undirected graph whose edges may carry a ``weight`` attribute. + personalization: Node -> restart weight; ``None`` means uniform. + weight: Edge attribute holding the weight. + alpha: Damping factor. + max_iter: Iteration cap. The last iterate is returned if it is hit — + retrieval degrades to a slightly less converged ranking rather than + raising, which is what the library does. + tol: Convergence tolerance; iteration stops below ``len(graph) * tol``. + + Returns: + Node -> PageRank mass, summing to ~1.0. Empty dict for an empty graph. + """ + nodes = list(graph.nodes) + node_count = len(nodes) + if node_count == 0: + return {} + + restart = _restart_vector(nodes, personalization) + + # Row-normalized transitions. An undirected edge is traversable from both + # endpoints, so each node normalizes over its own incident weights. + transitions: Dict[Any, List[Any]] = {} + for node in nodes: + neighbors = [] + total = 0.0 + for neighbor, data in graph[node].items(): + edge_weight = float(data.get(weight, 1.0) or 1.0) + if edge_weight <= 0: + continue + neighbors.append((neighbor, edge_weight)) + total += edge_weight + transitions[node] = ( + [(n, w / total) for n, w in neighbors] if total > 0 else [] + ) + + # A node with no usable edge is dangling: its mass would vanish each pass, + # so it is redistributed along the restart vector instead. + dangling = [node for node in nodes if not transitions[node]] + + ranks = {node: 1.0 / node_count for node in nodes} + for _ in range(max_iter): + previous = ranks + ranks = dict.fromkeys(nodes, 0.0) + leaked = alpha * sum(previous[node] for node in dangling) + for node in nodes: + share = alpha * previous[node] + for neighbor, transition in transitions[node]: + ranks[neighbor] += share * transition + for node in nodes: + ranks[node] += (leaked + 1.0 - alpha) * restart[node] + if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: + break + return ranks + + class GraphRAGRetriever(BaseRetriever): """Per-source PPR retriever; falls back to ClassicRAG when a source has no graph.""" @@ -117,7 +214,9 @@ class GraphRAGRetriever(BaseRetriever): if not any(personalization.values()): personalization = None - ranks = nx.pagerank(graph, personalization=personalization, weight="weight") + ranks = _personalized_pagerank( + graph, personalization=personalization, weight="weight" + ) return { node: rank * _idf(graph.nodes[node].get("doc_freq", 0)) for node, rank in ranks.items() diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index e3fe2fae..5f8cdf80 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -800,3 +800,125 @@ class TestGraphRAGBatching: assert rag._get_data() == [] mock_store_cls.assert_not_called() rag._classic._get_data.assert_not_called() + + +# ── Personalized PageRank without scipy ────────────────────────────────────── + + +@pytest.fixture +def _no_scipy(monkeypatch): + """Make ``import scipy`` fail, as it does in a default install. + + ``scipy`` is not a DocsGPT dependency — it only reaches this test env + through the optional docling extra. ``networkx.pagerank`` delegates to its + scipy implementation, so ranking must not go through it. + """ + import sys + + for name in [m for m in list(sys.modules) if m == "scipy" or m.startswith("scipy.")]: + monkeypatch.delitem(sys.modules, name) + monkeypatch.setitem(sys.modules, "scipy", None) + + +def _chain_graph(): + """Weighted chain a-b-c-d plus a heavier shortcut a-d.""" + import networkx as nx + + graph = nx.Graph() + graph.add_weighted_edges_from( + [("a", "b", 1.0), ("b", "c", 2.0), ("c", "d", 1.0), ("a", "d", 0.5)] + ) + return graph + + +@pytest.mark.unit +class TestPersonalizedPageRankWithoutScipy: + def test_ranking_runs_when_scipy_is_missing(self, _no_scipy): + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + ranks = _personalized_pagerank( + graph, personalization={"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0} + ) + + assert set(ranks) == {"a", "b", "c", "d"} + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + assert all(rank > 0 for rank in ranks.values()) + # Pinned from the parity test below, which runs the library + # implementation over the same graph while scipy is installed here. + assert sorted(ranks, key=ranks.get, reverse=True) == ["b", "a", "c", "d"] + # The seed outranks the node furthest from it along the heavy path. + assert ranks["a"] > ranks["d"] + + def test_matches_networkx_within_tolerance(self): + """Parity with the library implementation, while it is installed here.""" + import networkx as nx + + pytest.importorskip("scipy") + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + personalization = {"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0} + + ours = _personalized_pagerank(graph, personalization=personalization) + theirs = nx.pagerank(graph, personalization=personalization, weight="weight") + + for node in theirs: + assert ours[node] == pytest.approx(theirs[node], abs=1e-6) + + def test_uniform_personalization_when_none(self): + pytest.importorskip("scipy") + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + ours = _personalized_pagerank(graph, personalization=None) + theirs = nx.pagerank(graph, personalization=None, weight="weight") + + for node in theirs: + assert ours[node] == pytest.approx(theirs[node], abs=1e-6) + + def test_isolated_node_still_gets_mass(self): + """A node with no edges is dangling; its mass must not vanish.""" + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = nx.Graph() + graph.add_edge("a", "b", weight=1.0) + graph.add_node("lonely") + + ranks = _personalized_pagerank(graph, personalization=None) + + assert ranks["lonely"] > 0 + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + + def test_empty_graph_returns_empty(self): + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + assert _personalized_pagerank(nx.Graph(), personalization=None) == {} + + @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) + @patch("docsgpt.retriever.graph_rag.GraphStore") + @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) + def test_graph_retrieval_does_not_fall_back_without_scipy( + self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed, _no_scipy + ): + """The whole PPR path runs with scipy absent — no ClassicRAG fallback.""" + nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}] + edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}] + node_chunks = {"n1": ["c1"], "n2": ["c2"]} + chunk_texts = {"c1": "near", "c2": "far"} + seed_rows = [{"id": "n1", "distance": 0.0}] + store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows) + mock_store_cls.return_value = store + + rag = _make_retriever(chunks=2) + rag._classic_for_sources = Mock(side_effect=AssertionError("fell back")) + + docs = rag._get_data() + + assert [doc["text"] for doc in docs] == ["near", "far"] From 92e19ac177ef5bf170b8e63ecea5f6e0e4174f1d Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:40:20 +0100 Subject: [PATCH 02/27] fix(graphrag): dispatch extraction to the provider that serves the model MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _build_extraction_llm passed settings.LLM_PROVIDER with the resolved extraction model id. Those two disagree in any deployment that leaves the provider at its default: the model id comes from GRAPHRAG_EXTRACTION_MODEL or LLM_NAME, while the provider stays "docsgpt" — the hosted public endpoint, which does not serve it. The request is rejected, the shared fallback answers instead, and the graph gets built by a different model than the one configured, with nothing in the summary saying so. Resolve the provider from the model registry (owner-scoped, so per-user BYOM ids resolve too), fall back to settings.LLM_PROVIDER only when the model is unknown, and take the API key for the provider actually dispatched to rather than the generic settings.API_KEY. The effective provider is logged, since a silent swap was the whole failure mode. --- docsgpt/graphrag/extraction.py | 31 +++++++- tests/graphrag/test_extraction.py | 118 ++++++++++++++++++++++++++++++ 2 files changed, 147 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index beaf3d66..aa76189f 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -23,6 +23,10 @@ import logging import re from typing import Any, Callable, Dict, List, Optional +from docsgpt.core.model_utils import ( + get_api_key_for_provider, + get_provider_from_model_id, +) from docsgpt.core.settings import settings from docsgpt.llm.llm_creator import LLMCreator from docsgpt.storage.db.source_config import SourceConfig @@ -63,14 +67,37 @@ def _resolve_max_chunks(config: SourceConfig) -> int: return config.graph.max_chunks or settings.GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION +def _resolve_extraction_provider( + model_id: Optional[str], user: Optional[str] +) -> str: + """The provider that serves ``model_id``, else the deployment default. + + ``settings.LLM_PROVIDER`` is only a default (``docsgpt``, the hosted public + endpoint, out of the box). Dispatching the resolved extraction model + through it sends the request to a provider that does not serve that model: + the call is rejected, the shared fallback answers instead, and the graph is + built by a different model than the one configured — with nothing in the + summary to say so. ``user`` scopes the lookup so a per-user (BYOM) model id + resolves as well. + """ + provider = ( + get_provider_from_model_id(model_id, user_id=user) if model_id else None + ) + return provider or settings.LLM_PROVIDER + + def _build_extraction_llm( model_id: Optional[str], user: Optional[str], request_id: Optional[str] ): """Build the extraction LLM tagged for token-usage attribution to the owner.""" decoded_token = {"sub": user} if user else None + provider = _resolve_extraction_provider(model_id, user) + logger.info( + "Graph extraction dispatching model=%s via provider=%s", model_id, provider + ) llm = LLMCreator.create_llm( - settings.LLM_PROVIDER, - api_key=settings.API_KEY, + provider, + api_key=get_api_key_for_provider(provider), user_api_key=None, decoded_token=decoded_token, model_id=model_id, diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index ad2d670f..eaaf881d 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -403,6 +403,124 @@ class TestModelResolution: assert extraction_module._resolve_max_chunks(config) == 5 +@pytest.mark.unit +class TestExtractionProviderResolution: + """The extraction model decides the provider, not ``LLM_PROVIDER``. + + ``settings.LLM_PROVIDER`` is the deployment default (``docsgpt`` out of the + box, i.e. the hosted public endpoint). Dispatching the resolved extraction + model through it sends the call to a provider that never serves that model: + the request is rejected, the shared fallback answers instead, and the graph + is quietly built by a different model than the one configured. + """ + + def _capture_create_llm(self, monkeypatch, llm=None): + captured = {} + + def _create(provider, *args, **kwargs): + captured["provider"] = provider + captured["args"] = args + captured["kwargs"] = kwargs + return llm or _StubLLM([]) + + monkeypatch.setattr( + extraction_module.LLMCreator, "create_llm", staticmethod(_create) + ) + return captured + + def test_provider_comes_from_the_model_registry(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: "openai" + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "sk-openai" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("gpt-4o-mini", "owner-1", "req-1") + + assert captured["provider"] == "openai" + assert captured["kwargs"]["api_key"] == "sk-openai" + assert captured["kwargs"]["model_id"] == "gpt-4o-mini" + + def test_owner_scopes_the_registry_lookup(self, monkeypatch): + """A per-user (BYOM) model only resolves when the owner is passed.""" + seen = {} + + def _resolve(model_id, user_id=None): + seen["model_id"] = model_id + seen["user_id"] = user_id + return "anthropic" + + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", _resolve + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "k" + ) + self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("byom-uuid", "owner-7", "req-1") + + assert seen == {"model_id": "byom-uuid", "user_id": "owner-7"} + + def test_unknown_model_falls_back_to_the_configured_provider(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: None + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "fallback-key" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("mystery-model", "owner-1", "req-1") + + assert captured["provider"] == "docsgpt" + assert captured["kwargs"]["api_key"] == "fallback-key" + + def test_no_model_id_skips_the_lookup(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "openai") + calls = [] + monkeypatch.setattr( + extraction_module, + "get_provider_from_model_id", + lambda *a, **k: calls.append(a) or "anthropic", + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "k" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm(None, "owner-1", "req-1") + + assert calls == [] + assert captured["provider"] == "openai" + + def test_api_key_follows_the_resolved_provider(self, monkeypatch): + """The key must match the provider actually dispatched to.""" + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr(extraction_module.settings, "API_KEY", "generic-key") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: "anthropic" + ) + keyed_for = {} + + def _key(provider): + keyed_for["provider"] = provider + return "sk-anthropic" + + monkeypatch.setattr(extraction_module, "get_api_key_for_provider", _key) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("claude-x", "owner-1", "req-1") + + assert keyed_for["provider"] == "anthropic" + assert captured["kwargs"]["api_key"] == "sk-anthropic" + assert captured["kwargs"]["api_key"] != "generic-key" + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): From e3d819d9fdbbcd3f9eaffbec2b2a8cce5a822eac Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:44:23 +0100 Subject: [PATCH 03/27] fix(graphrag): stop losing chunks silently during a graph build MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three ways a chunk disappeared from a graph with no way to tell: A build checks out one connection and then spends minutes per chunk waiting on the model, so the connection idles long enough for the server or a pooler to drop it. The pool only validates a connection when it hands one out, and this one was handed out at the start of the build, so the next write raised "the connection is lost", the chunk was marked failed, and the build carried on a chunk short. apply_chunk and mark_chunk now reconnect and retry once; every statement they run is an idempotent upsert, so a replay cannot double-write. Only connection loss retries — a bad statement still surfaces. An unparseable model response marked the chunk failed and logged nothing at all, so failed_chunks was the only evidence and it named no chunk. Both failure modes now log the chunk id. The summary's node count summed per-chunk upserts, so an entity appearing in ten chunks counted ten times: it reported writes, not graph size. It now reports the distinct node count, falling back to the write count only if the count query fails. --- docsgpt/graphrag/extraction.py | 49 ++++++-- docsgpt/graphrag/store.py | 188 ++++++++++++++++++++---------- tests/graphrag/test_extraction.py | 103 ++++++++++++++++ tests/graphrag/test_store.py | 112 ++++++++++++++++++ 4 files changed, 379 insertions(+), 73 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index aa76189f..580b54cb 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -149,8 +149,15 @@ def _parse_extraction(raw: Any) -> Optional[Dict[str, List[Dict[str, Any]]]]: } -def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: - """Run exactly one extraction call for a chunk (gleanings off).""" +def _extract_chunk( + llm, text: str, chunk_id: Optional[str] = None +) -> Optional[Dict[str, List[Dict[str, Any]]]]: + """Run exactly one extraction call for a chunk (gleanings off). + + Both failure modes name the chunk: an unparseable response used to return + ``None`` silently, so a graph could come back short with nothing in the + logs to say which chunk was dropped or why. + """ messages = [ {"role": "system", "content": _SYSTEM_PROMPT}, {"role": "user", "content": f"\n{text}\n"}, @@ -161,9 +168,17 @@ def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: messages=messages, ) except Exception as exc: - logger.warning("Graph extraction call failed, skipping chunk: %s", exc) + logger.warning( + "Graph extraction call failed for chunk %s, skipping: %s", chunk_id, exc + ) return None - return _parse_extraction(response) + parsed = _parse_extraction(response) + if parsed is None: + logger.warning( + "Graph extraction returned unparseable output for chunk %s; marking it failed.", + chunk_id, + ) + return parsed def _coerce_weight(value: Any) -> float: @@ -205,7 +220,9 @@ def extract_graph_for_source( Returns: A summary ``{nodes, edges, chunks_processed, skipped_over_cap, - failed_chunks}``. + failed_chunks}``, where ``nodes`` is how many distinct nodes the + source's graph holds after the run — not how many upserts ran, which + counts the same entity once per chunk it appears in. """ from docsgpt.graphrag.store import GraphStore @@ -228,7 +245,7 @@ def extract_graph_for_source( _resolve_extraction_model(config), user, request_id ) - nodes = 0 + node_upserts = 0 edges = 0 chunks_processed = 0 failed_chunks = 0 @@ -242,7 +259,7 @@ def extract_graph_for_source( { "current": chunks_processed + failed_chunks, "total": total, - "nodes": nodes, + "nodes": node_upserts, "edges": edges, } ) @@ -257,7 +274,7 @@ def extract_graph_for_source( _report() continue - extracted = _extract_chunk(llm, text) + extracted = _extract_chunk(llm, text, chunk_id) if extracted is None: store.mark_chunk(source_id, chunk_id, "failed") failed_chunks += 1 @@ -271,7 +288,7 @@ def extract_graph_for_source( chunk_nodes, chunk_edges = store.apply_chunk( source_id, chunk_id, entities, relationships, name_embeddings ) - nodes += chunk_nodes + node_upserts += chunk_nodes edges += chunk_edges store.mark_chunk(source_id, chunk_id, "done") chunks_processed += 1 @@ -290,6 +307,20 @@ def extract_graph_for_source( except Exception as exc: logger.warning("set_node_degrees failed for source %s: %s", source_id, exc) + # Upserts are writes, not nodes: one entity seen in ten chunks is ten + # upserts and a single node, so the old count overstated every graph whose + # entities recur. Report what the graph holds, falling back to the write + # count only if the count query itself fails. + nodes = node_upserts + try: + nodes = store.count_nodes(source_id) + except Exception as exc: + logger.warning( + "count_nodes failed for source %s; reporting upserts instead: %s", + source_id, + exc, + ) + return { "nodes": nodes, "edges": edges, diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 818004c0..ceb0fc48 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -17,6 +17,7 @@ import logging import uuid from typing import Any, Dict, List, Optional +import psycopg from psycopg.types.json import Jsonb from docsgpt.core.settings import settings @@ -73,6 +74,25 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]: ) +def _is_connection_lost(exc: BaseException) -> bool: + """True when ``exc`` says the server connection went away, not that the SQL was bad. + + psycopg raises ``OperationalError`` ("the connection is lost") when the + socket dies under a statement and ``InterfaceError`` when the connection + object is already closed. Everything else — a bad statement, a constraint + violation — is a real failure that a retry would only repeat. + """ + return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError)) + + +def _safe_rollback(conn) -> None: + """Roll back, tolerating a connection too broken to roll back.""" + try: + conn.rollback() + except Exception as exc: + logging.debug("Rollback on a broken connection failed: %s", exc) + + class GraphStore: """Stores and queries a per-source knowledge graph in the pgvector DB.""" @@ -138,6 +158,37 @@ class GraphStore: self._pooled = False return self._connection + def _write_with_reconnect(self, operation): + """Run ``operation(conn)``, once more on a fresh connection if it was dead. + + A graph build holds one checked-out connection for the length of the + whole extraction and spends minutes per chunk waiting on the model, so + the connection idles long enough for the server (or a pooler) to drop + it. The pool only validates a connection when it is handed out, and + this one was handed out at the start of the build, so the next write + raises and its chunk is lost from the graph. Every statement here is an + idempotent upsert, so replaying one on a new connection cannot + double-write. + + Args: + operation: Callable taking the connection and doing one write. + + Returns: + Whatever ``operation`` returns. + """ + for attempt in (1, 2): + conn = self._get_connection() + try: + return operation(conn) + except Exception as exc: + if attempt == 2 or not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. @@ -499,56 +550,61 @@ class GraphStore: (not linked to the chunk), mirroring the per-call path. ``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees are not bumped here — the caller runs ``set_node_degrees`` once at the - end. Returns ``(nodes_upserted, edges_added)``. + end. Reconnects and retries once if the connection died while the + extraction was waiting on the model. Returns + ``(nodes_upserted, edges_added)``. """ self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - node_ids: Dict[str, str] = {} - edges_added = 0 - try: - for entity in entities: - normalized_name = entity["normalized_name"] - node_id = self._upsert_node( - cursor, - source_id, - entity["name"], - normalized_name, - entity.get("type"), - entity.get("description"), - name_embeddings.get(normalized_name), - ) - node_ids[normalized_name] = node_id - self._link_node_chunk(cursor, source_id, node_id, chunk_id) - for rel in relationships: - src_id = self._resolve_endpoint( - cursor, source_id, rel.get("source"), node_ids, name_embeddings - ) - dst_id = self._resolve_endpoint( - cursor, source_id, rel.get("target"), node_ids, name_embeddings - ) - if src_id is None or dst_id is None: - continue - self._add_edge( - cursor, - source_id, - src_id, - dst_id, - type=rel.get("type"), - description=rel.get("description"), - weight=float(rel.get("weight") or 1.0), - source_chunk_ids=[chunk_id], - ) - edges_added += 1 + def _write(conn): + cursor = conn.cursor() + node_ids: Dict[str, str] = {} + edges_added = 0 + try: + for entity in entities: + normalized_name = entity["normalized_name"] + node_id = self._upsert_node( + cursor, + source_id, + entity["name"], + normalized_name, + entity.get("type"), + entity.get("description"), + name_embeddings.get(normalized_name), + ) + node_ids[normalized_name] = node_id + self._link_node_chunk(cursor, source_id, node_id, chunk_id) - conn.commit() - return len(entities), edges_added - except Exception: - conn.rollback() - raise - finally: - cursor.close() + for rel in relationships: + src_id = self._resolve_endpoint( + cursor, source_id, rel.get("source"), node_ids, name_embeddings + ) + dst_id = self._resolve_endpoint( + cursor, source_id, rel.get("target"), node_ids, name_embeddings + ) + if src_id is None or dst_id is None: + continue + self._add_edge( + cursor, + source_id, + src_id, + dst_id, + type=rel.get("type"), + description=rel.get("description"), + weight=float(rel.get("weight") or 1.0), + source_chunk_ids=[chunk_id], + ) + edges_added += 1 + + conn.commit() + return len(entities), edges_added + except Exception: + _safe_rollback(conn) + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def _resolve_endpoint( self, @@ -1026,25 +1082,29 @@ class GraphStore: cursor.close() def mark_chunk(self, source_id: str, chunk_id: str, status: str): + """Record a chunk's extraction status, reconnecting once if the connection died.""" self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - try: - cursor.execute( - """ - INSERT INTO graph_ingest_progress (source_id, chunk_id, status) - VALUES (%s, %s, %s) - ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; - """, - (source_id, str(chunk_id), status), - ) - conn.commit() - except Exception as e: - conn.rollback() - logging.error(f"Error marking chunk: {e}") - raise - finally: - cursor.close() + + def _write(conn): + cursor = conn.cursor() + try: + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, %s) + ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id), status), + ) + conn.commit() + except Exception as e: + _safe_rollback(conn) + logging.error(f"Error marking chunk: {e}") + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def pending_chunks(self, source_id: str, all_chunk_ids: List[str]) -> List[str]: """Chunk ids from ``all_chunk_ids`` not yet marked ``done`` for the source.""" diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index eaaf881d..b6192450 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -521,6 +521,109 @@ class TestExtractionProviderResolution: assert captured["kwargs"]["api_key"] != "generic-key" +@pytest.mark.unit +class TestFailedChunksAreReported: + """Every dropped chunk has to leave a trace. + + A chunk whose extraction cannot be parsed is marked ``failed`` and skipped. + That path logged nothing at all, so a graph could come back short with the + summary's ``failed_chunks`` count as the only hint and no way to tell which + chunk, or why, from the logs. + """ + + def _fake_store(self, monkeypatch, chunk_ids): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = list(chunk_ids) + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + return store + + def test_unparseable_output_is_logged_with_the_chunk_id( + self, monkeypatch, caplog, stub_embedding + ): + import logging + + store = self._fake_store(monkeypatch, ["c1"]) + _install_stub_llm(monkeypatch, _StubLLM(["not json at all"])) + + with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c1", "some text")], + config=SourceConfig(), + request_id="req-1", + ) + + assert summary["failed_chunks"] == 1 + store.mark_chunk.assert_called_once() + assert store.mark_chunk.call_args.args[2] == "failed" + messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING] + assert any("c1" in message for message in messages), messages + + def test_llm_errors_still_name_the_chunk( + self, monkeypatch, caplog, stub_embedding + ): + import logging + + self._fake_store(monkeypatch, ["c7"]) + _install_stub_llm(monkeypatch, _StubLLM([RuntimeError("model exploded")])) + + with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): + extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c7", "some text")], + config=SourceConfig(), + request_id="req-1", + ) + + messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING] + assert any("c7" in message for message in messages), messages + + +@pytest.mark.integration +class TestSummaryNodeCount: + """``nodes`` must describe the graph, not the number of upserts.""" + + @pytest.fixture + def store(self, monkeypatch, postgresql): + store = _live_store(monkeypatch, postgresql.info) + yield store + store.close() + + def test_repeated_entity_counts_once( + self, store, monkeypatch, stub_embedding + ): + source_id = str(uuid.uuid4()) + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + _install_stub_llm(monkeypatch, _StubLLM([payload, payload])) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[_chunk("c1", "Ada one."), _chunk("c2", "Ada two.")], + config=SourceConfig(), + request_id="req-1", + ) + + # Two chunks upserted the same entity: one node in the graph. + assert store.count_nodes(source_id) == 1 + assert summary["nodes"] == 1 + assert summary["chunks_processed"] == 2 + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index d712e7fe..f2ce6abe 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -890,3 +890,115 @@ class TestCountNodesMany: store, _, _ = self._store_with_mock_conn([(source_id.lower(), 3)]) assert store.count_nodes_many([source_id]) == {source_id: 3} + + +@pytest.mark.unit +class TestWritesSurviveALostConnection: + """A graph build holds one pooled connection across its LLM calls. + + Extraction spends minutes per chunk waiting on a model, so the connection + sits idle between writes and the server (or a pooler) can drop it. The pool + only validates a connection at checkout, and this one was checked out once + at the start of the build, so the next write raises and the chunk is marked + ``failed`` — silently losing it from the graph. The write reconnects and + retries once instead; the statements are idempotent upserts, so a retry + cannot double-write. + """ + + def _store_with_connections(self, conns): + """Store that hands out ``conns`` in order, one per (re)connect.""" + store = GraphStore.__new__(GraphStore) + store._tables_ensured = True + store._connection = None + handed = [] + closed = [] + + def _get_connection(): + if store._connection is None: + store._connection = conns[len(handed)] + handed.append(store._connection) + return store._connection + + def _close(): + if store._connection is not None: + closed.append(store._connection) + store._connection = None + + store._get_connection = _get_connection + store.close = _close + return store, handed, closed + + @staticmethod + def _conn(execute_error=None): + cursor = MagicMock() + cursor.fetchone.return_value = [str(uuid.uuid4())] + cursor.fetchall.return_value = [] + if execute_error is not None: + cursor.execute.side_effect = execute_error + conn = MagicMock() + conn.cursor.return_value = cursor + return conn + + def test_mark_chunk_retries_on_a_dropped_connection(self): + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + alive = self._conn() + store, handed, closed = self._store_with_connections([dead, alive]) + + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [dead, alive] + assert closed == [dead] + alive.commit.assert_called_once() + + def test_apply_chunk_retries_on_a_dropped_connection(self): + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + alive = self._conn() + store, handed, closed = self._store_with_connections([dead, alive]) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + + nodes, edges = store.apply_chunk( + str(uuid.uuid4()), "c1", entities, [], {"ada": _embedding(0.5)} + ) + + assert (nodes, edges) == (1, 0) + assert handed == [dead, alive] + assert closed == [dead] + alive.commit.assert_called_once() + + def test_a_second_connection_failure_is_not_retried_again(self): + """One retry, not a loop: a genuinely unreachable DB still fails.""" + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + also_dead = self._conn(psycopg.OperationalError("the connection is lost")) + store, handed, _ = self._store_with_connections([dead, also_dead]) + + with pytest.raises(psycopg.OperationalError): + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [dead, also_dead] + + def test_a_query_error_is_not_retried(self): + """Only connection loss is retryable; a bad statement must surface.""" + import psycopg + + broken = self._conn(psycopg.ProgrammingError("syntax error")) + spare = self._conn() + store, handed, _ = self._store_with_connections([broken, spare]) + + with pytest.raises(psycopg.ProgrammingError): + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [broken] + broken.rollback.assert_called_once() From 6ea737e57b3fd2bf70ed24225ccadca7c37b125d Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:49:56 +0100 Subject: [PATCH 04/27] fix(tasks): a deferred duplicate stands down instead of failing the task MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A task that runs longer than the broker's visibility timeout is redelivered while its first run is still going. The idempotency lease correctly stops the duplicate from doing the work, but the duplicate then re-queued itself once per LEASE_TTL until celery ran out of retries and raised MaxRetriesExceededError — so a perfectly healthy long task (a large graph extraction is the one that found this) reported a task failure, with a traceback, while the real run was still making progress next to it. Catch the exhaustion and return a "deferred" result instead. A normal deferral still re-queues: only the give-up path changes, and the lease holder's dedup row is left untouched so its own completion still records. --- docsgpt/api/user/idempotency.py | 28 ++++++-- tests/api/user/test_idempotency_decorator.py | 69 ++++++++++++++++++++ 2 files changed, 93 insertions(+), 4 deletions(-) diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py index 1381f241..e38e8a73 100644 --- a/docsgpt/api/user/idempotency.py +++ b/docsgpt/api/user/idempotency.py @@ -9,6 +9,8 @@ import threading import uuid from typing import Any, Callable, Optional +from celery.exceptions import MaxRetriesExceededError + from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository from docsgpt.storage.db.session import db_readonly, db_session @@ -81,10 +83,28 @@ def with_idempotency( "idempotency: live lease held; deferring task=%s key=%s", task_name, key, ) - raise self.retry( - countdown=LEASE_TTL_SECONDS, - max_retries=LEASE_RETRY_MAX, - ) + try: + raise self.retry( + countdown=LEASE_TTL_SECONDS, + max_retries=LEASE_RETRY_MAX, + ) + except MaxRetriesExceededError: + # The holder is simply slower than LEASE_RETRY_MAX + # deferrals — a task that outruns the broker's visibility + # timeout is redelivered while its first run is still + # going. Standing down is the correct end state for the + # duplicate; raising here would report a failure for a + # task that is running normally somewhere else. + logger.info( + "idempotency: lease still held after %s deferrals; " + "leaving task=%s key=%s to its holder", + LEASE_RETRY_MAX, task_name, key, + ) + return { + "status": "deferred", + "reason": "another worker holds the lease", + "idempotency_key": key, + } if attempt > MAX_TASK_ATTEMPTS: logger.error( diff --git a/tests/api/user/test_idempotency_decorator.py b/tests/api/user/test_idempotency_decorator.py index 35dfb9dd..b1d26f93 100644 --- a/tests/api/user/test_idempotency_decorator.py +++ b/tests/api/user/test_idempotency_decorator.py @@ -335,6 +335,75 @@ class TestLiveLeaseDefersConcurrentRun: assert row[2] == "completed" +@pytest.mark.unit +class TestLeaseDeferralGivesUpQuietly: + """Deferral is bookkeeping, not failure. + + A task that outruns the broker's visibility timeout is redelivered while + the first worker is still running it. The lease keeps the duplicate from + doing the work, but the duplicate kept re-queueing itself until celery + exhausted ``LEASE_RETRY_MAX`` and raised ``MaxRetriesExceededError``, so a + healthy long task logged a task failure. The duplicate should stand down + instead and leave the run to the worker that holds the lease. + """ + + def _hold_lease(self, pg_conn, key): + from docsgpt.storage.db.repositories.idempotency import ( + IdempotencyRepository, + ) + + IdempotencyRepository(pg_conn).try_claim_lease( + key=key, task_name="thing", + task_id="t-worker-1", owner_id="worker-1", + ) + + def test_exhausted_retries_return_deferred_instead_of_raising(self, pg_conn): + from celery.exceptions import MaxRetriesExceededError + + from docsgpt.api.user.idempotency import with_idempotency + + self._hold_lease(pg_conn, "k-long-run") + invocations = {"count": 0} + + @with_idempotency(task_name="thing") + def task(self, idempotency_key=None): + invocations["count"] += 1 + return {"ran": True} + + # Celery raises this from ``self.retry`` once max_retries is hit. + worker2 = _fake_celery_self("t-worker-2") + worker2.retry.side_effect = MaxRetriesExceededError("out of retries") + + with _patch_decorator_db(pg_conn): + result = task(worker2, idempotency_key="k-long-run") + + assert result["status"] == "deferred" + # The lease holder is still running it; the duplicate did not. + assert invocations["count"] == 0 + # The holder's row is untouched — not failed, not completed. + row = _row_for(pg_conn, "k-long-run") + assert row[2] == "pending" + + def test_a_normal_retry_still_propagates(self, pg_conn): + """Only exhaustion stands down; the first deferrals must re-queue.""" + from docsgpt.api.user.idempotency import with_idempotency + + self._hold_lease(pg_conn, "k-busy-once") + + @with_idempotency(task_name="thing") + def task(self, idempotency_key=None): + return {"ran": True} + + class _RetrySignal(Exception): + pass + + worker2 = _fake_celery_self("t-worker-2") + worker2.retry.side_effect = _RetrySignal("retry scheduled") + + with _patch_decorator_db(pg_conn), pytest.raises(_RetrySignal): + task(worker2, idempotency_key="k-busy-once") + + @pytest.mark.unit class TestExceptionPathReleasesLease: """When ``fn`` raises, the lease is dropped so the next attempt From 9e9f130ed028b6e65e7db88b9707dbf88f8ac084 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:03:39 +0100 Subject: [PATCH 05/27] refactor(graphrag): spell the write retry out, and log a capped PPR run MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Review feedback on _write_with_reconnect: an explicit return inside the loop plus an implicit fall-through past it reads as a path that returns None. The retry is exactly two attempts, so say so — first attempt, reconnect on connection loss, second and final attempt — and every path now returns or raises. Also narrows the transitions hint to the (node, weight) tuples it holds, and logs at debug when the power iteration hits its iteration cap instead of returning the last iterate with no trace. --- docsgpt/graphrag/store.py | 29 +++++++++++++++++------------ docsgpt/retriever/graph_rag.py | 9 ++++++++- 2 files changed, 25 insertions(+), 13 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index ceb0fc48..116d3611 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -175,19 +175,24 @@ class GraphStore: Returns: Whatever ``operation`` returns. + + Raises: + Exception: Anything ``operation`` raises that is not connection + loss, and anything the single retry raises. """ - for attempt in (1, 2): - conn = self._get_connection() - try: - return operation(conn) - except Exception as exc: - if attempt == 2 or not _is_connection_lost(exc): - raise - logging.warning( - "Graph write lost its connection (%s); reconnecting and retrying once.", - exc, - ) - self.close() + try: + return operation(self._get_connection()) + except Exception as exc: + if not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + # Second and final attempt, on a connection freshly checked out by + # ``_get_connection``. A failure here belongs to the caller. + return operation(self._get_connection()) def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 5593a9b2..ef4b0ea5 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -103,7 +103,7 @@ def _personalized_pagerank( # Row-normalized transitions. An undirected edge is traversable from both # endpoints, so each node normalizes over its own incident weights. - transitions: Dict[Any, List[Any]] = {} + transitions: Dict[Any, List[tuple[Any, float]]] = {} for node in nodes: neighbors = [] total = 0.0 @@ -134,6 +134,13 @@ def _personalized_pagerank( ranks[node] += (leaked + 1.0 - alpha) * restart[node] if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: break + else: + logging.debug( + "Personalized PageRank hit its %s-iteration cap on a %s-node " + "subgraph; ranking with the last iterate.", + max_iter, + node_count, + ) return ranks From 3f774d813c5dd22e6d4c0fc6602093570ee590e2 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:31:19 +0100 Subject: [PATCH 06/27] =?UTF-8?q?fix(graphrag):=20review=20fixes=20?= =?UTF-8?q?=E2=80=94=20replay-safe=20writes,=20strict=20count,=20zero=20we?= =?UTF-8?q?ights?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three findings from review, all on code this branch introduced. apply_chunk was not replay-safe. commit() can report connection loss *after* Postgres committed, and the reconnect retry then replays the write: _upsert_node bumps doc_freq a second time and _add_edge inserts another row, since graph_edges has no uniqueness constraint for a logical edge. The chunk's graph_ingest_progress row is now written in the same transaction as the rows it describes, and a replay that finds it already "done" returns (0, 0) without touching the graph. Extraction drops its separate mark_chunk("done"): the checkpoint and the graph can no longer disagree. count_nodes swallows every query failure and answers 0, so extraction's "fall back to the write count" handler could never run — a failed count after a successful build reported an empty graph. count_nodes grows a strict mode that re-raises; retrieval keeps the swallow, which is what routes a source to ClassicRAG. A zero edge weight was read as a full-strength link: `or 1.0` rewrote an explicit 0 before the <= 0 filter. Only missing and null weights default now. The same coercion sat in _ppr_scores, where it would have kept the ranker's rule unreachable from the product path, so it is fixed there too. --- docsgpt/graphrag/extraction.py | 9 ++- docsgpt/graphrag/store.py | 52 ++++++++++++-- docsgpt/retriever/graph_rag.py | 14 +++- tests/graphrag/test_extraction.py | 36 ++++++++++ tests/graphrag/test_store.py | 112 ++++++++++++++++++++++++++++++ tests/retriever/test_graph_rag.py | 60 ++++++++++++++++ 6 files changed, 274 insertions(+), 9 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 580b54cb..1bdfb054 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -290,7 +290,9 @@ def extract_graph_for_source( ) node_upserts += chunk_nodes edges += chunk_edges - store.mark_chunk(source_id, chunk_id, "done") + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. chunks_processed += 1 except Exception as exc: logger.warning( @@ -311,9 +313,12 @@ def extract_graph_for_source( # upserts and a single node, so the old count overstated every graph whose # entities recur. Report what the graph holds, falling back to the write # count only if the count query itself fails. + # ``strict`` is what makes the fallback below reachable: the default + # count swallows query failures and answers 0, which would report a + # successful build as an empty graph. nodes = node_upserts try: - nodes = store.count_nodes(source_id) + nodes = store.count_nodes(source_id, strict=True) except Exception as exc: logger.warning( "count_nodes failed for source %s; reporting upserts instead: %s", diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 116d3611..acb35911 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -556,8 +556,12 @@ class GraphStore: ``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees are not bumped here — the caller runs ``set_node_degrees`` once at the end. Reconnects and retries once if the connection died while the - extraction was waiting on the model. Returns - ``(nodes_upserted, edges_added)``. + extraction was waiting on the model. + + The chunk's ``graph_ingest_progress`` row is written in this same + transaction, so the checkpoint and the rows it describes commit + together and a replay of an already-applied chunk returns ``(0, 0)`` + without touching the graph. Returns ``(nodes_upserted, edges_added)``. """ self._ensure_tables_once() @@ -566,6 +570,21 @@ class GraphStore: node_ids: Dict[str, str] = {} edges_added = 0 try: + # ``commit()`` can report connection loss *after* the server + # committed, and the retry then replays this write: doc_freq + # would be bumped twice and a second logical edge inserted + # (graph_edges has no uniqueness constraint). The progress row + # below is written in this transaction, so a replay sees it. + cursor.execute( + "SELECT status FROM graph_ingest_progress " + "WHERE source_id = %s AND chunk_id = %s;", + (source_id, str(chunk_id)), + ) + applied = cursor.fetchone() + if applied is not None and applied[0] == "done": + conn.rollback() + return 0, 0 + for entity in entities: normalized_name = entity["normalized_name"] node_id = self._upsert_node( @@ -601,6 +620,15 @@ class GraphStore: ) edges_added += 1 + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, 'done') + ON CONFLICT (source_id, chunk_id) + DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id)), + ) conn.commit() return len(entities), edges_added except Exception: @@ -671,8 +699,22 @@ class GraphStore: cursor.close() conn.rollback() - def count_nodes(self, source_id: str) -> int: - """Number of nodes for a source. Zero drives the ClassicRAG fallback.""" + def count_nodes(self, source_id: str, strict: bool = False) -> int: + """Number of nodes for a source. Zero drives the ClassicRAG fallback. + + Args: + source_id: Source whose nodes to count. + strict: Re-raise a query failure instead of reporting ``0``. + Retrieval wants the swallow — a broken count there just routes + the source to ClassicRAG — but a caller reporting how big a + graph is must not read a failed query as "the graph is empty". + + Returns: + int: The node count, or ``0`` when a query failure is swallowed. + + Raises: + Exception: The underlying query failure, when ``strict`` is set. + """ conn = self._get_connection() cursor = conn.cursor() try: @@ -683,6 +725,8 @@ class GraphStore: return int(cursor.fetchone()[0]) except Exception as e: logging.error(f"Error counting nodes: {e}") + if strict: + raise return 0 finally: cursor.close() diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index ef4b0ea5..81ed9bf0 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -108,7 +108,11 @@ def _personalized_pagerank( neighbors = [] total = 0.0 for neighbor, data in graph[node].items(): - edge_weight = float(data.get(weight, 1.0) or 1.0) + raw_weight = data.get(weight, 1.0) + # Default only a missing or null weight. ``or 1.0`` would also + # rewrite an explicit 0 — "these entities are not related" — into a + # full-strength transition, which changes the ranking. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) if edge_weight <= 0: continue neighbors.append((neighbor, edge_weight)) @@ -212,8 +216,12 @@ class GraphRAGRetriever(BaseRetriever): for edge in subgraph.get("edges", []): src, dst = edge["src_node_id"], edge["dst_node_id"] if src in graph and dst in graph: - weight = float(edge.get("weight") or 1.0) - graph.add_edge(src, dst, weight=weight) + raw_weight = edge.get("weight") + # Same rule the ranker applies: default only a missing or null + # weight. Coercing an explicit 0 to 1.0 here would make "these + # entities are not related" the strongest possible link. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + graph.add_edge(src, dst, weight=edge_weight) if graph.number_of_nodes() == 0: return {} diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index b6192450..48e567d4 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -624,6 +624,42 @@ class TestSummaryNodeCount: store.delete_by_source(source_id) +@pytest.mark.unit +class TestSummaryCountFailure: + """A broken count query must not be reported as an empty graph.""" + + def test_a_failed_count_reports_the_write_count( + self, monkeypatch, stub_embedding + ): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = ["c1"] + store.apply_chunk.return_value = (2, 1) + store.count_nodes.side_effect = RuntimeError("count query failed") + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + _install_stub_llm( + monkeypatch, + _StubLLM([_extraction_json([{"name": "Ada"}], [])]), + ) + + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c1", "Ada.")], + config=SourceConfig(), + request_id="req-1", + ) + + # Falls back to what was actually written, not to zero. + assert summary["nodes"] == 2 + # And it asked for a count that raises rather than one that returns 0, + # or the fallback above could never run. + assert store.count_nodes.call_args.kwargs.get("strict") is True + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index f2ce6abe..baa19604 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -1002,3 +1002,115 @@ class TestWritesSurviveALostConnection: assert handed == [broken] broken.rollback.assert_called_once() + + +@pytest.mark.unit +class TestCountNodesFailureModes: + """Retrieval wants a swallowed count; extraction wants to hear about it.""" + + def _store_with_failing_cursor(self): + store = GraphStore.__new__(GraphStore) + store._tables_ensured = True + cursor = MagicMock() + cursor.execute.side_effect = RuntimeError("relation does not exist") + conn = MagicMock() + conn.cursor.return_value = cursor + store._connection = conn + store._get_connection = lambda: conn + return store + + def test_default_reports_zero_to_drive_the_classic_fallback(self): + store = self._store_with_failing_cursor() + + assert store.count_nodes(str(uuid.uuid4())) == 0 + + def test_strict_surfaces_the_query_failure(self): + """A caller reporting graph size must not read a broken query as empty.""" + store = self._store_with_failing_cursor() + + with pytest.raises(RuntimeError): + store.count_nodes(str(uuid.uuid4()), strict=True) + + +@pytest.mark.integration +class TestApplyChunkIsReplaySafe: + """A retry after an ambiguous commit must not apply a chunk twice. + + ``_write_with_reconnect`` replays the write when the connection dies, and + ``commit()`` itself can raise connection loss *after* the server committed. + Replaying then bumps ``doc_freq`` a second time and inserts a second + logical edge (``graph_edges`` has no uniqueness constraint), so the chunk's + own progress row is written in the same transaction and short-circuits it. + """ + + @pytest.fixture + def store(self, postgresql): + store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) + try: + store._ensure_tables() + except Exception as exc: + pytest.skip(f"pgvector extension unavailable: {exc}") + yield store + store.close() + + def test_a_replayed_chunk_is_not_applied_twice(self, store): + source_id = str(uuid.uuid4()) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + relationships = [ + { + "source": "Ada", + "target": "Engine", + "type": "worked_on", + "description": "x", + "weight": 2.0, + } + ] + embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)} + try: + first = store.apply_chunk( + source_id, "c1", entities, relationships, embeddings + ) + replay = store.apply_chunk( + source_id, "c1", entities, relationships, embeddings + ) + + assert first == (1, 1) + assert replay == (0, 0) + node = store.get_node_by_normalized(source_id, "ada") + assert node["doc_freq"] == 1 + overview = store.get_graph_overview(source_id) + assert len(overview["edges"]) == 1 + # The write records its own progress, so the caller's checkpoint + # and the rows it describes commit together. + assert store.get_progress(source_id)["c1"] == "done" + finally: + store.delete_by_source(source_id) + + def test_a_different_chunk_still_applies(self, store): + """The guard is per chunk, not a blanket 'already saw this source'.""" + source_id = str(uuid.uuid4()) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + embeddings = {"ada": _embedding(0.1)} + try: + store.apply_chunk(source_id, "c1", entities, [], embeddings) + second = store.apply_chunk(source_id, "c2", entities, [], embeddings) + + assert second == (1, 0) + node = store.get_node_by_normalized(source_id, "ada") + assert node["doc_freq"] == 2 + finally: + store.delete_by_source(source_id) diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 5f8cdf80..8943174c 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -901,6 +901,66 @@ class TestPersonalizedPageRankWithoutScipy: assert _personalized_pagerank(nx.Graph(), personalization=None) == {} + def test_a_zero_weight_edge_is_not_traversable(self): + """Zero means "not related", not "use the default weight".""" + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = nx.Graph() + graph.add_edge("seed", "zero", weight=0.0) + graph.add_edge("seed", "real", weight=1.0) + + ranks = _personalized_pagerank( + graph, personalization={"seed": 1.0, "zero": 0.0, "real": 0.0} + ) + + # ``zero`` is reachable only across the zero-weight edge, so no mass + # walks to it; ``real`` is on a live edge and must outrank it. + assert ranks["real"] > ranks["zero"] + assert ranks["zero"] == pytest.approx(0.0, abs=1e-9) + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + + def test_stored_zero_weights_reach_the_ranker_intact(self): + """The subgraph builder must not coerce a stored 0 into a real edge. + + Without this the ranker's zero-weight rule is unreachable in + production: every 0 from ``graph_edges`` arrives as 1.0. + """ + subgraph = { + "nodes": [ + {"id": "seed", "doc_freq": 1}, + {"id": "zero", "doc_freq": 1}, + {"id": "real", "doc_freq": 1}, + ], + "edges": [ + {"src_node_id": "seed", "dst_node_id": "zero", "weight": 0}, + {"src_node_id": "seed", "dst_node_id": "real", "weight": 1.0}, + ], + } + # Called unbound with ``None`` for self: _ppr_scores reads no state. + scores = GraphRAGRetriever._ppr_scores(None, subgraph, {"seed": 1.0}) + + assert scores["real"] > scores["zero"] + assert scores["zero"] == pytest.approx(0.0, abs=1e-9) + + def test_missing_and_null_weights_default_to_one(self): + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + absent = nx.Graph() + absent.add_edge("a", "b") # no weight attribute at all + null = nx.Graph() + null.add_edge("a", "b", weight=None) + + personalization = {"a": 1.0, "b": 0.0} + from_absent = _personalized_pagerank(absent, personalization=personalization) + from_null = _personalized_pagerank(null, personalization=personalization) + + assert from_absent["b"] == pytest.approx(from_null["b"], abs=1e-9) + assert from_absent["b"] > 0 + @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) @patch("docsgpt.retriever.graph_rag.GraphStore") @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) From b7e7872bf7722002aac7515f8af5960a579d7ddf Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:51:06 +0100 Subject: [PATCH 07/27] fix(tasks): a deferred duplicate records no result, rather than success MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Returning a deferred marker traded one wrong signal for a worse one. A redelivery reuses the original task id — Context.as_execution_options carries task_id into the retry — so the duplicate's return marked the very id the client polls as SUCCESS. /api/task_status reports celery's state verbatim and the UI maps SUCCESS to "done", so the GraphRAG enable modal would announce a finished build, rendered from a payload with no counts, while the run holding the lease was still extracting. Raise Ignore instead: celery records no state for the duplicate, so the task id keeps whatever the holder sets and the poller keeps waiting. The autoretry wrapper re-raises Ignore ahead of autoretry_for, so the wider autoretry_for=(Exception,) on these tasks cannot turn it back into a retry. --- docsgpt/api/user/idempotency.py | 22 +++++++++++--------- tests/api/user/test_idempotency_decorator.py | 14 ++++++++----- 2 files changed, 21 insertions(+), 15 deletions(-) diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py index e38e8a73..1cfc1b80 100644 --- a/docsgpt/api/user/idempotency.py +++ b/docsgpt/api/user/idempotency.py @@ -9,7 +9,7 @@ import threading import uuid from typing import Any, Callable, Optional -from celery.exceptions import MaxRetriesExceededError +from celery.exceptions import Ignore, MaxRetriesExceededError from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository from docsgpt.storage.db.session import db_readonly, db_session @@ -90,21 +90,23 @@ def with_idempotency( ) except MaxRetriesExceededError: # The holder is simply slower than LEASE_RETRY_MAX - # deferrals — a task that outruns the broker's visibility + # deferrals: a task that outruns the broker's visibility # timeout is redelivered while its first run is still - # going. Standing down is the correct end state for the - # duplicate; raising here would report a failure for a - # task that is running normally somewhere else. + # going. Letting the exhaustion propagate would report a + # failure for a task that is running normally — but so + # would returning a value, only less visibly. A redelivery + # reuses the original task id (``Context`` carries + # ``task_id`` into the retry), so a return marks the very + # id the client polls SUCCESS, and ``/api/task_status`` + # hands that to the UI as a finished build. ``Ignore`` + # records no state at all, leaving the outcome to the run + # that actually holds the lease. logger.info( "idempotency: lease still held after %s deferrals; " "leaving task=%s key=%s to its holder", LEASE_RETRY_MAX, task_name, key, ) - return { - "status": "deferred", - "reason": "another worker holds the lease", - "idempotency_key": key, - } + raise Ignore() from None if attempt > MAX_TASK_ATTEMPTS: logger.error( diff --git a/tests/api/user/test_idempotency_decorator.py b/tests/api/user/test_idempotency_decorator.py index b1d26f93..cbed5f19 100644 --- a/tests/api/user/test_idempotency_decorator.py +++ b/tests/api/user/test_idempotency_decorator.py @@ -357,8 +357,8 @@ class TestLeaseDeferralGivesUpQuietly: task_id="t-worker-1", owner_id="worker-1", ) - def test_exhausted_retries_return_deferred_instead_of_raising(self, pg_conn): - from celery.exceptions import MaxRetriesExceededError + def test_exhausted_retries_stand_down_without_recording_a_result(self, pg_conn): + from celery.exceptions import Ignore, MaxRetriesExceededError from docsgpt.api.user.idempotency import with_idempotency @@ -374,10 +374,14 @@ class TestLeaseDeferralGivesUpQuietly: worker2 = _fake_celery_self("t-worker-2") worker2.retry.side_effect = MaxRetriesExceededError("out of retries") - with _patch_decorator_db(pg_conn): - result = task(worker2, idempotency_key="k-long-run") + # Ignore rather than a return value: a redelivery reuses the original + # task id, so returning would mark the id the client is polling + # SUCCESS — /api/task_status hands that straight to the UI, which + # would announce a finished (empty) build while the holder is still + # working. Ignore leaves the id's state to the holder. + with _patch_decorator_db(pg_conn), pytest.raises(Ignore): + task(worker2, idempotency_key="k-long-run") - assert result["status"] == "deferred" # The lease holder is still running it; the duplicate did not. assert invocations["count"] == 0 # The holder's row is untouched — not failed, not completed. From a83e1dc0afb95c8b868b2fe83075c968fc5bdde7 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 08/27] feat(graphrag): seed the walk from what entities are, and rank with passages and vector hits Graph retrieval tied plain vector search at best and never beat it. Measured across five corpora, the bottleneck was seeding, not the graph: the walk started from nodes whose embeddings were computed from bare entity names, and a whole question shares almost nothing with a name like "Quill". Extraction now embeds each node from "name (type): description" and each relationship as the fact it asserts ("Alder streams_to Quill: ..."), stored on a new nullable graph_edges.fact_embedding column that ensure_vector_schema adds in place. Entity names are canonicalised (case, punctuation, word breaks and a cautious plural) so "VECTOR_STORE" and "vector stores" land on one node. Extraction calls run concurrently (GRAPHRAG_EXTRACTION_WORKERS, default 8) while embedding and graph writes stay serial on the task thread, so ordering and idempotency are unchanged; that measured 8.4x faster with identical output. Retrieval gains per-source options, stored under retrieval.graph and read live at query time: - seed_strategy: start from matching entities (default) or matching relationships, which can reach an entity the question never names; - passage_nodes (on): walk the source's passages alongside entities, with PageRank damping 0.5 instead of 0.85; - blend_vector (on): fuse the graph ranking with the source's vector ranking by reciprocal rank. The defaults are the measured-best configuration. Through GraphRAGRetriever, the new seeding moved recall@4 from 0.41 to 0.68 on a multi-hop corpus and from 0.50 to 1.00 on the docs corpus, and regressed none of the corpora measured. Existing graphs keep name-only embeddings until rebuilt. --- docs/content/Deploying/Settings-Reference.mdx | 6 + docsgpt/core/settings/retrieval.py | 9 + docsgpt/graphrag/extraction.py | 169 +++++++-- docsgpt/graphrag/naming.py | 94 +++++ docsgpt/graphrag/store.py | 331 +++++++++++++++++- docsgpt/retriever/graph_rag.py | 286 +++++++++++++-- docsgpt/storage/db/source_config.py | 24 +- tests/graphrag/test_extraction.py | 213 +++++++++++ tests/graphrag/test_retriever_default_path.py | 109 ++++++ tests/graphrag/test_retriever_passages.py | 132 +++++++ tests/graphrag/test_retriever_seeding.py | 129 +++++++ tests/graphrag/test_store.py | 131 ++++++- tests/retriever/test_graph_rag.py | 23 ++ 13 files changed, 1577 insertions(+), 79 deletions(-) create mode 100644 docsgpt/graphrag/naming.py create mode 100644 tests/graphrag/test_retriever_default_path.py create mode 100644 tests/graphrag/test_retriever_passages.py create mode 100644 tests/graphrag/test_retriever_seeding.py diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 32181adb..3221867f 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -447,6 +447,12 @@ Type `int`, default `2000`, must be `>= 0`. Hard cap on chunks extracted per source (cost control); 0 extracts nothing. +### `GRAPHRAG_EXTRACTION_WORKERS` + +Type `int`, default `8`, must be `>= 1` and `<= 32`. + +Concurrent extraction calls during ingest. Model calls run in parallel while graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial. + ## Vector stores diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index 651d3f3d..53dbf56b 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -31,6 +31,15 @@ class RetrievalSettings(SettingsGroup): GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing." ) + GRAPHRAG_EXTRACTION_WORKERS: int = Field( + default=8, + ge=1, + le=32, + description=( + "Concurrent extraction calls during ingest. Model calls run in parallel while " + "graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial." + ), + ) @field_validator("VECTOR_STORE", mode="before") @classmethod diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 1bdfb054..db8ea27a 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -27,6 +27,7 @@ from docsgpt.core.model_utils import ( get_api_key_for_provider, get_provider_from_model_id, ) +from docsgpt.graphrag.naming import normalize_entity_name from docsgpt.core.settings import settings from docsgpt.llm.llm_creator import LLMCreator from docsgpt.storage.db.source_config import SourceConfig @@ -224,6 +225,8 @@ def extract_graph_for_source( source's graph holds after the run — not how many upserts ran, which counts the same entity once per chunk it appears in. """ + from concurrent.futures import ThreadPoolExecutor + from docsgpt.graphrag.store import GraphStore store = GraphStore() @@ -266,43 +269,85 @@ def extract_graph_for_source( except Exception as exc: logger.debug("graph progress callback failed: %s", exc) - for chunk, chunk_id in to_process: + def _prepare(item): + """One chunk's LLM extraction — the only step run concurrently. + + A chunk spends almost all of its time waiting on the model, so that is + what runs in the pool. Everything else stays on the calling thread: + graph writes, so transactions and the progress checkpoint are exactly + what they were serially, and embedding. Inside a Celery worker the + embeddings client decides to embed locally from the task on the + *current thread's* stack; a pool thread has none, so it would instead + dispatch an embed task to the worker and wait on it, which Celery + refuses inside a task — failing every chunk of the build. + """ + chunk, chunk_id = item text = _chunk_text(chunk) if not text: - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 - _report() - continue + return chunk_id, "empty", None extracted = _extract_chunk(llm, text, chunk_id) if extracted is None: - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() - continue - + return chunk_id, "failed", None try: entities = _build_entities(extracted["entities"]) relationships = _build_relationships(extracted["relationships"]) - name_embeddings = _embed_names(embedding, entities, relationships) - chunk_nodes, chunk_edges = store.apply_chunk( - source_id, chunk_id, entities, relationships, name_embeddings - ) - node_upserts += chunk_nodes - edges += chunk_edges - # ``apply_chunk`` marks the chunk done inside the transaction that - # writes its rows, so the checkpoint cannot disagree with the graph - # and a replayed write cannot apply the chunk twice. - chunks_processed += 1 except Exception as exc: logger.warning( - "Graph extraction write failed for chunk %s, skipping: %s", - chunk_id, - exc, + "Graph extraction failed for chunk %s, skipping: %s", chunk_id, exc ) - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() + return chunk_id, "failed", None + return chunk_id, "ok", (entities, relationships) + + workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1)) + pool = None + if workers > 1 and len(to_process) > 1: + pool = ThreadPoolExecutor(max_workers=workers) + # ``map`` yields in submission order, so chunks are still applied in the + # order they were given and a run stays reproducible. + prepared = pool.map(_prepare, to_process) + else: + prepared = (_prepare(item) for item in to_process) + + try: + for chunk_id, status, payload in prepared: + if status == "empty": + store.mark_chunk(source_id, chunk_id, "done") + chunks_processed += 1 + _report() + continue + if status == "failed": + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks += 1 + _report() + continue + + entities, relationships = payload + try: + # On this thread, not in the pool — see ``_prepare``. + name_embeddings = _embed_names(embedding, entities, relationships) + _embed_facts(embedding, relationships) + chunk_nodes, chunk_edges = store.apply_chunk( + source_id, chunk_id, entities, relationships, name_embeddings + ) + node_upserts += chunk_nodes + edges += chunk_edges + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. + chunks_processed += 1 + except Exception as exc: + logger.warning( + "Graph extraction embed/write failed for chunk %s, skipping: %s", + chunk_id, + exc, + ) + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks += 1 + _report() + finally: + if pool is not None: + pool.shutdown(wait=True) try: store.set_node_degrees(source_id) @@ -347,7 +392,7 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: entities.append( { "name": name, - "normalized_name": name.lower(), + "normalized_name": normalize_entity_name(name), "type": str(e.get("type") or "") or None, "description": str(e.get("description") or "") or None, } @@ -373,6 +418,70 @@ def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]: return relationships +def _fact_text(rel: Dict[str, Any]) -> str: + """A relationship rendered as the sentence it asserts. + + Embedded and stored on the edge so retrieval can match a question against + the *relation* rather than against entity names — the difference between + "which entity is this about" and "which fact answers this". + """ + source = str(rel.get("source") or "").strip() + target = str(rel.get("target") or "").strip() + if not source or not target: + return "" + relation = str(rel.get("type") or "related to").strip() or "related to" + text = f"{source} {relation} {target}" + description = str(rel.get("description") or "").strip() + return f"{text}: {description}" if description else text + + +def _embed_facts(embedding, relationships: List[Dict[str, Any]]) -> None: + """Attach a fact embedding to each relationship, in one batched call. + + Mutates the relationship dicts so the embedding travels with the edge into + ``apply_chunk`` without a second mapping to keep in step. Always on: it is + one extra batched call per chunk against an LLM call that already costs + far more, and it lets a source switch to relationship seeding at query time + without being rebuilt. + """ + pending = [(rel, _fact_text(rel)) for rel in relationships] + pending = [(rel, text) for rel, text in pending if text] + if not pending: + return + try: + vectors = embedding.embed_documents([text for _rel, text in pending]) + except Exception as exc: # noqa: BLE001 + # The graph is still correct without them; only fact seeding degrades. + logger.warning("Fact embedding failed, continuing without: %s", exc) + return + for (rel, _text), vector in zip(pending, vectors): + rel["fact_embedding"] = vector + + +def _seed_text(entity: Dict[str, Any]) -> str: + """The text a node's embedding is computed from. + + Retrieval seeds the graph walk by matching a whole question against these + embeddings, and a bare entity name is a poor thing to match a question + against — a question about what a service writes to shares almost no + surface with the name ``Quill``. Including the type and description gives + the match something to work with; measured across five corpora it moved + recall@4 by +0.07 to +0.50. + + Relationship endpoints keep their bare names: they arrive as strings with + no type or description attached. + """ + name = str(entity.get("name") or "").strip() + text = name + entity_type = str(entity.get("type") or "").strip() + if entity_type: + text += f" ({entity_type})" + description = str(entity.get("description") or "").strip() + if description: + text += f": {description}" + return text or name + + def _embed_names( embedding, entities: List[Dict[str, Any]], @@ -385,14 +494,16 @@ def _embed_names( """ name_by_norm: Dict[str, str] = {} for entity in entities: - name_by_norm.setdefault(entity["normalized_name"], entity["name"]) + name_by_norm.setdefault(entity["normalized_name"], _seed_text(entity)) for rel in relationships: for endpoint in (rel.get("source"), rel.get("target")): if endpoint is None: continue clean = str(endpoint).strip() if clean: - name_by_norm.setdefault(clean.lower(), clean) + # Same key the store resolves endpoints by, or the embedding + # computed here never reaches the node it was computed for. + name_by_norm.setdefault(normalize_entity_name(clean), clean) if not name_by_norm: return {} diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py new file mode 100644 index 00000000..3dbfa1cb --- /dev/null +++ b/docsgpt/graphrag/naming.py @@ -0,0 +1,94 @@ +"""Canonical entity naming for the per-source knowledge graph. + +Nodes are merged on ``normalized_name``, which has been ``name.lower()``. That +splits entities a reader would call the same thing: measured on the DocsGPT docs +corpus, ``agent``/``agents``, ``VECTOR_STORE``/``Vector store``/``vector stores``, +``Celery worker``/``Celery workers`` and ``.env file``/``env_file`` all landed as +separate nodes — 58 such collisions across 1,704 entities, with 75% of entities +appearing in exactly one chunk as a result. + +:func:`canonical_name` folds the differences that are purely orthographic: +case, surrounding punctuation, underscore/hyphen word breaks, and a *cautious* +plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``, +``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s" +would corrupt them into new entities rather than merge anything. + +Always on: every graph is built with canonical names. +""" + +from __future__ import annotations + +import re + +_PUNCT = re.compile(r"[^\w\s]+", re.UNICODE) +_UNDERSCORE = re.compile(r"[_\-]+") +_SPACE = re.compile(r"\s+") + +#: Words that end in "s" without being plural. Singularising these would invent +#: entities ("postgre", "kubernete") instead of merging existing ones. +_NOT_PLURAL = frozenset( + { + "postgres", "kubernetes", "https", "aws", "dns", "tls", "cors", "css", + "js", "sas", "gas", "ss", "class", "access", "process", "status", + "analysis", "basis", "axis", "https", "rss", "less", "express", + "redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos", + "always", "sometimes", "series", "docs", "ops", "devops", "sse", + } +) + + +def _singular(word: str) -> str: + """Best-effort singular of one word, biased hard towards leaving it alone. + + Only the endings that are unambiguous in this domain are touched: + ``-ies`` -> ``-y`` (``policies``), ``-ses``/``-xes``/``-zes``/``-ches``/ + ``-shes`` -> drop ``es`` (``indexes``, ``batches``), and a bare trailing + ``s`` on a word long enough to be safe. Everything in :data:`_NOT_PLURAL`, + and anything ending in ``ss``/``us``/``is``, is returned unchanged. + """ + if len(word) < 4 or word in _NOT_PLURAL: + return word + if word.endswith(("ss", "us", "is")): + return word + if word.endswith("ies") and len(word) > 4: + return word[:-3] + "y" + if word.endswith(("ses", "xes", "zes", "ches", "shes")): + return word[:-2] + if word.endswith("s"): + return word[:-1] + return word + + +def canonical_name(name: str) -> str: + """Merge key for an entity name. + + Args: + name: The entity name as the model wrote it. + + Returns: + A lowercase, punctuation-free, singularised key. Returns ``""`` for an + empty or punctuation-only name, which callers treat as "no entity". + + Examples: + ``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``; + ``.env file`` and ``env_file`` -> ``env file``; + ``postgres`` stays ``postgres``. + """ + if not name: + return "" + text = _UNDERSCORE.sub(" ", str(name)) + text = _PUNCT.sub(" ", text) + text = _SPACE.sub(" ", text).strip().lower() + if not text: + return "" + return " ".join(_singular(word) for word in text.split()) + + +def normalize_entity_name(name: str) -> str: + """The key an entity is merged on: its :func:`canonical_name`. + + Every graph the corpora were measured on was built this way, so it is the + only mode rather than a flag. A graph built before this used plain + ``lower()`` keys; re-extracting it merges onto these instead. + """ + return canonical_name(name) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index acb35911..134998ba 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -74,6 +74,16 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]: ) +def _pgvector_vector_column() -> str: + """Resolve the embedding column name from the same ``PGVectorStore`` defaults.""" + import inspect + + from docsgpt.vectorstore.pgvector import PGVectorStore + + params = inspect.signature(PGVectorStore.__init__).parameters + return _safe_identifier(params["vector_column"].default) + + def _is_connection_lost(exc: BaseException) -> bool: """True when ``exc`` says the server connection went away, not that the SQL was bad. @@ -253,7 +263,7 @@ class GraphStore: ) cursor.execute( - """ + f""" CREATE TABLE IF NOT EXISTS graph_edges ( id UUID PRIMARY KEY, source_id UUID NOT NULL, @@ -262,10 +272,18 @@ class GraphStore: type TEXT, description TEXT, weight REAL DEFAULT 1.0, - source_chunk_ids JSONB + source_chunk_ids JSONB, + fact_embedding vector({dimension}) ); """ ) + # ``CREATE TABLE IF NOT EXISTS`` is a no-op on a database that + # already has the table, so a column added after the fact needs its + # own statement or every existing deployment silently lacks it. + cursor.execute( + f"ALTER TABLE graph_edges " + f"ADD COLUMN IF NOT EXISTS fact_embedding vector({dimension});" + ) cursor.execute( """ @@ -453,19 +471,78 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge on an open cursor (no commit, no degree bump). + fact_embedding: Optional[List[float]] = None, + ) -> tuple[Optional[str], bool]: + """Write an edge on an open cursor (no commit, no degree bump). + + Returns ``(edge_id, created)``. Two shapes of noise are rejected here + rather than at read time, because once written neither is visible: + + * A self-loop feeds a node's PageRank mass straight back to itself. It + is dropped, reported as ``(None, False)``. + * A pair already related by the same type is *merged* rather than + inserted again. ``graph_edges`` carries no uniqueness constraint, so + re-extracting one relationship across many chunks otherwise writes a + row per chunk — a fifth of a real corpus's edges — inflating that + pair's traversal weight and spending the bounded subgraph fetch on + duplicates. The surviving row keeps the strongest weight seen and + every contributing chunk id. Callers that batch many edges run ``set_node_degrees`` once afterwards instead of bumping degree per edge. """ + if str(src_node_id) == str(dst_node_id): + return None, False + + cursor.execute( + """ + SELECT id + FROM graph_edges + WHERE source_id = %s AND src_node_id = %s AND dst_node_id = %s + AND type IS NOT DISTINCT FROM %s + LIMIT 1; + """, + (source_id, src_node_id, dst_node_id, type), + ) + existing = cursor.fetchone() + if existing: + edge_id = existing[0] + # The chunk ids are merged in SQL, against the row's own current + # value, rather than read here and written back: a read-modify-write + # would drop whatever a concurrent writer appended in between. + cursor.execute( + """ + UPDATE graph_edges + SET weight = GREATEST(COALESCE(weight, 0), %s), + description = COALESCE(description, %s), + -- Backfills the fact embedding for an edge first written + -- before fact embeddings were switched on. + fact_embedding = COALESCE(fact_embedding, %s::vector), + source_chunk_ids = COALESCE(source_chunk_ids, '[]'::jsonb) || ( + SELECT COALESCE(jsonb_agg(candidate), '[]'::jsonb) + FROM jsonb_array_elements(%s::jsonb) AS candidate + WHERE NOT COALESCE(source_chunk_ids, '[]'::jsonb) + @> jsonb_build_array(candidate) + ) + WHERE id = %s; + """, + ( + weight, + description, + fact_embedding, + Jsonb(list(source_chunk_ids or [])), + edge_id, + ), + ) + return str(edge_id), False + edge_id = str(uuid.uuid4()) cursor.execute( """ INSERT INTO graph_edges (id, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids) - VALUES (%s, %s, %s, %s, %s, %s, %s, %s); + weight, source_chunk_ids, fact_embedding) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s); """, ( edge_id, @@ -476,9 +553,10 @@ class GraphStore: description, weight, Jsonb(source_chunk_ids or []), + fact_embedding, ), ) - return edge_id + return edge_id, True def add_edge( self, @@ -489,21 +567,28 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge and bump the degree of both endpoints. Returns its id.""" + fact_embedding: Optional[List[float]] = None, + ) -> Optional[str]: + """Write an edge and bump the degree of both endpoints. Returns its id. + + Returns ``None`` for a self-loop, which is not written. A repeat of an + existing pair merges into that row and returns its id, leaving degree + alone — the endpoints gained no new neighbour. + """ self._ensure_tables_once() conn = self._get_connection() cursor = conn.cursor() try: - edge_id = self._add_edge( + edge_id, created = self._add_edge( cursor, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids, - ) - cursor.execute( - "UPDATE graph_nodes SET degree = degree + 1 " - "WHERE source_id = %s AND id IN (%s, %s);", - (source_id, src_node_id, dst_node_id), + weight, source_chunk_ids, fact_embedding, ) + if created: + cursor.execute( + "UPDATE graph_nodes SET degree = degree + 1 " + "WHERE source_id = %s AND id IN (%s, %s);", + (source_id, src_node_id, dst_node_id), + ) conn.commit() return edge_id except Exception as e: @@ -608,7 +693,7 @@ class GraphStore: ) if src_id is None or dst_id is None: continue - self._add_edge( + _, created = self._add_edge( cursor, source_id, src_id, @@ -617,8 +702,10 @@ class GraphStore: description=rel.get("description"), weight=float(rel.get("weight") or 1.0), source_chunk_ids=[chunk_id], + fact_embedding=rel.get("fact_embedding"), ) - edges_added += 1 + if created: + edges_added += 1 cursor.execute( """ @@ -653,7 +740,11 @@ class GraphStore: clean = str(name).strip() if not clean: return None - normalized_name = clean.lower() + from docsgpt.graphrag.naming import normalize_entity_name + + normalized_name = normalize_entity_name(clean) + if not normalized_name: + return None if normalized_name in node_ids: return node_ids[normalized_name] node_id = self._upsert_node( @@ -839,6 +930,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND (src_node_id = ANY(%s) OR dst_node_id = ANY(%s)) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, ( @@ -886,6 +978,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND src_node_id = ANY(%s) AND dst_node_id = ANY(%s) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, (source_id, node_id_list, node_id_list, MAX_SUBGRAPH_EDGES), @@ -999,6 +1092,206 @@ class GraphStore: cursor.close() conn.rollback() + def seed_nodes_from_facts( + self, + source_id: str, + query_embedding: List[float], + fact_limit: int = 5, + limit: int = 10, + ) -> List[Dict[str, Any]]: + """Seed nodes drawn from the *relationships* nearest the question. + + Name matching asks "which entity is this question about", which a + multi-document question cannot answer: the entity holding the answer is + named in another document, not in the question. A fact string carries + the relation — "Alder streams_to Quill: ..." — so a question about what + a service writes to can match the edge itself and seed the walk on both + of its endpoints, including the one nothing in the question names. + + Endpoints are weighted by fact score divided by the entity's + ``doc_freq``: an entity appearing in every chunk is a poor seed even + when it sits on a well-matched fact, and dividing by how widely it + occurs prefers the specific endpoint over the hub. + + Rows match :meth:`search_nodes_by_embedding`'s shape, so the caller's + seed weighting is unchanged. Returns nothing when the source has no + fact embeddings, which is the signal to fall back to name matching. + """ + if not query_embedding: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + WITH top_facts AS ( + SELECT src_node_id, dst_node_id, + 1 - (fact_embedding <=> %s::vector) AS score + FROM graph_edges + WHERE source_id = %s AND fact_embedding IS NOT NULL + ORDER BY fact_embedding <=> %s::vector + LIMIT %s + ) + SELECT n.id::text, n.name, n.description, + MAX(f.score / GREATEST(COALESCE(n.doc_freq, 1), 1)) AS weight + FROM top_facts f + JOIN graph_nodes n + ON n.id = f.src_node_id OR n.id = f.dst_node_id + WHERE n.source_id = %s + GROUP BY n.id, n.name, n.description + ORDER BY weight DESC + LIMIT %s; + """, + ( + query_embedding, + source_id, + query_embedding, + max(1, int(fact_limit)), + source_id, + max(1, int(limit)), + ), + ) + return [ + { + "id": row[0], + "name": row[1], + "description": row[2], + # The caller reads weight back as ``1 - distance``. + "distance": 1.0 - float(row[3] or 0.0), + } + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error seeding nodes from facts: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_relationships( + self, source_id: str, name: str, limit: int = 25 + ) -> List[Dict[str, Any]]: + """The relationships an entity takes part in, strongest first. + + This is the one thing a caller cannot get from vector search: which + *named* thing an entity is connected to. Matching is on the name rather + than a node id because the caller is an LLM holding a name it read in + the text, not an id. + """ + clean = (name or "").strip() + if not clean: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + SELECT s.name, e.type, d.name, e.description + FROM graph_edges e + JOIN graph_nodes s ON s.id = e.src_node_id + JOIN graph_nodes d ON d.id = e.dst_node_id + WHERE e.source_id = %s AND (s.name ILIKE %s OR d.name ILIKE %s) + ORDER BY e.weight DESC NULLS LAST + LIMIT %s; + """, + (source_id, f"%{clean}%", f"%{clean}%", max(1, int(limit))), + ) + return [ + {"source": row[0], "type": row[1], "target": row[2], "description": row[3]} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading relationships for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_pages( + self, source_id: str, name: str, limit: int = 4 + ) -> List[Dict[str, Any]]: + """Chunks an entity appears in, with the chunk it is *about* first. + + A plain substring match answers "Halvard" with pages that merely mention + Halvard, and an unordered ``LIMIT`` then decides which of those the + caller sees. Nodes whose name is the entity (or the entity plus a + qualifier the extractor appended, "Quill" -> "Quill Store") are + preferred, and among those the chunk whose text opens with the name + comes first; a substring match is the fallback so an unusual name still + resolves. + """ + clean = (name or "").strip() + if not clean: + return [] + table, text_col, metadata_col, source_col = _pgvector_identifiers() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + f""" + SELECT d.{metadata_col}, d.{text_col}, + (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + FROM graph_node_chunks gc + JOIN graph_nodes n ON n.id = gc.node_id + JOIN {table} d ON d.id::text = gc.chunk_id + WHERE gc.source_id = %s AND d.{source_col} = %s + AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) + GROUP BY d.{metadata_col}, d.{text_col}, is_subject + ORDER BY is_subject DESC, (d.{text_col} ILIKE %s) DESC + LIMIT %s; + """, + ( + clean.lower(), f"{clean.lower()} %", + source_id, source_id, + clean.lower(), f"{clean.lower()} %", f"%{clean}%", + f"{clean}%", + max(1, int(limit)), + ), + ) + return [ + {"metadata": row[0] or {}, "text": row[1] or ""} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading pages for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def chunk_similarities( + self, source_id: str, chunk_ids: List[str], query_embedding: List[float] + ) -> Dict[str, float]: + """Cosine similarity between the query and specific chunks of a source. + + Passage nodes need their own relevance to claim a share of the walk's + restart mass, and that number lives in the co-located pgvector table — + the same one :meth:`get_chunk_texts` reads. Restricted to the chunk ids + the subgraph actually reached, so this never scans the whole source. + """ + if not chunk_ids or not query_embedding: + return {} + table, _text_col, _metadata_col, source_col = _pgvector_identifiers() + vector_col = _pgvector_vector_column() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + f""" + SELECT id::text, 1 - ({vector_col} <=> %s::vector) + FROM {table} + WHERE {source_col} = %s AND id::text = ANY(%s); + """, + (query_embedding, source_id, [str(c) for c in chunk_ids]), + ) + return {row[0]: float(row[1]) for row in cursor.fetchall()} + except Exception as e: + logging.error(f"Error scoring chunks against the query: {e}") + return {} + finally: + cursor.close() + conn.rollback() + def get_chunk_texts( self, source_id: str, diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 81ed9bf0..818f4f07 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -32,6 +32,7 @@ from docsgpt.graphrag.store import GraphStore from docsgpt.retriever.base import BaseRetriever from docsgpt.retriever.classic_rag import ClassicRAG from docsgpt.retriever.labels import labels_from_metadata +from docsgpt.storage.db.source_config import GraphRetrievalConfig from docsgpt.utils import num_tokens_from_string from docsgpt.vectorstore.base import get_embeddings @@ -39,11 +40,27 @@ SEED_NODES = 10 SUBGRAPH_HOPS = 1 +PASSAGE_NODE_WEIGHT = 0.05 +FACT_SEED_FACTS = 5 +RRF_K = 60 + +# PageRank damping per ranking mode — each is the value that mode was measured +# at. Lower keeps mass nearer the seeds; with passages in the walk 0.5 measured +# better, while entity-only ranking was measured at the conventional 0.85. +DAMPING_WITH_PASSAGES = 0.5 +DAMPING_ENTITIES_ONLY = 0.85 + + def _idf(doc_freq: Any) -> float: """Node-specificity weight: rarer entities (low ``doc_freq``) score higher.""" return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _damping(passage_nodes: bool) -> float: + """PageRank damping for a ranking mode: the value that mode was measured at.""" + return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY + + def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]: """Normalized restart distribution over ``nodes``. @@ -210,18 +227,9 @@ class GraphRAGRetriever(BaseRetriever): After PPR, each node's mass is scaled by ``1/log(2 + doc_freq)`` so a high-degree hub contributes less than a specific entity at equal mass. """ - graph = nx.Graph() - for node in subgraph.get("nodes", []): - graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) - for edge in subgraph.get("edges", []): - src, dst = edge["src_node_id"], edge["dst_node_id"] - if src in graph and dst in graph: - raw_weight = edge.get("weight") - # Same rule the ranker applies: default only a missing or null - # weight. Coercing an explicit 0 to 1.0 here would make "these - # entities are not related" the strongest possible link. - edge_weight = 1.0 if raw_weight is None else float(raw_weight) - graph.add_edge(src, dst, weight=edge_weight) + # Through the class, not ``self``: this method reads no instance state, + # and callers (and tests) rely on being able to invoke it unbound. + graph = GraphRAGRetriever._subgraph_graph(subgraph) if graph.number_of_nodes() == 0: return {} @@ -230,13 +238,33 @@ class GraphRAGRetriever(BaseRetriever): personalization = None ranks = _personalized_pagerank( - graph, personalization=personalization, weight="weight" + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=False), ) return { node: rank * _idf(graph.nodes[node].get("doc_freq", 0)) for node, rank in ranks.items() } + @staticmethod + def _subgraph_graph(subgraph) -> "nx.Graph": + """The fetched subgraph as a weighted undirected graph.""" + graph = nx.Graph() + for node in subgraph.get("nodes", []): + graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) + for edge in subgraph.get("edges", []): + src, dst = edge["src_node_id"], edge["dst_node_id"] + if src in graph and dst in graph: + raw_weight = edge.get("weight") + # Default only a missing or null weight. Coercing an explicit 0 + # to 1.0 would make "these entities are not related" the + # strongest possible link. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + graph.add_edge(src, dst, weight=edge_weight) + return graph + def _rank_chunks(self, store, source_id, node_scores) -> List[str]: """Score chunks by summed (PPR mass x IDF) of their linked nodes; top candidates. @@ -254,6 +282,74 @@ class GraphRAGRetriever(BaseRetriever): candidates = max(self.chunks * 2, self.chunks + 5) return ranked[: max(1, candidates)] + def _rank_chunks_with_passages( + self, store, source_id, subgraph, seeds, query_embedding + ) -> List[str]: + """Rank chunks by walking a graph that contains the chunks themselves. + + :meth:`_rank_chunks` reads a chunk's score *off* its entities, summing + their PPR mass — so a chunk touching many mid-scoring generic entities + outranks one touching the few entities the question is about. Putting + the chunks in the walk instead, each joined to its own entities and + carrying a small share of the restart mass proportional to its own + vector similarity, makes a chunk reachable both ways: by being about the + question, and by being connected to what is. Graph retrieval then + contains vector retrieval rather than competing with it. + + Only an improvement when the seeds are good: measured across five + corpora it helped alongside richer seed embeddings and *hurt* with + bare-name seeds (0.73 -> 0.57 on one corpus). Graphs are now always + built with the richer seed text; one built before that change should be + rebuilt before this is relied on. + """ + node_ids = [node["id"] for node in subgraph.get("nodes", [])] + chunk_links = store.get_chunk_ids_for_nodes(source_id, node_ids) + candidate_ids = sorted({c for chunks in chunk_links.values() for c in chunks}) + if not candidate_ids: + return [] + + graph = self._subgraph_graph(subgraph) + similarities = store.chunk_similarities( + source_id, candidate_ids, query_embedding + ) + # Normalised so the passage share is a fixed fraction of the restart + # mass rather than whatever absolute cosine this embedding model emits. + scores = [similarities.get(c, 0.0) for c in candidate_ids] + low, high = (min(scores), max(scores)) if scores else (0.0, 0.0) + spread = high - low + + personalization = dict(seeds) + passage_of: Dict[str, str] = {} + for chunk_id in candidate_ids: + linked = [n for n, chunks in chunk_links.items() if chunk_id in chunks] + linked = [n for n in linked if n in graph] + if not linked: + continue + passage_node = f"chunk::{chunk_id}" + passage_of[passage_node] = chunk_id + for node in linked: + graph.add_edge(passage_node, node, weight=1.0) + similarity = similarities.get(chunk_id, 0.0) + normalized = (similarity - low) / spread if spread > 0 else 0.0 + personalization[passage_node] = normalized * PASSAGE_NODE_WEIGHT + + if graph.number_of_nodes() == 0 or not any(personalization.values()): + return [] + + ranks = _personalized_pagerank( + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=True), + ) + chunk_scores = { + chunk_id: ranks.get(passage_node, 0.0) + for passage_node, chunk_id in passage_of.items() + } + ranked = sorted(chunk_scores, key=lambda c: chunk_scores[c], reverse=True) + 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. @@ -273,6 +369,111 @@ class GraphRAGRetriever(BaseRetriever): base = self.base_chunks if self.base_chunks is not None else self.chunks return max(1, base // max(1, len(self.vectorstores))) + def _vector_ranking(self, source_id, query_embedding: List[float]) -> List[tuple]: + """The source's own vector ranking, as ``(text, metadata)`` in score order. + + Used only by the hybrid path. Vector hits carry no row id, so the fused + ranking is keyed on the chunk text itself — the one identifier both + rankings share — and the metadata travels with it so a hit the graph + never surfaced can still be emitted as a document. + """ + from docsgpt.vectorstore.vector_creator import VectorCreator + + store = None + try: + store = VectorCreator.create_vectorstore( + settings.VECTOR_STORE, source_id, settings.EMBEDDINGS_KEY + ) + hits = store.search( + self._classic._get_rephrased_question(), + k=max(self.chunks * 4, 20), + query_vector=query_embedding, + ) + except Exception as e: + logging.error( + "GraphRAG hybrid: vector ranking failed for %s: %s", source_id, e + ) + return [] + finally: + close = getattr(store, "close", None) + if close is not None: + try: + close() + except Exception as e: + logging.debug("Error closing hybrid vector store: %s", e) + ranked = [] + for hit in hits: + text = getattr(hit, "page_content", None) + metadata = getattr(hit, "metadata", None) + if text is None and isinstance(hit, dict): + text = hit.get("text") or hit.get("page_content") + metadata = hit.get("metadata") + if text: + ranked.append((text, metadata or {})) + return ranked + + @staticmethod + def _rrf_order(rankings: List[List[str]], k: int) -> Dict[str, float]: + """Reciprocal rank fusion over ranked lists of the same key type. + + Rank-based on purpose: PPR mass and cosine similarity are not on + comparable scales, and normalising either one invents a calibration + that does not exist. + """ + scores: Dict[str, float] = {} + for ranking in rankings: + for position, key in enumerate(ranking): + scores[key] = scores.get(key, 0.0) + 1.0 / (k + position + 1) + return scores + + def _graph_options(self, source_id) -> GraphRetrievalConfig: + """This source's graph retrieval options, or the recommended defaults. + + Options travel on the per-source retrieval config the Dispatcher hands + over. A request that carries no per-source detail gets the defaults, + which are the measured-best configuration rather than a neutral one. + """ + cfg = (getattr(self, "per_source_retrieval", None) or {}).get(source_id) + options = cfg.get("graph") if isinstance(cfg, dict) else getattr(cfg, "graph", None) + if isinstance(options, GraphRetrievalConfig): + return options + try: + return GraphRetrievalConfig.model_validate(options or {}) + except Exception: + return GraphRetrievalConfig() + + def _seed_rows( + self, store, source_id, query_embedding: List[float] + ) -> List[Dict[str, Any]]: + """The nodes the walk restarts from, per the source's ``seed_strategy``. + + Seeding decides more than ranking does: a walk that starts on the wrong + nodes cannot be rescued downstream. + + ``entities`` + Cosine NN over entity embeddings, built from each entity's name, + type and description so a whole question has something to match. + The default: best or tied-best on every corpus measured. + ``relationships`` + Cosine NN over relationship sentences ("A streams_to B: ..."), + seeding both endpoints of the best-matching facts. The only way to + start on an entity the question never names; strongest on + chain-structured content, weaker on ordinary prose. + + Relationship seeding falls back to entity matching for a source with no + fact embeddings (one built before they were recorded), so it still + retrieves rather than returning nothing. + """ + if self._graph_options(source_id).seed_strategy == "relationships": + by_fact = store.seed_nodes_from_facts( + source_id, query_embedding, fact_limit=FACT_SEED_FACTS, limit=SEED_NODES + ) + if by_fact: + return by_fact + return store.search_nodes_by_embedding( + source_id, query_embedding, k=SEED_NODES + ) + def _graph_docs_for_source( self, store, source_id, query_embedding: List[float] ) -> List[Dict[str, Any]]: @@ -284,9 +485,7 @@ class GraphRAGRetriever(BaseRetriever): query_embedding: Embedding of the rephrased question, computed once by the caller for the whole retrieval. """ - seed_rows = store.search_nodes_by_embedding( - source_id, query_embedding, k=SEED_NODES - ) + seed_rows = self._seed_rows(store, source_id, query_embedding) if not seed_rows: return [] @@ -300,26 +499,63 @@ class GraphRAGRetriever(BaseRetriever): for row in seed_rows } + options = self._graph_options(source_id) subgraph = store.get_subgraph(source_id, seed_ids, hops=SUBGRAPH_HOPS) - node_scores = self._ppr_scores(subgraph, seeds) - if not node_scores: + if options.passage_nodes: + chunk_ids = self._rank_chunks_with_passages( + store, source_id, subgraph, seeds, query_embedding + ) + else: + node_scores = self._ppr_scores(subgraph, seeds) + if not node_scores: + return [] + chunk_ids = self._rank_chunks(store, source_id, node_scores) + if not chunk_ids: return [] - chunk_ids = self._rank_chunks(store, source_id, node_scores) chunk_data = store.get_chunk_texts(source_id, chunk_ids) + # ``(text, metadata)`` in rank order. Chunk ids stop being the currency + # here: a hit contributed by the vector ranking has no graph chunk id, + # and keying on ids is what made an earlier version of this fusion able + # only to reorder the graph's own candidates. + candidates: List[tuple] = [] + for chunk_id in chunk_ids: + chunk = chunk_data.get(chunk_id) + text = chunk.get("text") if chunk else None + if text: + candidates.append((text, chunk.get("metadata"))) + + if options.blend_vector: + # The graph ranks by how much PPR mass landed on a chunk's + # entities, which says nothing about whether the chunk is about the + # question. Fusing with the source's own vector ranking keeps the + # graph's reach while letting plain relevance back in — including + # chunks the graph never surfaced, which is where most of the value + # is: no reordering can rescue a question whose answer the graph + # missed entirely. + vector_hits = self._vector_ranking(source_id, query_embedding) + if vector_hits: + metadata_by_text = {text: meta for text, meta in candidates} + for text, meta in vector_hits: + metadata_by_text.setdefault(text, meta) + fused = self._rrf_order( + [[t for t, _ in candidates], [t for t, _ in vector_hits]], + RRF_K, + ) + candidates = [ + (text, metadata_by_text.get(text)) + for text in sorted(fused, key=lambda t: fused[t], reverse=True) + ] + 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: + for text, metadata in candidates: if len(docs) >= source_top_k: break - chunk = chunk_data.get(chunk_id) - text = chunk.get("text") if chunk else None - if not text: - continue - labels = labels_from_metadata(chunk.get("metadata"), text, source_id) + labels = labels_from_metadata(metadata, text, source_id) doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}") if cumulative_tokens + doc_tokens >= token_budget: break diff --git a/docsgpt/storage/db/source_config.py b/docsgpt/storage/db/source_config.py index 20873b46..8f8d6209 100644 --- a/docsgpt/storage/db/source_config.py +++ b/docsgpt/storage/db/source_config.py @@ -12,7 +12,7 @@ reproduces today's chunking byte-for-byte. from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import BaseModel, ConfigDict, field_validator, model_validator @@ -90,6 +90,27 @@ class ChunkingConfig(BaseModel): duplicate_headers: bool = False +class GraphRetrievalConfig(BaseModel): + """How the graph retriever walks a graphrag source (live; no re-ingest). + + The defaults are the configuration that measured best across the corpora + tested rather than a neutral starting point: seed from entity matches, put + the passages in the walk, and blend with the source's own vector ranking. + """ + + model_config = ConfigDict(extra="forbid") + + # Where the walk starts: entities whose descriptions match the question, or + # relationships ("A streams_to B") that do. Relationships can start the + # walk on an entity the question never names. + seed_strategy: Literal["entities", "relationships"] = "entities" + # Chunks join the walk as nodes, so a passage is reachable both by being + # about the question and by being connected to what is. + passage_nodes: bool = True + # Fuse the graph ranking with plain vector search by reciprocal rank. + blend_vector: bool = True + + class RetrievalConfig(BaseModel): """Query-time retrieval knobs (live; no re-ingest needed).""" @@ -102,6 +123,7 @@ class RetrievalConfig(BaseModel): rephrase_query: bool = True # toggle ClassicRAG._rephrase_query side-call reranker: Optional[dict] = None # reserved: future cross-encoder/LLM reorder prescreen: Optional[dict] = None # None = off; else PreScreenConfig dict (D12) + graph: GraphRetrievalConfig = GraphRetrievalConfig() # graphrag retriever only @field_validator("chunks") @classmethod diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 48e567d4..eda8dd7c 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -138,6 +138,129 @@ def _extraction_json(entities, relationships): return json.dumps({"entities": entities, "relationships": relationships}) +class TestFactText: + """A relationship rendered as the sentence it asserts. + + This is what fact seeding matches a question against, so it has to read as + a claim rather than as three fields concatenated. + """ + + def test_renders_the_relationship_as_a_sentence(self): + text = extraction_module._fact_text( + { + "source": "Alder", + "target": "Quill", + "type": "streams_to", + "description": "Alder streams audit events to Quill.", + } + ) + + assert text == "Alder streams_to Quill: Alder streams audit events to Quill." + + def test_omits_an_absent_description(self): + text = extraction_module._fact_text( + {"source": "Alder", "target": "Quill", "type": "streams_to"} + ) + + assert text == "Alder streams_to Quill" + + def test_defaults_a_missing_relation(self): + text = extraction_module._fact_text({"source": "Alder", "target": "Quill"}) + + assert text == "Alder related to Quill" + + @pytest.mark.parametrize( + "rel", + [ + {"source": "Alder", "target": ""}, + {"source": "", "target": "Quill"}, + {}, + ], + ) + def test_an_edge_without_both_endpoints_has_no_fact(self, rel): + assert extraction_module._fact_text(rel) == "" + + +class TestEmbedFacts: + """Fact embeddings are always recorded, so a source can switch to + relationship seeding at query time without being rebuilt.""" + + def _relationships(self): + return [{"source": "Alder", "target": "Quill", "type": "streams_to"}] + + def test_attaches_one_embedding_per_fact_in_a_single_call(self): + relationships = self._relationships() + [{"source": "", "target": "Nowhere"}] + calls = [] + + class _Embedding: + def embed_documents(self, texts): + calls.append(texts) + return [[0.5] * 4 for _ in texts] + + extraction_module._embed_facts(_Embedding(), relationships) + + # One batched call, and the endpoint-less relationship is skipped + # rather than embedded as an empty string. + assert calls == [["Alder streams_to Quill"]] + assert relationships[0]["fact_embedding"] == [0.5] * 4 + assert "fact_embedding" not in relationships[1] + + def test_survives_an_embedding_failure(self): + """The graph is still correct without fact embeddings — only + relationship seeding degrades, and it falls back to entities — so a + failure here must not fail the chunk.""" + relationships = self._relationships() + + class _Embedding: + def embed_documents(self, texts): + raise RuntimeError("embeddings down") + + extraction_module._embed_facts(_Embedding(), relationships) + + assert "fact_embedding" not in relationships[0] + + +class TestSeedText: + """What a node's embedding is computed from. + + Retrieval matches a whole question against these embeddings, so what goes + into them decides what the graph walk can start from. + """ + + def _entity(self): + return { + "name": "Quill", + "normalized_name": "quill", + "type": "store", + "description": "A write-ahead store.", + } + + def test_includes_type_and_description(self): + assert ( + extraction_module._seed_text(self._entity()) + == "Quill (store): A write-ahead store." + ) + + def test_falls_back_to_the_name_when_fields_are_missing(self): + assert extraction_module._seed_text({"name": "Quill"}) == "Quill" + + def test_embedded_text_is_keyed_by_the_normalized_name(self): + """The richer text must reach ``embed_documents``, keyed by the same + normalized name the store resolves nodes by — otherwise the embedding + is computed for a node it never reaches.""" + captured = {} + + class _Embedding: + def embed_documents(self, texts): + captured["texts"] = texts + return [[0.0] * 4 for _ in texts] + + result = extraction_module._embed_names(_Embedding(), [self._entity()], []) + + assert captured["texts"] == ["Quill (store): A write-ahead store."] + assert set(result) == {"quill"} + + @pytest.mark.integration class TestExtractionLive: @pytest.fixture @@ -193,6 +316,96 @@ class TestExtractionLive: finally: store.delete_by_source(source_id) + def test_parallel_workers_process_every_chunk_once( + self, store, source_id, monkeypatch, stub_embedding + ): + """Running the model calls concurrently must not change what gets written. + + Extraction spends nearly all of a chunk's time waiting on the model, so + the calls run in a pool while every graph write stays on the calling + thread. Six chunks share one entity here: whatever order the pool + finishes in, that entity is upserted once, each chunk is linked, and all + six are marked processed. + """ + from docsgpt.core.settings import settings + + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + llm = _StubLLM([payload] * 6) + _install_stub_llm(monkeypatch, llm) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[ + _chunk(f"c{i}", f"Ada appears here, take {i}.") for i in range(6) + ], + config=SourceConfig(), + request_id="req-parallel", + ) + + assert summary["chunks_processed"] == 6 + assert summary["failed_chunks"] == 0 + assert summary["nodes"] == 1 + assert len(llm.gen_calls) == 6 + + node = store.get_node_by_normalized(source_id, "ada") + assert node is not None + mapping = store.get_chunk_ids_for_nodes(source_id, [node["id"]]) + assert sorted(mapping[node["id"]]) == [f"c{i}" for i in range(6)] + finally: + store.delete_by_source(source_id) + + def test_embedding_runs_on_the_calling_thread( + self, store, source_id, monkeypatch, stub_embedding + ): + """Only the LLM call may run in the extraction pool, never embedding. + + Inside a Celery worker the embeddings client decides to embed locally + from the task on the *current thread's* stack. A pool thread has none, + so from there it dispatches an embed task to the worker and waits on + it — which Celery refuses inside a task, so every chunk of a graph + build failed. + """ + import threading + + from docsgpt.core.settings import settings + + caller = threading.current_thread() + seen = [] + real_embed_names = extraction_module._embed_names + + def _recording_embed_names(*args, **kwargs): + seen.append(threading.current_thread()) + return real_embed_names(*args, **kwargs) + + monkeypatch.setattr(extraction_module, "_embed_names", _recording_embed_names) + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + _install_stub_llm(monkeypatch, _StubLLM([payload] * 4)) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(4)], + config=SourceConfig(), + request_id="req-thread", + ) + + assert summary["failed_chunks"] == 0 + assert len(seen) == 4 + assert all(thread is caller for thread in seen) + finally: + store.delete_by_source(source_id) + def test_same_entity_across_chunks_merges( self, store, source_id, monkeypatch, stub_embedding ): diff --git a/tests/graphrag/test_retriever_default_path.py b/tests/graphrag/test_retriever_default_path.py new file mode 100644 index 00000000..c11a0f0b --- /dev/null +++ b/tests/graphrag/test_retriever_default_path.py @@ -0,0 +1,109 @@ +"""The graph retriever's default path, end to end through ``_graph_docs_for_source``. + +The shipped defaults — seed from entities, walk the passages, blend with vector +search — are the configuration that measured best, so they are what most graph +sources run. This drives that whole path with a store that returns real values, +and checks each per-source option actually switches its stage off. +""" + +from __future__ import annotations + +from docsgpt.retriever.graph_rag import GraphRAGRetriever +from docsgpt.storage.db.source_config import RetrievalConfig + +TEXTS = { + "c-alder": "Alder streams audit events to Quill.", + "c-quill": "Quill is compacted every six hours.", +} +VECTOR_ONLY = "A passage only plain vector search found." + + +class _Store: + """A two-entity chain: the question matches Alder, the answer is on Quill.""" + + def __init__(self): + self.calls: list[str] = [] + + def search_nodes_by_embedding(self, source_id, query_embedding, k=10): + return [{"id": "alder", "name": "Alder", "distance": 0.1}] + + def get_subgraph(self, source_id, node_ids, hops=1): + return { + "nodes": [{"id": "alder", "doc_freq": 1}, {"id": "quill", "doc_freq": 1}], + "edges": [{"src_node_id": "alder", "dst_node_id": "quill", "weight": 1.0}], + } + + def get_chunk_ids_for_nodes(self, source_id, node_ids): + return {"alder": ["c-alder"], "quill": ["c-quill"]} + + def chunk_similarities(self, source_id, chunk_ids, query_embedding): + self.calls.append("chunk_similarities") + return {"c-alder": 0.9, "c-quill": 0.2} + + def get_chunk_texts(self, source_id, chunk_ids): + return { + c: {"text": TEXTS[c], "metadata": {"title": c}} + for c in chunk_ids + if c in TEXTS + } + + +def _retriever(per_source=None): + """A retriever without its constructor (which builds a ClassicRAG).""" + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = 3 + retriever.base_chunks = None + retriever.doc_token_limit = 50000 + retriever.vectorstores = ["src"] + retriever.per_source_retrieval = per_source or {} + retriever.vector_calls = 0 + + def _vector_ranking(source_id, query_embedding): + retriever.vector_calls += 1 + return [(VECTOR_ONLY, {"title": "vector"})] + + retriever._vector_ranking = _vector_ranking + return retriever + + +def _texts(docs): + return [doc["text"] for doc in docs] + + +class TestDefaultPath: + def test_walks_passages_and_blends_in_vector_hits(self): + store = _Store() + retriever = _retriever() + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + # The answer sits one edge away from the seed: the walk reached it. + assert TEXTS["c-quill"] in _texts(docs) + # A hit only vector search found is blended in, not lost. + assert VECTOR_ONLY in _texts(docs) + assert store.calls == ["chunk_similarities"] + assert retriever.vector_calls == 1 + + +class TestPerSourceOptions: + def test_passage_walk_can_be_switched_off(self): + store = _Store() + retriever = _retriever( + {"src": RetrievalConfig(chunks=3, graph={"passage_nodes": False})} + ) + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + assert "chunk_similarities" not in store.calls + assert TEXTS["c-quill"] in _texts(docs) + + def test_vector_blending_can_be_switched_off(self): + store = _Store() + retriever = _retriever( + {"src": RetrievalConfig(chunks=3, graph={"blend_vector": False})} + ) + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + assert retriever.vector_calls == 0 + assert VECTOR_ONLY not in _texts(docs) diff --git a/tests/graphrag/test_retriever_passages.py b/tests/graphrag/test_retriever_passages.py new file mode 100644 index 00000000..fcc561ff --- /dev/null +++ b/tests/graphrag/test_retriever_passages.py @@ -0,0 +1,132 @@ +"""Chunks as nodes in the walk, and the damping that decides how far mass spreads. + +``_rank_chunks`` reads a chunk's score off its entities by summing their PPR +mass, which rewards a chunk for touching *many* entities rather than the right +ones. The passage-node path puts the chunks in the graph instead, so a chunk is +reachable both by being about the question and by being connected to what is. + +These tests use a stub store: the ranking is graph arithmetic, and pinning it +against a real database would measure Postgres rather than the ranking. +""" + +from __future__ import annotations + +import pytest + +from docsgpt.retriever.graph_rag import GraphRAGRetriever, _damping + + +class _StubStore: + """The two reads the passage path makes, and nothing else.""" + + def __init__(self, chunk_links, similarities): + self._chunk_links = chunk_links + self._similarities = similarities + + def get_chunk_ids_for_nodes(self, source_id, node_ids): + return {n: c for n, c in self._chunk_links.items() if n in set(node_ids)} + + def chunk_similarities(self, source_id, chunk_ids, query_embedding): + return {c: self._similarities.get(c, 0.0) for c in chunk_ids} + + +def _retriever(chunks=2): + """A retriever without its constructor — which builds a ClassicRAG, opens + settings-driven collaborators, and has nothing to do with ranking.""" + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = chunks + return retriever + + +def _subgraph(): + return { + "nodes": [ + {"id": "a", "doc_freq": 1}, + {"id": "b", "doc_freq": 1}, + {"id": "hub", "doc_freq": 40}, + ], + "edges": [ + {"src_node_id": "a", "dst_node_id": "hub", "weight": 1.0}, + {"src_node_id": "b", "dst_node_id": "hub", "weight": 1.0}, + ], + } + + +class TestDamping: + """Each ranking mode runs at the damping it was measured at.""" + + def test_passage_walk_keeps_mass_near_the_seeds(self): + assert _damping(passage_nodes=True) == 0.5 + + def test_entity_only_ranking_keeps_the_conventional_value(self): + assert _damping(passage_nodes=False) == 0.85 + + +class TestPassageNodes: + def test_ranks_the_chunk_the_question_matches(self, monkeypatch): + """Two chunks are equally connected; only their own relevance differs, + so the more relevant one must win.""" + store = _StubStore( + chunk_links={"a": ["c1"], "b": ["c2"]}, + similarities={"c1": 0.1, "c2": 0.9}, + ) + + ranked = _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0, "b": 1.0}, [0.0] * 4 + ) + + assert ranked[0] == "c2" + + def test_a_chunk_reached_only_through_the_graph_still_ranks(self, monkeypatch): + """The point of the walk: a chunk with no similarity of its own is + still reachable through the entity the seeds point at.""" + store = _StubStore( + chunk_links={"a": ["c1"], "b": ["c2"]}, + similarities={"c1": 0.0, "c2": 0.0}, + ) + + ranked = _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + + assert set(ranked) == {"c1", "c2"} + + def test_no_linked_chunks_returns_nothing(self): + store = _StubStore(chunk_links={}, similarities={}) + + assert ( + _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + == [] + ) + + def test_over_fetches_past_the_chunk_budget(self, monkeypatch): + """Same contract as ``_rank_chunks``: candidates exceed the budget so + chunks with missing text cannot drop the final count below it.""" + links = {"a": [f"c{i}" for i in range(10)]} + store = _StubStore( + chunk_links=links, + similarities={f"c{i}": i / 10 for i in range(10)}, + ) + + ranked = _retriever(chunks=2)._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + + assert len(ranked) == max(2 * 2, 2 + 5) + + +class TestChunkSimilaritiesGuard: + """The store call the passage path depends on short-circuits before it + touches a connection, so an empty subgraph costs no query.""" + + @pytest.mark.parametrize( + "chunk_ids,embedding", [([], [0.1]), (["c1"], []), ([], [])] + ) + def test_empty_inputs_return_empty(self, chunk_ids, embedding): + from docsgpt.graphrag.store import GraphStore + + store = object.__new__(GraphStore) + + assert store.chunk_similarities("src", chunk_ids, embedding) == {} diff --git a/tests/graphrag/test_retriever_seeding.py b/tests/graphrag/test_retriever_seeding.py new file mode 100644 index 00000000..0918aa4c --- /dev/null +++ b/tests/graphrag/test_retriever_seeding.py @@ -0,0 +1,129 @@ +"""Where the graph walk starts, per the source's graph retrieval options. + +Seeding decides more than ranking does — a walk that starts on the wrong nodes +cannot be rescued downstream. The options are per source and live (no +re-ingest), carried on the per-source retrieval config the Dispatcher hands the +retriever, so both the dispatch and how the options are resolved are pinned. + +The fallback matters most: relationship seeding reads fact embeddings written +at ingest, and a source built before they were recorded has none. It must keep +retrieving through entity matching rather than returning nothing. +""" + +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from docsgpt.retriever.graph_rag import GraphRAGRetriever +from docsgpt.storage.db.source_config import GraphRetrievalConfig, RetrievalConfig + +ENTITY_ROWS = [{"id": "n1", "name": "Quill", "distance": 0.2}] +FACT_ROWS = [ + {"id": "f1", "name": "Alder", "distance": 0.1}, + {"id": "f2", "name": "Quill", "distance": 0.1}, +] + + +class _StubStore: + def __init__(self, fact_rows=None): + self.fact_rows = list(FACT_ROWS) if fact_rows is None else fact_rows + self.calls: list[str] = [] + + def seed_nodes_from_facts(self, source_id, query_embedding, fact_limit=5, limit=10): + self.calls.append("facts") + return self.fact_rows + + def search_nodes_by_embedding(self, source_id, query_embedding, k=10): + self.calls.append("entities") + return list(ENTITY_ROWS) + + +def _retriever(per_source=None): + """A retriever without its constructor, which builds a ClassicRAG.""" + retriever = object.__new__(GraphRAGRetriever) + if per_source is not None: + retriever.per_source_retrieval = per_source + return retriever + + +def _relationships_config(): + return RetrievalConfig(graph={"seed_strategy": "relationships"}) + + +class TestDefaults: + def test_measured_best_configuration_is_the_default(self): + options = GraphRetrievalConfig() + + assert options.seed_strategy == "entities" + assert options.passage_nodes is True + assert options.blend_vector is True + + def test_a_source_with_no_per_source_config_gets_the_defaults(self): + assert _retriever()._graph_options("src") == GraphRetrievalConfig() + + def test_existing_retrieval_configs_validate_without_the_new_block(self): + """Source configs saved before this existed carry no ``graph`` key.""" + config = RetrievalConfig.model_validate({"retriever": "graphrag"}) + + assert config.graph == GraphRetrievalConfig() + + @pytest.mark.parametrize( + "bad", [{"seed_strategy": "vector"}, {"seed_strategy": "union"}, {"damping": 0.5}] + ) + def test_retired_and_unknown_options_are_rejected(self, bad): + with pytest.raises(ValidationError): + GraphRetrievalConfig.model_validate(bad) + + +class TestEntitySeeding: + def test_seeds_from_entities_and_never_reads_facts(self): + store = _StubStore() + + rows = _retriever()._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["n1"] + assert store.calls == ["entities"] + + +class TestRelationshipSeeding: + def test_seeds_from_facts_when_the_source_asks_for_it(self): + store = _StubStore() + retriever = _retriever({"src": _relationships_config()}) + + rows = retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["f1", "f2"] + assert store.calls == ["facts"] + + def test_reads_the_option_from_a_plain_dict_config_too(self): + store = _StubStore() + retriever = _retriever({"src": {"graph": {"seed_strategy": "relationships"}}}) + + retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert store.calls == ["facts"] + + def test_falls_back_to_entities_without_fact_embeddings(self): + """A source built before fact embeddings were recorded still retrieves.""" + store = _StubStore(fact_rows=[]) + retriever = _retriever({"src": _relationships_config()}) + + rows = retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["n1"] + assert store.calls == ["facts", "entities"] + + def test_options_are_per_source(self): + store = _StubStore() + retriever = _retriever({"other": _relationships_config()}) + + retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert store.calls == ["entities"] + + def test_a_malformed_stored_option_falls_back_to_the_defaults(self): + """A bad value must not take graph retrieval down with it.""" + retriever = _retriever({"src": {"graph": {"seed_strategy": "nonsense"}}}) + + assert retriever._graph_options("src") == GraphRetrievalConfig() diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index baa19604..02a315db 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -157,6 +157,122 @@ class TestGraphStoreLive: finally: store.delete_by_source(source_id) + def test_seed_nodes_from_facts_returns_both_endpoints_of_the_match( + self, store, source_id + ): + """Fact seeding's whole point: the question matches the *relationship*, + and both of its endpoints become seeds — including the one the question + never names.""" + try: + alder = store.upsert_node(source_id, "Alder", "alder", "service", "d") + quill = store.upsert_node(source_id, "Quill", "quill", "store", "d") + birch = store.upsert_node(source_id, "Birch", "birch", "service", "d") + ridge = store.upsert_node(source_id, "Ridge", "ridge", "store", "d") + store.add_edge( + source_id, alder, quill, "streams_to", "Alder streams to Quill", + 1.0, ["c1"], fact_embedding=_embedding(1.0), + ) + store.add_edge( + source_id, birch, ridge, "streams_to", "Birch streams to Ridge", + 1.0, ["c2"], fact_embedding=_embedding(-1.0), + ) + + rows = store.seed_nodes_from_facts( + source_id, _embedding(1.0), fact_limit=1, limit=10 + ) + + assert {row["name"] for row in rows} == {"Alder", "Quill"} + assert all(row["distance"] <= 1.0 for row in rows) + finally: + store.delete_by_source(source_id) + + def test_seed_nodes_from_facts_is_empty_without_fact_embeddings( + self, store, source_id + ): + """A source ingested before fact embeddings existed returns nothing, + which is the signal the retriever falls back to name matching on.""" + try: + a = store.upsert_node(source_id, "A", "a") + b = store.upsert_node(source_id, "B", "b") + store.add_edge(source_id, a, b, "rel") + + assert store.seed_nodes_from_facts(source_id, _embedding(1.0)) == [] + finally: + store.delete_by_source(source_id) + + def test_add_edge_skips_self_loops(self, store, source_id): + """A relationship whose endpoints resolve to one node is noise. + + A self-loop feeds a node's PageRank mass straight back to itself, and a + real extraction produced 121 of them on a 98-page corpus. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "desc a") + assert ( + store.add_edge(source_id, a, a, "related", "a relates to a", 1.0, ["c1"]) + is None + ) + assert store.get_subgraph(source_id, [a], hops=1)["edges"] == [] + finally: + store.delete_by_source(source_id) + + def test_add_edge_merges_a_repeated_pair(self, store, source_id): + """The same relationship seen in many chunks is one edge, not many rows. + + ``graph_edges`` carries no uniqueness constraint, so re-extracting a + relationship used to insert a row per chunk — 19.9% of a real corpus's + edges — inflating traversal weight and wasting the subgraph fetch + budget. The surviving row keeps the strongest weight and both chunk ids. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "desc a") + b = store.upsert_node(source_id, "B", "b", "thing", "desc b") + first = store.add_edge(source_id, a, b, "related", "d", 2.0, ["chunk-1"]) + second = store.add_edge(source_id, a, b, "related", "d", 5.0, ["chunk-2"]) + + assert second == first + edges = store.get_subgraph(source_id, [a, b], hops=1)["edges"] + assert len(edges) == 1 + assert float(edges[0]["weight"]) == 5.0 + + # Both chunks are still recorded as evidence for the merged edge. + conn = store._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + "SELECT source_chunk_ids FROM graph_edges WHERE id = %s;", (first,) + ) + chunk_ids = cursor.fetchone()[0] + finally: + cursor.close() + conn.rollback() + assert sorted(chunk_ids) == ["chunk-1", "chunk-2"] + finally: + store.delete_by_source(source_id) + + def test_get_subgraph_keeps_the_heaviest_edges_when_capped( + self, store, source_id, monkeypatch + ): + """A capped fetch must drop the weakest edges, not an arbitrary subset. + + The cap is applied with ``LIMIT``; without an ordering Postgres is free + to return any rows at all, so a dense graph silently retrieves a random + neighbourhood. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "d") + b = store.upsert_node(source_id, "B", "b", "thing", "d") + c = store.upsert_node(source_id, "C", "c", "thing", "d") + store.add_edge(source_id, a, b, "light", "d", 1.0, ["c1"]) + store.add_edge(source_id, a, c, "heavy", "d", 9.0, ["c1"]) + + monkeypatch.setattr(store_module, "MAX_SUBGRAPH_EDGES", 1) + edges = store.get_subgraph(source_id, [a], hops=1)["edges"] + + assert [e["type"] for e in edges] == ["heavy"] + finally: + store.delete_by_source(source_id) + def test_apply_chunk_writes_nodes_links_and_edges(self, store, source_id): """One transactional write: entities linked to the chunk, edges added, and a bare relationship endpoint upserted but not chunk-linked.""" @@ -203,18 +319,23 @@ class TestGraphStoreLive: store.delete_by_source(source_id) def test_self_loop_degree_agrees_across_paths(self, store, source_id): - """``add_edge``'s incremental +1 and ``set_node_degrees`` recompute must - agree on a self-loop (count it once).""" + """``add_edge``'s incremental bump and ``set_node_degrees`` recompute must + agree on a self-loop. + + They now agree on zero rather than one: the self-loop is rejected at + write time, so neither path has an edge to count. The property under + test is that the two paths agree, not the number they agree on. + """ try: node = store.upsert_node(source_id, "Solo", "solo") - store.add_edge(source_id, node, node, "self") + assert store.add_edge(source_id, node, node, "self") is None incremental = store.get_node_by_normalized(source_id, "solo")["degree"] - assert incremental == 1 + assert incremental == 0 store.set_node_degrees(source_id) recomputed = store.get_node_by_normalized(source_id, "solo")["degree"] - assert recomputed == 1 + assert recomputed == incremental == 0 finally: store.delete_by_source(source_id) diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 8943174c..e4d1bf46 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -45,6 +45,29 @@ def _patch_embed(monkeypatch): ) +@pytest.fixture(autouse=True) +def _entity_only_ranking(monkeypatch): + """Pin the ranking path these tests were written for. + + Everything here exercises entity-only PPR ranking without vector blending, + driven through ``MagicMock`` stores. The shipped default now walks the + passages and blends with vector search — covered end to end in + ``tests/graphrag/test_retriever_default_path.py`` with a store that returns + real values. Pinning keeps each test here asserting what it was written to + assert, rather than whatever a mock happens to return on a path it never set + up. + """ + from docsgpt.storage.db.source_config import GraphRetrievalConfig + + monkeypatch.setattr( + GraphRAGRetriever, + "_graph_options", + lambda self, source_id: GraphRetrievalConfig( + passage_nodes=False, blend_vector=False + ), + ) + + # ── Fallback to ClassicRAG ──────────────────────────────────────────────────── From 6b9b1933371b1e9bcc8e32ac3b3b037d918ec901 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 09/27] feat(agents): let agents walk a source's graph when they can search it A one-shot graph ranking diffuses over the whole neighbourhood; a question whose answer sits two hops away is better served by following the edges. The graph_search tool gives an agent search_entities, get_relationships and read_entity_pages over its graph sources. On a multi-hop corpus where the bridging entity is never named, answers went from 0/10 with classic vector retrieval to 10/10 with the tool, end to end through /stream. The tool has no setting of its own. It is offered exactly where the agent can already search: agentic and research agents always, a classic agent only for sources exposed as a search tool. A graph source left at prefetch is used for ranking only. Tests now pin GRAPHRAG_ENABLED to its shipped default, as CI has: with a dev .env enabling it, every agent test's graph check read the developer's real database and left a pool to it behind. --- docsgpt/agents/agentic_agent.py | 2 + docsgpt/agents/classic_agent.py | 2 + docsgpt/agents/research_agent.py | 2 + docsgpt/agents/tools/graph_search.py | 277 +++++++++++++++++++++++ tests/conftest.py | 15 ++ tests/graphrag/test_graph_search_tool.py | 172 ++++++++++++++ 6 files changed, 470 insertions(+) create mode 100644 docsgpt/agents/tools/graph_search.py create mode 100644 tests/graphrag/test_graph_search_tool.py diff --git a/docsgpt/agents/agentic_agent.py b/docsgpt/agents/agentic_agent.py index b83c485d..0b87485c 100644 --- a/docsgpt/agents/agentic_agent.py +++ b/docsgpt/agents/agentic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -33,6 +34,7 @@ class AgenticAgent(BaseAgent): ) -> Generator[Dict, None, None]: tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/classic_agent.py b/docsgpt/agents/classic_agent.py index 2bc25130..5956ed3c 100644 --- a/docsgpt/agents/classic_agent.py +++ b/docsgpt/agents/classic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -38,6 +39,7 @@ class ClassicAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() if self.retriever_config: add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py index 1d19a7a9..ee92507a 100644 --- a/docsgpt/agents/research_agent.py +++ b/docsgpt/agents/research_agent.py @@ -6,6 +6,7 @@ from typing import Dict, Generator, List, Optional from docsgpt.agents.base import BaseAgent from docsgpt.agents.tool_executor import ToolExecutor +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import ( INTERNAL_TOOL_ID, add_internal_search_tool, @@ -277,6 +278,7 @@ class ResearchAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py new file mode 100644 index 00000000..c2c15a7b --- /dev/null +++ b/docsgpt/agents/tools/graph_search.py @@ -0,0 +1,277 @@ +"""Let the model search the knowledge graph itself, one edge at a time. + +Graph retrieval normally runs as a ranker: seed a walk from the question, +diffuse mass over a subgraph, hand back the highest-scoring chunks. Measured +across five corpora that never beat plain vector search, because a question +whose answer lives two documents away has nothing in it for the seeding step to +match — the bridging entity is named in the *first* document, not the question. + +Exposing the graph as tools removes the guess. The model can look up the +service, read which store it names, then fetch that store's page: the chain +followed deliberately rather than approximated by a diffusion. On a corpus built +so that vector search cannot shortcut the chain, this took two-hop answers from +1/8 to 8/8, against 0.40 for vector and 0.47 for one-shot graph retrieval. + +It is not a general win, and is deliberately not a default. On ordinary prose +documentation it *lost* to plain vector search (0.50 against 0.90): it answers +well when a question names an entity and wanders when the question is a task +description. It also costs several model round-trips per answer instead of one. +So it is offered only where a source owner has already chosen search over +prefetch — the per-source exposure setting, or an agentic/research agent — and +suits content that is genuinely chain-structured: runbooks, service catalogues, +infrastructure inventories. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, List, Optional + +from docsgpt.agents.tools.base import Tool +from docsgpt.core.settings import settings + +logger = logging.getLogger(__name__) + +GRAPH_TOOL_ID = "graph_search" +MAX_PAGE_CHARS = 1500 + + +class GraphSearchTool(Tool): + """Entity lookup, relationship traversal and page reads over a source's graph.""" + + internal = True + + def __init__(self, config: Dict): + self.config = config or {} + self._store = None + self.retrieved_docs: List[Dict] = [] + + # -- plumbing ------------------------------------------------------------ + def _sources(self) -> List[str]: + source = self.config.get("source") or {} + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + return [str(s) for s in active if s] + + def _get_store(self): + if self._store is None: + from docsgpt.graphrag.store import GraphStore + + self._store = GraphStore() + return self._store + + def _embed(self, text: str) -> Optional[List[float]]: + try: + from docsgpt.vectorstore.base import get_embeddings + + return get_embeddings().embed_query(text) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool could not embed the query: {e}") + return None + + # -- actions ------------------------------------------------------------- + def execute_action(self, action_name: str, **kwargs): + if not settings.GRAPHRAG_ENABLED: + return "The knowledge graph is not enabled for this deployment." + if not self._sources(): + return "No graph-backed sources are configured." + try: + if action_name == "search_entities": + return self._search_entities(**kwargs) + if action_name == "get_relationships": + return self._get_relationships(**kwargs) + if action_name == "read_entity_pages": + return self._read_entity_pages(**kwargs) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True) + return "The graph lookup failed." + return f"Unknown action: {action_name}" + + def _search_entities(self, **kwargs) -> str: + query = str(kwargs.get("query") or "").strip() + if not query: + return "Error: 'query' parameter is required." + limit = max(1, min(int(kwargs.get("k") or 8), 25)) + + embedding = self._embed(query) + if embedding is None: + return "Entity search is unavailable." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for row in store.search_nodes_by_embedding(source_id, embedding, k=limit): + similarity = 1.0 - float(row.get("distance") or 0.0) + description = (row.get("description") or "").strip() + suffix = f" — {description[:160]}" if description else "" + lines.append(f"- {row['name']} (match {similarity:.2f}){suffix}") + if not lines: + return f"No entities found for {query!r}." + return "Entities:\n" + "\n".join(lines[:limit]) + + def _get_relationships(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for edge in store.entity_relationships(source_id, entity): + relation = edge.get("type") or "related to" + lines.append(f"- {edge['source']} --{relation}--> {edge['target']}") + if not lines: + return ( + f"No relationships found for {entity!r}. Try search_entities first " + "to get the exact name used in the graph." + ) + return f"Relationships for {entity!r}:\n" + "\n".join(lines) + + def _read_entity_pages(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + parts: List[str] = [] + for source_id in self._sources(): + for page in store.entity_pages(source_id, entity): + metadata = page.get("metadata") or {} + title = ( + metadata.get("file_path") + or metadata.get("title") + or metadata.get("source") + or "document" + ) + text = (page.get("text") or "")[:MAX_PAGE_CHARS] + doc = {"title": title, "text": text, "source": metadata.get("source", "")} + if doc not in self.retrieved_docs: + self.retrieved_docs.append(doc) + parts.append(f"--- {title} ---\n{text}") + if not parts: + return f"No documents mention {entity!r}." + return "\n\n".join(parts) + + # -- metadata ------------------------------------------------------------ + def get_actions_metadata(self): + return [ + { + "name": "search_entities", + "description": ( + "Find named things in the knowledge graph — services, components, " + "settings, people — whose names resemble a query. Use this first to " + "learn the exact name the graph uses before asking for its " + "relationships." + ), + "parameters": { + "properties": { + "query": { + "type": "string", + "description": "What to look for, e.g. a service or component name.", + "filled_by_llm": True, + "required": True, + }, + "k": { + "type": "integer", + "description": "How many entities to return (default 8).", + "filled_by_llm": True, + "required": False, + }, + } + }, + }, + { + "name": "get_relationships", + "description": ( + "List what an entity is connected to, as 'source --relation--> target'. " + "This is how you answer a question about something the question does not " + "name: look up what it points at, then read that thing's pages." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + { + "name": "read_entity_pages", + "description": ( + "Read the documentation an entity appears in, the page it is about " + "first. Use this once you know which entity holds the answer." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + ] + + def get_config_requirements(self): + return {} + + +def build_graph_tool_entry() -> Dict: + """The synthetic ``tools_dict`` entry for the graph tool.""" + tool = GraphSearchTool({}) + actions = [] + for action in tool.get_actions_metadata(): + entry = dict(action) + entry["active"] = True + actions.append(entry) + return {"name": "graph_search", "actions": actions} + + +def sources_have_graph(source: Dict) -> bool: + """Whether any active source actually has a graph to search.""" + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + if not active: + return False + try: + from docsgpt.graphrag.store import GraphStore + + counts = GraphStore().count_nodes_many([str(a) for a in active]) + return any(count > 0 for count in counts.values()) + except Exception as e: # noqa: BLE001 + logger.debug(f"Could not check for graphs: {e}") + return False + + +def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None: + """Add the graph tool when the agent's search-tool sources include a graph. + + No setting of its own: ``retriever_config`` already carries exactly the + sources the agent may *search* — the ones a source owner exposed as a + search tool, or every source for an agentic/research agent — so the graph + tool follows that same per-source exposure choice. A graph source left at + ``prefetch`` in a classic agent is used for ranking only. + """ + if not settings.GRAPHRAG_ENABLED: + return + source = retriever_config.get("source") or {} + if not source.get("active_docs") or not sources_have_graph(source): + return + + entry = build_graph_tool_entry() + # The executor resolves tools by ``id``; this one is synthetic (no DB row). + entry["id"] = GRAPH_TOOL_ID + entry["config"] = {"source": source} + tools_dict[GRAPH_TOOL_ID] = entry + + +def build_graph_tool_config(source: Dict, **_ignored: Any) -> Dict: + """Config for :class:`GraphSearchTool` — it only needs the source ids.""" + return {"source": source} diff --git a/tests/conftest.py b/tests/conftest.py index 2d2e338b..876e6d89 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -183,6 +183,21 @@ def _no_worker_delegation(monkeypatch): monkeypatch.setattr("docsgpt.cache._pubsub_redis_creation_failed", True) +@pytest.fixture(autouse=True) +def _graphrag_off_by_default(monkeypatch): + """Run with GraphRAG at its shipped default (off), as CI does. + + Every agent that gets a search tool checks its sources for a graph, and + that check reads the configured vector database. A dev ``.env`` enabling + GraphRAG sent unrelated agent tests to the developer's real database and + left a pool to it in ``pgconn._POOLS``, failing a live test that asserts it + owns the only pool. Tests that exercise GraphRAG turn it on themselves. + """ + from docsgpt.core.settings import settings + + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False, raising=False) + + @pytest.fixture def mock_llm(): llm = Mock() diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py new file mode 100644 index 00000000..2f63cf13 --- /dev/null +++ b/tests/graphrag/test_graph_search_tool.py @@ -0,0 +1,172 @@ +"""The graph exposed to an agent as callable tools. + +Ranking with a graph never beat vector search in measurement; letting a model +*follow* an edge did, on content where the answer is two documents away. These +tests cover the contract that makes that possible — the tool must return the +relationships verbatim enough for the model to read a name out of them, and +must refuse clearly rather than silently when it has nothing to offer. +""" + +from __future__ import annotations + +from docsgpt.agents.tools.graph_search import ( + GRAPH_TOOL_ID, + GraphSearchTool, + add_graph_search_tool, + build_graph_tool_entry, +) +from docsgpt.core.settings import settings + +SOURCE = {"active_docs": ["src-1"]} + + +class _StubStore: + def __init__(self, nodes=None, relationships=None, pages=None): + self._nodes = nodes or [] + self._relationships = relationships or [] + self._pages = pages or [] + + def search_nodes_by_embedding(self, source_id, embedding, k=10): + return self._nodes[:k] + + def entity_relationships(self, source_id, name, limit=25): + return self._relationships + + def entity_pages(self, source_id, name, limit=4): + return self._pages + + +def _tool(monkeypatch, store, enabled=True): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled) + tool = GraphSearchTool({"source": SOURCE}) + tool._store = store + monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1]) + return tool + + +class TestGating: + def test_reports_when_graphs_are_disabled(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(), enabled=False) + + assert "not enabled" in tool.execute_action("search_entities", query="x") + + def test_reports_when_no_sources_are_configured(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + tool = GraphSearchTool({"source": {"active_docs": []}}) + + assert "No graph-backed sources" in tool.execute_action( + "search_entities", query="x" + ) + + def test_unknown_action_is_named(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "Unknown action" in tool.execute_action("wander") + + +class TestActions: + def test_search_entities_lists_names_with_match_strength(self, monkeypatch): + tool = _tool( + monkeypatch, + _StubStore(nodes=[{"name": "Quill", "distance": 0.2, "description": "A store."}]), + ) + + result = tool.execute_action("search_entities", query="quill") + + assert "Quill" in result + assert "0.80" in result + + def test_search_entities_requires_a_query(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "required" in tool.execute_action("search_entities", query=" ") + + def test_relationships_are_rendered_as_triples(self, monkeypatch): + """The model reads the *target* out of this line to take its next step, + so the target name has to survive rendering intact.""" + tool = _tool( + monkeypatch, + _StubStore( + relationships=[ + {"source": "Alder", "type": "streams_to", "target": "Quill", "description": ""} + ] + ), + ) + + result = tool.execute_action("get_relationships", entity="Alder") + + assert "Alder --streams_to--> Quill" in result + + def test_missing_relationships_suggest_the_next_step(self, monkeypatch): + """A dead end should point at search_entities rather than stop the agent.""" + tool = _tool(monkeypatch, _StubStore(relationships=[])) + + result = tool.execute_action("get_relationships", entity="Nope") + + assert "search_entities" in result + + def test_pages_are_titled_truncated_and_recorded(self, monkeypatch): + tool = _tool( + monkeypatch, + _StubStore( + pages=[{"metadata": {"file_path": "quill-store.md"}, "text": "x" * 5000}] + ), + ) + + result = tool.execute_action("read_entity_pages", entity="Quill") + + assert "--- quill-store.md ---" in result + assert len(result) < 3000 + # Accumulated so the answer can cite what the walk actually read. + assert tool.retrieved_docs[0]["title"] == "quill-store.md" + + def test_pages_absent_is_stated_plainly(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(pages=[])) + + assert "No documents" in tool.execute_action("read_entity_pages", entity="Quill") + + +class TestWiring: + def test_entry_exposes_every_action(self): + entry = build_graph_tool_entry() + + assert {a["name"] for a in entry["actions"]} == { + "search_entities", + "get_relationships", + "read_entity_pages", + } + assert all(action["active"] for action in entry["actions"]) + + def test_not_added_when_graphs_are_disabled(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools == {} + + def test_not_added_when_the_sources_have_no_graph(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools == {} + + def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID + assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE From 5e67f8c9276fefbf449c63092dba42306825f65a Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 10/27] feat(settings): graph retrieval options in a graph source's retrieval settings Exposes the three per-source graph options in the source's retrieval settings, shown only when the retriever is graphrag: where the walk starts (entities or relationships), whether passages join the walk, and whether vector hits are blended in. Defaults match the backend's measured-best configuration and are filled in for sources saved before the options existed. A note points at the "search tool" exposure, which is what offers the graph to an agent. --- frontend/src/locale/de.json | 13 +++ frontend/src/locale/en.json | 13 +++ frontend/src/locale/es.json | 13 +++ frontend/src/locale/jp.json | 13 +++ frontend/src/locale/ru.json | 13 +++ frontend/src/locale/zh-TW.json | 13 +++ frontend/src/locale/zh.json | 13 +++ frontend/src/models/misc.ts | 12 ++ .../components/RetrievalOptions.test.tsx | 43 +++++++ .../settings/components/RetrievalOptions.tsx | 110 ++++++++++++++++++ 10 files changed, 256 insertions(+) diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index e64ed52f..b914fbfd 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -268,6 +268,19 @@ }, "exposureHint": "Lade diese Quelle vorab in den Prompt oder lass den Agenten sie bei Bedarf als Werkzeug durchsuchen." }, + "graphRetrieval": { + "title": "Graph-Abruf", + "tag": "ohne Neuimport", + "seedStrategy": "Suche beginnt bei", + "seedStrategyHint": "Entitäten eignen sich für die meisten Dokumente. Beziehungen erreichen auch Entitäten, die in der Frage nicht vorkommen – ideal für Inhalte, die beschreiben, wie Dinge zusammenhängen.", + "seedEntities": "Entitäten (empfohlen)", + "seedRelationships": "Beziehungen", + "passageNodes": "Textabschnitte in die Suche einbeziehen", + "passageNodesHint": "Ein Abschnitt wird gefunden, wenn er zur Frage passt oder mit etwas Passendem verbunden ist. Am besten bei Graphen, die mit dieser Version erstellt wurden.", + "blendVector": "Mit Vektorsuche kombinieren", + "blendVectorHint": "Ergänzt Ergebnisse der Vektorsuche, damit kein Abschnitt verloren geht, den der Graph übersieht.", + "agentToolHint": "Agenten können diesen Beziehungen auch selbst folgen, wenn die Bereitstellung dieser Quelle „Suchwerkzeug auf Abruf“ ist oder ein agentischer Agent sie nutzt." + }, "prescreen": { "enable": "LLM-Vorfilterung aktivieren", "warning": "Ruft eine größere Kandidatenmenge ab und filtert sie mit einem LLM. Das erhöht Latenz und Kosten pro Anfrage.", diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 03f0f0dd..ba67a1cb 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -272,6 +272,19 @@ }, "exposureHint": "Pre-fetch this source into the prompt, or let the agent search it on demand as a tool." }, + "graphRetrieval": { + "title": "Graph retrieval", + "tag": "no re-ingest", + "seedStrategy": "Start the walk from", + "seedStrategyHint": "Entities suit most documents. Relationships can reach an entity the question never names, and suit content that describes how things connect.", + "seedEntities": "Entities (recommended)", + "seedRelationships": "Relationships", + "passageNodes": "Include passages in the walk", + "passageNodesHint": "Lets a passage be found both by matching the question and by being connected to what does. Works best on graphs built with this version.", + "blendVector": "Blend with vector search", + "blendVectorHint": "Adds plain vector search results, so a passage the graph misses is not lost.", + "agentToolHint": "Agents can also follow these relationships themselves when this source's exposure is “On-demand search tool”, or when an agentic agent uses it." + }, "prescreen": { "enable": "Enable LLM prescreen", "warning": "Fetches a larger candidate set and uses an LLM to filter it. This adds query-time latency and cost.", diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index c7b4f794..f30cadd7 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -268,6 +268,19 @@ }, "exposureHint": "Precarga esta fuente en el prompt, o deja que el agente la busque bajo demanda como herramienta." }, + "graphRetrieval": { + "title": "Recuperación por grafo", + "tag": "sin reingesta", + "seedStrategy": "Iniciar el recorrido desde", + "seedStrategyHint": "Las entidades funcionan para la mayoría de documentos. Las relaciones pueden llegar a una entidad que la pregunta no menciona; son ideales para contenido que describe cómo se conectan las cosas.", + "seedEntities": "Entidades (recomendado)", + "seedRelationships": "Relaciones", + "passageNodes": "Incluir fragmentos en el recorrido", + "passageNodesHint": "Un fragmento puede encontrarse por coincidir con la pregunta o por estar conectado con lo que coincide. Funciona mejor en grafos creados con esta versión.", + "blendVector": "Combinar con búsqueda vectorial", + "blendVectorHint": "Añade resultados de la búsqueda vectorial para no perder fragmentos que el grafo pase por alto.", + "agentToolHint": "Los agentes también pueden seguir estas relaciones por sí mismos cuando la exposición de esta fuente es «Herramienta de búsqueda bajo demanda» o cuando la usa un agente agéntico." + }, "prescreen": { "enable": "Habilitar preselección con LLM", "warning": "Obtiene un conjunto de candidatos más grande y usa un LLM para filtrarlo. Esto añade latencia y costo por consulta.", diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index dc022fe8..e475f954 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -268,6 +268,19 @@ }, "exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。" }, + "graphRetrieval": { + "title": "グラフ検索", + "tag": "再取り込み不要", + "seedStrategy": "探索の開始点", + "seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。", + "seedEntities": "エンティティ(推奨)", + "seedRelationships": "リレーション", + "passageNodes": "パッセージを探索に含める", + "passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。", + "blendVector": "ベクトル検索と組み合わせる", + "blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。", + "agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。" + }, "prescreen": { "enable": "LLMプリスクリーニングを有効にする", "warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。", diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index dd7aef1e..ba064d54 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -268,6 +268,19 @@ }, "exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту." }, + "graphRetrieval": { + "title": "Поиск по графу", + "tag": "без повторной загрузки", + "seedStrategy": "Начинать обход с", + "seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.", + "seedEntities": "Сущностей (рекомендуется)", + "seedRelationships": "Связей", + "passageNodes": "Включать фрагменты в обход", + "passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.", + "blendVector": "Сочетать с векторным поиском", + "blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.", + "agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент." + }, "prescreen": { "enable": "Включить предварительный отбор LLM", "warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.", diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 16236cec..c8c3e7e1 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -268,6 +268,19 @@ }, "exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。" }, + "graphRetrieval": { + "title": "圖譜檢索", + "tag": "無需重新匯入", + "seedStrategy": "走訪起點", + "seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。", + "seedEntities": "實體(建議)", + "seedRelationships": "關係", + "passageNodes": "將段落納入走訪", + "passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。", + "blendVector": "與向量檢索結合", + "blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。", + "agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。" + }, "prescreen": { "enable": "啟用 LLM 預篩選", "warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。", diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 5b4cf5a6..40e8fa0e 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -268,6 +268,19 @@ }, "exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。" }, + "graphRetrieval": { + "title": "图谱检索", + "tag": "无需重新导入", + "seedStrategy": "遍历起点", + "seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。", + "seedEntities": "实体(推荐)", + "seedRelationships": "关系", + "passageNodes": "将段落纳入遍历", + "passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。", + "blendVector": "与向量检索结合", + "blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。", + "agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。" + }, "prescreen": { "enable": "启用 LLM 预筛选", "warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。", diff --git a/frontend/src/models/misc.ts b/frontend/src/models/misc.ts index fc8b0782..c0c5efb2 100644 --- a/frontend/src/models/misc.ts +++ b/frontend/src/models/misc.ts @@ -32,6 +32,17 @@ export type SourcePrescreenConfig = { max_keep?: number; // default 8, <= candidate_k }; +// Where the graph walk starts: matching entities, or matching relationships +// ("A streams_to B"), which can reach an entity the question never names. +export type GraphSeedStrategy = 'entities' | 'relationships'; + +// Query-time graph retrieval knobs (graphrag only; live, no re-ingest). +export type SourceGraphRetrievalConfig = { + seed_strategy?: GraphSeedStrategy; // default 'entities' + passage_nodes?: boolean; // default true + blend_vector?: boolean; // default true +}; + // Query-time retrieval knobs (live; no re-ingest needed). export type SourceRetrievalConfig = { retriever?: string; // default 'classic' (only option for now) @@ -40,6 +51,7 @@ export type SourceRetrievalConfig = { score_threshold?: number | null; // default null rephrase_query?: boolean; // default true prescreen?: SourcePrescreenConfig | null; // null = off + graph?: SourceGraphRetrievalConfig; // graphrag retriever only }; // Ingest-time GraphRAG extraction knobs (only used when kind === 'graphrag'). diff --git a/frontend/src/settings/components/RetrievalOptions.test.tsx b/frontend/src/settings/components/RetrievalOptions.test.tsx index 92142b23..264c1794 100644 --- a/frontend/src/settings/components/RetrievalOptions.test.tsx +++ b/frontend/src/settings/components/RetrievalOptions.test.tsx @@ -146,6 +146,11 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { batch_size: 5, max_keep: 10, }, + graph: { + seed_strategy: 'relationships', + passage_nodes: false, + blend_vector: false, + }, }, graph: { extraction_model: null, @@ -157,6 +162,44 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { }); }); +describe('graph retrieval options', () => { + it('defaults to the measured-best configuration', () => { + expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('fills the defaults for a source saved before the options existed', () => { + const opts = configToOptions({ retrieval: { retriever: 'graphrag' } }); + expect(opts.retrieval.graph).toEqual( + DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph, + ); + }); + + it('honors stored options and fills only the missing ones', () => { + const opts = configToOptions({ + retrieval: { graph: { seed_strategy: 'relationships' } }, + }); + expect(opts.retrieval.graph).toEqual({ + seed_strategy: 'relationships', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('writes the options into the retrieval block', () => { + const v = clone(DEFAULT_RETRIEVAL_OPTIONS); + v.retrieval.graph.blend_vector = false; + expect(optionsToConfig(v).retrieval?.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: false, + }); + }); +}); + describe('isPrescreenConfigValid', () => { const withPrescreen = ( chunks: number, diff --git a/frontend/src/settings/components/RetrievalOptions.tsx b/frontend/src/settings/components/RetrievalOptions.tsx index e2723358..b9dec5d6 100644 --- a/frontend/src/settings/components/RetrievalOptions.tsx +++ b/frontend/src/settings/components/RetrievalOptions.tsx @@ -16,6 +16,7 @@ import { import { Switch } from '../../components/ui/switch'; import type { ChunkingStrategy, + GraphSeedStrategy, RetrievalExposure, SourceConfig, } from '../../models/misc'; @@ -59,6 +60,11 @@ export type RetrievalOptionsValue = { batch_size: number; max_keep: number; }; + graph: { + seed_strategy: GraphSeedStrategy; + passage_nodes: boolean; + blend_vector: boolean; + }; }; graph: { extraction_model: string | null; @@ -85,6 +91,12 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = { enabled: false, ...DEFAULT_PRESCREEN, }, + // The configuration that measured best across the corpora tested. + graph: { + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }, }, graph: { extraction_model: null, @@ -204,6 +216,7 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { const chunking = config?.chunking ?? {}; const retrieval = config?.retrieval ?? {}; const prescreen = retrieval.prescreen ?? null; + const retrievalGraph = retrieval.graph ?? {}; const graph = config?.graph ?? {}; const d = DEFAULT_RETRIEVAL_OPTIONS; return { @@ -228,6 +241,14 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { batch_size: prescreen?.batch_size ?? DEFAULT_PRESCREEN.batch_size, max_keep: prescreen?.max_keep ?? DEFAULT_PRESCREEN.max_keep, }, + graph: { + seed_strategy: + retrievalGraph.seed_strategy ?? d.retrieval.graph.seed_strategy, + passage_nodes: + retrievalGraph.passage_nodes ?? d.retrieval.graph.passage_nodes, + blend_vector: + retrievalGraph.blend_vector ?? d.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: graph.extraction_model ?? d.graph.extraction_model, @@ -271,6 +292,11 @@ export function optionsToConfig(value: RetrievalOptionsValue): SourceConfig { max_keep: ps.max_keep, } : null, + graph: { + seed_strategy: value.retrieval.graph.seed_strategy, + passage_nodes: value.retrieval.graph.passage_nodes, + blend_vector: value.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: value.graph.extraction_model?.trim() @@ -399,6 +425,12 @@ export default function RetrievalOptions({ }); }; + const setGraphRetrieval = ( + patch: Partial, + ) => { + setRetrieval({ graph: { ...value.retrieval.graph, ...patch } }); + }; + const modelOptions = useMemo(() => { const builtin: Model[] = []; const user: Model[] = []; @@ -611,6 +643,84 @@ export default function RetrievalOptions({ )} + {/* Graph retrieval group (graphrag only; live, so shown when testing too) */} + {isGraphRAG && ( +
+ +

+ {tr('graphRetrieval.agentToolHint')} +

+ +
+ + + + + + + setGraphRetrieval({ passage_nodes: checked }) + } + /> + + + + + setGraphRetrieval({ blend_vector: checked }) + } + /> + +
+
+ )} + {/* Graph extraction group (graphrag only; re-ingest required to apply) */} {isGraphRAG && !queryOnly && (
From 15bdda8554cc332230310d2f7fdfa9f9e9af80eb Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:42 +0100 Subject: [PATCH 11/27] fix(worker): know you are in a worker from any thread, not only the task's Celery records the executing task on the thread that runs it, so a thread that task starts sees none. The embeddings client and read_document both decided "am I in a worker?" from that alone, and from any other thread took the web-process branch: dispatch to the worker they were running in and block on the result. Celery refuses that get() ("Never call result.get() within a task!"), so the embed failed and latched the 30s dispatch cooldown for every caller after it; with joins allowed, read_document would instead wait on a parsing queue only its own busy process serves. Threads inside tasks are not hypothetical: per-source retrieval fans out to a pool, so a scheduled or webhook agent searching several sources embedded from pool threads. Graph extraction did too, which failed every chunk of a build. in_worker() in celery_init answers for the whole process. Celery's task_join_will_block is process-wide and set for every blocking pool (prefork, solo, threads) -- exactly the condition under which dispatch-and-wait goes wrong; eventlet/gevent leave it unset, so the task's own thread still counts through current_worker_task. Verified with real workers on each blocking pool: from a thread a task started, the old check dispatched and hit the error, the new one embedded locally. --- docsgpt/agents/tools/read_document.py | 9 ++-- docsgpt/celery_init.py | 23 ++++++++ docsgpt/vectorstore/embeddings_delegated.py | 11 ++-- tests/agents/tools/test_read_document_tool.py | 12 ++--- tests/test_celery.py | 52 +++++++++++++++++++ .../vectorstore/test_embeddings_delegated.py | 26 ++++++++++ 6 files changed, 117 insertions(+), 16 deletions(-) diff --git a/docsgpt/agents/tools/read_document.py b/docsgpt/agents/tools/read_document.py index ab9d2294..94e488df 100644 --- a/docsgpt/agents/tools/read_document.py +++ b/docsgpt/agents/tools/read_document.py @@ -17,8 +17,6 @@ import signal import threading from typing import Any, Callable, Dict, List, Optional -from celery import current_task - from docsgpt.agents.tools.artifact_ref import resolve_artifact_id from docsgpt.agents.tools.attachment_bridge import ( AttachmentBridgeError, @@ -26,6 +24,7 @@ from docsgpt.agents.tools.attachment_bridge import ( match_attachment, ) from docsgpt.agents.tools.base import Tool +from docsgpt.celery_init import in_worker from docsgpt.core.json_schema_utils import ( JsonSchemaValidationError, normalize_json_schema_payload, @@ -229,9 +228,9 @@ class ReadDocumentTool(Tool): # (floored at DOCUMENT_PARSE_TIMEOUT). timeout = parse_timeout_for_size(self._input_size) - # ``current_task`` is a Celery proxy: truthy only while this runs inside a worker task, - # falsy in the web process (the bare proxy is NOT identity-None, so test truthiness). - if current_task: + # Process-wide, not the thread-local ``current_task``: a thread a task starts has no + # task of its own, and dispatching from there is the self-deadlock described above. + if in_worker(): from docsgpt.worker import run_parse_document try: diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py index 4272eac1..5e11d3e6 100644 --- a/docsgpt/celery_init.py +++ b/docsgpt/celery_init.py @@ -172,6 +172,29 @@ def _run_version_check(*args, **kwargs): celery = make_celery() celery.config_from_object("docsgpt.celeryconfig") + +def in_worker() -> bool: + """True anywhere in a Celery worker process, on any thread. + + ``current_worker_task`` alone is not enough: Celery records the executing + task on the thread that runs it, so a thread the task starts sees none and + would take the web-process branch — dispatching to the worker it is running + in and blocking on the result. Celery refuses that ``get()`` ("Never call + result.get() within a task!"), or, where joins are allowed, it waits on a + queue only this busy process serves. + + ``task_join_will_block`` is process-wide and set for every blocking pool + (prefork, solo, threads) — exactly the condition under which dispatching + and waiting goes wrong. eventlet/gevent leave it unset, so the task's own + thread still counts through ``current_worker_task``. + + Returns: + bool: Whether this call is running inside a worker process. + """ + from celery.result import task_join_will_block + + return task_join_will_block() or celery.current_worker_task is not None + #: Task-name prefix the package carried before the rename to ``docsgpt``. diff --git a/docsgpt/vectorstore/embeddings_delegated.py b/docsgpt/vectorstore/embeddings_delegated.py index 75f730b8..bc184f34 100644 --- a/docsgpt/vectorstore/embeddings_delegated.py +++ b/docsgpt/vectorstore/embeddings_delegated.py @@ -10,8 +10,9 @@ Celery and the vector comes back. The API pays a broker round trip per query and no resident model. Inside a worker there is nothing to delegate to -- dispatching would queue work -behind the task already running and wait on itself -- so a call made while a -task is executing runs locally, on a model this process loads once and caches. +behind the task already running and wait on itself -- so a call made anywhere in +a worker process, including from a thread a task started, runs locally, on a +model this process loads once and caches. ``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side. Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding @@ -79,11 +80,11 @@ def _forget(result) -> None: def _in_worker() -> bool: - """True when a Celery task is executing in this process.""" + """True anywhere in a Celery worker process -- on any thread, not only the task's.""" try: - from docsgpt.celery_init import celery + from docsgpt.celery_init import in_worker - return celery.current_worker_task is not None + return in_worker() except Exception: return False diff --git a/tests/agents/tools/test_read_document_tool.py b/tests/agents/tools/test_read_document_tool.py index 8913495d..f0e7b177 100644 --- a/tests/agents/tools/test_read_document_tool.py +++ b/tests/agents/tools/test_read_document_tool.py @@ -357,9 +357,9 @@ def test_malformed_json_schema_rejected_before_enqueue(monkeypatch): @pytest.mark.unit def test_dispatch_inline_when_in_worker(monkeypatch): _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - # Inside a worker current_task is truthy -> parse inline, never enqueue (else the + # Inside a worker -> parse inline, never enqueue (else the # parsing queue self-deadlocks the worker that also serves it). - monkeypatch.setattr(rd, "current_task", object()) + monkeypatch.setattr(rd, "in_worker", lambda: True) import docsgpt.api.user.tasks as tasks monkeypatch.setattr( @@ -387,8 +387,8 @@ def test_dispatch_inline_when_in_worker(monkeypatch): @pytest.mark.unit def test_dispatch_enqueues_when_not_in_worker(monkeypatch): _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - # Web process: current_task falsy -> dispatch to the parsing queue, never inline. - monkeypatch.setattr(rd, "current_task", None) + # Web process -> dispatch to the parsing queue, never inline. + monkeypatch.setattr(rd, "in_worker", lambda: False) captured = _patch_task(monkeypatch, payload={"status": "ok", "content": "queued", "truncated": False}) import docsgpt.worker as worker @@ -414,9 +414,9 @@ _TIMED_OUT = "document parsing timed out after" def _inline(monkeypatch, run_parse, *, timeout=0.2) -> ReadDocumentTool: - """Drive the inline branch (current_task truthy) with a patched parse window.""" + """Drive the inline (in-worker) branch with a patched parse window.""" _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - monkeypatch.setattr(rd, "current_task", object()) + monkeypatch.setattr(rd, "in_worker", lambda: True) import docsgpt.api.user.tasks as tasks monkeypatch.setattr( diff --git a/tests/test_celery.py b/tests/test_celery.py index c5b692df..5a3e66be 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -276,3 +276,55 @@ class TestReclaimIsSkippedForEmbeds: from docsgpt.vectorstore.embeddings_delegated import EMBED_TASK assert EMBED_TASK in _NO_RECLAIM_TASKS + + +@pytest.mark.unit +class TestInWorker: + """Whether code runs inside a worker must not depend on which thread asks. + + Celery records the executing task on the thread that runs it, so a thread + that task starts sees no task at all. Code deciding "am I in the worker?" + from that alone takes the web-process branch there: it dispatches to the + worker it is running in and blocks on the result, which Celery refuses + ("Never call result.get() within a task!") or, where joins are allowed, + waits on a queue only this busy process serves. + """ + + @staticmethod + def _ask_from_a_new_thread(): + import threading + + from docsgpt.celery_init import in_worker + + seen = [] + thread = threading.Thread(target=lambda: seen.append(in_worker())) + thread.start() + thread.join() + return seen[0] + + def test_false_outside_a_worker(self): + from docsgpt.celery_init import in_worker + + assert in_worker() is False + assert self._ask_from_a_new_thread() is False + + def test_true_on_a_thread_started_inside_a_worker(self): + # Blocking pools (prefork, solo, threads) mark the whole process as one + # where joining a task would block; ``denied_join_result`` sets exactly + # that flag. + from celery.result import denied_join_result + + with denied_join_result(): + assert self._ask_from_a_new_thread() is True + + def test_true_on_the_task_thread_of_a_non_blocking_pool(self): + # eventlet/gevent pools leave the process flag unset; the thread + # running the task still knows it is in one. + from unittest.mock import PropertyMock + + from docsgpt.celery_init import celery, in_worker + + with patch.object( + type(celery), "current_worker_task", new_callable=PropertyMock, return_value=object() + ): + assert in_worker() is True diff --git a/tests/vectorstore/test_embeddings_delegated.py b/tests/vectorstore/test_embeddings_delegated.py index 060717ea..b9759510 100644 --- a/tests/vectorstore/test_embeddings_delegated.py +++ b/tests/vectorstore/test_embeddings_delegated.py @@ -72,6 +72,32 @@ class TestInsideAWorker: assert vector == [1.0, 2.0] celery.send_task.assert_not_called() + def test_a_thread_started_inside_the_worker_embeds_locally(self): + """The task's own thread is not the only one in a worker. + + Graph extraction and per-source retrieval both fan out to thread pools + inside tasks. The check used to read the task off the current thread + only, so from those threads it dispatched to the worker it was running + in -- and Celery refuses that ``get()`` inside a worker, failing the + call and latching the 30s dispatch cooldown for every caller after it. + """ + from celery.result import denied_join_result + + from docsgpt.celery_init import celery + + local = MagicMock() + local.embed_documents.return_value = [[1.0, 2.0]] + client = DelegatedEmbeddings("some/model") + vectors = [] + with denied_join_result(): + with patch("docsgpt.vectorstore.base.build_local_embeddings", return_value=local): + with patch.object(celery, "send_task") as send_task: + thread = threading.Thread(target=lambda: vectors.append(client.embed_query("hi"))) + thread.start() + thread.join() + assert vectors == [[1.0, 2.0]] + send_task.assert_not_called() + def test_the_local_model_is_built_once(self): local = MagicMock() local.embed_documents.return_value = [[1.0]] From 5285bb4115e050fa10051a7db522de3bfd5dd70c Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:42 +0100 Subject: [PATCH 12/27] fix(graphrag): give a failed chunk one more attempt before giving up A chunk whose extraction failed was marked failed and the build moved on. The checkpoint treats failed chunks as pending, but nothing ever reran the build, so one transient error -- a provider hiccup, a single response that did not parse -- left a permanent hole in the graph until someone rebuilt the whole source. Failed chunks now get one more attempt after the rest of the build, so a burst of rate limiting has time to pass. Extraction errors, unparseable responses and failed writes are all retried; a chunk is marked failed only when its retry fails too, so it costs at most two calls. The chunks given up on are logged by id, and progress still ends at the total. --- docsgpt/graphrag/extraction.py | 129 ++++++++++++--------- tests/graphrag/test_extraction.py | 181 ++++++++++++++++++++++++++++-- 2 files changed, 247 insertions(+), 63 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index db8ea27a..9539f4fd 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -170,13 +170,13 @@ def _extract_chunk( ) except Exception as exc: logger.warning( - "Graph extraction call failed for chunk %s, skipping: %s", chunk_id, exc + "Graph extraction call failed for chunk %s: %s", chunk_id, exc ) return None parsed = _parse_extraction(response) if parsed is None: logger.warning( - "Graph extraction returned unparseable output for chunk %s; marking it failed.", + "Graph extraction returned unparseable output for chunk %s.", chunk_id, ) return parsed @@ -203,8 +203,10 @@ def extract_graph_for_source( Resumable and idempotent: chunks already marked ``done`` are skipped via the ``graph_ingest_progress`` checkpoint, so a retry never re-extracts (and never re-bills). Processes at most the resolved chunk cap; excess chunks are - reported under ``skipped_over_cap``. A malformed response or an LLM error on - a single chunk marks it ``failed`` and continues — the pipeline never crashes. + reported under ``skipped_over_cap``. A malformed response, an LLM error or a + failed write on a single chunk is retried once after the rest of the build; + a chunk that fails again is marked ``failed`` and the run continues — the + pipeline never crashes. Each chunk is written in a single transaction with one batched embedding call (entity + relationship-endpoint names together). @@ -273,13 +275,9 @@ def extract_graph_for_source( """One chunk's LLM extraction — the only step run concurrently. A chunk spends almost all of its time waiting on the model, so that is - what runs in the pool. Everything else stays on the calling thread: - graph writes, so transactions and the progress checkpoint are exactly - what they were serially, and embedding. Inside a Celery worker the - embeddings client decides to embed locally from the task on the - *current thread's* stack; a pool thread has none, so it would instead - dispatch an embed task to the worker and wait on it, which Celery - refuses inside a task — failing every chunk of the build. + what runs in the pool. Graph writes and embedding stay on the calling + thread, so transactions and the progress checkpoint are exactly what + they were serially and the pool never touches the embeddings client. """ chunk, chunk_id = item text = _chunk_text(chunk) @@ -294,56 +292,85 @@ def extract_graph_for_source( relationships = _build_relationships(extracted["relationships"]) except Exception as exc: logger.warning( - "Graph extraction failed for chunk %s, skipping: %s", chunk_id, exc + "Graph extraction failed for chunk %s: %s", chunk_id, exc ) return chunk_id, "failed", None return chunk_id, "ok", (entities, relationships) + def _write(chunk_id, status, payload) -> bool: + """Apply one prepared chunk to the graph; False when it did not land.""" + nonlocal node_upserts, edges, chunks_processed + if status == "empty": + store.mark_chunk(source_id, chunk_id, "done") + chunks_processed += 1 + return True + if status == "failed": + return False + + entities, relationships = payload + try: + name_embeddings = _embed_names(embedding, entities, relationships) + _embed_facts(embedding, relationships) + chunk_nodes, chunk_edges = store.apply_chunk( + source_id, chunk_id, entities, relationships, name_embeddings + ) + except Exception as exc: + logger.warning( + "Graph extraction embed/write failed for chunk %s: %s", chunk_id, exc + ) + return False + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. + node_upserts += chunk_nodes + edges += chunk_edges + chunks_processed += 1 + return True + workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1)) pool = None if workers > 1 and len(to_process) > 1: pool = ThreadPoolExecutor(max_workers=workers) - # ``map`` yields in submission order, so chunks are still applied in the - # order they were given and a run stays reproducible. - prepared = pool.map(_prepare, to_process) - else: - prepared = (_prepare(item) for item in to_process) + + def _pass(items): + """Extract and write ``items``; return the ones that did not land.""" + if pool is not None: + # ``map`` yields in submission order, so chunks are still applied in + # the order they were given and a run stays reproducible. + prepared = pool.map(_prepare, items) + else: + prepared = (_prepare(item) for item in items) + missed = [] + for item, (chunk_id, status, payload) in zip(items, prepared): + if not _write(chunk_id, status, payload): + missed.append(item) + _report() + return missed try: - for chunk_id, status, payload in prepared: - if status == "empty": - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 - _report() - continue - if status == "failed": - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() - continue - - entities, relationships = payload - try: - # On this thread, not in the pool — see ``_prepare``. - name_embeddings = _embed_names(embedding, entities, relationships) - _embed_facts(embedding, relationships) - chunk_nodes, chunk_edges = store.apply_chunk( - source_id, chunk_id, entities, relationships, name_embeddings - ) - node_upserts += chunk_nodes - edges += chunk_edges - # ``apply_chunk`` marks the chunk done inside the transaction that - # writes its rows, so the checkpoint cannot disagree with the graph - # and a replayed write cannot apply the chunk twice. - chunks_processed += 1 - except Exception as exc: - logger.warning( - "Graph extraction embed/write failed for chunk %s, skipping: %s", - chunk_id, - exc, - ) - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 + missed = _pass(to_process) + if missed: + # A failure is usually transient — a provider error, one response + # that did not parse — and the checkpoint only picks it up on a + # rerun nothing schedules. One more attempt, after the rest of the + # build so a burst of rate limiting has passed, and no more: a + # chunk that cannot be extracted costs at most two calls. + logger.info( + "Graph extraction retrying %d failed chunk(s) for source %s", + len(missed), + source_id, + ) + missed = _pass(missed) + for _, chunk_id in missed: + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks = len(missed) + if missed: + logger.warning( + "Graph extraction gave up on %d chunk(s) for source %s after a retry: %s", + failed_chunks, + source_id, + ", ".join(str(chunk_id) for _, chunk_id in missed), + ) _report() finally: if pool is not None: diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index eda8dd7c..0573cb7c 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -88,6 +88,37 @@ class _StubLLM: return response +class _ScriptedLLM: + """Stub LLM answering per chunk, so results do not depend on call order. + + The extraction pool runs calls concurrently, so a stub that hands out + responses in call order gives each chunk whichever response its thread + happened to grab first. ``script`` maps a chunk's text to the responses + for that chunk, consumed one per call. + """ + + def __init__(self, script): + self._script = {text: list(responses) for text, responses in script.items()} + self.model_id = "stub-model" + self.calls = [] + self._token_usage_source = None + self._request_id = None + + def gen(self, model=None, messages=None, **kwargs): + text = messages[-1]["content"].removeprefix("\n").removesuffix("\n") + self.calls.append(text) + responses = self._script.get(text) + if not responses: + raise AssertionError(f"unexpected extraction call for {text!r}") + response = responses.pop(0) + if isinstance(response, Exception): + raise response + return response + + def calls_for(self, text): + return self.calls.count(text) + + class _StubEmbedding: """Stub embeddings model producing deterministic fixed-dim vectors.""" @@ -365,11 +396,12 @@ class TestExtractionLive: ): """Only the LLM call may run in the extraction pool, never embedding. - Inside a Celery worker the embeddings client decides to embed locally - from the task on the *current thread's* stack. A pool thread has none, - so from there it dispatches an embed task to the worker and waits on - it — which Celery refuses inside a task, so every chunk of a graph - build failed. + Embedding from the pool is what broke every graph build inside a + worker: the embeddings client used to decide "embed locally" from the + task on the *current thread's* stack, so a pool thread dispatched to + the worker instead and Celery refused the wait. ``in_worker`` is + process-wide now, but the pool still has no reason to touch the + embeddings client — it exists to overlap model latency. """ import threading @@ -505,11 +537,13 @@ class TestExtractionLive: entities=[{"name": "Ada", "type": "person", "description": "d"}], relationships=[], ) - llm = _StubLLM([ - "not json at all", - RuntimeError("model exploded"), - good, - ]) + # Each failing chunk fails its retry too; one that recovers on + # retry is covered in ``TestFailedChunksAreRetried``. + llm = _ScriptedLLM({ + "garbage": ["not json at all", "still not json"], + "boom": [RuntimeError("model exploded"), RuntimeError("model exploded again")], + "Ada.": [good], + }) _install_stub_llm(monkeypatch, llm) summary = extract_graph_for_source( @@ -762,7 +796,7 @@ class TestFailedChunksAreReported: import logging store = self._fake_store(monkeypatch, ["c1"]) - _install_stub_llm(monkeypatch, _StubLLM(["not json at all"])) + _install_stub_llm(monkeypatch, _StubLLM(["not json at all", "still not json"])) with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): summary = extract_graph_for_source( @@ -785,7 +819,10 @@ class TestFailedChunksAreReported: import logging self._fake_store(monkeypatch, ["c7"]) - _install_stub_llm(monkeypatch, _StubLLM([RuntimeError("model exploded")])) + _install_stub_llm( + monkeypatch, + _StubLLM([RuntimeError("model exploded"), RuntimeError("model exploded again")]), + ) with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): extract_graph_for_source( @@ -800,6 +837,126 @@ class TestFailedChunksAreReported: assert any("c7" in message for message in messages), messages +@pytest.mark.unit +class TestFailedChunksAreRetried: + """A chunk that fails once gets one more attempt before the build ends. + + Failures are recorded as ``failed`` and the checkpoint treats them as + pending, but nothing ever ran the build again, so a single transient error + — a provider hiccup, one response that did not parse — left a permanent + hole in the graph until someone rebuilt the whole source. Retries run + after the rest of the build, which gives a burst of rate limiting time to + pass, and are bounded at one per chunk so a chunk that can never be + extracted costs at most two calls. + """ + + GOOD = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + + def _fake_store(self, monkeypatch, chunk_ids): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = list(chunk_ids) + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + return store + + def _run(self, chunks, progress=None): + return extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=chunks, + config=SourceConfig(), + request_id="req-retry", + progress_cb=progress, + ) + + @staticmethod + def _marked_failed(store): + return [c.args[1] for c in store.mark_chunk.call_args_list if c.args[2] == "failed"] + + def test_a_transient_failure_is_retried_and_written(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + llm = _ScriptedLLM({"flaky": [RuntimeError("rate limited"), self.GOOD]}) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "flaky")]) + + assert summary["failed_chunks"] == 0 + assert summary["chunks_processed"] == 1 + assert store.apply_chunk.call_args.args[1] == "c1" + assert self._marked_failed(store) == [] + + def test_an_unparseable_response_is_retried(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + _install_stub_llm(monkeypatch, _ScriptedLLM({"odd": ["not json", self.GOOD]})) + + summary = self._run([_chunk("c1", "odd")]) + + assert summary["failed_chunks"] == 0 + assert self._marked_failed(store) == [] + + def test_a_chunk_that_fails_again_is_marked_failed_once(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + llm = _ScriptedLLM({"broken": ["not json", "still not json"]}) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "broken")]) + + assert summary["failed_chunks"] == 1 + assert summary["chunks_processed"] == 0 + assert llm.calls_for("broken") == 2 + assert self._marked_failed(store) == ["c1"] + + def test_only_failed_chunks_are_retried(self, monkeypatch, stub_embedding): + from docsgpt.core.settings import settings + + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + self._fake_store(monkeypatch, ["c1", "c2", "c3"]) + llm = _ScriptedLLM({ + "one": [self.GOOD], + "two": [RuntimeError("timeout"), self.GOOD], + "three": [self.GOOD], + }) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "one"), _chunk("c2", "two"), _chunk("c3", "three")]) + + assert summary["chunks_processed"] == 3 + assert summary["failed_chunks"] == 0 + assert (llm.calls_for("one"), llm.calls_for("two"), llm.calls_for("three")) == (1, 2, 1) + + def test_a_failed_write_is_retried(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + store.apply_chunk.side_effect = [RuntimeError("write failed"), (1, 0)] + _install_stub_llm(monkeypatch, _ScriptedLLM({"text": [self.GOOD, self.GOOD]})) + + summary = self._run([_chunk("c1", "text")]) + + assert summary["failed_chunks"] == 0 + assert summary["chunks_processed"] == 1 + assert self._marked_failed(store) == [] + + def test_progress_ends_at_the_total(self, monkeypatch, stub_embedding): + self._fake_store(monkeypatch, ["c1", "c2"]) + _install_stub_llm(monkeypatch, _ScriptedLLM({ + "fine": [self.GOOD], + "broken": ["not json", "still not json"], + })) + events = [] + + self._run([_chunk("c1", "fine"), _chunk("c2", "broken")], progress=events.append) + + assert all(e["current"] <= e["total"] for e in events) + assert events[-1]["current"] == events[-1]["total"] == 2 + + @pytest.mark.integration class TestSummaryNodeCount: """``nodes`` must describe the graph, not the number of upserts.""" From 022cf69b491048b74c84fff29d9b53b334a3d653 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 13/27] fix(graphrag): fold an entity's singular and plural onto one key canonical_name dropped "es" from every -ches/-ses/-zes plural, so "caches" became "cach" while "cache" stayed "cache" -- the singular and plural landed on two nodes, which is the split the function exists to prevent. The same rule split the words documentation uses most: databases/database, responses/ response, releases/release, sizes/size. An "-es" plural cannot say whether it is cache + "s" or batch + "es", so instead of guessing, both sides now meet at the stem: the singular endings -che/-she/-se/-ze/-xe drop their "e" the way the plurals drop "es", and a singular's -ie folds to -y as -ies already did (cookie/cookies). The key is a merge key that is never shown, so it only has to agree, not be a word. alias, canvas, atlas and bias join the words that only look plural. Found by the naming tests CI was missing: the module had no direct tests. --- docsgpt/graphrag/naming.py | 39 ++++++++++++---- tests/graphrag/test_naming.py | 88 +++++++++++++++++++++++++++++++++++ 2 files changed, 118 insertions(+), 9 deletions(-) create mode 100644 tests/graphrag/test_naming.py diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py index 3dbfa1cb..2a173eea 100644 --- a/docsgpt/graphrag/naming.py +++ b/docsgpt/graphrag/naming.py @@ -13,6 +13,9 @@ plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``, ``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s" would corrupt them into new entities rather than merge anything. +The result is a merge key, never shown to anyone, so it only has to be the +same for a word's singular and plural — not to be a word itself. + Always on: every graph is built with canonical names. """ @@ -33,18 +36,30 @@ _NOT_PLURAL = frozenset( "analysis", "basis", "axis", "https", "rss", "less", "express", "redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos", "always", "sometimes", "series", "docs", "ops", "devops", "sse", + "alias", "canvas", "atlas", "bias", "pandas", } ) +#: Plural endings that drop ``es``, and the singular endings that meet them. +#: ``caches`` cannot say whether it is ``cache`` + "s" or ``cach`` + "es" +#: (as ``batches`` is ``batch`` + "es"), so rather than guess, both +#: ``caches`` and ``cache`` fold to ``cach`` — as ``databases``/``database`` +#: fold to ``databas``. Hardly any real word differs from one of these singulars +#: by its final "e" alone, so the fold merges next to nothing it should not. +_ES_PLURAL = ("ches", "shes", "ses", "zes", "xes") +_E_SINGULAR = ("che", "she", "se", "ze", "xe") + def _singular(word: str) -> str: - """Best-effort singular of one word, biased hard towards leaving it alone. + """Fold one word so its singular and plural share a key, else leave it alone. - Only the endings that are unambiguous in this domain are touched: - ``-ies`` -> ``-y`` (``policies``), ``-ses``/``-xes``/``-zes``/``-ches``/ - ``-shes`` -> drop ``es`` (``indexes``, ``batches``), and a bare trailing - ``s`` on a word long enough to be safe. Everything in :data:`_NOT_PLURAL`, - and anything ending in ``ss``/``us``/``is``, is returned unchanged. + ``-ies`` and a singular's ``-ie`` both fold to ``-y`` (``policies``, + ``cookies``/``cookie``). The ``-es`` endings in :data:`_ES_PLURAL` drop + ``es`` and the singular endings in :data:`_E_SINGULAR` drop their ``e``, so + both sides of an ambiguous plural meet (``caches``/``cache`` -> ``cach``). + Otherwise a bare trailing ``s`` is dropped on a word long enough to be + safe. Everything in :data:`_NOT_PLURAL`, and anything ending in + ``ss``/``us``/``is``, is returned unchanged. """ if len(word) < 4 or word in _NOT_PLURAL: return word @@ -52,8 +67,12 @@ def _singular(word: str) -> str: return word if word.endswith("ies") and len(word) > 4: return word[:-3] + "y" - if word.endswith(("ses", "xes", "zes", "ches", "shes")): + if word.endswith(_ES_PLURAL): return word[:-2] + if word.endswith("ie") and len(word) > 4: + return word[:-2] + "y" + if word.endswith(_E_SINGULAR): + return word[:-1] if word.endswith("s"): return word[:-1] return word @@ -66,12 +85,14 @@ def canonical_name(name: str) -> str: name: The entity name as the model wrote it. Returns: - A lowercase, punctuation-free, singularised key. Returns ``""`` for an - empty or punctuation-only name, which callers treat as "no entity". + A lowercase, punctuation-free key shared by a name's singular and + plural. Returns ``""`` for an empty or punctuation-only name, which + callers treat as "no entity". Examples: ``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``; ``.env file`` and ``env_file`` -> ``env file``; + ``cache`` and ``caches`` -> ``cach``; ``postgres`` stays ``postgres``. """ if not name: diff --git a/tests/graphrag/test_naming.py b/tests/graphrag/test_naming.py new file mode 100644 index 00000000..ff14f36a --- /dev/null +++ b/tests/graphrag/test_naming.py @@ -0,0 +1,88 @@ +"""Tests for canonical entity naming (the key graph nodes are merged on). + +Two failure directions matter. Too little folding splits one entity across +nodes ("agent" / "agents", "VECTOR_STORE" / "vector stores"), so the walk never +connects what the text connects. Too much folding invents entities: stripping +the "s" off ``postgres`` or ``redis`` would merge nothing and create a node no +chunk ever named. +""" + +from __future__ import annotations + +import pytest + +from docsgpt.graphrag.naming import canonical_name, normalize_entity_name + + +@pytest.mark.unit +class TestCanonicalName: + @pytest.mark.parametrize( + "variants, key", + [ + (["VECTOR_STORE", "Vector store", "vector stores", "vector-stores"], "vector store"), + ([".env file", "env_file", "ENV FILE"], "env file"), + (["Celery worker", "Celery workers"], "celery worker"), + (["agent", "Agents", "agents!"], "agent"), + ], + ) + def test_orthographic_variants_share_one_key(self, variants, key): + assert {canonical_name(v) for v in variants} == {key} + + @pytest.mark.parametrize( + "singular, plural", + [ + ("policy", "policies"), + ("index", "indexes"), + ("batch", "batches"), + ("hash", "hashes"), + ("class", "classes"), + ("process", "processes"), + ("status", "statuses"), + ("bus", "buses"), + ("alias", "aliases"), + ("document", "documents"), + ("service", "services"), + # Singulars ending in "e" whose plural also ends in "-es": the + # plural alone cannot say whether to drop "s" or "es". + ("cache", "caches"), + ("database", "databases"), + ("response", "responses"), + ("release", "releases"), + ("case", "cases"), + ("size", "sizes"), + ("cookie", "cookies"), + ], + ) + def test_singular_and_plural_share_one_key(self, singular, plural): + assert canonical_name(singular) == canonical_name(plural) + + def test_the_key_need_not_be_a_word(self): + # It is a merge key, never shown: "cache" and "caches" meet at the + # stem an "-es" plural cannot see past, rather than guessing a form. + assert canonical_name("caches") == "cach" + assert canonical_name("batches") == "batch" + + @pytest.mark.parametrize( + "word", + ["postgres", "kubernetes", "redis", "https", "status", "analysis", "access", "docs", "series"], + ) + def test_words_that_only_look_plural_are_left_alone(self, word): + assert canonical_name(word) == word + + @pytest.mark.parametrize("word", ["class", "corpus", "thesis"]) + def test_ss_us_is_endings_are_never_stripped(self, word): + assert canonical_name(word) == word + + @pytest.mark.parametrize("word", ["aws", "ids", "ops"]) + def test_short_words_are_left_alone(self, word): + assert canonical_name(word) == word + + def test_each_word_of_a_phrase_is_folded(self): + assert canonical_name("Postgres Replicas") == "postgres replica" + + @pytest.mark.parametrize("name", [None, "", " ", "!!!", "--_--"]) + def test_a_name_with_nothing_left_is_no_entity(self, name): + assert canonical_name(name) == "" + + def test_normalize_entity_name_is_the_canonical_key(self): + assert normalize_entity_name("Vector Stores") == canonical_name("Vector Stores") From a20402a6dc09acb3a51e31f05437cf54be208b35 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 14/27] fix(graphrag): give each extraction thread its own LLM Provider-reported usage is kept on the LLM instance (_last_usage) and claimed by whichever call finishes next. With GRAPHRAG_EXTRACTION_WORKERS > 1 (the default is 8) every extraction thread shared one instance, so a call could claim another call's provider counts while its own fell back to the estimate: token_usage rows, and the cost they bill, could be attributed to the wrong call and summed wrong. Each pool thread now builds its own extraction LLM on first use. The calling thread's instance is still built up front, so a misconfigured model fails the run before any chunk is touched. --- docsgpt/graphrag/extraction.py | 24 +++++++++++--- tests/graphrag/test_extraction.py | 55 +++++++++++++++++++++++++++++++ 2 files changed, 75 insertions(+), 4 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 9539f4fd..5ddbe1e1 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -227,6 +227,7 @@ def extract_graph_for_source( source's graph holds after the run — not how many upserts ran, which counts the same entity once per chunk it appears in. """ + import threading from concurrent.futures import ThreadPoolExecutor from docsgpt.graphrag.store import GraphStore @@ -246,9 +247,24 @@ def extract_graph_for_source( embedding = get_embeddings() - llm = _build_extraction_llm( - _resolve_extraction_model(config), user, request_id - ) + model_id = _resolve_extraction_model(config) + # Built here first so a misconfigured model fails the run before any + # chunk is touched; this instance serves the calling thread. + thread_llm = threading.local() + thread_llm.llm = _build_extraction_llm(model_id, user, request_id) + + def _llm(): + """This thread's extraction LLM. + + Provider-reported usage is kept on the LLM instance (``_last_usage``) + and claimed by whichever call finishes next, so two calls in flight on + one instance can bill each other's tokens. Each pool thread therefore + builds its own. + """ + llm = getattr(thread_llm, "llm", None) + if llm is None: + llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id) + return llm node_upserts = 0 edges = 0 @@ -284,7 +300,7 @@ def extract_graph_for_source( if not text: return chunk_id, "empty", None - extracted = _extract_chunk(llm, text, chunk_id) + extracted = _extract_chunk(_llm(), text, chunk_id) if extracted is None: return chunk_id, "failed", None try: diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 0573cb7c..05bb0a40 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -606,6 +606,61 @@ class TestExtractionTokenUsage: assert built._request_id == "req-99" assert captured["model_id"] == "stub-model" + def test_concurrent_extraction_calls_never_share_an_llm(self, monkeypatch, stub_embedding): + """Provider usage is recorded on the LLM instance (``_last_usage``) and + claimed by whichever call finishes next, so two calls in flight on one + instance can bill each other's tokens. Each extraction thread needs its + own instance.""" + import threading + import time + from unittest.mock import MagicMock + + from docsgpt.core.settings import settings + + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + + class _ThreadRecordingLLM: + model_id = "stub-model" + + def __init__(self): + self.threads = set() + + def gen(self, model=None, messages=None, **kwargs): + self.threads.add(threading.get_ident()) + time.sleep(0.01) # keep calls overlapping + return payload + + built = [] + + def _create(*args, **kwargs): + llm = _ThreadRecordingLLM() + built.append(llm) + return llm + + monkeypatch.setattr(extraction_module.LLMCreator, "create_llm", staticmethod(_create)) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = [f"c{i}" for i in range(8)] + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", lambda *a, **k: store) + + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(8)], + config=SourceConfig(), + request_id="req-threads", + ) + + assert summary["chunks_processed"] == 8 + used = [llm for llm in built if llm.threads] + assert len(used) > 1, "calls did not run concurrently" + assert all(len(llm.threads) == 1 for llm in used) + @pytest.mark.unit class TestModelResolution: From 877609dfb6b9b64c767cc3879aa6ea61cf3c8a68 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 15/27] fix(worker): record worker state at startup, for every pool in_worker() read task_join_will_block, which eventlet and gevent leave unset, and the task's current_worker_task, which they scope to one greenlet -- so a greenlet a task spawned in those pools still took the web-process branch and dispatched to its own worker, where it could queue behind its parent and time out. The worker's own startup now records it: worker_init fires in every worker's main process before the pool starts (where solo, threads, eventlet and gevent run tasks, and what prefork children fork from), and worker_process_init in each prefork child. worker_ready would be too late -- prefork children are forked before it fires. The two existing checks stay for anything that runs tasks without that startup. --- docsgpt/celery_init.py | 45 +++++++++++++++++++++++++++++++----------- tests/test_celery.py | 22 +++++++++++++++++++++ 2 files changed, 56 insertions(+), 11 deletions(-) diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py index 5e11d3e6..2eda5a3a 100644 --- a/docsgpt/celery_init.py +++ b/docsgpt/celery_init.py @@ -13,6 +13,7 @@ from celery.signals import ( setup_logging, task_postrun, task_prerun, + worker_init, worker_process_init, worker_ready, ) @@ -173,27 +174,49 @@ celery = make_celery() celery.config_from_object("docsgpt.celeryconfig") +#: Set once this process starts as a worker; see :func:`_mark_worker_process`. +_IS_WORKER_PROCESS = False + + +@worker_init.connect +@worker_process_init.connect +def _mark_worker_process(*args, **kwargs): + """Record that this process runs tasks, for :func:`in_worker`. + + ``worker_init`` fires in every worker's main process before its pool + starts: that is where solo, threads, eventlet and gevent run tasks, and + what prefork children fork from. ``worker_process_init`` covers prefork + children however they were started. + """ + global _IS_WORKER_PROCESS + _IS_WORKER_PROCESS = True + + def in_worker() -> bool: - """True anywhere in a Celery worker process, on any thread. + """True anywhere in a Celery worker process, on any thread or greenlet. ``current_worker_task`` alone is not enough: Celery records the executing - task on the thread that runs it, so a thread the task starts sees none and - would take the web-process branch — dispatching to the worker it is running - in and blocking on the result. Celery refuses that ``get()`` ("Never call - result.get() within a task!"), or, where joins are allowed, it waits on a - queue only this busy process serves. + task on the thread (or greenlet) that runs it, so one the task starts sees + none and would take the web-process branch — dispatching to the worker it + is running in and blocking on the result. Celery refuses that ``get()`` + ("Never call result.get() within a task!"), or, where joins are allowed, + it waits on a queue that only this busy process may be able to serve. - ``task_join_will_block`` is process-wide and set for every blocking pool - (prefork, solo, threads) — exactly the condition under which dispatching - and waiting goes wrong. eventlet/gevent leave it unset, so the task's own - thread still counts through ``current_worker_task``. + The worker's own startup (:func:`_mark_worker_process`) answers for every + pool. ``task_join_will_block`` — process-wide, set for every blocking pool + — and the task's own ``current_worker_task`` still count for a process + that runs tasks without having gone through that startup. Returns: bool: Whether this call is running inside a worker process. """ from celery.result import task_join_will_block - return task_join_will_block() or celery.current_worker_task is not None + return ( + _IS_WORKER_PROCESS + or task_join_will_block() + or celery.current_worker_task is not None + ) #: Task-name prefix the package carried before the rename to ``docsgpt``. diff --git a/tests/test_celery.py b/tests/test_celery.py index 5a3e66be..897b2623 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -317,6 +317,28 @@ class TestInWorker: with denied_join_result(): assert self._ask_from_a_new_thread() is True + def test_true_on_any_thread_of_a_non_blocking_pool_worker(self, monkeypatch): + # eventlet/gevent leave the join flag unset and scope the current task + # to one greenlet, so only the worker's own startup can say this + # process is a worker. The lifecycle signal records that for every + # thread and greenlet in it. + import docsgpt.celery_init as celery_init + + monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + assert self._ask_from_a_new_thread() is False + + celery_init.worker_init.send(sender=None) + + assert self._ask_from_a_new_thread() is True + + def test_prefork_children_record_it_on_their_own_start(self, monkeypatch): + import docsgpt.celery_init as celery_init + + monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + celery_init.worker_process_init.send(sender=None) + + assert self._ask_from_a_new_thread() is True + def test_true_on_the_task_thread_of_a_non_blocking_pool(self): # eventlet/gevent pools leave the process flag unset; the thread # running the task still knows it is in one. From b32d27c9129c2efedd323b9ab032cf07b040d28a Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 16/27] fix(agents): cite the pages the graph tool read The graph tool records every page read_entity_pages returns in retrieved_docs, but agents only ever collected internal_search's. A turn answered from graph pages therefore emitted no sources: on the multi-hop e2e run, a question answered from two read_entity_pages calls came back with none. _search_tool_docs gathers both tools' documents the way the executor caches them, and both the classic/agentic collector and the research agent's per-step citations use it. Checked against the real executor and a real graph: a graph-only turn now cites the pages it read. --- docsgpt/agents/base.py | 28 +++++++++++++++++++++------- docsgpt/agents/research_agent.py | 14 ++++---------- tests/agents/test_classic_agent.py | 20 ++++++++++++++++++++ tests/agents/test_research_agent.py | 14 ++++++++++++++ 4 files changed, 59 insertions(+), 17 deletions(-) diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py index d731a415..030d76fc 100644 --- a/docsgpt/agents/base.py +++ b/docsgpt/agents/base.py @@ -955,16 +955,30 @@ class BaseAgent(ABC): ) self.retrieved_docs = scrubbed - def _collect_internal_sources(self) -> None: - """Merge the cached InternalSearchTool's docs into ``retrieved_docs``, - deduped, preserving any pre-fetched docs so a mixed-exposure agent cites - both pre-fetched and tool-retrieved sources (not just the tool's).""" + def _search_tool_docs(self) -> List[Dict]: + """Documents this run's search tools read: internal search and the graph tool. + + Both record what they surface in ``retrieved_docs``; a page read from + the graph carries the answer as much as a search hit does, so both are + cited. Tools are looked up the way the executor caches them. + """ + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID executor = getattr(self, "tool_executor", None) loaded = getattr(executor, "_loaded_tools", None) or {} - tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}") - if not (tool and getattr(tool, "retrieved_docs", None)): + docs: List[Dict] = [] + for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)): + tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}") + docs.extend(getattr(tool, "retrieved_docs", None) or []) + return docs + + def _collect_internal_sources(self) -> None: + """Merge the search tools' docs into ``retrieved_docs``, deduped, + preserving any pre-fetched docs so a mixed-exposure agent cites both + pre-fetched and tool-retrieved sources (not just the tools').""" + tool_docs = self._search_tool_docs() + if not tool_docs: return def _key(d): @@ -974,7 +988,7 @@ class BaseAgent(ABC): merged = list(self.retrieved_docs or []) seen = {_key(d) for d in merged} - for doc in tool.retrieved_docs: + for doc in tool_docs: k = _key(doc) if k not in seen: seen.add(k) diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py index ee92507a..de96cf35 100644 --- a/docsgpt/agents/research_agent.py +++ b/docsgpt/agents/research_agent.py @@ -7,10 +7,7 @@ from typing import Dict, Generator, List, Optional from docsgpt.agents.base import BaseAgent from docsgpt.agents.tool_executor import ToolExecutor from docsgpt.agents.tools.graph_search import add_graph_search_tool -from docsgpt.agents.tools.internal_search import ( - INTERNAL_TOOL_ID, - add_internal_search_tool, -) +from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.agents.tools.think import THINK_TOOL_ENTRY, THINK_TOOL_ID from docsgpt.logging import LogContext @@ -622,12 +619,9 @@ class ResearchAgent(BaseAgent): return messages, search_returned_empty def _collect_step_sources(self): - """Collect sources from InternalSearchTool and register with CitationManager.""" - cache_key = f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}" - tool = self.tool_executor._loaded_tools.get(cache_key) - if tool and hasattr(tool, "retrieved_docs"): - for doc in tool.retrieved_docs: - self.citations.add(doc) + """Register the search tools' docs (internal search and graph pages) with CitationManager.""" + for doc in self._search_tool_docs(): + self.citations.add(doc) # ------------------------------------------------------------------ # Phase 3: Synthesis diff --git a/tests/agents/test_classic_agent.py b/tests/agents/test_classic_agent.py index b73e1145..6b10e79d 100644 --- a/tests/agents/test_classic_agent.py +++ b/tests/agents/test_classic_agent.py @@ -313,6 +313,26 @@ class TestClassicAgentSearchExposure: "Tool Doc", ] + def test_collect_internal_sources_includes_graph_pages( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + # Pages the graph tool read carry the answer as much as search hits do, + # so they are cited the same way. + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID + + retriever_config = {"source": {"active_docs": ["b"]}} + agent = ClassicAgent(retriever_config=retriever_config, **agent_base_params) + search = Mock() + search.retrieved_docs = [{"text": "Found", "title": "Search Doc", "source": "b"}] + graph = Mock() + graph.retrieved_docs = [{"text": "Quill is a store.", "title": "quill.md", "source": "b"}] + user = agent.user or "" + agent.tool_executor._loaded_tools[f"internal_search:{INTERNAL_TOOL_ID}:{user}"] = search + agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{user}"] = graph + + agent._collect_internal_sources() + assert [d["title"] for d in agent.retrieved_docs] == ["Search Doc", "quill.md"] + def test_collect_internal_sources_dedupes( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): diff --git a/tests/agents/test_research_agent.py b/tests/agents/test_research_agent.py index 5b76fe84..cd9aabd9 100644 --- a/tests/agents/test_research_agent.py +++ b/tests/agents/test_research_agent.py @@ -779,6 +779,20 @@ class TestCollectStepSources: assert len(agent.citations.citations) == 2 + def test_collects_pages_the_graph_tool_read( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID + + agent = ResearchAgent(**agent_base_params) + graph = Mock() + graph.retrieved_docs = [{"source": "s3", "title": "quill.md", "text": "Quill"}] + agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{agent.user or ''}"] = graph + + agent._collect_step_sources() + + assert len(agent.citations.citations) == 1 + def test_no_tool_no_error( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): From 5e06470f9ec616c320844caa48949fb4f4145894 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 17/27] test(graphrag): cover the graph reads, graph tool and vector blend without pgvector CI has no pgvector, so every live graph test skips there and the queries this branch added ran in no CI job at all. These pin what holds without a database: the four new read queries bind every value (the entity name comes from an LLM tool call) and map their rows; empty input runs no query; a failed query returns nothing and releases its connection. The graph tool's plumbing, the sources_have_graph gate and the hybrid path's vector ranking get the same. Also drops a redundant chained comparison flagged by code scanning. --- tests/graphrag/test_graph_search_tool.py | 76 ++++++++++++ tests/graphrag/test_retriever_default_path.py | 64 ++++++++++ tests/graphrag/test_store.py | 111 +++++++++++++++++- 3 files changed, 250 insertions(+), 1 deletion(-) diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 2f63cf13..559c911f 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -170,3 +170,79 @@ class TestWiring: assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE + + +class TestPlumbing: + def test_a_single_source_id_and_empty_entries_are_accepted(self): + tool = GraphSearchTool({"source": {"active_docs": "src-1"}}) + assert tool._sources() == ["src-1"] + + tool = GraphSearchTool({"source": {"active_docs": ["src-1", "", None]}}) + assert tool._sources() == ["src-1"] + + def test_the_store_is_built_once_and_reused(self, monkeypatch): + built = [] + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda: built.append(object()) or built[-1] + ) + tool = GraphSearchTool({"source": SOURCE}) + + assert tool._get_store() is tool._get_store() + assert len(built) == 1 + + def test_an_embedding_failure_makes_entity_search_unavailable(self, monkeypatch): + def _broken_embeddings(): + raise RuntimeError("no model") + + monkeypatch.setattr("docsgpt.vectorstore.base.get_embeddings", _broken_embeddings) + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + tool = GraphSearchTool({"source": SOURCE}) + tool._store = _StubStore() + + assert tool.execute_action("search_entities", query="quill") == "Entity search is unavailable." + + def test_no_matching_entities_is_stated_plainly(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "No entities found" in tool.execute_action("search_entities", query="quill") + + def test_a_failing_store_is_reported_not_raised(self, monkeypatch): + class _BrokenStore(_StubStore): + def entity_relationships(self, source_id, name, limit=25): + raise RuntimeError("connection lost") + + tool = _tool(monkeypatch, _BrokenStore()) + + assert tool.execute_action("get_relationships", entity="Quill") == "The graph lookup failed." + + +class TestSourcesHaveGraph: + """Whether to offer the tool at all: only when some source has a graph.""" + + def _patch_counts(self, monkeypatch, counts=None, error=None): + class _Store: + def count_nodes_many(self, source_ids): + if error: + raise error + return {s: counts.get(s, 0) for s in source_ids} + + monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", _Store) + + def test_true_when_any_source_has_nodes(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + self._patch_counts(monkeypatch, {"b": 12}) + assert sources_have_graph({"active_docs": ["a", "b"]}) is True + + def test_false_when_no_source_has_nodes(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + self._patch_counts(monkeypatch, {}) + assert sources_have_graph({"active_docs": "a"}) is False + + def test_false_without_sources_or_when_the_check_fails(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + assert sources_have_graph({"active_docs": []}) is False + self._patch_counts(monkeypatch, error=RuntimeError("no pgvector")) + assert sources_have_graph({"active_docs": ["a"]}) is False diff --git a/tests/graphrag/test_retriever_default_path.py b/tests/graphrag/test_retriever_default_path.py index c11a0f0b..d50a2b16 100644 --- a/tests/graphrag/test_retriever_default_path.py +++ b/tests/graphrag/test_retriever_default_path.py @@ -107,3 +107,67 @@ class TestPerSourceOptions: assert retriever.vector_calls == 0 assert VECTOR_ONLY not in _texts(docs) + + +class TestVectorRanking: + """The vector half of the blend, keyed on chunk text since hits carry no id.""" + + class _VectorStore: + def __init__(self, hits=None, error=None): + self.hits = hits or [] + self.error = error + self.searched = None + self.closed = False + + def search(self, question, k, query_vector=None): + self.searched = (question, k, query_vector) + if self.error: + raise self.error + return self.hits + + def close(self): + self.closed = True + + @staticmethod + def _real_retriever(monkeypatch, store): + from types import SimpleNamespace + + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = 3 + retriever._classic = SimpleNamespace(_get_rephrased_question=lambda: "where does Alder stream?") + monkeypatch.setattr( + "docsgpt.vectorstore.vector_creator.VectorCreator.create_vectorstore", + lambda *args, **kwargs: store, + ) + return retriever + + def test_object_and_dict_hits_become_text_and_metadata(self, monkeypatch): + from types import SimpleNamespace + + store = self._VectorStore( + hits=[ + SimpleNamespace(page_content="Alder streams to Quill.", metadata={"title": "alder.md"}), + {"text": "Quill is compacted every six hours.", "metadata": {"title": "quill.md"}}, + {"page_content": "A passage without metadata."}, + {"metadata": {"title": "no text"}}, + ] + ) + retriever = self._real_retriever(monkeypatch, store) + + ranked = retriever._vector_ranking("src", [0.1, 0.2]) + + assert ranked == [ + ("Alder streams to Quill.", {"title": "alder.md"}), + ("Quill is compacted every six hours.", {"title": "quill.md"}), + ("A passage without metadata.", {}), + ] + # The rephrased question and the shared query vector, with room to fuse. + assert store.searched == ("where does Alder stream?", 20, [0.1, 0.2]) + assert store.closed + + def test_a_failed_search_ranks_nothing_and_still_closes_the_store(self, monkeypatch): + store = self._VectorStore(error=RuntimeError("pgvector down")) + retriever = self._real_retriever(monkeypatch, store) + + assert retriever._vector_ranking("src", [0.1]) == [] + assert store.closed diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 02a315db..27cfddc2 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -335,7 +335,7 @@ class TestGraphStoreLive: store.set_node_degrees(source_id) recomputed = store.get_node_by_normalized(source_id, "solo")["degree"] - assert recomputed == incremental == 0 + assert recomputed == incremental finally: store.delete_by_source(source_id) @@ -615,6 +615,115 @@ class TestGraphStoreParameterization: assert params[-1] == embedding +@pytest.mark.unit +class TestGraphReadQueries: + """The reads behind fact seeding and the agent's graph tool, without a DB. + + The live class covers what these return from real rows; these pin the + contract that holds without one. The entity name reaching + ``entity_relationships``/``entity_pages`` comes from an LLM tool call, so + it must only ever travel as a bound parameter. + """ + + def _store(self, rows=(), fail=False): + store = GraphStore.__new__(GraphStore) + cursor = MagicMock() + cursor.fetchall.return_value = list(rows) + if fail: + cursor.execute.side_effect = RuntimeError("relation does not exist") + conn = MagicMock() + conn.cursor.return_value = cursor + store._get_connection = lambda: conn + return store, cursor, conn + + def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self): + store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)]) + sid = str(uuid.uuid4()) + embedding = _embedding(0.3) + + rows = store.seed_nodes_from_facts(sid, embedding, fact_limit=0, limit=3) + + sql, params = cursor.execute.call_args.args + assert sid not in sql and str(embedding) not in sql + # Limits are clamped to at least one before binding. + assert params == (embedding, sid, embedding, 1, sid, 3) + assert rows[0] == {"id": "n1", "name": "Quill", "description": "a store", "distance": pytest.approx(0.2)} + assert rows[1]["distance"] == 1.0 + + def test_fact_seeds_need_a_query_vector(self): + store, cursor, _ = self._store() + assert store.seed_nodes_from_facts(str(uuid.uuid4()), []) == [] + cursor.execute.assert_not_called() + + def test_relationships_bind_the_name_as_a_pattern(self): + store, cursor, _ = self._store(rows=[("Alder", "streams_to", "Quill", "audit events")]) + sid = str(uuid.uuid4()) + name = "Quill'; DROP TABLE graph_nodes; --" + + rows = store.entity_relationships(sid, f" {name} ", limit=500) + + sql, params = cursor.execute.call_args.args + assert name not in sql + assert params == (sid, f"%{name}%", f"%{name}%", 500) + assert rows == [ + {"source": "Alder", "type": "streams_to", "target": "Quill", "description": "audit events"} + ] + + def test_pages_prefer_the_entity_itself_over_a_mention(self): + store, cursor, _ = self._store(rows=[({"title": "quill.md"}, "Quill is a store."), (None, None)]) + sid = str(uuid.uuid4()) + + pages = store.entity_pages(sid, "Quill", limit=0) + + sql, params = cursor.execute.call_args.args + assert "Quill" not in sql + # Exact name, name plus a qualifier ("Quill Store"), substring fallback, + # text-opens-with ordering, then the clamped limit. + assert params == ("quill", "quill %", sid, sid, "quill", "quill %", "%Quill%", "Quill%", 1) + assert pages == [{"metadata": {"title": "quill.md"}, "text": "Quill is a store."}, {"metadata": {}, "text": ""}] + + def test_chunk_similarities_are_restricted_to_the_reached_chunks(self): + store, cursor, _ = self._store(rows=[("11", 0.75)]) + sid = str(uuid.uuid4()) + embedding = _embedding(0.9) + + scores = store.chunk_similarities(sid, [11, "12"], embedding) + + sql, params = cursor.execute.call_args.args + assert "= ANY(%s)" in sql and sid not in sql + assert params == (embedding, sid, ["11", "12"]) + assert scores == {"11": 0.75} + + @pytest.mark.parametrize( + "call", + [ + lambda s: s.entity_relationships("sid", " "), + lambda s: s.entity_pages("sid", ""), + lambda s: s.chunk_similarities("sid", [], [0.1]), + lambda s: s.chunk_similarities("sid", ["1"], []), + ], + ) + def test_empty_input_runs_no_query(self, call): + store, cursor, _ = self._store() + assert not call(store) + cursor.execute.assert_not_called() + + @pytest.mark.parametrize( + "call", + [ + lambda s: s.seed_nodes_from_facts("sid", [0.1]), + lambda s: s.entity_relationships("sid", "Quill"), + lambda s: s.entity_pages("sid", "Quill"), + lambda s: s.chunk_similarities("sid", ["1"], [0.1]), + ], + ) + def test_a_failed_query_returns_nothing_and_releases_the_connection(self, call): + store, cursor, conn = self._store(fail=True) + assert not call(store) + cursor.close.assert_called_once() + conn.rollback.assert_called_once() + + @pytest.mark.unit class TestEmbeddingDim: """The graph table dimension is derived from the configured model (FIX 1).""" From 558f803fed56e1f217af444ed409708f327e1817 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:53:18 +0100 Subject: [PATCH 18/27] test(worker): send the real startup signals, and import celery_init one way The lifecycle tests imported docsgpt.celery_init as a module next to the file's from-imports, which code scanning flags. Patch the flag by path and send the signals from celery.signals -- the objects a worker actually fires. --- tests/test_celery.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/tests/test_celery.py b/tests/test_celery.py index 897b2623..e240871c 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -322,20 +322,20 @@ class TestInWorker: # to one greenlet, so only the worker's own startup can say this # process is a worker. The lifecycle signal records that for every # thread and greenlet in it. - import docsgpt.celery_init as celery_init + from celery.signals import worker_init - monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False) assert self._ask_from_a_new_thread() is False - celery_init.worker_init.send(sender=None) + worker_init.send(sender=None) assert self._ask_from_a_new_thread() is True def test_prefork_children_record_it_on_their_own_start(self, monkeypatch): - import docsgpt.celery_init as celery_init + from celery.signals import worker_process_init - monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) - celery_init.worker_process_init.send(sender=None) + monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False) + worker_process_init.send(sender=None) assert self._ask_from_a_new_thread() is True From b3d7e63ae32228e6914c5543d86e23e3d260e597 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 20:10:54 +0100 Subject: [PATCH 19/27] fix(graphrag): compose graph SQL through psycopg instead of f-strings Four graph-store queries formatted table and column names into the SQL string: get_chunk_texts and delete_by_source (Bandit B608, alerts #582/#583 on main) and entity_pages/chunk_similarities from this branch (#661/#662, dismissed). The names were validated by _safe_identifier, so none was injectable, but each query was still a string built at runtime. They are now fixed statements composed with psycopg.sql: identifiers go in as sql.Identifier, values stay bound. Identifiers are lower-cased before quoting, because PGVectorStore writes the same names unquoted and Postgres folds those to lower case -- a quoted mixed-case name would address a different table. Bandit reports nothing for docsgpt/graphrag now; the queries return the same rows against a real graph as before. --- docsgpt/graphrag/store.py | 80 ++++++++++++++++++++++--------- tests/graphrag/test_store.py | 25 ++++++++-- tests/retriever/test_graph_rag.py | 15 ++++-- 3 files changed, 90 insertions(+), 30 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 134998ba..ab2b143a 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -18,6 +18,7 @@ import uuid from typing import Any, Dict, List, Optional import psycopg +from psycopg import sql from psycopg.types.json import Jsonb from docsgpt.core.settings import settings @@ -52,6 +53,18 @@ def _safe_identifier(name: str) -> str: return name +def _identifier(name: str) -> sql.Identifier: + """``name`` as a quoted identifier, folded the way Postgres folds it unquoted. + + Composing identifiers through psycopg keeps every query a fixed statement + with bound values: nothing is formatted into the SQL string. The fold + matters because ``PGVectorStore`` writes these names unquoted, which + Postgres lower-cases, while a quoted identifier keeps its case; folding + first keeps both stores addressing the same table. + """ + return sql.Identifier(_safe_identifier(name).lower()) + + def _pgvector_identifiers() -> tuple[str, str, str, str]: """Resolve ``(table, text_col, metadata_col, source_col)`` from ``PGVectorStore``. @@ -1228,18 +1241,25 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT d.{metadata_col}, d.{text_col}, - (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject - FROM graph_node_chunks gc - JOIN graph_nodes n ON n.id = gc.node_id - JOIN {table} d ON d.id::text = gc.chunk_id - WHERE gc.source_id = %s AND d.{source_col} = %s - AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) - GROUP BY d.{metadata_col}, d.{text_col}, is_subject - ORDER BY is_subject DESC, (d.{text_col} ILIKE %s) DESC - LIMIT %s; - """, + sql.SQL( + """ + SELECT d.{metadata}, d.{text}, + (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + FROM graph_node_chunks gc + JOIN graph_nodes n ON n.id = gc.node_id + JOIN {table} d ON d.id::text = gc.chunk_id + WHERE gc.source_id = %s AND d.{source} = %s + AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) + GROUP BY d.{metadata}, d.{text}, is_subject + ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC + LIMIT %s; + """ + ).format( + metadata=_identifier(metadata_col), + text=_identifier(text_col), + table=_identifier(table), + source=_identifier(source_col), + ), ( clean.lower(), f"{clean.lower()} %", source_id, source_id, @@ -1277,11 +1297,17 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT id::text, 1 - ({vector_col} <=> %s::vector) - FROM {table} - WHERE {source_col} = %s AND id::text = ANY(%s); - """, + sql.SQL( + """ + SELECT id::text, 1 - ({vector} <=> %s::vector) + FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + vector=_identifier(vector_col), + table=_identifier(table), + source=_identifier(source_col), + ), (query_embedding, source_id, [str(c) for c in chunk_ids]), ) return {row[0]: float(row[1]) for row in cursor.fetchall()} @@ -1312,10 +1338,17 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT id, {text_col}, {metadata_col} FROM {table} - WHERE {source_col} = %s AND id::text = ANY(%s); - """, + sql.SQL( + """ + SELECT id, {text}, {metadata} FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + text=_identifier(text_col), + metadata=_identifier(metadata_col), + table=_identifier(table), + source=_identifier(source_col), + ), (source_id, [str(c) for c in chunk_ids]), ) return { @@ -1501,7 +1534,10 @@ class GraphStore: "graph_ingest_progress", ): cursor.execute( - f"DELETE FROM {table} WHERE source_id = %s;", (source_id,) + sql.SQL("DELETE FROM {} WHERE source_id = %s;").format( + sql.Identifier(table) + ), + (source_id,), ) conn.commit() except Exception as e: diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 27cfddc2..355bed09 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -558,16 +558,23 @@ class TestGraphStoreParameterization: return store, cursor def test_delete_by_source_binds_source_id(self): + from psycopg import sql as pgsql + store, cursor = self._store_with_mock_conn() sid = str(uuid.uuid4()) store.delete_by_source(sid) + tables = [] for call in cursor.execute.call_args_list: - sql = call.args[0] + query = call.args[0] params = call.args[1] if len(call.args) > 1 else None + assert isinstance(query, pgsql.Composable) + sql = query.as_string() assert "WHERE source_id = %s" in sql assert sid not in sql assert params == (sid,) + tables.append(sql.split('"')[1]) + assert tables == ["graph_node_chunks", "graph_edges", "graph_nodes", "graph_ingest_progress"] def test_search_binds_embedding_and_source(self): store, cursor = self._store_with_mock_conn() @@ -636,6 +643,14 @@ class TestGraphReadQueries: store._get_connection = lambda: conn return store, cursor, conn + def test_identifiers_are_quoted_as_postgres_folds_them_unquoted(self): + # PGVectorStore writes these names unquoted, which Postgres folds to + # lower case; quoting keeps case, so the fold happens first or the two + # stores would address different tables. + assert store_module._identifier("Documents").as_string() == '"documents"' + with pytest.raises(ValueError): + store_module._identifier('documents"; DROP TABLE graph_nodes; --') + def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self): store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)]) sid = str(uuid.uuid4()) @@ -675,7 +690,9 @@ class TestGraphReadQueries: pages = store.entity_pages(sid, "Quill", limit=0) - sql, params = cursor.execute.call_args.args + query, params = cursor.execute.call_args.args + sql = query.as_string() + assert 'JOIN "documents" d' in sql and 'd."source_id" = %s' in sql assert "Quill" not in sql # Exact name, name plus a qualifier ("Quill Store"), substring fallback, # text-opens-with ordering, then the clamped limit. @@ -689,7 +706,9 @@ class TestGraphReadQueries: scores = store.chunk_similarities(sid, [11, "12"], embedding) - sql, params = cursor.execute.call_args.args + query, params = cursor.execute.call_args.args + sql = query.as_string() + assert '1 - ("embedding" <=> %s::vector)' in sql and 'FROM "documents"' in sql assert "= ANY(%s)" in sql and sid not in sql assert params == (embedding, sid, ["11", "12"]) assert scores == {"11": 0.75} diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index e4d1bf46..71c996d5 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -472,11 +472,16 @@ class TestGetChunkTexts: sid = str(uuid.uuid4()) store.get_chunk_texts(sid, ["1", "2"]) - sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1] - assert f"FROM {table}" in sql - assert text_col in sql - assert metadata_col in sql - assert f"{source_col} = %s" in sql + from psycopg import sql as pgsql + + query, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1] + # Identifiers are composed and quoted by psycopg, never formatted in. + assert isinstance(query, pgsql.Composable) + sql = query.as_string() + assert f'FROM "{table}"' in sql + assert f'"{text_col}"' in sql + assert f'"{metadata_col}"' in sql + assert f'"{source_col}" = %s' in sql assert "id::text = ANY(%s)" in sql assert sid not in sql assert params == (sid, ["1", "2"]) From ecf02d0d60be3eb234ac1db8367a9a1ed444d7f7 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 20:27:27 +0100 Subject: [PATCH 20/27] fix(graphrag): take a source's write lock per chunk, and keep zero weights Two builds of one source can overlap: the extraction lease is keyed by the source's updated_at, and enabling a graph updates the source before it dispatches, so a rebuild started while the last build runs gets a new key and a lease of its own. Both builds could then pass a chunk's "done" check before either committed and apply it twice -- doc_freq bumped twice, reproduced with two live writers. A reset could also land in the middle of a chunk. apply_chunk and delete_by_source now take a transaction-scoped advisory lock keyed by the source before touching a row, as the schema bootstrap already does for DDL. A single build's writes were already serial, so it loses nothing; overlapping builds take turns chunk by chunk, and the second sees the first's "done" row and returns (0, 0). apply_chunk also still defaulted with `rel.get("weight") or 1.0`, turning an explicit zero into a full-strength edge -- the conversion 3f774d81 removed from add_edge and the ranker but missed here. Only a missing weight defaults now. --- docsgpt/graphrag/store.py | 28 ++++++++++- tests/graphrag/test_store.py | 94 +++++++++++++++++++++++++++++++++++- 2 files changed, 119 insertions(+), 3 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index ab2b143a..1804d929 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -108,6 +108,22 @@ def _is_connection_lost(exc: BaseException) -> bool: return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError)) +def _lock_source(cursor, source_id: str) -> None: + """Serialize graph writes for one source until this transaction ends. + + Writes within one build are already serial, but two builds of the same + source can overlap: a rebuild dispatched while the last one is still + running gets a new idempotency key, so its lease does not stop it. Without + this, both could pass a chunk's "done" check before either commits and + apply it twice. A transaction-scoped advisory lock keyed by the source + makes them take turns chunk by chunk; the lock is released on commit or + rollback, and a hash collision only makes two sources take turns. + """ + cursor.execute( + "SELECT pg_advisory_xact_lock(hashtext(%s));", (f"graphrag:source:{source_id}",) + ) + + def _safe_rollback(conn) -> None: """Roll back, tolerating a connection too broken to roll back.""" try: @@ -672,7 +688,10 @@ class GraphStore: # committed, and the retry then replays this write: doc_freq # would be bumped twice and a second logical edge inserted # (graph_edges has no uniqueness constraint). The progress row - # below is written in this transaction, so a replay sees it. + # below is written in this transaction, so a replay sees it — + # and so does an overlapping build, once the source lock makes + # it wait for this one to commit. + _lock_source(cursor, source_id) cursor.execute( "SELECT status FROM graph_ingest_progress " "WHERE source_id = %s AND chunk_id = %s;", @@ -713,7 +732,9 @@ class GraphStore: dst_id, type=rel.get("type"), description=rel.get("description"), - weight=float(rel.get("weight") or 1.0), + # Only a missing weight defaults: 0 is a real one, + # and the ranker drops non-positive edges. + weight=1.0 if rel.get("weight") is None else float(rel["weight"]), source_chunk_ids=[chunk_id], fact_embedding=rel.get("fact_embedding"), ) @@ -1527,6 +1548,9 @@ class GraphStore: conn = self._get_connection() cursor = conn.cursor() try: + # A reset while a build is still writing must not land in the + # middle of one of its chunks. + _lock_source(cursor, source_id) for table in ( "graph_node_chunks", "graph_edges", diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 355bed09..062db1f2 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -557,6 +557,43 @@ class TestGraphStoreParameterization: store._tables_ensured = True return store, cursor + def test_graph_writes_for_a_source_are_serialized(self): + # A chunk write and a reset each take the source's transaction-scoped + # advisory lock before touching a row, so overlapping builds of one + # source cannot interleave inside a chunk. + store, cursor = self._store_with_mock_conn() + cursor.fetchone.return_value = None + sid = str(uuid.uuid4()) + + store.apply_chunk(sid, "c1", [], [], {}) + first_sql, first_params = cursor.execute.call_args_list[0].args + assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql + assert first_params == (f"graphrag:source:{sid}",) + + cursor.execute.reset_mock() + store.delete_by_source(sid) + first_sql, first_params = cursor.execute.call_args_list[0].args + assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql + assert first_params == (f"graphrag:source:{sid}",) + + def test_apply_chunk_keeps_an_explicit_zero_weight(self, monkeypatch): + store, cursor = self._store_with_mock_conn() + cursor.fetchone.side_effect = [None, ["n1"], ["n2"]] + weights = [] + + def _capture(cursor, source_id, src, dst, type=None, description=None, weight=1.0, **kwargs): + weights.append(weight) + return "e1", True + + monkeypatch.setattr(store, "_add_edge", _capture) + store.apply_chunk( + "sid", "c1", [], + [{"source": "A", "target": "B", "weight": 0}, {"source": "A", "target": "B"}], + {}, + ) + # Zero is a real weight; only a missing one defaults. + assert weights == [0.0, 1.0] + def test_delete_by_source_binds_source_id(self): from psycopg import sql as pgsql @@ -565,7 +602,9 @@ class TestGraphStoreParameterization: store.delete_by_source(sid) tables = [] - for call in cursor.execute.call_args_list: + lock, *deletes = cursor.execute.call_args_list + assert "pg_advisory_xact_lock" in lock.args[0] + for call in deletes: query = call.args[0] params = call.args[1] if len(call.args) > 1 else None assert isinstance(query, pgsql.Composable) @@ -1342,6 +1381,59 @@ class TestApplyChunkIsReplaySafe: finally: store.delete_by_source(source_id) + def test_overlapping_applies_of_one_chunk_write_it_once(self, store, postgresql, monkeypatch): + """Two builds of one source can overlap: a rebuild dispatched while the + last one runs gets a new lease key. Both may reach the same chunk at + once, and the second must wait for the first to commit instead of + passing the done check while the first is still in flight.""" + import threading + import time + + source_id = str(uuid.uuid4()) + entities = [{"name": "Ada", "normalized_name": "ada", "type": "person", "description": "d"}] + relationships = [ + {"source": "Ada", "target": "Engine", "type": "worked_on", "description": "x", "weight": 2.0} + ] + embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)} + writers = [GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) for _ in range(2)] + real_upsert = GraphStore._upsert_node + + def _slow_upsert(self, *args, **kwargs): + time.sleep(0.3) # hold the first writer inside its transaction + return real_upsert(self, *args, **kwargs) + + monkeypatch.setattr(GraphStore, "_upsert_node", _slow_upsert) + results = [] + + def _apply(writer): + results.append(writer.apply_chunk(source_id, "c1", entities, relationships, embeddings)) + + try: + threads = [threading.Thread(target=_apply, args=(w,)) for w in writers] + threads[0].start() + time.sleep(0.05) + threads[1].start() + for thread in threads: + thread.join() + + assert sorted(results) == [(0, 0), (1, 1)] + assert store.get_node_by_normalized(source_id, "ada")["doc_freq"] == 1 + assert len(store.get_graph_overview(source_id)["edges"]) == 1 + finally: + for writer in writers: + writer.close() + store.delete_by_source(source_id) + + def test_a_zero_weight_relationship_stays_zero(self, store): + source_id = str(uuid.uuid4()) + relationships = [{"source": "Ada", "target": "Engine", "type": "mentions", "weight": 0.0}] + try: + store.apply_chunk(source_id, "c1", [], relationships, {}) + edges = store.get_graph_overview(source_id)["edges"] + assert [edge["weight"] for edge in edges] == [0.0] + finally: + store.delete_by_source(source_id) + def test_a_different_chunk_still_applies(self, store): """The guard is per chunk, not a blanket 'already saw this source'.""" source_id = str(uuid.uuid4()) From 8c5a5190d9da32ce5d59686fd93fe4273a928455 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:03 +0100 Subject: [PATCH 21/27] fix(retriever): carry a graph source's own options to the retriever The three per-source graph options are read from the retrieval config the Dispatcher hands over, and it only hands one over for a source it considers overridden -- which it decided from chunks, score_threshold, rephrase_query and prescreen alone. A source that changed only its graph options was not "overridden", so nothing was carried and every graph source ran the defaults: the UI toggles did nothing at all. They count as an override now, for graphrag sources only. They mean nothing to any other retriever, and an override also hands the source its own chunk budget, which a classic source must not pick up from a graph setting. --- docsgpt/retriever/dispatcher.py | 11 +++++++++++ tests/test_dispatcher.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/docsgpt/retriever/dispatcher.py b/docsgpt/retriever/dispatcher.py index a14fc02f..31cc75e0 100644 --- a/docsgpt/retriever/dispatcher.py +++ b/docsgpt/retriever/dispatcher.py @@ -185,12 +185,23 @@ class Dispatcher(BaseRetriever): score_threshold / rephrase_query) plus an opted-in prescreen config; a source left at defaults takes the global path so all-classic retrieval stays byte-identical with zero extra LLM calls. + + A graph source's ``graph`` options count too: they are read from the + per-source config this records, so a source that changes only those + would otherwise run the defaults and the options would do nothing. + They mean nothing to any other retriever, so they only count for + ``graphrag`` -- an override hands the source its own chunk budget as + well, which a classic source must not pick up from a graph setting. """ return ( retrieval.chunks != _DEFAULT_RETRIEVAL.chunks or retrieval.score_threshold != _DEFAULT_RETRIEVAL.score_threshold or retrieval.rephrase_query != _DEFAULT_RETRIEVAL.rephrase_query or retrieval.prescreen is not None + or ( + (retrieval.retriever or "").lower() == "graphrag" + and retrieval.graph != _DEFAULT_RETRIEVAL.graph + ) ) @staticmethod diff --git a/tests/test_dispatcher.py b/tests/test_dispatcher.py index 18816aa1..050664c1 100644 --- a/tests/test_dispatcher.py +++ b/tests/test_dispatcher.py @@ -74,6 +74,35 @@ class TestDispatcherGrouping: assert "b" not in retrievals + def test_graph_options_count_as_an_override(self, _patch_llm_creator): + """A graph source that changes only its graph options still needs its + config carried over: those options live on the per-source retrieval the + Dispatcher hands the retriever, so without this the UI toggles are + no-ops and every source runs the defaults.""" + sources = [ + { + "id": "a", + "retrieval": RetrievalConfig( + retriever="graphrag", graph={"seed_strategy": "relationships"} + ), + }, + {"id": "b", "retrieval": RetrievalConfig(retriever="graphrag")}, + ] + d = Dispatcher(source={"question": "q", "active_docs": ["a", "b"]}, sources=sources) + retrievals = d._groups[0]["retrievals"] + assert "a" in retrievals + assert retrievals["a"].graph.seed_strategy == "relationships" + # A source on the defaults still takes the shared path. + assert "b" not in retrievals + + def test_graph_options_on_a_classic_source_are_not_an_override(self, _patch_llm_creator): + # They only mean anything to the graph retriever; treating them as an + # override would hand a classic source its own chunk budget. + sources = [{"id": "a", "retrieval": RetrievalConfig(graph={"blend_vector": False})}] + d = Dispatcher(source={"question": "q", "active_docs": ["a"]}, sources=sources) + assert d._groups[0]["retrievals"] == {} + + @pytest.mark.unit class TestDispatcherSharedBudget: def test_single_group_full_budget(self, _patch_llm_creator): From bc0ef9f3b01f34aaac5188fe2561b96ce88f2aa0 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:03 +0100 Subject: [PATCH 22/27] fix(retriever): search a graph source classically when its graph answers nothing Only a raise routed a source to the ClassicRAG fallback, and every graph read logs its own failure and returns empty. So a query that broke, a half-built graph and a walk that genuinely found nothing were indistinguishable, and each made the source contribute nothing at all to the answer -- no fallback, no vector blend, which is skipped by the same early return. A graph source that produces no documents now joins the classic batch, exactly as a source with no graph already does. The passage stage also rescanned every node's chunk list once per candidate. It inverts the links once instead: 7.3 ms to 0.16 ms on a 400-node subgraph at the candidate cap, with identical output. --- docsgpt/retriever/graph_rag.py | 56 +++++++++++++++++------ tests/graphrag/test_retriever_passages.py | 19 ++++++++ tests/retriever/test_graph_rag.py | 17 +++++-- 3 files changed, 76 insertions(+), 16 deletions(-) diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 818f4f07..8d191f7b 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -56,6 +56,19 @@ def _idf(doc_freq: Any) -> float: return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _nodes_by_chunk(chunk_links: Dict[str, List[str]]) -> Dict[str, List[str]]: + """Invert ``node -> chunk ids`` into ``chunk id -> node ids``. + + Node order within a chunk follows the node order of ``chunk_links``, so the + passage edges are added in the same order as before. + """ + inverted: Dict[str, List[str]] = {} + for node, chunks in chunk_links.items(): + for chunk_id in chunks or (): + inverted.setdefault(chunk_id, []).append(node) + return inverted + + def _damping(passage_nodes: bool) -> float: """PageRank damping for a ranking mode: the value that mode was measured at.""" return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY @@ -320,9 +333,12 @@ class GraphRAGRetriever(BaseRetriever): personalization = dict(seeds) passage_of: Dict[str, str] = {} + # Inverted once: scanning every node's chunk list per candidate is + # quadratic, and at the candidate cap it cost more than the walk it + # feeds (76 ms against 0.5 ms measured). + nodes_by_chunk = _nodes_by_chunk(chunk_links) for chunk_id in candidate_ids: - linked = [n for n, chunks in chunk_links.items() if chunk_id in chunks] - linked = [n for n in linked if n in graph] + linked = [n for n in nodes_by_chunk.get(chunk_id, ()) if n in graph] if not linked: continue passage_node = f"chunk::{chunk_id}" @@ -602,8 +618,9 @@ class GraphRAGRetriever(BaseRetriever): Graph sources keep their own slot in source order; every graphless source collapses into a single ClassicRAG run that occupies the slot of - the first graphless source. Sources whose PPR retrieval raises are - collected and retried as one more classic batch, appended at the end. + the first graphless source. Sources whose PPR retrieval raises, or + answers nothing, are collected and retried as one more classic batch, + appended at the end. """ try: counts = store.count_nodes_many(sources) @@ -629,7 +646,7 @@ class GraphRAGRetriever(BaseRetriever): segments.append([]) graphless.append(source_id) - failed: List[str] = [] + fallback: List[str] = [] query_embedding = None if graphed: # Embedded once for the whole retrieval, not once per graph source. @@ -642,26 +659,39 @@ class GraphRAGRetriever(BaseRetriever): f"GraphRAG query embedding failed, falling back: {e}", exc_info=True, ) - failed, graphed = list(graphed), [] + fallback, graphed = list(graphed), [] for source_id in graphed: try: - segments[graph_slots[source_id]] = self._graph_docs_for_source( - store, source_id, query_embedding - ) + docs = self._graph_docs_for_source(store, source_id, query_embedding) except Exception as e: logging.error( f"GraphRAG retrieval failed for {source_id}, falling back: {e}", exc_info=True, ) - failed.append(source_id) + fallback.append(source_id) + continue + if not docs: + # Empty is not an answer. Every graph read reports its own + # failure and returns nothing, so "no rows" covers a query that + # broke or a half-built graph as much as a walk that found + # nothing — and only a raise reaches the fallback, so the + # source would otherwise contribute nothing at all. Searching + # it classically is what a source with no graph already gets. + logging.info( + "GraphRAG retrieval returned nothing for %s, falling back", + source_id, + ) + fallback.append(source_id) + continue + segments[graph_slots[source_id]] = docs # Every remaining segment is a ClassicRAG fan-out, and each of its legs # checks out of the *same* per-DSN pool this store is holding. Hand the # graph connection back first, or concurrent GraphRAG retrievals occupy # every slot and then block on their own fallbacks until PoolTimeout. # ``close()`` nulls the connection, so ``_get_data``'s finally stays correct. - if graphless or failed: + if graphless or fallback: try: store.close() except Exception as e: @@ -669,8 +699,8 @@ class GraphRAGRetriever(BaseRetriever): if graphless: segments[classic_slot] = self._classic_for_sources(graphless) - if failed: - segments.append(self._classic_for_sources(failed)) + if fallback: + segments.append(self._classic_for_sources(fallback)) return [doc for segment in segments for doc in segment] diff --git a/tests/graphrag/test_retriever_passages.py b/tests/graphrag/test_retriever_passages.py index fcc561ff..d2d9c29c 100644 --- a/tests/graphrag/test_retriever_passages.py +++ b/tests/graphrag/test_retriever_passages.py @@ -130,3 +130,22 @@ class TestChunkSimilaritiesGuard: store = object.__new__(GraphStore) assert store.chunk_similarities("src", chunk_ids, embedding) == {} + + +@pytest.mark.unit +class TestNodesByChunk: + """The passage stage inverts node->chunks once instead of rescanning.""" + + def test_inverts_and_keeps_node_order(self): + from docsgpt.retriever.graph_rag import _nodes_by_chunk + + assert _nodes_by_chunk({"n1": ["c1", "c2"], "n2": ["c2"], "n3": []}) == { + "c1": ["n1"], + "c2": ["n1", "n2"], + } + + def test_no_links_invert_to_nothing(self): + from docsgpt.retriever.graph_rag import _nodes_by_chunk + + assert _nodes_by_chunk({}) == {} + assert _nodes_by_chunk({"n1": None}) == {} diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 71c996d5..5203eee1 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -180,7 +180,9 @@ class TestGraphRAGPoolDiscipline: mock_store_cls.return_value = store rag = _make_retriever() - with patch.object(rag, "_graph_docs_for_source", return_value=[]): + # A real result: an empty one now falls back like a failure does. + graph_docs = [{"title": "g", "text": "graph text", "source": "src1", "filename": "g"}] + with patch.object(rag, "_graph_docs_for_source", return_value=graph_docs): with patch.object(rag, "_classic_for_sources") as classic: rag._get_data() @@ -355,15 +357,24 @@ class TestGraphRAGHappyPath: @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) @patch("docsgpt.retriever.graph_rag.GraphStore") @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) - def test_no_seeds_returns_empty( + def test_a_graph_that_answers_nothing_falls_back_to_classic( self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed ): + """Empty is not an answer. Every graph read swallows its own errors and + returns nothing, so "no rows" covers a broken query as much as a walk + that found nothing — and the source would contribute nothing at all, + with no fallback, because only a raise routes one to ClassicRAG.""" store = _store_with_graph([], [], {}, {}, []) store.count_nodes_many.side_effect = lambda ids: {s: 5 for s in ids} mock_store_cls.return_value = store rag = _make_retriever() - assert rag._get_data() == [] + seen = _recording_classic(rag, [_CLASSIC_DOC]) + + docs = rag._get_data() + + assert seen == [["src1"]] + assert [doc["text"] for doc in docs] == ["classic"] # ── IDF down-weighting ──────────────────────────────────────────────────────── From ccd8eb612ffb2476fe59cd435266b20b885eb9a4 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:04 +0100 Subject: [PATCH 23/27] fix(agents): hand the graph tool's connection back, and gate it like the rest Three faults in the graph tool, all on the agent's path: The store was cached on the tool, and the executor caches the tool for the whole agent run -- so one pgvector pooled connection stayed checked out across every LLM round trip of that run, minutes at a time, and enough concurrent runs exhaust the pool. GraphRAGRetriever releases its store before falling back for this reason. The tool now releases it at the end of each action. It gated on GRAPHRAG_ENABLED where everything else asks graphrag_available(), which also requires the pgvector store. Under any other vector store the graph tables are not the ones the sources were ingested into, but the tool was still offered and still queried Postgres. Pages were labelled by hand rather than through labels_from_metadata, which exists so citation labels match across retrievers. A page read by the tool and the same chunk retrieved by internal_search are one document, and citations key on (source, title) -- so the research agent gave that document two citation numbers. The recorded doc also keeps the full chunk text now, so it dedupes against the retriever's copy; only what the model reads is truncated. --- docsgpt/agents/tools/graph_search.py | 48 ++++++++--- tests/graphrag/test_graph_search_tool.py | 104 ++++++++++++++++++++++- 2 files changed, 138 insertions(+), 14 deletions(-) diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py index c2c15a7b..b370949e 100644 --- a/docsgpt/agents/tools/graph_search.py +++ b/docsgpt/agents/tools/graph_search.py @@ -28,7 +28,8 @@ import logging from typing import Any, Dict, List, Optional from docsgpt.agents.tools.base import Tool -from docsgpt.core.settings import settings +from docsgpt.graphrag import graphrag_available +from docsgpt.retriever.labels import labels_from_metadata logger = logging.getLogger(__name__) @@ -61,6 +62,24 @@ class GraphSearchTool(Tool): self._store = GraphStore() return self._store + def _release_store(self) -> None: + """Hand the pooled connection back at the end of an action. + + The executor caches this tool for the whole agent run, so a store kept + between actions pins one connection of the shared pgvector pool across + every LLM round trip of that run -- minutes at a time, and enough + concurrent runs exhaust the pool. ``GraphRAGRetriever`` releases its + store before falling back for the same reason. Checking one back out + costs a pool acquire. + """ + store, self._store = self._store, None + if store is None: + return + try: + store.close() + except Exception as exc: # noqa: BLE001 -- releasing must not fail an action + logger.debug(f"Graph tool could not release its store: {exc}") + def _embed(self, text: str) -> Optional[List[float]]: try: from docsgpt.vectorstore.base import get_embeddings @@ -72,7 +91,9 @@ class GraphSearchTool(Tool): # -- actions ------------------------------------------------------------- def execute_action(self, action_name: str, **kwargs): - if not settings.GRAPHRAG_ENABLED: + # The graph lives in the pgvector store, so the flag alone is not + # enough: under another vector store there is no graph to read. + if not graphrag_available(): return "The knowledge graph is not enabled for this deployment." if not self._sources(): return "No graph-backed sources are configured." @@ -86,6 +107,8 @@ class GraphSearchTool(Tool): except Exception as e: # noqa: BLE001 logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True) return "The graph lookup failed." + finally: + self._release_store() return f"Unknown action: {action_name}" def _search_entities(self, **kwargs) -> str: @@ -137,18 +160,17 @@ class GraphSearchTool(Tool): parts: List[str] = [] for source_id in self._sources(): for page in store.entity_pages(source_id, entity): - metadata = page.get("metadata") or {} - title = ( - metadata.get("file_path") - or metadata.get("title") - or metadata.get("source") - or "document" - ) - text = (page.get("text") or "")[:MAX_PAGE_CHARS] - doc = {"title": title, "text": text, "source": metadata.get("source", "")} + text = page.get("text") or "" + # The retrievers' own labelling: a page read here and the same + # chunk retrieved by internal_search are one document, and + # citations key on (source, title). Labelling it differently + # gives that document two citation numbers. + labels = labels_from_metadata(page.get("metadata"), text, source_id) + doc = {**labels, "text": text} if doc not in self.retrieved_docs: self.retrieved_docs.append(doc) - parts.append(f"--- {title} ---\n{text}") + header = labels["filename"] or labels["title"] + parts.append(f"--- {header} ---\n{text[:MAX_PAGE_CHARS]}") if not parts: return f"No documents mention {entity!r}." return "\n\n".join(parts) @@ -259,7 +281,7 @@ def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None: tool follows that same per-source exposure choice. A graph source left at ``prefetch`` in a classic agent is used for ranking only. """ - if not settings.GRAPHRAG_ENABLED: + if not graphrag_available(): return source = retriever_config.get("source") or {} if not source.get("active_docs") or not sources_have_graph(source): diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 559c911f..465a1c25 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -9,6 +9,8 @@ must refuse clearly rather than silently when it has nothing to offer. from __future__ import annotations +import pytest + from docsgpt.agents.tools.graph_search import ( GRAPH_TOOL_ID, GraphSearchTool, @@ -38,6 +40,7 @@ class _StubStore: def _tool(monkeypatch, store, enabled=True): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") tool = GraphSearchTool({"source": SOURCE}) tool._store = store monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1]) @@ -52,6 +55,7 @@ class TestGating: def test_reports_when_no_sources_are_configured(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") tool = GraphSearchTool({"source": {"active_docs": []}}) assert "No graph-backed sources" in tool.execute_action( @@ -109,16 +113,49 @@ class TestActions: tool = _tool( monkeypatch, _StubStore( - pages=[{"metadata": {"file_path": "quill-store.md"}, "text": "x" * 5000}] + pages=[ + { + "metadata": {"title": "quill-store.md", "source": "quill-store.md"}, + "text": "x" * 5000, + } + ] ), ) result = tool.execute_action("read_entity_pages", entity="Quill") assert "--- quill-store.md ---" in result + # Only what the model reads is truncated. assert len(result) < 3000 # Accumulated so the answer can cite what the walk actually read. assert tool.retrieved_docs[0]["title"] == "quill-store.md" + assert len(tool.retrieved_docs[0]["text"]) == 5000 + + def test_page_labels_match_what_the_retrievers_record(self, monkeypatch): + """A page read here and the same chunk retrieved by internal_search are + one document. The citation manager keys on (source, title), so labels + derived differently give the same document two citation numbers.""" + from docsgpt.retriever.labels import labels_from_metadata + + metadata = {"title": "Quill Store", "source": "quill-store.md"} + text = "Quill is a write-ahead store." + tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": metadata, "text": text}])) + + tool.execute_action("read_entity_pages", entity="Quill") + + expected = labels_from_metadata(metadata, text, "src-1") + doc = tool.retrieved_docs[0] + assert {k: doc[k] for k in ("title", "source", "filename")} == expected + # The full chunk text, so the doc dedupes against the retriever's copy; + # only what the model reads is truncated. + assert doc["text"] == text + + def test_a_page_with_no_metadata_falls_back_to_its_source_id(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": {}, "text": "body"}])) + + tool.execute_action("read_entity_pages", entity="Quill") + + assert tool.retrieved_docs[0]["source"] == "src-1" def test_pages_absent_is_stated_plainly(self, monkeypatch): tool = _tool(monkeypatch, _StubStore(pages=[])) @@ -150,6 +187,7 @@ class TestWiring: def test_not_added_when_the_sources_have_no_graph(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") monkeypatch.setattr( "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False ) @@ -161,6 +199,7 @@ class TestWiring: def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") monkeypatch.setattr( "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True ) @@ -246,3 +285,66 @@ class TestSourcesHaveGraph: assert sources_have_graph({"active_docs": []}) is False self._patch_counts(monkeypatch, error=RuntimeError("no pgvector")) assert sources_have_graph({"active_docs": ["a"]}) is False + + +class TestGraphsMustBeAvailable: + """The graph lives in the pgvector store, so the flag alone is not enough. + + With another vector store configured the graph tables are not the ones the + sources were ingested into; everything else in the app asks + ``graphrag_available()``, which requires both. + """ + + def test_the_tool_is_not_offered_without_pgvector(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "faiss") + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", + lambda source: pytest.fail("must not reach the database"), + ) + tools = {} + add_graph_search_tool(tools, {"source": SOURCE}) + assert tools == {} + + def test_actions_report_it_rather_than_querying(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "faiss") + tool = GraphSearchTool({"source": SOURCE}) + tool._store = _StubStore(nodes=[{"name": "Quill", "distance": 0.1}]) + + assert "not enabled" in tool.execute_action("search_entities", query="quill") + + +class TestPooledConnection: + """The tool is cached for the whole agent run; its connection must not be.""" + + class _ClosingStore(_StubStore): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.closed = 0 + + def close(self): + self.closed += 1 + + def test_the_connection_goes_back_after_each_action(self, monkeypatch): + store = self._ClosingStore(relationships=[{"source": "A", "target": "B", "type": "r"}]) + tool = _tool(monkeypatch, store) + + tool.execute_action("get_relationships", entity="A") + + # Held open, one pooled connection would be pinned across every LLM + # round trip of the run. + assert store.closed == 1 + assert tool._store is None + + def test_a_failing_action_still_releases_it(self, monkeypatch): + class _Broken(self._ClosingStore): + def entity_relationships(self, source_id, name, limit=25): + raise RuntimeError("connection lost") + + store = _Broken() + tool = _tool(monkeypatch, store) + + assert tool.execute_action("get_relationships", entity="A") == "The graph lookup failed." + assert store.closed == 1 + assert tool._store is None From 2b6d4d509e13d4be7d667b3d0010f4a77615b583 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:04 +0100 Subject: [PATCH 24/27] fix(graphrag): return a chunk once from entity_pages The subject flag was in the GROUP BY, so a chunk two matching entities link -- one the page is about, one merely mentioned in it -- came back as two identical pages and spent the caller's page budget twice on the same text. It is aggregated with bool_or now, which is what the ordering wanted anyway. Covered by a live test against a pgvector-shaped documents table: the graph tables alone cannot answer this query, so nothing exercised it before. --- docsgpt/graphrag/store.py | 8 +++-- tests/graphrag/test_store.py | 63 ++++++++++++++++++++++++++++++++++++ 2 files changed, 69 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 1804d929..0b73d031 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -1265,13 +1265,17 @@ class GraphStore: sql.SQL( """ SELECT d.{metadata}, d.{text}, - (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + bool_or(lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject FROM graph_node_chunks gc JOIN graph_nodes n ON n.id = gc.node_id JOIN {table} d ON d.id::text = gc.chunk_id WHERE gc.source_id = %s AND d.{source} = %s AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) - GROUP BY d.{metadata}, d.{text}, is_subject + -- One page per chunk. Grouping on the subject flag as well + -- split a chunk two entities link -- one naming it, one + -- merely mentioned -- into two identical pages, spending + -- the caller's page budget twice on the same text. + GROUP BY d.{metadata}, d.{text} ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC LIMIT %s; """ diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 062db1f2..51002d18 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -19,6 +19,7 @@ import uuid from unittest.mock import MagicMock, patch import pytest +from psycopg.types.json import Jsonb import docsgpt.graphrag.store as store_module from docsgpt.vectorstore import pgconn @@ -661,6 +662,68 @@ class TestGraphStoreParameterization: assert params[-1] == embedding +@pytest.mark.integration +class TestEntityPagesLive: + """``entity_pages`` against a real pgvector-shaped table. + + The graph tables alone cannot answer it: the rows it returns live in the + documents table the sources were ingested into, so the test creates a + minimal one with the same column names ``PGVectorStore`` uses. + """ + + @pytest.fixture + def store(self, postgresql): + store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) + try: + store._ensure_tables() + except Exception as exc: + pytest.skip(f"pgvector extension unavailable: {exc}") + conn = store._get_connection() + cursor = conn.cursor() + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS documents ( + id SERIAL PRIMARY KEY, + text TEXT, + metadata JSONB, + source_id TEXT + ); + """ + ) + conn.commit() + cursor.close() + yield store + store.close() + + def test_a_page_linked_by_two_entities_is_returned_once(self, store): + """One chunk, two nodes whose names both match: an exact hit and a + mention. They differ only in whether the page is *about* the entity, so + grouping on that flag returned the same page twice and spent a quarter + of the page budget on it.""" + source_id = str(uuid.uuid4()) + conn = store._get_connection() + cursor = conn.cursor() + cursor.execute( + "INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;", + ("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id), + ) + chunk_id = str(cursor.fetchone()[0]) + conn.commit() + cursor.close() + try: + subject = store.upsert_node(source_id, "Quill", "quill") + mention = store.upsert_node(source_id, "Legacy Quill", "legacy quill") + store.link_node_chunk(source_id, subject, chunk_id) + store.link_node_chunk(source_id, mention, chunk_id) + + pages = store.entity_pages(source_id, "Quill", limit=4) + + assert [page["text"] for page in pages] == ["Quill is a write-ahead store."] + assert pages[0]["metadata"] == {"title": "quill.md"} + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestGraphReadQueries: """The reads behind fact seeding and the agent's graph tool, without a DB. From ef5b718f371da4180c1667fe4aaf0ce7ef1ecb81 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:05 +0100 Subject: [PATCH 25/27] fix(graphrag): drop entities whose name normalizes to nothing ``canonical_name`` answers "" for a punctuation-only name, which callers are meant to read as "no entity" -- ``_resolve_endpoint`` already does. Entity extraction did not, and nodes merge on that key, so every such entity in a source collapsed onto one shared node that belonged to none of them. --- docsgpt/graphrag/extraction.py | 10 ++++++++-- tests/graphrag/test_extraction.py | 21 +++++++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 5ddbe1e1..daa8f910 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -424,7 +424,7 @@ def extract_graph_for_source( def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: - """Normalize the LLM's entity dicts (drop nameless ones).""" + """Normalize the LLM's entity dicts (drop the ones with no usable name).""" entities = [] for e in raw_entities: if not isinstance(e, dict): @@ -432,10 +432,16 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: name = str(e.get("name", "")).strip() if not name: continue + normalized_name = normalize_entity_name(name) + if not normalized_name: + # A punctuation-only name normalizes to nothing, and nodes merge on + # that key: keeping it collapses every such entity onto one shared + # node. The relationship side already drops them. + continue entities.append( { "name": name, - "normalized_name": normalize_entity_name(name), + "normalized_name": normalized_name, "type": str(e.get("type") or "") or None, "description": str(e.get("description") or "") or None, } diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 05bb0a40..1125184f 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -1085,6 +1085,27 @@ class TestSummaryCountFailure: assert store.count_nodes.call_args.kwargs.get("strict") is True +@pytest.mark.unit +class TestEntityNormalization: + """An entity whose name normalizes to nothing is not an entity. + + ``canonical_name`` answers "" for a punctuation-only name, and nodes are + merged on that key, so keeping them collapses every such entity onto one + shared node. ``_resolve_endpoint`` already drops them on the relationship + side. + """ + + @pytest.mark.parametrize("name", ["!!!", "--", "?", " *** "]) + def test_a_name_that_normalizes_to_nothing_is_dropped(self, name): + assert extraction_module._build_entities([{"name": name}]) == [] + + def test_real_names_survive(self): + built = extraction_module._build_entities( + [{"name": "Quill Store"}, {"name": "!!!"}, {"name": "Alder"}] + ) + assert [e["normalized_name"] for e in built] == ["quill store", "alder"] + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): From 692edcdd4ea24ee053b21e75ba0dddc4760a87e5 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:54:55 +0100 Subject: [PATCH 26/27] test(agents): give the graph tool tests the vector store they need The tool now asks graphrag_available(), which wants pgvector as well as the flag. One case set only the flag and passed locally off a dev .env, then failed on CI's faiss default. An autouse fixture sets it for the module, so a case that forgets fails for its own reason; the two about a different vector store override it themselves. --- tests/graphrag/test_graph_search_tool.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 465a1c25..ca2f1330 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -22,6 +22,17 @@ from docsgpt.core.settings import settings SOURCE = {"active_docs": ["src-1"]} +@pytest.fixture(autouse=True) +def _graph_store_is_pgvector(monkeypatch): + """The graph only exists under pgvector, and CI's default is faiss. + + Set here rather than per test so a case that forgets it fails for its own + reason instead of the gate; the cases about a different vector store + override it in the test body. + """ + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") + + class _StubStore: def __init__(self, nodes=None, relationships=None, pages=None): self._nodes = nodes or [] From 63d93fc4b9ba5baccd268ffadd4169b79a431a53 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:54:55 +0100 Subject: [PATCH 27/27] test(graphrag): pin that identical pages collapse into one Review asked whether two document rows with the same text and metadata should stay two pages. They should not: the caller gets four pages to hand a model, and a crawl that ingested the same text twice would spend two of them on it. The rows differ only by an id the model never sees. Pinned either way now. --- tests/graphrag/test_store.py | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 51002d18..d9de877a 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -724,6 +724,34 @@ class TestEntityPagesLive: store.delete_by_source(source_id) + def test_two_chunks_with_identical_text_collapse_into_one_page(self, store): + """Deliberate: the caller gets at most four pages to hand a model, and a + crawl that ingested the same text twice would spend two of them saying + the same thing. The rows differ only by an id the model never sees.""" + source_id = str(uuid.uuid4()) + conn = store._get_connection() + cursor = conn.cursor() + chunk_ids = [] + for _ in range(2): + cursor.execute( + "INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;", + ("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id), + ) + chunk_ids.append(str(cursor.fetchone()[0])) + conn.commit() + cursor.close() + try: + node = store.upsert_node(source_id, "Quill", "quill") + for chunk_id in chunk_ids: + store.link_node_chunk(source_id, node, chunk_id) + + pages = store.entity_pages(source_id, "Quill", limit=4) + + assert [page["text"] for page in pages] == ["Quill is a write-ahead store."] + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestGraphReadQueries: """The reads behind fact seeding and the agent's graph tool, without a DB.