mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
fix(graphrag): take a source's write lock per chunk, and keep zero weights
Two builds of one source can overlap: the extraction lease is keyed by the
source's updated_at, and enabling a graph updates the source before it
dispatches, so a rebuild started while the last build runs gets a new key and
a lease of its own. Both builds could then pass a chunk's "done" check before
either committed and apply it twice -- doc_freq bumped twice, reproduced with
two live writers. A reset could also land in the middle of a chunk.
apply_chunk and delete_by_source now take a transaction-scoped advisory lock
keyed by the source before touching a row, as the schema bootstrap already
does for DDL. A single build's writes were already serial, so it loses
nothing; overlapping builds take turns chunk by chunk, and the second sees the
first's "done" row and returns (0, 0).
apply_chunk also still defaulted with `rel.get("weight") or 1.0`, turning an
explicit zero into a full-strength edge -- the conversion 3f774d81 removed from
add_edge and the ranker but missed here. Only a missing weight defaults now.
This commit is contained in:
1 parent
b3d7e63ae3
commit
ecf02d0d60
2 files changed
+119
-3
No files matched your search
@@ -108,6 +108,22 @@ def _is_connection_lost(exc: BaseException) -> bool:
|
||||
return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError))
|
||||
|
||||
|
||||
def _lock_source(cursor, source_id: str) -> None:
|
||||
"""Serialize graph writes for one source until this transaction ends.
|
||||
|
||||
Writes within one build are already serial, but two builds of the same
|
||||
source can overlap: a rebuild dispatched while the last one is still
|
||||
running gets a new idempotency key, so its lease does not stop it. Without
|
||||
this, both could pass a chunk's "done" check before either commits and
|
||||
apply it twice. A transaction-scoped advisory lock keyed by the source
|
||||
makes them take turns chunk by chunk; the lock is released on commit or
|
||||
rollback, and a hash collision only makes two sources take turns.
|
||||
"""
|
||||
cursor.execute(
|
||||
"SELECT pg_advisory_xact_lock(hashtext(%s));", (f"graphrag:source:{source_id}",)
|
||||
)
|
||||
|
||||
|
||||
def _safe_rollback(conn) -> None:
|
||||
"""Roll back, tolerating a connection too broken to roll back."""
|
||||
try:
|
||||
@@ -672,7 +688,10 @@ class GraphStore:
|
||||
# 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.
|
||||
# below is written in this transaction, so a replay sees it —
|
||||
# and so does an overlapping build, once the source lock makes
|
||||
# it wait for this one to commit.
|
||||
_lock_source(cursor, source_id)
|
||||
cursor.execute(
|
||||
"SELECT status FROM graph_ingest_progress "
|
||||
"WHERE source_id = %s AND chunk_id = %s;",
|
||||
@@ -713,7 +732,9 @@ class GraphStore:
|
||||
dst_id,
|
||||
type=rel.get("type"),
|
||||
description=rel.get("description"),
|
||||
weight=float(rel.get("weight") or 1.0),
|
||||
# Only a missing weight defaults: 0 is a real one,
|
||||
# and the ranker drops non-positive edges.
|
||||
weight=1.0 if rel.get("weight") is None else float(rel["weight"]),
|
||||
source_chunk_ids=[chunk_id],
|
||||
fact_embedding=rel.get("fact_embedding"),
|
||||
)
|
||||
@@ -1527,6 +1548,9 @@ class GraphStore:
|
||||
conn = self._get_connection()
|
||||
cursor = conn.cursor()
|
||||
try:
|
||||
# A reset while a build is still writing must not land in the
|
||||
# middle of one of its chunks.
|
||||
_lock_source(cursor, source_id)
|
||||
for table in (
|
||||
"graph_node_chunks",
|
||||
"graph_edges",
|
||||
|
||||
@@ -557,6 +557,43 @@ class TestGraphStoreParameterization:
|
||||
store._tables_ensured = True
|
||||
return store, cursor
|
||||
|
||||
def test_graph_writes_for_a_source_are_serialized(self):
|
||||
# A chunk write and a reset each take the source's transaction-scoped
|
||||
# advisory lock before touching a row, so overlapping builds of one
|
||||
# source cannot interleave inside a chunk.
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
cursor.fetchone.return_value = None
|
||||
sid = str(uuid.uuid4())
|
||||
|
||||
store.apply_chunk(sid, "c1", [], [], {})
|
||||
first_sql, first_params = cursor.execute.call_args_list[0].args
|
||||
assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql
|
||||
assert first_params == (f"graphrag:source:{sid}",)
|
||||
|
||||
cursor.execute.reset_mock()
|
||||
store.delete_by_source(sid)
|
||||
first_sql, first_params = cursor.execute.call_args_list[0].args
|
||||
assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql
|
||||
assert first_params == (f"graphrag:source:{sid}",)
|
||||
|
||||
def test_apply_chunk_keeps_an_explicit_zero_weight(self, monkeypatch):
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
cursor.fetchone.side_effect = [None, ["n1"], ["n2"]]
|
||||
weights = []
|
||||
|
||||
def _capture(cursor, source_id, src, dst, type=None, description=None, weight=1.0, **kwargs):
|
||||
weights.append(weight)
|
||||
return "e1", True
|
||||
|
||||
monkeypatch.setattr(store, "_add_edge", _capture)
|
||||
store.apply_chunk(
|
||||
"sid", "c1", [],
|
||||
[{"source": "A", "target": "B", "weight": 0}, {"source": "A", "target": "B"}],
|
||||
{},
|
||||
)
|
||||
# Zero is a real weight; only a missing one defaults.
|
||||
assert weights == [0.0, 1.0]
|
||||
|
||||
def test_delete_by_source_binds_source_id(self):
|
||||
from psycopg import sql as pgsql
|
||||
|
||||
@@ -565,7 +602,9 @@ class TestGraphStoreParameterization:
|
||||
store.delete_by_source(sid)
|
||||
|
||||
tables = []
|
||||
for call in cursor.execute.call_args_list:
|
||||
lock, *deletes = cursor.execute.call_args_list
|
||||
assert "pg_advisory_xact_lock" in lock.args[0]
|
||||
for call in deletes:
|
||||
query = call.args[0]
|
||||
params = call.args[1] if len(call.args) > 1 else None
|
||||
assert isinstance(query, pgsql.Composable)
|
||||
@@ -1342,6 +1381,59 @@ class TestApplyChunkIsReplaySafe:
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_overlapping_applies_of_one_chunk_write_it_once(self, store, postgresql, monkeypatch):
|
||||
"""Two builds of one source can overlap: a rebuild dispatched while the
|
||||
last one runs gets a new lease key. Both may reach the same chunk at
|
||||
once, and the second must wait for the first to commit instead of
|
||||
passing the done check while the first is still in flight."""
|
||||
import threading
|
||||
import time
|
||||
|
||||
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)}
|
||||
writers = [GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) for _ in range(2)]
|
||||
real_upsert = GraphStore._upsert_node
|
||||
|
||||
def _slow_upsert(self, *args, **kwargs):
|
||||
time.sleep(0.3) # hold the first writer inside its transaction
|
||||
return real_upsert(self, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(GraphStore, "_upsert_node", _slow_upsert)
|
||||
results = []
|
||||
|
||||
def _apply(writer):
|
||||
results.append(writer.apply_chunk(source_id, "c1", entities, relationships, embeddings))
|
||||
|
||||
try:
|
||||
threads = [threading.Thread(target=_apply, args=(w,)) for w in writers]
|
||||
threads[0].start()
|
||||
time.sleep(0.05)
|
||||
threads[1].start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
assert sorted(results) == [(0, 0), (1, 1)]
|
||||
assert store.get_node_by_normalized(source_id, "ada")["doc_freq"] == 1
|
||||
assert len(store.get_graph_overview(source_id)["edges"]) == 1
|
||||
finally:
|
||||
for writer in writers:
|
||||
writer.close()
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_a_zero_weight_relationship_stays_zero(self, store):
|
||||
source_id = str(uuid.uuid4())
|
||||
relationships = [{"source": "Ada", "target": "Engine", "type": "mentions", "weight": 0.0}]
|
||||
try:
|
||||
store.apply_chunk(source_id, "c1", [], relationships, {})
|
||||
edges = store.get_graph_overview(source_id)["edges"]
|
||||
assert [edge["weight"] for edge in edges] == [0.0]
|
||||
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())
|
||||
|
||||
Reference in new issue
Block a user