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)