mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 18:13:03 +00:00
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:
1 parent
9e9f130ed0
commit
3f774d813c
6 files changed
+274
-9
No files matched your search
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user