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()