mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 14:12:58 +00:00
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.
1458 lines
56 KiB
Python
1458 lines
56 KiB
Python
"""Tests for the GraphRAG GraphStore (on-demand tables in the pgvector DB).
|
|
|
|
Two layers:
|
|
|
|
* A live-pg integration test that exercises the real DDL + SQL against the
|
|
pgvector store DB (same connection-string source as ``PGVectorStore``). It
|
|
uses a unique temp ``source_id`` and tears down every row it creates.
|
|
* A mock-cursor test that asserts the parameterized SQL shapes — ``source_id``
|
|
and embeddings are bound params, never interpolated.
|
|
|
|
The embedding dimension is mocked everywhere so the suite never loads the real
|
|
SentenceTransformer model: the live store creates ``TEST_EMBEDDING_DIM`` vectors
|
|
and the helpers build matching ones.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import uuid
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
import docsgpt.graphrag.store as store_module
|
|
from docsgpt.vectorstore import pgconn
|
|
from docsgpt.vectorstore import pgvector as pgvector_module
|
|
|
|
GraphStore = store_module.GraphStore
|
|
|
|
TEST_EMBEDDING_DIM = 8
|
|
|
|
POOL_DSN = "postgresql://u:p@localhost/graphpool"
|
|
|
|
_REAL_EMBEDDING_DIM = GraphStore._embedding_dim
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _mock_embedding_dim(monkeypatch):
|
|
monkeypatch.setattr(
|
|
GraphStore, "_embedding_dim", lambda self: TEST_EMBEDDING_DIM
|
|
)
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _close_pools():
|
|
"""Never leak a pool into another test; an ephemeral DSN dies with its DB."""
|
|
yield
|
|
for dsn, pool in list(pgconn._POOLS.items()):
|
|
try:
|
|
pool.close()
|
|
except Exception:
|
|
pass
|
|
pgconn._POOLS.pop(dsn, None)
|
|
|
|
|
|
def _ephemeral_dsn(info) -> str:
|
|
"""libpq DSN for the ephemeral pytest-postgresql database.
|
|
|
|
Deliberately not the operator's ``POSTGRES_URI``: these tests create and
|
|
drop graph tables, and the dev database is not theirs to rewrite.
|
|
"""
|
|
password = f":{info.password}" if info.password else ""
|
|
return (
|
|
f"postgresql://{info.user}{password}@{info.host}:{info.port}/{info.dbname}"
|
|
)
|
|
|
|
|
|
def _embedding(seed: float) -> list:
|
|
vec = [0.0] * TEST_EMBEDDING_DIM
|
|
vec[0] = seed
|
|
return vec
|
|
|
|
|
|
@pytest.mark.integration
|
|
class TestGraphStoreLive:
|
|
@pytest.fixture
|
|
def store(self, postgresql):
|
|
"""Graph store on a fresh ephemeral database.
|
|
|
|
Construction no longer creates tables, so the fixture calls
|
|
``_ensure_tables`` explicitly — exactly what ``ensure_vector_schema``
|
|
does at boot in production — and read-before-write tests still pass.
|
|
"""
|
|
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()
|
|
|
|
@pytest.fixture
|
|
def source_id(self):
|
|
return str(uuid.uuid4())
|
|
|
|
def test_ensure_tables_idempotent(self, store):
|
|
store._ensure_tables()
|
|
store._ensure_tables()
|
|
|
|
def test_upsert_node_merges_by_normalized_name(self, store, source_id):
|
|
try:
|
|
first = store.upsert_node(
|
|
source_id=source_id,
|
|
name="Ada Lovelace",
|
|
normalized_name="ada lovelace",
|
|
type="person",
|
|
description="A mathematician.",
|
|
name_embedding=_embedding(1.0),
|
|
)
|
|
second = store.upsert_node(
|
|
source_id=source_id,
|
|
name="Ada Lovelace",
|
|
normalized_name="ada lovelace",
|
|
type="person",
|
|
description="Wrote the first algorithm.",
|
|
)
|
|
assert first == second
|
|
|
|
node = store.get_node_by_normalized(source_id, "ada lovelace")
|
|
assert node is not None
|
|
assert node["id"] == first
|
|
assert node["doc_freq"] == 2
|
|
assert "mathematician" in node["description"]
|
|
assert "first algorithm" in node["description"]
|
|
|
|
duplicate = store.upsert_node(
|
|
source_id=source_id,
|
|
name="Ada Lovelace",
|
|
normalized_name="ada lovelace",
|
|
description="Wrote the first algorithm.",
|
|
)
|
|
assert duplicate == first
|
|
node = store.get_node_by_normalized(source_id, "ada lovelace")
|
|
assert node["description"].count("first algorithm") == 1
|
|
|
|
assert store.count_nodes(source_id) == 1
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_add_edge_and_link_chunk(self, store, source_id):
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a", "thing", "desc a")
|
|
b = store.upsert_node(source_id, "B", "b", "thing", "desc b")
|
|
store.add_edge(
|
|
source_id, a, b, "related", "a relates to b", 2.0, ["chunk-1"]
|
|
)
|
|
store.link_node_chunk(source_id, a, "chunk-1")
|
|
store.link_node_chunk(source_id, a, "chunk-1")
|
|
store.link_node_chunk(source_id, b, "chunk-1")
|
|
|
|
mapping = store.get_chunk_ids_for_nodes(source_id, [a, b])
|
|
assert mapping[a] == ["chunk-1"]
|
|
assert mapping[b] == ["chunk-1"]
|
|
|
|
store.set_node_degrees(source_id)
|
|
node_a = store.get_node_by_normalized(source_id, "a")
|
|
assert node_a["degree"] == 1
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_seed_nodes_from_facts_returns_both_endpoints_of_the_match(
|
|
self, store, source_id
|
|
):
|
|
"""Fact seeding's whole point: the question matches the *relationship*,
|
|
and both of its endpoints become seeds — including the one the question
|
|
never names."""
|
|
try:
|
|
alder = store.upsert_node(source_id, "Alder", "alder", "service", "d")
|
|
quill = store.upsert_node(source_id, "Quill", "quill", "store", "d")
|
|
birch = store.upsert_node(source_id, "Birch", "birch", "service", "d")
|
|
ridge = store.upsert_node(source_id, "Ridge", "ridge", "store", "d")
|
|
store.add_edge(
|
|
source_id, alder, quill, "streams_to", "Alder streams to Quill",
|
|
1.0, ["c1"], fact_embedding=_embedding(1.0),
|
|
)
|
|
store.add_edge(
|
|
source_id, birch, ridge, "streams_to", "Birch streams to Ridge",
|
|
1.0, ["c2"], fact_embedding=_embedding(-1.0),
|
|
)
|
|
|
|
rows = store.seed_nodes_from_facts(
|
|
source_id, _embedding(1.0), fact_limit=1, limit=10
|
|
)
|
|
|
|
assert {row["name"] for row in rows} == {"Alder", "Quill"}
|
|
assert all(row["distance"] <= 1.0 for row in rows)
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_seed_nodes_from_facts_is_empty_without_fact_embeddings(
|
|
self, store, source_id
|
|
):
|
|
"""A source ingested before fact embeddings existed returns nothing,
|
|
which is the signal the retriever falls back to name matching on."""
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a")
|
|
b = store.upsert_node(source_id, "B", "b")
|
|
store.add_edge(source_id, a, b, "rel")
|
|
|
|
assert store.seed_nodes_from_facts(source_id, _embedding(1.0)) == []
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_add_edge_skips_self_loops(self, store, source_id):
|
|
"""A relationship whose endpoints resolve to one node is noise.
|
|
|
|
A self-loop feeds a node's PageRank mass straight back to itself, and a
|
|
real extraction produced 121 of them on a 98-page corpus.
|
|
"""
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a", "thing", "desc a")
|
|
assert (
|
|
store.add_edge(source_id, a, a, "related", "a relates to a", 1.0, ["c1"])
|
|
is None
|
|
)
|
|
assert store.get_subgraph(source_id, [a], hops=1)["edges"] == []
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_add_edge_merges_a_repeated_pair(self, store, source_id):
|
|
"""The same relationship seen in many chunks is one edge, not many rows.
|
|
|
|
``graph_edges`` carries no uniqueness constraint, so re-extracting a
|
|
relationship used to insert a row per chunk — 19.9% of a real corpus's
|
|
edges — inflating traversal weight and wasting the subgraph fetch
|
|
budget. The surviving row keeps the strongest weight and both chunk ids.
|
|
"""
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a", "thing", "desc a")
|
|
b = store.upsert_node(source_id, "B", "b", "thing", "desc b")
|
|
first = store.add_edge(source_id, a, b, "related", "d", 2.0, ["chunk-1"])
|
|
second = store.add_edge(source_id, a, b, "related", "d", 5.0, ["chunk-2"])
|
|
|
|
assert second == first
|
|
edges = store.get_subgraph(source_id, [a, b], hops=1)["edges"]
|
|
assert len(edges) == 1
|
|
assert float(edges[0]["weight"]) == 5.0
|
|
|
|
# Both chunks are still recorded as evidence for the merged edge.
|
|
conn = store._get_connection()
|
|
cursor = conn.cursor()
|
|
try:
|
|
cursor.execute(
|
|
"SELECT source_chunk_ids FROM graph_edges WHERE id = %s;", (first,)
|
|
)
|
|
chunk_ids = cursor.fetchone()[0]
|
|
finally:
|
|
cursor.close()
|
|
conn.rollback()
|
|
assert sorted(chunk_ids) == ["chunk-1", "chunk-2"]
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_get_subgraph_keeps_the_heaviest_edges_when_capped(
|
|
self, store, source_id, monkeypatch
|
|
):
|
|
"""A capped fetch must drop the weakest edges, not an arbitrary subset.
|
|
|
|
The cap is applied with ``LIMIT``; without an ordering Postgres is free
|
|
to return any rows at all, so a dense graph silently retrieves a random
|
|
neighbourhood.
|
|
"""
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a", "thing", "d")
|
|
b = store.upsert_node(source_id, "B", "b", "thing", "d")
|
|
c = store.upsert_node(source_id, "C", "c", "thing", "d")
|
|
store.add_edge(source_id, a, b, "light", "d", 1.0, ["c1"])
|
|
store.add_edge(source_id, a, c, "heavy", "d", 9.0, ["c1"])
|
|
|
|
monkeypatch.setattr(store_module, "MAX_SUBGRAPH_EDGES", 1)
|
|
edges = store.get_subgraph(source_id, [a], hops=1)["edges"]
|
|
|
|
assert [e["type"] for e in edges] == ["heavy"]
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_apply_chunk_writes_nodes_links_and_edges(self, store, source_id):
|
|
"""One transactional write: entities linked to the chunk, edges added,
|
|
and a bare relationship endpoint upserted but not chunk-linked."""
|
|
try:
|
|
entities = [
|
|
{"name": "Ada", "normalized_name": "ada", "type": "person",
|
|
"description": "mathematician"},
|
|
{"name": "Engine", "normalized_name": "engine", "type": "machine",
|
|
"description": None},
|
|
]
|
|
relationships = [
|
|
{"source": "Ada", "target": "Engine", "type": "designed",
|
|
"description": "Ada designed the Engine", "weight": 2.0},
|
|
# 'Babbage' is only an endpoint — upserted edge-only.
|
|
{"source": "Babbage", "target": "Engine", "type": "built",
|
|
"description": None, "weight": 1.0},
|
|
]
|
|
name_embeddings = {
|
|
"ada": [0.1] * store._embedding_dim(),
|
|
"engine": [0.2] * store._embedding_dim(),
|
|
"babbage": [0.3] * store._embedding_dim(),
|
|
}
|
|
|
|
nodes, edges = store.apply_chunk(
|
|
source_id, "c1", entities, relationships, name_embeddings
|
|
)
|
|
assert nodes == 2 # only entities are counted
|
|
assert edges == 2
|
|
|
|
ada = store.get_node_by_normalized(source_id, "ada")
|
|
engine = store.get_node_by_normalized(source_id, "engine")
|
|
babbage = store.get_node_by_normalized(source_id, "babbage")
|
|
assert ada is not None and engine is not None
|
|
assert babbage is not None # endpoint upserted
|
|
|
|
mapping = store.get_chunk_ids_for_nodes(
|
|
source_id, [ada["id"], engine["id"], babbage["id"]]
|
|
)
|
|
assert mapping[ada["id"]] == ["c1"]
|
|
assert mapping[engine["id"]] == ["c1"]
|
|
# Bare endpoint is not linked to the chunk.
|
|
assert babbage["id"] not in mapping
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_self_loop_degree_agrees_across_paths(self, store, source_id):
|
|
"""``add_edge``'s incremental bump and ``set_node_degrees`` recompute must
|
|
agree on a self-loop.
|
|
|
|
They now agree on zero rather than one: the self-loop is rejected at
|
|
write time, so neither path has an edge to count. The property under
|
|
test is that the two paths agree, not the number they agree on.
|
|
"""
|
|
try:
|
|
node = store.upsert_node(source_id, "Solo", "solo")
|
|
assert store.add_edge(source_id, node, node, "self") is None
|
|
|
|
incremental = store.get_node_by_normalized(source_id, "solo")["degree"]
|
|
assert incremental == 0
|
|
|
|
store.set_node_degrees(source_id)
|
|
recomputed = store.get_node_by_normalized(source_id, "solo")["degree"]
|
|
assert recomputed == incremental
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_search_nodes_by_embedding(self, store, source_id):
|
|
try:
|
|
near = store.upsert_node(
|
|
source_id, "Near", "near", "thing", "d", _embedding(1.0)
|
|
)
|
|
store.upsert_node(
|
|
source_id, "Far", "far", "thing", "d", _embedding(-1.0)
|
|
)
|
|
results = store.search_nodes_by_embedding(source_id, _embedding(1.0), k=2)
|
|
assert len(results) == 2
|
|
assert results[0]["id"] == near
|
|
assert results[0]["distance"] <= results[1]["distance"]
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_get_subgraph_bounded(self, store, source_id):
|
|
try:
|
|
a = store.upsert_node(source_id, "A", "a")
|
|
b = store.upsert_node(source_id, "B", "b")
|
|
c = store.upsert_node(source_id, "C", "c")
|
|
store.add_edge(source_id, a, b, "rel")
|
|
store.add_edge(source_id, b, c, "rel")
|
|
|
|
one_hop = store.get_subgraph(source_id, [a], hops=1)
|
|
node_ids = {n["id"] for n in one_hop["nodes"]}
|
|
assert a in node_ids and b in node_ids
|
|
assert c not in node_ids
|
|
|
|
two_hop = store.get_subgraph(source_id, [a], hops=2)
|
|
node_ids = {n["id"] for n in two_hop["nodes"]}
|
|
assert {a, b, c} <= node_ids
|
|
assert len(two_hop["edges"]) >= 2
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_get_subgraph_frontier_truncation_is_deterministic(
|
|
self, store, source_id, monkeypatch
|
|
):
|
|
"""Bounded expansion must pick the same neighbors run-to-run so PPR (G5)
|
|
is reproducible."""
|
|
try:
|
|
hub = store.upsert_node(source_id, "Hub", "hub")
|
|
leaves = []
|
|
for i in range(6):
|
|
leaf = store.upsert_node(source_id, f"L{i}", f"l{i}")
|
|
store.add_edge(source_id, hub, leaf, "rel")
|
|
leaves.append(leaf)
|
|
|
|
monkeypatch.setattr(store_module, "MAX_SUBGRAPH_NODES", 4)
|
|
|
|
first = {n["id"] for n in store.get_subgraph(source_id, [hub])["nodes"]}
|
|
second = {n["id"] for n in store.get_subgraph(source_id, [hub])["nodes"]}
|
|
assert first == second
|
|
assert len(first) == 4
|
|
|
|
kept_leaves = sorted(leaves)[:3]
|
|
assert first == {hub, *kept_leaves}
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_get_graph_overview_bounded_by_degree(self, store, source_id):
|
|
try:
|
|
hub = store.upsert_node(source_id, "Hub", "hub")
|
|
leaves = [
|
|
store.upsert_node(source_id, f"L{i}", f"l{i}") for i in range(4)
|
|
]
|
|
for leaf in leaves:
|
|
store.add_edge(source_id, hub, leaf, "rel")
|
|
store.set_node_degrees(source_id)
|
|
|
|
overview = store.get_graph_overview(source_id, limit=3)
|
|
node_ids = [n["id"] for n in overview["nodes"]]
|
|
assert len(node_ids) == 3
|
|
# The hub has the highest degree, so it must lead the bounded set.
|
|
assert node_ids[0] == hub
|
|
# Edges only connect nodes that survived the limit.
|
|
for edge in overview["edges"]:
|
|
assert edge["source"] in node_ids
|
|
assert edge["target"] in node_ids
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_get_graph_overview_empty_source(self, store, source_id):
|
|
overview = store.get_graph_overview(source_id)
|
|
assert overview == {"nodes": [], "edges": []}
|
|
|
|
def test_get_node_detail_with_linked_chunks(self, store, source_id):
|
|
try:
|
|
node = store.upsert_node(
|
|
source_id, "Ada", "ada", "person", "A mathematician."
|
|
)
|
|
store.link_node_chunk(source_id, node, "chunk-1")
|
|
|
|
detail = store.get_node_detail(source_id, node)
|
|
assert detail is not None
|
|
assert detail["name"] == "Ada"
|
|
assert detail["description"] == "A mathematician."
|
|
chunk_ids = [c["chunk_id"] for c in detail["chunks"]]
|
|
assert "chunk-1" in chunk_ids
|
|
|
|
assert store.get_node_detail(source_id, str(uuid.uuid4())) is None
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_checkpoint_pending_and_mark(self, store, source_id):
|
|
try:
|
|
all_chunks = ["c1", "c2", "c3"]
|
|
assert store.pending_chunks(source_id, all_chunks) == all_chunks
|
|
|
|
store.mark_chunk(source_id, "c1", "done")
|
|
store.mark_chunk(source_id, "c2", "pending")
|
|
assert store.pending_chunks(source_id, all_chunks) == ["c2", "c3"]
|
|
|
|
store.mark_chunk(source_id, "c2", "done")
|
|
assert store.pending_chunks(source_id, all_chunks) == ["c3"]
|
|
|
|
progress = store.get_progress(source_id)
|
|
assert progress["c1"] == "done"
|
|
assert progress["c2"] == "done"
|
|
finally:
|
|
store.delete_by_source(source_id)
|
|
|
|
def test_count_nodes_many_batches_and_zero_fills(self, store):
|
|
"""One query for N sources; a source with no graph still gets an entry."""
|
|
a, b, c = (str(uuid.uuid4()) for _ in range(3))
|
|
try:
|
|
store.upsert_node(a, "A1", "a1")
|
|
store.upsert_node(a, "A2", "a2")
|
|
store.upsert_node(b, "B1", "b1")
|
|
|
|
counts = store.count_nodes_many([a, b, c])
|
|
|
|
assert counts == {a: 2, b: 1, c: 0}
|
|
# Agrees with the per-source query it replaces.
|
|
assert [store.count_nodes(s) for s in (a, b, c)] == [2, 1, 0]
|
|
assert store.count_nodes_many([]) == {}
|
|
finally:
|
|
store.delete_by_source(a)
|
|
store.delete_by_source(b)
|
|
|
|
def test_pooled_connection_is_returned_to_the_shared_pool(self, store):
|
|
"""The live store borrows from the shared pool and gives the socket back."""
|
|
source_id = str(uuid.uuid4())
|
|
assert store.count_nodes_many([source_id]) == {source_id: 0}
|
|
assert store._pooled is True
|
|
assert list(pgconn._POOLS) == [store._connection_string]
|
|
|
|
pool = pgconn._POOLS[store._connection_string]
|
|
store.close()
|
|
|
|
assert store._connection is None
|
|
stats = pool.get_stats()
|
|
assert stats["pool_available"] == stats["pool_size"]
|
|
|
|
def test_the_vector_store_reuses_the_graph_store_pool(self, store, postgresql):
|
|
"""Same DSN, one pool: the graph store does not double the connections."""
|
|
from docsgpt.vectorstore.pgvector import PGVectorStore
|
|
|
|
stub = MagicMock()
|
|
stub.dimension = TEST_EMBEDDING_DIM
|
|
stub.embed_query.return_value = [0.0] * TEST_EMBEDDING_DIM
|
|
with patch(
|
|
"docsgpt.vectorstore.base.BaseVectorStore._get_embeddings",
|
|
return_value=stub,
|
|
):
|
|
vector_store = PGVectorStore(
|
|
source_id="live-source", connection_string=store._connection_string
|
|
)
|
|
try:
|
|
store.count_nodes_many([str(uuid.uuid4())])
|
|
vector_store._get_connection()
|
|
|
|
assert list(pgconn._POOLS) == [store._connection_string]
|
|
assert vector_store._pooled is True
|
|
finally:
|
|
vector_store.close()
|
|
|
|
def test_delete_by_source_isolation(self, store):
|
|
keep = str(uuid.uuid4())
|
|
drop = str(uuid.uuid4())
|
|
try:
|
|
k = store.upsert_node(keep, "K", "k")
|
|
d = store.upsert_node(drop, "D", "d")
|
|
store.add_edge(keep, k, k, "self")
|
|
store.add_edge(drop, d, d, "self")
|
|
store.link_node_chunk(keep, k, "kc")
|
|
store.link_node_chunk(drop, d, "dc")
|
|
store.mark_chunk(keep, "kc", "done")
|
|
store.mark_chunk(drop, "dc", "done")
|
|
|
|
store.delete_by_source(drop)
|
|
|
|
assert store.count_nodes(drop) == 0
|
|
assert store.get_progress(drop) == {}
|
|
assert store.count_nodes(keep) == 1
|
|
assert store.get_progress(keep) == {"kc": "done"}
|
|
finally:
|
|
store.delete_by_source(keep)
|
|
store.delete_by_source(drop)
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestGraphStoreParameterization:
|
|
"""Asserts SQL is parameterized without touching a real DB."""
|
|
|
|
def _store_with_mock_conn(self):
|
|
store = GraphStore.__new__(GraphStore)
|
|
cursor = MagicMock()
|
|
cursor.fetchone.return_value = [str(uuid.uuid4())]
|
|
cursor.fetchall.return_value = []
|
|
conn = MagicMock()
|
|
conn.cursor.return_value = cursor
|
|
store._connection = conn
|
|
store._get_connection = lambda: conn
|
|
# Boot owns the schema; the write-path safety net has its own tests.
|
|
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
|
|
|
|
store, cursor = self._store_with_mock_conn()
|
|
sid = str(uuid.uuid4())
|
|
store.delete_by_source(sid)
|
|
|
|
tables = []
|
|
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)
|
|
sql = query.as_string()
|
|
assert "WHERE source_id = %s" in sql
|
|
assert sid not in sql
|
|
assert params == (sid,)
|
|
tables.append(sql.split('"')[1])
|
|
assert tables == ["graph_node_chunks", "graph_edges", "graph_nodes", "graph_ingest_progress"]
|
|
|
|
def test_search_binds_embedding_and_source(self):
|
|
store, cursor = self._store_with_mock_conn()
|
|
sid = str(uuid.uuid4())
|
|
embedding = _embedding(0.5)
|
|
store.search_nodes_by_embedding(sid, embedding, k=5)
|
|
|
|
sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
|
|
assert "%s::vector" in sql
|
|
assert "source_id = %s" in sql
|
|
assert sid not in sql
|
|
assert str(embedding) not in sql
|
|
assert params == (embedding, sid, embedding, 5)
|
|
|
|
def test_graph_overview_binds_source_and_clamps_limit(self):
|
|
from docsgpt.graphrag.store import GRAPH_OVERVIEW_MAX_LIMIT
|
|
|
|
store, cursor = self._store_with_mock_conn()
|
|
cursor.fetchall.return_value = []
|
|
sid = str(uuid.uuid4())
|
|
|
|
store.get_graph_overview(sid, limit=10_000)
|
|
|
|
sql, params = (
|
|
cursor.execute.call_args.args[0],
|
|
cursor.execute.call_args.args[1],
|
|
)
|
|
assert "source_id = %s" in sql
|
|
assert sid not in sql
|
|
# An empty node fetch short-circuits; only the node query ran, and the
|
|
# limit is clamped to the hard cap before binding.
|
|
assert params == (sid, GRAPH_OVERVIEW_MAX_LIMIT)
|
|
|
|
def test_upsert_node_binds_all_values(self):
|
|
store, cursor = self._store_with_mock_conn()
|
|
sid = str(uuid.uuid4())
|
|
embedding = _embedding(0.1)
|
|
store.upsert_node(sid, "Name", "name", "type", "desc", embedding)
|
|
|
|
sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
|
|
assert "ON CONFLICT (source_id, normalized_name) DO UPDATE" in sql
|
|
assert sid not in sql
|
|
assert "name" not in [t for t in sql.split() if t == sid]
|
|
assert params[1] == sid
|
|
assert params[-1] == embedding
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestGraphReadQueries:
|
|
"""The reads behind fact seeding and the agent's graph tool, without a DB.
|
|
|
|
The live class covers what these return from real rows; these pin the
|
|
contract that holds without one. The entity name reaching
|
|
``entity_relationships``/``entity_pages`` comes from an LLM tool call, so
|
|
it must only ever travel as a bound parameter.
|
|
"""
|
|
|
|
def _store(self, rows=(), fail=False):
|
|
store = GraphStore.__new__(GraphStore)
|
|
cursor = MagicMock()
|
|
cursor.fetchall.return_value = list(rows)
|
|
if fail:
|
|
cursor.execute.side_effect = RuntimeError("relation does not exist")
|
|
conn = MagicMock()
|
|
conn.cursor.return_value = cursor
|
|
store._get_connection = lambda: conn
|
|
return store, cursor, conn
|
|
|
|
def test_identifiers_are_quoted_as_postgres_folds_them_unquoted(self):
|
|
# PGVectorStore writes these names unquoted, which Postgres folds to
|
|
# lower case; quoting keeps case, so the fold happens first or the two
|
|
# stores would address different tables.
|
|
assert store_module._identifier("Documents").as_string() == '"documents"'
|
|
with pytest.raises(ValueError):
|
|
store_module._identifier('documents"; DROP TABLE graph_nodes; --')
|
|
|
|
def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self):
|
|
store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)])
|
|
sid = str(uuid.uuid4())
|
|
embedding = _embedding(0.3)
|
|
|
|
rows = store.seed_nodes_from_facts(sid, embedding, fact_limit=0, limit=3)
|
|
|
|
sql, params = cursor.execute.call_args.args
|
|
assert sid not in sql and str(embedding) not in sql
|
|
# Limits are clamped to at least one before binding.
|
|
assert params == (embedding, sid, embedding, 1, sid, 3)
|
|
assert rows[0] == {"id": "n1", "name": "Quill", "description": "a store", "distance": pytest.approx(0.2)}
|
|
assert rows[1]["distance"] == 1.0
|
|
|
|
def test_fact_seeds_need_a_query_vector(self):
|
|
store, cursor, _ = self._store()
|
|
assert store.seed_nodes_from_facts(str(uuid.uuid4()), []) == []
|
|
cursor.execute.assert_not_called()
|
|
|
|
def test_relationships_bind_the_name_as_a_pattern(self):
|
|
store, cursor, _ = self._store(rows=[("Alder", "streams_to", "Quill", "audit events")])
|
|
sid = str(uuid.uuid4())
|
|
name = "Quill'; DROP TABLE graph_nodes; --"
|
|
|
|
rows = store.entity_relationships(sid, f" {name} ", limit=500)
|
|
|
|
sql, params = cursor.execute.call_args.args
|
|
assert name not in sql
|
|
assert params == (sid, f"%{name}%", f"%{name}%", 500)
|
|
assert rows == [
|
|
{"source": "Alder", "type": "streams_to", "target": "Quill", "description": "audit events"}
|
|
]
|
|
|
|
def test_pages_prefer_the_entity_itself_over_a_mention(self):
|
|
store, cursor, _ = self._store(rows=[({"title": "quill.md"}, "Quill is a store."), (None, None)])
|
|
sid = str(uuid.uuid4())
|
|
|
|
pages = store.entity_pages(sid, "Quill", limit=0)
|
|
|
|
query, params = cursor.execute.call_args.args
|
|
sql = query.as_string()
|
|
assert 'JOIN "documents" d' in sql and 'd."source_id" = %s' in sql
|
|
assert "Quill" not in sql
|
|
# Exact name, name plus a qualifier ("Quill Store"), substring fallback,
|
|
# text-opens-with ordering, then the clamped limit.
|
|
assert params == ("quill", "quill %", sid, sid, "quill", "quill %", "%Quill%", "Quill%", 1)
|
|
assert pages == [{"metadata": {"title": "quill.md"}, "text": "Quill is a store."}, {"metadata": {}, "text": ""}]
|
|
|
|
def test_chunk_similarities_are_restricted_to_the_reached_chunks(self):
|
|
store, cursor, _ = self._store(rows=[("11", 0.75)])
|
|
sid = str(uuid.uuid4())
|
|
embedding = _embedding(0.9)
|
|
|
|
scores = store.chunk_similarities(sid, [11, "12"], embedding)
|
|
|
|
query, params = cursor.execute.call_args.args
|
|
sql = query.as_string()
|
|
assert '1 - ("embedding" <=> %s::vector)' in sql and 'FROM "documents"' in sql
|
|
assert "= ANY(%s)" in sql and sid not in sql
|
|
assert params == (embedding, sid, ["11", "12"])
|
|
assert scores == {"11": 0.75}
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda s: s.entity_relationships("sid", " "),
|
|
lambda s: s.entity_pages("sid", ""),
|
|
lambda s: s.chunk_similarities("sid", [], [0.1]),
|
|
lambda s: s.chunk_similarities("sid", ["1"], []),
|
|
],
|
|
)
|
|
def test_empty_input_runs_no_query(self, call):
|
|
store, cursor, _ = self._store()
|
|
assert not call(store)
|
|
cursor.execute.assert_not_called()
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda s: s.seed_nodes_from_facts("sid", [0.1]),
|
|
lambda s: s.entity_relationships("sid", "Quill"),
|
|
lambda s: s.entity_pages("sid", "Quill"),
|
|
lambda s: s.chunk_similarities("sid", ["1"], [0.1]),
|
|
],
|
|
)
|
|
def test_a_failed_query_returns_nothing_and_releases_the_connection(self, call):
|
|
store, cursor, conn = self._store(fail=True)
|
|
assert not call(store)
|
|
cursor.close.assert_called_once()
|
|
conn.rollback.assert_called_once()
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestEmbeddingDim:
|
|
"""The graph table dimension is derived from the configured model (FIX 1)."""
|
|
|
|
def test_uses_configured_model_dimension(self, monkeypatch):
|
|
from docsgpt.vectorstore import base as base_module
|
|
|
|
monkeypatch.setattr(base_module.settings, "EMBEDDINGS_BASE_URL", None)
|
|
fake_embedding = MagicMock()
|
|
fake_embedding.dimension = 1536
|
|
monkeypatch.setattr(
|
|
base_module.EmbeddingsSingleton,
|
|
"get_instance",
|
|
staticmethod(lambda *a, **k: fake_embedding),
|
|
)
|
|
monkeypatch.setattr(GraphStore, "_embedding_dim", _REAL_EMBEDDING_DIM)
|
|
|
|
store = GraphStore.__new__(GraphStore)
|
|
assert store._embedding_dim() == 1536
|
|
|
|
def test_none_dimension_falls_back_to_default(self, monkeypatch):
|
|
"""A remote model outside the registry reports ``None``, not nothing.
|
|
|
|
``getattr`` with a default cannot catch that -- the attribute exists --
|
|
so the width reached the DDL as ``vector(None)``.
|
|
"""
|
|
from docsgpt.vectorstore import base as base_module
|
|
|
|
monkeypatch.setattr(base_module.settings, "EMBEDDINGS_BASE_URL", None)
|
|
fake_embedding = MagicMock()
|
|
fake_embedding.dimension = None
|
|
monkeypatch.setattr(
|
|
base_module.EmbeddingsSingleton,
|
|
"get_instance",
|
|
staticmethod(lambda *a, **k: fake_embedding),
|
|
)
|
|
monkeypatch.setattr(GraphStore, "_embedding_dim", _REAL_EMBEDDING_DIM)
|
|
|
|
store = GraphStore.__new__(GraphStore)
|
|
assert store._embedding_dim() == store_module.DEFAULT_NAME_EMBEDDING_DIM
|
|
|
|
def test_falls_back_to_default_dimension(self, monkeypatch):
|
|
from docsgpt.vectorstore import base as base_module
|
|
|
|
monkeypatch.setattr(base_module.settings, "EMBEDDINGS_BASE_URL", None)
|
|
fake_embedding = object()
|
|
monkeypatch.setattr(
|
|
base_module.EmbeddingsSingleton,
|
|
"get_instance",
|
|
staticmethod(lambda *a, **k: fake_embedding),
|
|
)
|
|
monkeypatch.setattr(GraphStore, "_embedding_dim", _REAL_EMBEDDING_DIM)
|
|
|
|
store = GraphStore.__new__(GraphStore)
|
|
assert store._embedding_dim() == store_module.DEFAULT_NAME_EMBEDDING_DIM
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestEmbeddingDimResolution:
|
|
def test_uses_shared_resolver(self, monkeypatch):
|
|
"""The dimension probe must not build its own embeddings instance."""
|
|
from unittest.mock import patch
|
|
|
|
fake_embedding = MagicMock()
|
|
fake_embedding.dimension = 1536
|
|
monkeypatch.setattr(GraphStore, "_embedding_dim", _REAL_EMBEDDING_DIM)
|
|
|
|
with patch(
|
|
"docsgpt.vectorstore.base.get_embeddings",
|
|
return_value=fake_embedding,
|
|
) as mock_resolver:
|
|
store = GraphStore.__new__(GraphStore)
|
|
assert store._embedding_dim() == 1536
|
|
|
|
mock_resolver.assert_called_once_with()
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestGraphSchemaIsBootOwned:
|
|
"""Construction must not run DDL: reads happen once per query, per source."""
|
|
|
|
def _mock_store(self):
|
|
store = GraphStore.__new__(GraphStore)
|
|
cursor = MagicMock()
|
|
cursor.fetchone.return_value = [str(uuid.uuid4())]
|
|
cursor.fetchall.return_value = []
|
|
conn = MagicMock()
|
|
conn.cursor.return_value = cursor
|
|
store._connection = conn
|
|
store._get_connection = lambda: conn
|
|
store._tables_ensured = False
|
|
store._ensure_tables = MagicMock()
|
|
return store, cursor
|
|
|
|
def test_init_opens_no_connection_and_creates_no_tables(self):
|
|
from unittest.mock import patch
|
|
|
|
with patch.dict(
|
|
"sys.modules",
|
|
{
|
|
"psycopg": MagicMock(),
|
|
"pgvector": MagicMock(),
|
|
"pgvector.psycopg": MagicMock(),
|
|
},
|
|
), patch.object(GraphStore, "_ensure_tables") as ensure, patch.object(
|
|
GraphStore, "_get_connection"
|
|
) as get_conn:
|
|
store = GraphStore(connection_string="postgresql://u:p@localhost/db")
|
|
|
|
ensure.assert_not_called()
|
|
get_conn.assert_not_called()
|
|
assert store._tables_ensured is False
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda s: s.upsert_node("sid", "N", "n"),
|
|
lambda s: s.add_edge("sid", "a", "b"),
|
|
lambda s: s.link_node_chunk("sid", "n", "c1"),
|
|
lambda s: s.apply_chunk("sid", "c1", [], [], {}),
|
|
lambda s: s.set_node_degrees("sid"),
|
|
lambda s: s.mark_chunk("sid", "c1", "done"),
|
|
lambda s: s.delete_by_source("sid"),
|
|
],
|
|
ids=[
|
|
"upsert_node",
|
|
"add_edge",
|
|
"link_node_chunk",
|
|
"apply_chunk",
|
|
"set_node_degrees",
|
|
"mark_chunk",
|
|
"delete_by_source",
|
|
],
|
|
)
|
|
def test_writes_ensure_tables_once(self, call):
|
|
store, _ = self._mock_store()
|
|
|
|
call(store)
|
|
store._tables_ensured = True # what the real _ensure_tables_once sets
|
|
call(store)
|
|
|
|
assert store._ensure_tables.call_count == 1
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda s: s.count_nodes("sid"),
|
|
lambda s: s.count_nodes_many(["sid"]),
|
|
lambda s: s.get_node_by_normalized("sid", "n"),
|
|
lambda s: s.search_nodes_by_embedding("sid", _embedding(1.0)),
|
|
lambda s: s.get_subgraph("sid", ["n"]),
|
|
lambda s: s.get_graph_overview("sid"),
|
|
lambda s: s.get_chunk_ids_for_nodes("sid", ["n"]),
|
|
lambda s: s.pending_chunks("sid", ["c1"]),
|
|
lambda s: s.get_progress("sid"),
|
|
],
|
|
ids=[
|
|
"count_nodes",
|
|
"count_nodes_many",
|
|
"get_node_by_normalized",
|
|
"search_nodes_by_embedding",
|
|
"get_subgraph",
|
|
"get_graph_overview",
|
|
"get_chunk_ids_for_nodes",
|
|
"pending_chunks",
|
|
"get_progress",
|
|
],
|
|
)
|
|
def test_reads_never_create_tables(self, call):
|
|
store, _ = self._mock_store()
|
|
|
|
call(store)
|
|
|
|
store._ensure_tables.assert_not_called()
|
|
|
|
def test_create_schema_emits_the_ddl_without_committing(self):
|
|
conn, cursor = MagicMock(), MagicMock()
|
|
conn.cursor.return_value = cursor
|
|
|
|
GraphStore.create_schema(conn, dimension=8)
|
|
|
|
statements = " ".join(str(c) for c in cursor.execute.call_args_list)
|
|
assert "CREATE EXTENSION IF NOT EXISTS vector" in statements
|
|
for table in (
|
|
"graph_nodes",
|
|
"graph_edges",
|
|
"graph_node_chunks",
|
|
"graph_ingest_progress",
|
|
):
|
|
assert f"CREATE TABLE IF NOT EXISTS {table}" in statements
|
|
assert "name_embedding vector(8)" in statements
|
|
assert statements.count("CREATE INDEX IF NOT EXISTS") == 5
|
|
conn.commit.assert_not_called()
|
|
|
|
def test_ensure_tables_locks_then_commits(self):
|
|
store, cursor = self._mock_store()
|
|
del store._ensure_tables # exercise the real method
|
|
|
|
store._ensure_tables()
|
|
|
|
statements = " ".join(str(c) for c in cursor.execute.call_args_list)
|
|
assert "pg_advisory_xact_lock" in statements
|
|
assert "CREATE TABLE IF NOT EXISTS graph_nodes" in statements
|
|
store._connection.commit.assert_called_once()
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestGraphStorePooling:
|
|
"""The graph store borrows from the same per-DSN pool as ``PGVectorStore``."""
|
|
|
|
def _store(self, dsn=POOL_DSN, pool_max_size=4):
|
|
store = GraphStore.__new__(GraphStore)
|
|
store._connection_string = dsn
|
|
store._connection = None
|
|
store._pooled = False
|
|
store._pool_max_size = pool_max_size
|
|
store._psycopg = MagicMock()
|
|
store._register_vector = MagicMock()
|
|
store._tables_ensured = True
|
|
return store
|
|
|
|
def _fake_pool(self):
|
|
pooled_conn = MagicMock()
|
|
pooled_conn.closed = False
|
|
pool = MagicMock()
|
|
pool.getconn.return_value = pooled_conn
|
|
return pool, pooled_conn
|
|
|
|
def test_get_connection_checks_out_of_the_pool(self, monkeypatch):
|
|
pool, pooled_conn = self._fake_pool()
|
|
monkeypatch.setattr(store_module.pgconn, "pool_for", lambda dsn, n: pool)
|
|
store = self._store()
|
|
|
|
conn = store._get_connection()
|
|
|
|
assert conn is pooled_conn
|
|
assert store._pooled is True
|
|
pool.getconn.assert_called_once()
|
|
store._psycopg.connect.assert_not_called()
|
|
# The pool's ``configure`` hook already registered the adapters.
|
|
store._register_vector.assert_not_called()
|
|
|
|
def test_close_rolls_back_and_returns_the_connection(self, monkeypatch):
|
|
pool, pooled_conn = self._fake_pool()
|
|
monkeypatch.setattr(store_module.pgconn, "pool_for", lambda dsn, n: pool)
|
|
monkeypatch.setitem(pgconn._POOLS, POOL_DSN, pool)
|
|
store = self._store()
|
|
store._get_connection()
|
|
pooled_conn.info.transaction_status.name = "INTRANS"
|
|
|
|
store.close()
|
|
|
|
pooled_conn.rollback.assert_called_once()
|
|
pool.putconn.assert_called_once_with(pooled_conn)
|
|
pooled_conn.close.assert_not_called()
|
|
assert store._connection is None
|
|
|
|
def test_close_does_not_roll_back_an_idle_connection(self, monkeypatch):
|
|
pool, pooled_conn = self._fake_pool()
|
|
monkeypatch.setattr(store_module.pgconn, "pool_for", lambda dsn, n: pool)
|
|
monkeypatch.setitem(pgconn._POOLS, POOL_DSN, pool)
|
|
store = self._store()
|
|
store._get_connection()
|
|
pooled_conn.info.transaction_status.name = "IDLE"
|
|
|
|
store.close()
|
|
|
|
pooled_conn.rollback.assert_not_called()
|
|
pool.putconn.assert_called_once_with(pooled_conn)
|
|
|
|
def test_a_dead_pooled_connection_is_returned_before_being_replaced(
|
|
self, monkeypatch
|
|
):
|
|
# Same contract as ``PGVectorStore``: a connection that dies while this
|
|
# store holds it must go back to the pool, or the slot is lost for the
|
|
# life of the process. Extraction holds one store across the whole
|
|
# per-chunk LLM loop, which is exactly when a backend gets reaped.
|
|
pool, pooled_conn = self._fake_pool()
|
|
monkeypatch.setattr(store_module.pgconn, "pool_for", lambda dsn, n: pool)
|
|
monkeypatch.setitem(pgconn._POOLS, POOL_DSN, pool)
|
|
store = self._store()
|
|
store._get_connection()
|
|
replacement = MagicMock()
|
|
replacement.closed = False
|
|
pool.getconn.return_value = replacement
|
|
|
|
pooled_conn.closed = True
|
|
conn = store._get_connection()
|
|
|
|
assert conn is replacement
|
|
pool.putconn.assert_called_once_with(pooled_conn)
|
|
assert pool.getconn.call_count == 2
|
|
|
|
def test_legacy_path_connects_directly_and_closes(self, monkeypatch):
|
|
def _never(dsn, n):
|
|
raise AssertionError("pooling is off; no pool must be built")
|
|
|
|
monkeypatch.setattr(store_module.pgconn, "pool_for", _never)
|
|
store = self._store(pool_max_size=0)
|
|
direct = MagicMock()
|
|
direct.closed = False
|
|
store._psycopg.connect.return_value = direct
|
|
|
|
conn = store._get_connection()
|
|
|
|
assert conn is direct
|
|
assert store._pooled is False
|
|
store._register_vector.assert_called_once_with(direct)
|
|
|
|
store.close()
|
|
direct.close.assert_called_once()
|
|
|
|
def test_del_never_raises(self):
|
|
store = self._store()
|
|
broken = MagicMock()
|
|
broken.closed = False
|
|
broken.close.side_effect = RuntimeError("already gone")
|
|
store._connection = broken
|
|
|
|
store.__del__() # must not propagate
|
|
|
|
def test_the_graph_store_and_the_vector_store_share_one_pool(self):
|
|
"""One DSN, one pool object — reached from either module."""
|
|
pool, _ = self._fake_pool()
|
|
store = self._store()
|
|
|
|
with patch("psycopg_pool.ConnectionPool", return_value=pool) as pool_cls:
|
|
store._get_connection()
|
|
# ``PGVectorStore``'s own entry point resolves to the same object.
|
|
assert pgvector_module._pool_for(POOL_DSN, 4) is pool
|
|
|
|
assert pool_cls.call_count == 1
|
|
assert pgconn._POOLS[POOL_DSN] is pool
|
|
assert pgvector_module._POOLS is pgconn._POOLS
|
|
|
|
@pytest.mark.parametrize(
|
|
"value,expected",
|
|
[(0, 0), (2, 2), (None, 8), ("4", 8), (True, 8), (-1, 8)],
|
|
)
|
|
def test_pool_size_is_resolved_defensively(self, monkeypatch, value, expected):
|
|
monkeypatch.setattr(
|
|
store_module.settings, "PGVECTOR_POOL_MAX_SIZE", value, raising=False
|
|
)
|
|
assert store_module._resolve_pool_max_size() == expected
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestCountNodesMany:
|
|
"""One ``ANY(%s)`` query replaces the retriever's per-source count fan-out."""
|
|
|
|
def _store_with_mock_conn(self, rows):
|
|
store = GraphStore.__new__(GraphStore)
|
|
cursor = MagicMock()
|
|
cursor.fetchall.return_value = rows
|
|
conn = MagicMock()
|
|
conn.cursor.return_value = cursor
|
|
store._connection = conn
|
|
store._get_connection = lambda: conn
|
|
store._tables_ensured = True
|
|
return store, cursor, conn
|
|
|
|
def test_binds_the_ids_as_one_array_and_zero_fills(self):
|
|
store, cursor, _ = self._store_with_mock_conn([("a", 2), ("b", 1)])
|
|
|
|
counts = store.count_nodes_many(["a", "b", "c"])
|
|
|
|
assert counts == {"a": 2, "b": 1, "c": 0}
|
|
sql, params = (
|
|
cursor.execute.call_args.args[0],
|
|
cursor.execute.call_args.args[1],
|
|
)
|
|
assert "source_id = ANY(%s)" in sql
|
|
assert "GROUP BY source_id" in sql
|
|
assert cursor.execute.call_count == 1
|
|
assert params == (["a", "b", "c"],)
|
|
|
|
def test_empty_input_short_circuits(self):
|
|
store, cursor, _ = self._store_with_mock_conn([])
|
|
|
|
assert store.count_nodes_many([]) == {}
|
|
assert store.count_nodes_many([None, ""]) == {}
|
|
cursor.execute.assert_not_called()
|
|
|
|
def test_a_failed_query_reports_every_source_as_graphless(self):
|
|
store, cursor, conn = self._store_with_mock_conn([])
|
|
cursor.execute.side_effect = RuntimeError("no such table")
|
|
|
|
assert store.count_nodes_many(["a", "b"]) == {"a": 0, "b": 0}
|
|
conn.rollback.assert_called_once()
|
|
|
|
def test_the_callers_id_spelling_is_preserved(self):
|
|
"""Postgres returns canonical lowercase UUID text; keys must still match."""
|
|
source_id = str(uuid.uuid4()).upper()
|
|
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()
|
|
|
|
|
|
@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_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())
|
|
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)
|