Files
DocsGPT/tests/graphrag/test_store.py
T
2026-09-28 17:35:16 +04:00

2163 lines
84 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
from psycopg.types.json import Jsonb
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 _seed_typed_graph(self, store, source_id):
"""Six nodes whose types spell ``person`` three ways, plus edges for degree."""
ids = {
"ada": store.upsert_node(source_id, "Ada Lovelace", "ada lovelace", "Person"),
"alan": store.upsert_node(source_id, "Alan Turing", "alan turing", "Person"),
"grace": store.upsert_node(source_id, "Grace Hopper", "grace hopper", "PERSON"),
"babbage": store.upsert_node(source_id, "Babbage", "babbage", "per son"),
"acme": store.upsert_node(source_id, "Acme_Corp", "acme_corp", "org"),
"sure": store.upsert_node(source_id, "100%_Sure", "100%_sure"),
}
# Degrees: ada 3, alan 2, acme 2, grace 1, babbage 0, sure 0.
store.add_edge(source_id, ids["ada"], ids["alan"], "knows")
store.add_edge(source_id, ids["ada"], ids["acme"], "works_at")
store.add_edge(source_id, ids["grace"], ids["ada"], "cites")
store.add_edge(source_id, ids["alan"], ids["acme"], "works_at")
store.set_node_degrees(source_id)
return ids
def test_list_nodes_filters_pages_and_counts(self, store, source_id):
try:
ids = self._seed_typed_graph(store, source_id)
everything = store.list_nodes(source_id)
assert everything["total"] == 6
ordered = [n["id"] for n in everything["nodes"]]
assert ordered[0] == ids["ada"]
# Degree DESC, then id: alan and acme tie on 2.
assert set(ordered[1:3]) == {ids["alan"], ids["acme"]}
assert ordered[1:3] == sorted(ordered[1:3])
first = everything["nodes"][0]
assert set(first) == {"id", "name", "type", "degree", "doc_freq"}
assert first["degree"] == 3 and first["doc_freq"] == 1
by_name = store.list_nodes(source_id, query="ADA")
assert [n["id"] for n in by_name["nodes"]] == [ids["ada"]]
assert by_name["total"] == 1
# ``%`` and ``_`` are literals, not wildcards.
assert [n["id"] for n in store.list_nodes(source_id, query="%")["nodes"]] == [
ids["sure"]
]
underscored = store.list_nodes(source_id, query="_")
assert {n["id"] for n in underscored["nodes"]} == {ids["acme"], ids["sure"]}
people = store.list_nodes(source_id, type_key="person")
assert people["total"] == 4
assert {n["id"] for n in people["nodes"]} == {
ids["ada"], ids["alan"], ids["grace"], ids["babbage"]
}
untyped = store.list_nodes(source_id, type_key="")
assert [n["id"] for n in untyped["nodes"]] == [ids["sure"]]
both = store.list_nodes(source_id, query="a", type_key="person")
assert both["total"] == 4
page = store.list_nodes(source_id, type_key="person", offset=1, limit=2)
assert page["total"] == 4
assert [n["id"] for n in page["nodes"]] == [
n["id"] for n in people["nodes"][1:3]
]
beyond = store.list_nodes(source_id, offset=50, limit=10)
assert beyond == {"nodes": [], "total": 6}
finally:
store.delete_by_source(source_id)
def test_node_type_facets_fold_spellings_into_one_key(self, store, source_id):
try:
self._seed_typed_graph(store, source_id)
facets = store.node_type_facets(source_id)
assert facets == [
{"key": "person", "label": "Person", "count": 4},
{"key": "org", "label": "org", "count": 1},
{"key": "", "label": None, "count": 1},
]
assert store.node_type_facets(str(uuid.uuid4())) == []
finally:
store.delete_by_source(source_id)
def test_count_edges(self, store, source_id):
try:
assert store.count_edges(source_id) == 0
self._seed_typed_graph(store, source_id)
assert store.count_edges(source_id) == 4
assert store.count_edges(str(uuid.uuid4())) == 0
finally:
store.delete_by_source(source_id)
def test_node_detail_lists_relationships_in_both_directions(self, store, source_id):
try:
ids = self._seed_typed_graph(store, source_id)
# ``add_edge`` drops self-loops, so plant one directly: the reader
# must skip it too.
conn = store._get_connection()
with conn.cursor() as cursor:
cursor.execute(
"INSERT INTO graph_edges (id, source_id, src_node_id, dst_node_id, type) "
"VALUES (%s, %s, %s, %s, 'self');",
(str(uuid.uuid4()), source_id, ids["ada"], ids["ada"]),
)
conn.commit()
detail = store.get_node_detail(source_id, ids["ada"])
assert detail["name"] == "Ada Lovelace"
assert "chunks" in detail
rels = detail["relationships"]
assert len(rels) == 3
# Neighbour degree DESC, then neighbour id: alan and acme tie on 2.
assert [r["id"] for r in rels[:2]] == sorted([ids["alan"], ids["acme"]])
by_id = {r["id"]: r for r in rels}
assert by_id[ids["alan"]] == {
"id": ids["alan"], "name": "Alan Turing", "type": "Person",
"degree": 2, "edge_type": "knows", "direction": "out",
}
assert by_id[ids["acme"]]["edge_type"] == "works_at"
assert by_id[ids["acme"]]["direction"] == "out"
assert rels[2] == {
"id": ids["grace"], "name": "Grace Hopper", "type": "PERSON",
"degree": 1, "edge_type": "cites", "direction": "in",
}
assert detail["relationships_total"] == 3
lonely = store.get_node_detail(source_id, ids["babbage"])
assert lonely["relationships"] == []
assert lonely["relationships_total"] == 0
finally:
store.delete_by_source(source_id)
def test_node_detail_counts_every_relationship_past_the_cap(
self, store, source_id, monkeypatch
):
try:
ids = self._seed_typed_graph(store, source_id)
conn = store._get_connection()
with conn.cursor() as cursor:
cursor.execute(
"INSERT INTO graph_edges (id, source_id, src_node_id, dst_node_id, type) "
"VALUES (%s, %s, %s, %s, 'self');",
(str(uuid.uuid4()), source_id, ids["ada"], ids["ada"]),
)
conn.commit()
monkeypatch.setattr(store_module, "MAX_NODE_RELATIONSHIPS", 2)
detail = store.get_node_detail(source_id, ids["ada"])
# The list stops at the cap; the total counts both directions and
# still skips the self-loop.
assert len(detail["relationships"]) == 2
assert detail["relationships_total"] == 3
finally:
store.delete_by_source(source_id)
def test_remap_chunk_moves_every_link_to_the_new_id(self, store, source_id):
try:
ada = store.upsert_node(source_id, "Ada", "ada")
alan = store.upsert_node(source_id, "Alan", "alan")
store.link_node_chunk(source_id, ada, "old")
store.link_node_chunk(source_id, alan, "keep")
store.add_edge(
source_id, ada, alan, "knows", source_chunk_ids=["keep", "old"]
)
store.mark_chunk(source_id, "old", "done")
other = str(uuid.uuid4())
other_node = store.upsert_node(other, "Ada", "ada")
store.link_node_chunk(other, other_node, "old")
store.remap_chunk(source_id, "old", "new")
links = store.get_chunk_ids_for_nodes(source_id, [ada, alan])
assert links == {ada: ["new"], alan: ["keep"]}
conn = store._get_connection()
with conn.cursor() as cursor:
cursor.execute(
"SELECT source_chunk_ids FROM graph_edges WHERE source_id = %s;",
(source_id,),
)
assert cursor.fetchone()[0] == ["keep", "new"]
conn.rollback()
assert store.pending_chunks(source_id, ["old", "new"]) == ["old"]
# Another source's link to the same chunk id is left alone.
assert store.get_chunk_ids_for_nodes(other, [other_node]) == {
other_node: ["old"]
}
finally:
store.delete_by_source(source_id)
store.delete_by_source(other)
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_count_edges_binds_source(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchone.return_value = [7]
sid = str(uuid.uuid4())
assert store.count_edges(sid) == 7
sql, params = cursor.execute.call_args.args
assert "FROM graph_edges" in sql
assert "source_id = %s" in sql
assert sid not in sql
assert params == (sid,)
def test_count_edges_failure_reports_zero(self):
store, cursor = self._store_with_mock_conn()
cursor.execute.side_effect = RuntimeError("relation does not exist")
assert store.count_edges(str(uuid.uuid4())) == 0
def test_count_edges_has_no_strict_mode(self):
import inspect
assert "strict" not in inspect.signature(GraphStore.count_edges).parameters
def test_list_nodes_binds_values_and_escapes_the_pattern(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchone.return_value = [0]
cursor.fetchall.side_effect = [[("Person",), ("person",), ("Org",)], []]
sid = str(uuid.uuid4())
query = "50%_a\\b"
result = store.list_nodes(
sid, query=query, type_key="person", offset=40, limit=10_000
)
assert result == {"nodes": [], "total": 0}
pattern = "%50\\%\\_a\\\\b%"
types_sql, types_params = cursor.execute.call_args_list[0].args
assert "DISTINCT type" in types_sql
assert types_params == (sid,)
for call in cursor.execute.call_args_list[1:]:
sql, params = call.args
assert sid not in sql
assert query not in sql
assert "person" not in sql
assert "ILIKE %s" in sql
assert "type = ANY(%s)" in sql
assert params[:3] == (sid, pattern, ["Person", "person"])
count_sql, count_params = cursor.execute.call_args_list[1].args
assert "count(*)" in count_sql
page_sql, page_params = cursor.execute.call_args_list[-1].args
assert "ORDER BY degree DESC, id" in page_sql
# The page size is clamped to the hard cap before binding.
assert page_params[3:] == (store_module.GRAPH_NODE_LIST_MAX_LIMIT, 40)
def test_list_nodes_without_filters_binds_only_the_source(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchone.return_value = [0]
sid = str(uuid.uuid4())
store.list_nodes(sid, query=" ", offset=-5, limit=0)
count_sql, count_params = cursor.execute.call_args_list[0].args
assert "ILIKE" not in count_sql
assert "regexp_replace" not in count_sql
assert count_params == (sid,)
_, page_params = cursor.execute.call_args_list[-1].args
assert page_params == (sid, 1, 0)
def test_list_nodes_failure_propagates(self):
store, cursor = self._store_with_mock_conn()
cursor.execute.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError):
store.list_nodes("sid")
cursor.close.assert_called_once()
def test_list_nodes_empty_graph_returns_empty(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchone.return_value = [0]
cursor.fetchall.return_value = []
assert store.list_nodes("sid") == {"nodes": [], "total": 0}
def test_node_type_facets_bind_source(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchall.return_value = [("Person", 4), (None, 1)]
sid = str(uuid.uuid4())
facets = store.node_type_facets(sid)
sql, params = cursor.execute.call_args.args
assert sid not in sql
assert params == (sid,)
assert facets == [
{"key": "person", "label": "Person", "count": 4},
{"key": "", "label": None, "count": 1},
]
def test_node_type_facets_failure_propagates(self):
store, cursor = self._store_with_mock_conn()
cursor.execute.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError):
store.node_type_facets("sid")
cursor.close.assert_called_once()
def test_node_type_facets_empty_graph_returns_empty(self):
store, cursor = self._store_with_mock_conn()
cursor.fetchall.return_value = []
assert store.node_type_facets("sid") == []
def test_node_detail_relationships_bind_ids_and_cap(self):
store, cursor = self._store_with_mock_conn()
sid = str(uuid.uuid4())
nid = str(uuid.uuid4())
other = str(uuid.uuid4())
cursor.fetchone.return_value = (nid, "Ada", "person", "d", 1, 1)
cursor.fetchall.side_effect = [
[(other, "Alan", "person", 2, "knows", "out")], # relationships
[], # chunk ids
]
detail = store.get_node_detail(sid, nid)
assert detail["relationships"] == [
{"id": other, "name": "Alan", "type": "person", "degree": 2,
"edge_type": "knows", "direction": "out"},
]
rel_call = next(
c for c in cursor.execute.call_args_list if "direction" in c.args[0]
)
sql, params = rel_call.args
assert sid not in sql and nid not in sql
assert "src_node_id <> e.dst_node_id" in sql
assert sid in params and nid in params
assert params[-1] == store_module.MAX_NODE_RELATIONSHIPS == 300
# Under the cap the list is the whole set: no COUNT query.
assert detail["relationships_total"] == 1
assert not any(
"COUNT(*)" in c.args[0] for c in cursor.execute.call_args_list
)
def test_node_detail_counts_relationships_when_capped(self, monkeypatch):
store, cursor = self._store_with_mock_conn()
monkeypatch.setattr(store_module, "MAX_NODE_RELATIONSHIPS", 2)
sid = str(uuid.uuid4())
nid = str(uuid.uuid4())
rel = (str(uuid.uuid4()), "Alan", "person", 2, "knows", "out")
cursor.fetchone.side_effect = [
(nid, "Ada", "person", "d", 1, 1), # the node
(1204,), # relationships total
]
cursor.fetchall.side_effect = [[rel, rel], []]
detail = store.get_node_detail(sid, nid)
assert len(detail["relationships"]) == 2
assert detail["relationships_total"] == 1204
count_sql, count_params = next(
c.args for c in cursor.execute.call_args_list if "COUNT(*)" in c.args[0]
)
assert sid not in count_sql and nid not in count_sql
assert "src_node_id <> e.dst_node_id" in count_sql
assert "LIMIT" not in count_sql
assert sid in count_params and nid in count_params
def test_node_detail_total_falls_back_when_the_count_fails(self, monkeypatch):
store, cursor = self._store_with_mock_conn()
monkeypatch.setattr(store_module, "MAX_NODE_RELATIONSHIPS", 1)
sid = str(uuid.uuid4())
nid = str(uuid.uuid4())
cursor.fetchone.return_value = (nid, "Ada", "person", "d", 1, 1)
cursor.fetchall.side_effect = [
[(str(uuid.uuid4()), "Alan", "person", 2, "knows", "out")],
[],
]
def execute(sql, params=None):
if "COUNT(*)" in sql:
raise RuntimeError("boom")
cursor.execute.side_effect = execute
detail = store.get_node_detail(sid, nid)
assert detail["relationships_total"] == 1
assert len(detail["relationships"]) == 1
def test_remap_chunk_binds_ids_under_the_source_lock(self):
store, cursor = self._store_with_mock_conn()
store._write_with_reconnect = lambda write: write(store._connection)
sid = str(uuid.uuid4())
store.remap_chunk(sid, "old-id", "new-id")
calls = cursor.execute.call_args_list
assert "pg_advisory_xact_lock" in calls[0].args[0]
tables = ["graph_node_chunks", "graph_edges", "graph_ingest_progress"]
for table in tables:
sql, params = next(c.args for c in calls if table in c.args[0])
assert sid not in sql and "old-id" not in sql and "new-id" not in sql
assert {sid, "old-id", "new-id"} <= set(params)
store._connection.commit.assert_called_once()
def test_node_detail_survives_a_relationships_failure(self):
# A failed relationships read degrades to an empty list; the node and
# its chunks still come back instead of the whole detail 404ing.
store, cursor = self._store_with_mock_conn()
conn = store._connection
sid = str(uuid.uuid4())
nid = str(uuid.uuid4())
cursor.fetchone.return_value = (nid, "Ada", "person", "d", 1, 1)
def execute(sql, params=None):
if "direction" in sql:
raise RuntimeError("boom")
cursor.execute.side_effect = execute
cursor.fetchall.return_value = []
detail = store.get_node_detail(sid, nid)
assert detail is not None
assert detail["name"] == "Ada"
assert detail["relationships"] == []
assert detail["relationships_total"] == 0
assert detail["chunks"] == []
# The aborted transaction is rolled back before the chunk read reuses it.
assert conn.rollback.call_count >= 2
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.integration
class TestEntityPagesLive:
"""``entity_pages`` against a real pgvector-shaped table.
The graph tables alone cannot answer it: the rows it returns live in the
documents table the sources were ingested into, so the test creates a
minimal one with the same column names ``PGVectorStore`` uses.
"""
@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}")
conn = store._get_connection()
cursor = conn.cursor()
cursor.execute(
"""
CREATE TABLE IF NOT EXISTS documents (
id SERIAL PRIMARY KEY,
text TEXT,
metadata JSONB,
source_id TEXT
);
"""
)
conn.commit()
cursor.close()
yield store
store.close()
def test_a_page_linked_by_two_entities_is_returned_once(self, store):
"""One chunk, two nodes whose names both match: an exact hit and a
mention. They differ only in whether the page is *about* the entity, so
grouping on that flag returned the same page twice and spent a quarter
of the page budget on it."""
source_id = str(uuid.uuid4())
conn = store._get_connection()
cursor = conn.cursor()
cursor.execute(
"INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;",
("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id),
)
chunk_id = str(cursor.fetchone()[0])
conn.commit()
cursor.close()
try:
subject = store.upsert_node(source_id, "Quill", "quill")
mention = store.upsert_node(source_id, "Legacy Quill", "legacy quill")
store.link_node_chunk(source_id, subject, chunk_id)
store.link_node_chunk(source_id, mention, chunk_id)
pages = store.entity_pages(source_id, "Quill", limit=4)
assert [page["text"] for page in pages] == ["Quill is a write-ahead store."]
assert pages[0]["metadata"] == {"title": "quill.md"}
finally:
store.delete_by_source(source_id)
def test_two_chunks_with_identical_text_collapse_into_one_page(self, store):
"""Deliberate: the caller gets at most four pages to hand a model, and a
crawl that ingested the same text twice would spend two of them saying
the same thing. The rows differ only by an id the model never sees."""
source_id = str(uuid.uuid4())
conn = store._get_connection()
cursor = conn.cursor()
chunk_ids = []
for _ in range(2):
cursor.execute(
"INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;",
("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id),
)
chunk_ids.append(str(cursor.fetchone()[0]))
conn.commit()
cursor.close()
try:
node = store.upsert_node(source_id, "Quill", "quill")
for chunk_id in chunk_ids:
store.link_node_chunk(source_id, node, chunk_id)
pages = store.entity_pages(source_id, "Quill", limit=4)
assert [page["text"] for page in pages] == ["Quill is a write-ahead store."]
finally:
store.delete_by_source(source_id)
@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"),
lambda s: s.count_edges("sid"),
# list_nodes now surfaces query errors, so feed it a numeric count row.
lambda s: (
setattr(s._get_connection().cursor(), "fetchone", lambda: [0]),
s.list_nodes("sid", query="a", type_key="person"),
),
lambda s: s.node_type_facets("sid"),
lambda s: s.get_node_detail("sid", "n"),
],
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",
"count_edges",
"list_nodes",
"node_type_facets",
"get_node_detail",
],
)
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)
@pytest.mark.unit
class TestGraphTypeKey:
"""The Python type key must match the frontend's."""
@pytest.mark.parametrize(
"raw, key",
[
("Person", "person"),
("PERSON", "person"),
("per son", "person"),
("Org-Unit_2", "orgunit2"),
("", ""),
(None, ""),
("---", ""),
],
)
def test_key(self, raw, key):
assert store_module.graph_type_key(raw) == key
# Non-ASCII types fold the same way the frontend's ``\p{L}\p{N}`` does.
assert store_module.graph_type_key("Ком-пания") == "компания"
assert store_module.graph_type_key("人 物") == "人物"
@pytest.mark.integration
class TestTypeKeysOnACLocaleDatabase:
"""Type keys must not depend on the database's ``LC_CTYPE``.
Managed Postgres commonly runs ``C``: there ``lower()`` leaves Cyrillic
alone and ``[:alnum:]`` matches ASCII only, so a key computed in SQL folds
``Компания`` to ``""``. Only the columns these reads touch are created, so
the test runs without the pgvector extension.
"""
@pytest.fixture
def c_store(self, postgresql):
import psycopg
base = _ephemeral_dsn(postgresql.info)
dbname = f"graph_c_{uuid.uuid4().hex[:8]}"
with psycopg.connect(base, autocommit=True) as admin:
admin.execute(
f"CREATE DATABASE {dbname} TEMPLATE template0 ENCODING 'UTF8' "
"LC_COLLATE 'C' LC_CTYPE 'C';"
)
dsn = base.rsplit("/", 1)[0] + f"/{dbname}"
with psycopg.connect(dsn, autocommit=True) as conn:
conn.execute(
"""
CREATE TABLE graph_nodes (
id UUID PRIMARY KEY, source_id UUID NOT NULL, name TEXT,
normalized_name TEXT, type TEXT, description TEXT,
degree INT DEFAULT 0, doc_freq INT DEFAULT 0
);
"""
)
store = GraphStore(connection_string=dsn)
store._pool_max_size = 0
try:
yield store, dsn
finally:
store.close()
with psycopg.connect(base, autocommit=True) as admin:
admin.execute(f"DROP DATABASE IF EXISTS {dbname} WITH (FORCE);")
def test_non_ascii_types_facet_and_filter(self, c_store):
import psycopg
store, dsn = c_store
sid, other = str(uuid.uuid4()), str(uuid.uuid4())
rows = [
(sid, "a", "Компания", 5),
(sid, "b", "Компания", 4),
(sid, "c", "компания", 3),
(sid, "d", "人物", 2),
(sid, "e", "人 物", 1),
(sid, "f", None, 0),
(sid, "g", "---", 0),
(sid, "h", "Person", 0),
(other, "x", "Компания", 9),
]
with psycopg.connect(dsn, autocommit=True) as conn:
for source, name, type_, degree in rows:
conn.execute(
"INSERT INTO graph_nodes (id, source_id, name, normalized_name, "
"type, degree) VALUES (%s, %s, %s, %s, %s, %s);",
(str(uuid.uuid4()), source, name, name, type_, degree),
)
assert store.node_type_facets(sid) == [
{"key": "компания", "label": "Компания", "count": 3},
{"key": "人物", "label": "人 物", "count": 2},
{"key": "", "label": None, "count": 2},
{"key": "person", "label": "Person", "count": 1},
]
companies = store.list_nodes(sid, type_key="компания")
assert companies["total"] == 3
assert [n["name"] for n in companies["nodes"]] == ["a", "b", "c"]
assert store.list_nodes(sid, type_key="人物")["total"] == 2
untyped = store.list_nodes(sid, type_key="")
assert sorted(n["name"] for n in untyped["nodes"]) == ["f", "g"]
assert store.list_nodes(sid, type_key="person", query="h")["total"] == 1
assert store.list_nodes(sid, type_key="missing") == {"nodes": [], "total": 0}
@pytest.mark.unit
class TestTypeKeyFoldingInPython:
"""Facets and the type filter fold raw types with ``graph_type_key``."""
def _store(self):
store = GraphStore.__new__(GraphStore)
cursor = MagicMock()
conn = MagicMock()
conn.cursor.return_value = cursor
store._connection = conn
store._get_connection = lambda: conn
store._tables_ensured = True
return store, cursor
def test_facets_group_raw_types_and_pick_the_label(self):
store, cursor = self._store()
cursor.fetchall.return_value = [
("Компания", 2), ("компания", 2), ("КОМПАНИЯ", 1),
(None, 1), ("---", 3), ("人物", 5),
]
sid = str(uuid.uuid4())
facets = store.node_type_facets(sid)
sql_text, params = cursor.execute.call_args.args
assert "GROUP BY type" in sql_text
assert "regexp_replace" not in sql_text
assert params == (sid,)
# Most frequent spelling wins; a tie goes to the alphabetically first.
assert facets == [
{"key": "компания", "label": "Компания", "count": 5},
{"key": "人物", "label": "人物", "count": 5},
{"key": "", "label": None, "count": 4},
]
def test_type_filter_binds_the_matching_raw_types(self):
store, cursor = self._store()
cursor.fetchall.side_effect = [
[("Компания",), ("компания",), ("Person",), (None,)],
[],
]
cursor.fetchone.return_value = [3]
sid = str(uuid.uuid4())
result = store.list_nodes(sid, type_key="компания")
assert result == {"nodes": [], "total": 3}
distinct_sql, distinct_params = cursor.execute.call_args_list[0].args
assert "DISTINCT type" in distinct_sql
assert distinct_params == (sid,)
count_sql, count_params = cursor.execute.call_args_list[1].args
assert "type = ANY(%s)" in count_sql
assert "IS NULL" not in count_sql
assert count_params == (sid, ["Компания", "компания"])
def test_untyped_filter_includes_null_and_types_that_fold_to_empty(self):
store, cursor = self._store()
cursor.fetchall.side_effect = [[("---",), ("Person",), (None,)], []]
cursor.fetchone.return_value = [2]
sid = str(uuid.uuid4())
store.list_nodes(sid, type_key="")
count_sql, count_params = cursor.execute.call_args_list[1].args
assert "type IS NULL OR type = ANY(%s)" in count_sql
assert count_params == (sid, ["---"])
def test_unknown_type_key_short_circuits(self):
store, cursor = self._store()
cursor.fetchall.return_value = [("Person",)]
assert store.list_nodes("sid", type_key="org") == {"nodes": [], "total": 0}
assert cursor.execute.call_count == 1