fix(graphrag): review fixes — replay-safe writes, strict count, zero weights

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.
This commit is contained in:
Alex committed 2026-09-17 16:31:19 +01:00
1 parent 9e9f130ed0
commit 3f774d813c
6 files changed
+274 -9

No files matched your search

+7 -2
View File
@@ -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",
+48 -4
View File
@@ -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()
+11 -3
View File
@@ -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 {}
+36
View File
@@ -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):
+112
View File
@@ -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)
+60
View File
@@ -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)