mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 10:13:06 +00:00
2163 lines
84 KiB
Python
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
|