fix(graphrag): stop losing chunks silently during a graph build

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.
This commit is contained in:
Alex committed 2026-09-17 15:44:23 +01:00
1 parent 92e19ac177
commit e3d819d9fd
4 files changed
+379 -73

No files matched your search

+40 -9
View File
@@ -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"<chunk>\n{text}\n</chunk>"},
@@ -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,
+124 -64
View File
@@ -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."""
+103
View File
@@ -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):
+112
View File
@@ -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()