diff --git a/datamind/capabilities/graph/providers/networkx_store.py b/datamind/capabilities/graph/providers/networkx_store.py index 136b37e..d19b631 100644 --- a/datamind/capabilities/graph/providers/networkx_store.py +++ b/datamind/capabilities/graph/providers/networkx_store.py @@ -15,6 +15,8 @@ import difflib import json import os +import tempfile +import threading from pathlib import Path from typing import Any, Sequence @@ -41,6 +43,10 @@ def __init__( self._path.parent.mkdir(parents=True, exist_ok=True) self._g: nx.MultiDiGraph = nx.MultiDiGraph() self._dirty = False + self._mutation_revision = 0 + self._state_lock = threading.RLock() + self._persist_lock = asyncio.Lock() + self._persist_worker_lock = threading.Lock() if autoload and self._path.exists(): self._load() _log.info( @@ -78,115 +84,155 @@ def _load(self) -> None: ) async def persist(self) -> None: - if not self._dirty: - return - def _run() -> None: - doc = { - "nodes": [ - { - "id": nid, - "label": d.get("label", nid), - "type": d.get("type", "entity"), - "props": {k: v for k, v in d.items() if k not in {"label", "type"}}, - } - for nid, d in self._g.nodes(data=True) - ], - "edges": [ - { - "src": u, - "dst": v, - "key": key, - "rel": d.get("relation", "related"), - "w": float(d.get("weight", 1.0)), - "props": { - k: val - for k, val in d.items() - if k not in {"relation", "weight"} - }, - } - for u, v, key, d in self._g.edges(keys=True, data=True) - ], - } - temporary = self._path.with_name(f".{self._path.name}.{os.getpid()}.tmp") - temporary.write_text( - json.dumps(doc, ensure_ascii=False, indent=2), - encoding="utf-8", - ) + async with self._persist_lock: + await asyncio.to_thread(self._persist_sync) + + def _persist_sync(self) -> None: + # Cancelling to_thread's awaiter does not stop its worker. Serialize + # capture, replacement, and dirty-state bookkeeping in the worker too, + # so a cancelled save cannot overwrite a later successful save. + with self._persist_worker_lock: + with self._state_lock: + if not self._dirty: + return + captured_revision = self._mutation_revision + doc = self._document_locked() + + self._write_document(doc) + + with self._state_lock: + # A mutation may have happened while the captured document + # was being written. In that case the newer state is not on + # disk yet and must remain dirty for the next persist call. + if self._mutation_revision == captured_revision: + self._dirty = False + + def _document_locked(self) -> dict[str, Any]: + """Build a detached JSON document while holding ``_state_lock``.""" + return { + "nodes": [ + { + "id": nid, + "label": d.get("label", nid), + "type": d.get("type", "entity"), + "props": {k: v for k, v in d.items() if k not in {"label", "type"}}, + } + for nid, d in self._g.nodes(data=True) + ], + "edges": [ + { + "src": u, + "dst": v, + "key": key, + "rel": d.get("relation", "related"), + "w": float(d.get("weight", 1.0)), + "props": { + k: val + for k, val in d.items() + if k not in {"relation", "weight"} + }, + } + for u, v, key, d in self._g.edges(keys=True, data=True) + ], + } + + def _write_document(self, doc: dict[str, Any]) -> None: + """Atomically write one detached document through a unique temp file.""" + fd, temporary = tempfile.mkstemp( + prefix=f".{self._path.name}.", suffix=".tmp", dir=self._path.parent, + ) + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(doc, handle, ensure_ascii=False, indent=2) os.replace(temporary, self._path) - await asyncio.to_thread(_run) - self._dirty = False + except BaseException: + try: + os.unlink(temporary) + except FileNotFoundError: + pass + raise # ------------------------------------------------------------ mutation async def upsert_triples(self, triples: Sequence[GraphTriple]) -> None: - for t in triples: - # Nodes - for side, node_id, node_type in ( - ("subject", t.subject, t.subject_type), - ("object", t.object, t.object_type), - ): - if not self._g.has_node(node_id): - self._g.add_node( - node_id, - label=node_id, - type=node_type, - ) - # Profile snapshots and runtime writes have separate identities; - # exact duplicates within either origin overwrite deterministically. - profile_managed = bool((t.properties or {}).get("_profile_managed")) - origin = "profile" if profile_managed else (t.source or "runtime") - self._g.add_edge( - t.subject, - t.object, - key=f"{t.relation}\x1f{origin}", - relation=t.relation, - weight=float(t.confidence), - source=t.source, - **{f"p_{k}": v for k, v in (t.properties or {}).items()}, - ) - self._dirty = True + with self._state_lock: + for t in triples: + # Nodes + for side, node_id, node_type in ( + ("subject", t.subject, t.subject_type), + ("object", t.object, t.object_type), + ): + if not self._g.has_node(node_id): + self._g.add_node( + node_id, + label=node_id, + type=node_type, + ) + # Profile snapshots and runtime writes have separate identities; + # exact duplicates within either origin overwrite deterministically. + profile_managed = bool((t.properties or {}).get("_profile_managed")) + origin = "profile" if profile_managed else (t.source or "runtime") + self._g.add_edge( + t.subject, + t.object, + key=f"{t.relation}\x1f{origin}", + relation=t.relation, + weight=float(t.confidence), + source=t.source, + **{f"p_{k}": v for k, v in (t.properties or {}).items()}, + ) + self._mutation_revision += 1 + self._dirty = True async def reconcile_profile_triples(self, triples: Sequence[GraphTriple]) -> None: """Replace only edges managed by the profile snapshot.""" - stale = [ - (u, v, key) - for u, v, key, data in self._g.edges(keys=True, data=True) - if data.get("p__profile_managed") is True - ] - self._g.remove_edges_from(stale) + with self._state_lock: + stale = [ + (u, v, key) + for u, v, key, data in self._g.edges(keys=True, data=True) + if data.get("p__profile_managed") is True + ] + self._g.remove_edges_from(stale) await self.upsert_triples(triples) # Drop now-orphaned profile nodes without touching runtime nodes. - self._g.remove_nodes_from(list(nx.isolates(self._g))) + with self._state_lock: + self._g.remove_nodes_from(list(nx.isolates(self._g))) async def reconcile_source_triples( self, source: str, triples: Sequence[GraphTriple] ) -> None: """Replace edges produced from one source file while preserving others.""" - stale = [ - (u, v, key) - for u, v, key, data in self._g.edges(keys=True, data=True) - if data.get("p__source_path") == source - ] - self._g.remove_edges_from(stale) + with self._state_lock: + stale = [ + (u, v, key) + for u, v, key, data in self._g.edges(keys=True, data=True) + if data.get("p__source_path") == source + ] + self._g.remove_edges_from(stale) await self.upsert_triples(triples) - self._g.remove_nodes_from(list(nx.isolates(self._g))) + with self._state_lock: + self._g.remove_nodes_from(list(nx.isolates(self._g))) async def reconcile_lineage_triples( self, root: str, triples: Sequence[GraphTriple] ) -> None: """Replace all deterministic lineage edges for one workspace root.""" - stale = [ - (u, v, key) - for u, v, key, data in self._g.edges(keys=True, data=True) - if data.get("p__lineage_root") == root - ] - self._g.remove_edges_from(stale) + with self._state_lock: + stale = [ + (u, v, key) + for u, v, key, data in self._g.edges(keys=True, data=True) + if data.get("p__lineage_root") == root + ] + self._g.remove_edges_from(stale) await self.upsert_triples(triples) - self._g.remove_nodes_from(list(nx.isolates(self._g))) + with self._state_lock: + self._g.remove_nodes_from(list(nx.isolates(self._g))) async def reset(self) -> None: - self._g = nx.MultiDiGraph() - self._dirty = True + with self._state_lock: + self._g = nx.MultiDiGraph() + self._mutation_revision += 1 + self._dirty = True # ------------------------------------------------------------- lookup diff --git a/datamind/tests/test_graph_persistence_races.py b/datamind/tests/test_graph_persistence_races.py new file mode 100644 index 0000000..eae022d --- /dev/null +++ b/datamind/tests/test_graph_persistence_races.py @@ -0,0 +1,98 @@ +"""Deterministic persistence-race coverage for the NetworkX store.""" +from __future__ import annotations + +import asyncio +import json +import threading + +import pytest + +from datamind.capabilities.graph.providers.networkx_store import NetworkXGraphStore +from datamind.core.protocols import GraphTriple + + +def triple(subject: str, relation: str, object_: str) -> GraphTriple: + return GraphTriple(subject=subject, relation=relation, object=object_) + + +@pytest.mark.asyncio +async def test_mutation_during_save_stays_dirty_until_new_revision_is_written(tmp_path): + path = tmp_path / "graph.json" + store = NetworkXGraphStore(persist_path=path) + await store.upsert_triples([triple("A", "old", "B")]) + + started = threading.Event() + release = threading.Event() + original_write = store._write_document + + def blocked_write(document): + started.set() + assert release.wait(5) + original_write(document) + + store._write_document = blocked_write + first = asyncio.create_task(store.persist()) + await asyncio.to_thread(started.wait, 5) + + await store.upsert_triples([triple("B", "new", "C")]) + release.set() + await first + + assert store._dirty is True + store._write_document = original_write + await store.persist() + + saved = json.loads(path.read_text(encoding="utf-8")) + assert {edge["rel"] for edge in saved["edges"]} == {"old", "new"} + assert NetworkXGraphStore(persist_path=path).stats()["edges"] == 2 + + +@pytest.mark.asyncio +async def test_overlapping_persists_are_serialized_and_keep_the_latest_document(tmp_path): + path = tmp_path / "graph.json" + store = NetworkXGraphStore(persist_path=path) + await store.upsert_triples([triple("A", "old", "B")]) + + started = threading.Event() + release = threading.Event() + original_write = store._write_document + writes = [] + + def blocked_write(document): + writes.append(document) + if len(writes) == 1: + started.set() + assert release.wait(5) + original_write(document) + + store._write_document = blocked_write + first = asyncio.create_task(store.persist()) + await asyncio.to_thread(started.wait, 5) + await store.upsert_triples([triple("B", "new", "C")]) + second = asyncio.create_task(store.persist()) + release.set() + await asyncio.gather(first, second) + + assert len(writes) == 2 + saved = json.loads(path.read_text(encoding="utf-8")) + assert {edge["rel"] for edge in saved["edges"]} == {"old", "new"} + assert store._dirty is False + + +@pytest.mark.asyncio +async def test_write_failure_keeps_store_dirty_for_retry(tmp_path): + store = NetworkXGraphStore(persist_path=tmp_path / "graph.json") + await store.upsert_triples([triple("A", "r", "B")]) + original_write = store._write_document + + def fail_write(document): + raise OSError("disk full") + + store._write_document = fail_write + with pytest.raises(OSError, match="disk full"): + await store.persist() + assert store._dirty is True + + store._write_document = original_write + await store.persist() + assert store._dirty is False diff --git a/datamind/tests/test_review_cancelled_persist.py b/datamind/tests/test_review_cancelled_persist.py new file mode 100644 index 0000000..e77ce1d --- /dev/null +++ b/datamind/tests/test_review_cancelled_persist.py @@ -0,0 +1,58 @@ +import asyncio +import json +import threading + +import pytest + +from datamind.capabilities.graph.providers.networkx_store import NetworkXGraphStore +from datamind.core.protocols import GraphTriple + + +@pytest.mark.asyncio +async def test_cancelled_save_cannot_overwrite_a_later_successful_save(tmp_path): + path = tmp_path / 'graph.json' + store = NetworkXGraphStore(persist_path=path) + await store.upsert_triples([GraphTriple(subject='A', relation='old', object='B')]) + started = threading.Event() + release = threading.Event() + finished = threading.Event() + original_write = store._write_document + count = 0 + + def controlled_write(doc): + nonlocal count + count += 1 + if count == 1: + started.set() + try: + assert release.wait(10) + original_write(doc) + finally: + finished.set() + else: + original_write(doc) + + store._write_document = controlled_write + first = asyncio.create_task(store.persist()) + second = None + try: + assert await asyncio.to_thread(started.wait, 5) + first.cancel() + with pytest.raises(asyncio.CancelledError): + await first + await store.upsert_triples([GraphTriple(subject='B', relation='new', object='C')]) + second = asyncio.create_task(store.persist()) + # Allow an incorrectly unlocked second writer to finish before the + # first is released; a correctly serialized writer remains pending. + await asyncio.wait({second}, timeout=0.2) + finally: + release.set() + assert await asyncio.to_thread(finished.wait, 5) + if second is not None: + await second + # Cancellation of to_thread does not stop its worker. The old writer must + # never replace the newer committed document after the async lock releases. + assert len(json.loads(path.read_text())['edges']) == 2 + assert store._dirty is False + await store.persist() + assert len(json.loads(path.read_text())['edges']) == 2