feat(graphrag): seed the walk from what entities are, and rank with passages and vector hits

Graph retrieval tied plain vector search at best and never beat it. Measured
across five corpora, the bottleneck was seeding, not the graph: the walk
started from nodes whose embeddings were computed from bare entity names, and
a whole question shares almost nothing with a name like "Quill".

Extraction now embeds each node from "name (type): description" and each
relationship as the fact it asserts ("Alder streams_to Quill: ..."), stored on
a new nullable graph_edges.fact_embedding column that ensure_vector_schema adds
in place. Entity names are canonicalised (case, punctuation, word breaks and a
cautious plural) so "VECTOR_STORE" and "vector stores" land on one node. Extraction calls run
concurrently (GRAPHRAG_EXTRACTION_WORKERS, default 8) while embedding and graph
writes stay serial on the task thread, so ordering and idempotency are
unchanged; that measured 8.4x faster with identical output.

Retrieval gains per-source options, stored under retrieval.graph and read live
at query time:

- seed_strategy: start from matching entities (default) or matching
  relationships, which can reach an entity the question never names;
- passage_nodes (on): walk the source's passages alongside entities, with
  PageRank damping 0.5 instead of 0.85;
- blend_vector (on): fuse the graph ranking with the source's vector ranking
  by reciprocal rank.

The defaults are the measured-best configuration. Through GraphRAGRetriever,
the new seeding moved recall@4 from 0.41 to 0.68 on a multi-hop corpus and
from 0.50 to 1.00 on the docs corpus, and regressed none of the corpora
measured. Existing graphs keep name-only embeddings until rebuilt.
This commit is contained in:
Alex committed 2026-09-19 14:07:41 +01:00
1 parent b7e7872bf7
commit a83e1dc0af
13 files changed
+1577 -79

No files matched your search

@@ -447,6 +447,12 @@ Type `int`, default `2000`, must be `>= 0`.
Hard cap on chunks extracted per source (cost control); 0 extracts nothing.
### `GRAPHRAG_EXTRACTION_WORKERS`
Type `int`, default `8`, must be `>= 1` and `<= 32`.
Concurrent extraction calls during ingest. Model calls run in parallel while graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial.
## Vector stores
+9
View File
@@ -31,6 +31,15 @@ class RetrievalSettings(SettingsGroup):
GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field(
default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing."
)
GRAPHRAG_EXTRACTION_WORKERS: int = Field(
default=8,
ge=1,
le=32,
description=(
"Concurrent extraction calls during ingest. Model calls run in parallel while "
"graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial."
),
)
@field_validator("VECTOR_STORE", mode="before")
@classmethod
+140 -29
View File
@@ -27,6 +27,7 @@ from docsgpt.core.model_utils import (
get_api_key_for_provider,
get_provider_from_model_id,
)
from docsgpt.graphrag.naming import normalize_entity_name
from docsgpt.core.settings import settings
from docsgpt.llm.llm_creator import LLMCreator
from docsgpt.storage.db.source_config import SourceConfig
@@ -224,6 +225,8 @@ def extract_graph_for_source(
source's graph holds after the run — not how many upserts ran, which
counts the same entity once per chunk it appears in.
"""
from concurrent.futures import ThreadPoolExecutor
from docsgpt.graphrag.store import GraphStore
store = GraphStore()
@@ -266,43 +269,85 @@ def extract_graph_for_source(
except Exception as exc:
logger.debug("graph progress callback failed: %s", exc)
for chunk, chunk_id in to_process:
def _prepare(item):
"""One chunk's LLM extraction — the only step run concurrently.
A chunk spends almost all of its time waiting on the model, so that is
what runs in the pool. Everything else stays on the calling thread:
graph writes, so transactions and the progress checkpoint are exactly
what they were serially, and embedding. Inside a Celery worker the
embeddings client decides to embed locally from the task on the
*current thread's* stack; a pool thread has none, so it would instead
dispatch an embed task to the worker and wait on it, which Celery
refuses inside a task — failing every chunk of the build.
"""
chunk, chunk_id = item
text = _chunk_text(chunk)
if not text:
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
_report()
continue
return chunk_id, "empty", None
extracted = _extract_chunk(llm, text, chunk_id)
if extracted is None:
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
continue
return chunk_id, "failed", None
try:
entities = _build_entities(extracted["entities"])
relationships = _build_relationships(extracted["relationships"])
name_embeddings = _embed_names(embedding, entities, relationships)
chunk_nodes, chunk_edges = store.apply_chunk(
source_id, chunk_id, entities, relationships, name_embeddings
)
node_upserts += chunk_nodes
edges += chunk_edges
# ``apply_chunk`` marks the chunk done inside the transaction that
# writes its rows, so the checkpoint cannot disagree with the graph
# and a replayed write cannot apply the chunk twice.
chunks_processed += 1
except Exception as exc:
logger.warning(
"Graph extraction write failed for chunk %s, skipping: %s",
chunk_id,
exc,
"Graph extraction failed for chunk %s, skipping: %s", chunk_id, exc
)
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
return chunk_id, "failed", None
return chunk_id, "ok", (entities, relationships)
workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1))
pool = None
if workers > 1 and len(to_process) > 1:
pool = ThreadPoolExecutor(max_workers=workers)
# ``map`` yields in submission order, so chunks are still applied in the
# order they were given and a run stays reproducible.
prepared = pool.map(_prepare, to_process)
else:
prepared = (_prepare(item) for item in to_process)
try:
for chunk_id, status, payload in prepared:
if status == "empty":
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
_report()
continue
if status == "failed":
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
continue
entities, relationships = payload
try:
# On this thread, not in the pool — see ``_prepare``.
name_embeddings = _embed_names(embedding, entities, relationships)
_embed_facts(embedding, relationships)
chunk_nodes, chunk_edges = store.apply_chunk(
source_id, chunk_id, entities, relationships, name_embeddings
)
node_upserts += chunk_nodes
edges += chunk_edges
# ``apply_chunk`` marks the chunk done inside the transaction that
# writes its rows, so the checkpoint cannot disagree with the graph
# and a replayed write cannot apply the chunk twice.
chunks_processed += 1
except Exception as exc:
logger.warning(
"Graph extraction embed/write failed for chunk %s, skipping: %s",
chunk_id,
exc,
)
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
finally:
if pool is not None:
pool.shutdown(wait=True)
try:
store.set_node_degrees(source_id)
@@ -347,7 +392,7 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
entities.append(
{
"name": name,
"normalized_name": name.lower(),
"normalized_name": normalize_entity_name(name),
"type": str(e.get("type") or "") or None,
"description": str(e.get("description") or "") or None,
}
@@ -373,6 +418,70 @@ def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]:
return relationships
def _fact_text(rel: Dict[str, Any]) -> str:
"""A relationship rendered as the sentence it asserts.
Embedded and stored on the edge so retrieval can match a question against
the *relation* rather than against entity names — the difference between
"which entity is this about" and "which fact answers this".
"""
source = str(rel.get("source") or "").strip()
target = str(rel.get("target") or "").strip()
if not source or not target:
return ""
relation = str(rel.get("type") or "related to").strip() or "related to"
text = f"{source} {relation} {target}"
description = str(rel.get("description") or "").strip()
return f"{text}: {description}" if description else text
def _embed_facts(embedding, relationships: List[Dict[str, Any]]) -> None:
"""Attach a fact embedding to each relationship, in one batched call.
Mutates the relationship dicts so the embedding travels with the edge into
``apply_chunk`` without a second mapping to keep in step. Always on: it is
one extra batched call per chunk against an LLM call that already costs
far more, and it lets a source switch to relationship seeding at query time
without being rebuilt.
"""
pending = [(rel, _fact_text(rel)) for rel in relationships]
pending = [(rel, text) for rel, text in pending if text]
if not pending:
return
try:
vectors = embedding.embed_documents([text for _rel, text in pending])
except Exception as exc: # noqa: BLE001
# The graph is still correct without them; only fact seeding degrades.
logger.warning("Fact embedding failed, continuing without: %s", exc)
return
for (rel, _text), vector in zip(pending, vectors):
rel["fact_embedding"] = vector
def _seed_text(entity: Dict[str, Any]) -> str:
"""The text a node's embedding is computed from.
Retrieval seeds the graph walk by matching a whole question against these
embeddings, and a bare entity name is a poor thing to match a question
against — a question about what a service writes to shares almost no
surface with the name ``Quill``. Including the type and description gives
the match something to work with; measured across five corpora it moved
recall@4 by +0.07 to +0.50.
Relationship endpoints keep their bare names: they arrive as strings with
no type or description attached.
"""
name = str(entity.get("name") or "").strip()
text = name
entity_type = str(entity.get("type") or "").strip()
if entity_type:
text += f" ({entity_type})"
description = str(entity.get("description") or "").strip()
if description:
text += f": {description}"
return text or name
def _embed_names(
embedding,
entities: List[Dict[str, Any]],
@@ -385,14 +494,16 @@ def _embed_names(
"""
name_by_norm: Dict[str, str] = {}
for entity in entities:
name_by_norm.setdefault(entity["normalized_name"], entity["name"])
name_by_norm.setdefault(entity["normalized_name"], _seed_text(entity))
for rel in relationships:
for endpoint in (rel.get("source"), rel.get("target")):
if endpoint is None:
continue
clean = str(endpoint).strip()
if clean:
name_by_norm.setdefault(clean.lower(), clean)
# Same key the store resolves endpoints by, or the embedding
# computed here never reaches the node it was computed for.
name_by_norm.setdefault(normalize_entity_name(clean), clean)
if not name_by_norm:
return {}
+94
View File
@@ -0,0 +1,94 @@
"""Canonical entity naming for the per-source knowledge graph.
Nodes are merged on ``normalized_name``, which has been ``name.lower()``. That
splits entities a reader would call the same thing: measured on the DocsGPT docs
corpus, ``agent``/``agents``, ``VECTOR_STORE``/``Vector store``/``vector stores``,
``Celery worker``/``Celery workers`` and ``.env file``/``env_file`` all landed as
separate nodes — 58 such collisions across 1,704 entities, with 75% of entities
appearing in exactly one chunk as a result.
:func:`canonical_name` folds the differences that are purely orthographic:
case, surrounding punctuation, underscore/hyphen word breaks, and a *cautious*
plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``,
``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s"
would corrupt them into new entities rather than merge anything.
Always on: every graph is built with canonical names.
"""
from __future__ import annotations
import re
_PUNCT = re.compile(r"[^\w\s]+", re.UNICODE)
_UNDERSCORE = re.compile(r"[_\-]+")
_SPACE = re.compile(r"\s+")
#: Words that end in "s" without being plural. Singularising these would invent
#: entities ("postgre", "kubernete") instead of merging existing ones.
_NOT_PLURAL = frozenset(
{
"postgres", "kubernetes", "https", "aws", "dns", "tls", "cors", "css",
"js", "sas", "gas", "ss", "class", "access", "process", "status",
"analysis", "basis", "axis", "https", "rss", "less", "express",
"redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos",
"always", "sometimes", "series", "docs", "ops", "devops", "sse",
}
)
def _singular(word: str) -> str:
"""Best-effort singular of one word, biased hard towards leaving it alone.
Only the endings that are unambiguous in this domain are touched:
``-ies`` -> ``-y`` (``policies``), ``-ses``/``-xes``/``-zes``/``-ches``/
``-shes`` -> drop ``es`` (``indexes``, ``batches``), and a bare trailing
``s`` on a word long enough to be safe. Everything in :data:`_NOT_PLURAL`,
and anything ending in ``ss``/``us``/``is``, is returned unchanged.
"""
if len(word) < 4 or word in _NOT_PLURAL:
return word
if word.endswith(("ss", "us", "is")):
return word
if word.endswith("ies") and len(word) > 4:
return word[:-3] + "y"
if word.endswith(("ses", "xes", "zes", "ches", "shes")):
return word[:-2]
if word.endswith("s"):
return word[:-1]
return word
def canonical_name(name: str) -> str:
"""Merge key for an entity name.
Args:
name: The entity name as the model wrote it.
Returns:
A lowercase, punctuation-free, singularised key. Returns ``""`` for an
empty or punctuation-only name, which callers treat as "no entity".
Examples:
``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``;
``.env file`` and ``env_file`` -> ``env file``;
``postgres`` stays ``postgres``.
"""
if not name:
return ""
text = _UNDERSCORE.sub(" ", str(name))
text = _PUNCT.sub(" ", text)
text = _SPACE.sub(" ", text).strip().lower()
if not text:
return ""
return " ".join(_singular(word) for word in text.split())
def normalize_entity_name(name: str) -> str:
"""The key an entity is merged on: its :func:`canonical_name`.
Every graph the corpora were measured on was built this way, so it is the
only mode rather than a flag. A graph built before this used plain
``lower()`` keys; re-extracting it merges onto these instead.
"""
return canonical_name(name)
+312 -19
View File
@@ -74,6 +74,16 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]:
)
def _pgvector_vector_column() -> str:
"""Resolve the embedding column name from the same ``PGVectorStore`` defaults."""
import inspect
from docsgpt.vectorstore.pgvector import PGVectorStore
params = inspect.signature(PGVectorStore.__init__).parameters
return _safe_identifier(params["vector_column"].default)
def _is_connection_lost(exc: BaseException) -> bool:
"""True when ``exc`` says the server connection went away, not that the SQL was bad.
@@ -253,7 +263,7 @@ class GraphStore:
)
cursor.execute(
"""
f"""
CREATE TABLE IF NOT EXISTS graph_edges (
id UUID PRIMARY KEY,
source_id UUID NOT NULL,
@@ -262,10 +272,18 @@ class GraphStore:
type TEXT,
description TEXT,
weight REAL DEFAULT 1.0,
source_chunk_ids JSONB
source_chunk_ids JSONB,
fact_embedding vector({dimension})
);
"""
)
# ``CREATE TABLE IF NOT EXISTS`` is a no-op on a database that
# already has the table, so a column added after the fact needs its
# own statement or every existing deployment silently lacks it.
cursor.execute(
f"ALTER TABLE graph_edges "
f"ADD COLUMN IF NOT EXISTS fact_embedding vector({dimension});"
)
cursor.execute(
"""
@@ -453,19 +471,78 @@ class GraphStore:
description: Optional[str] = None,
weight: float = 1.0,
source_chunk_ids: Optional[List[str]] = None,
) -> str:
"""Insert an edge on an open cursor (no commit, no degree bump).
fact_embedding: Optional[List[float]] = None,
) -> tuple[Optional[str], bool]:
"""Write an edge on an open cursor (no commit, no degree bump).
Returns ``(edge_id, created)``. Two shapes of noise are rejected here
rather than at read time, because once written neither is visible:
* A self-loop feeds a node's PageRank mass straight back to itself. It
is dropped, reported as ``(None, False)``.
* A pair already related by the same type is *merged* rather than
inserted again. ``graph_edges`` carries no uniqueness constraint, so
re-extracting one relationship across many chunks otherwise writes a
row per chunk — a fifth of a real corpus's edges — inflating that
pair's traversal weight and spending the bounded subgraph fetch on
duplicates. The surviving row keeps the strongest weight seen and
every contributing chunk id.
Callers that batch many edges run ``set_node_degrees`` once afterwards
instead of bumping degree per edge.
"""
if str(src_node_id) == str(dst_node_id):
return None, False
cursor.execute(
"""
SELECT id
FROM graph_edges
WHERE source_id = %s AND src_node_id = %s AND dst_node_id = %s
AND type IS NOT DISTINCT FROM %s
LIMIT 1;
""",
(source_id, src_node_id, dst_node_id, type),
)
existing = cursor.fetchone()
if existing:
edge_id = existing[0]
# The chunk ids are merged in SQL, against the row's own current
# value, rather than read here and written back: a read-modify-write
# would drop whatever a concurrent writer appended in between.
cursor.execute(
"""
UPDATE graph_edges
SET weight = GREATEST(COALESCE(weight, 0), %s),
description = COALESCE(description, %s),
-- Backfills the fact embedding for an edge first written
-- before fact embeddings were switched on.
fact_embedding = COALESCE(fact_embedding, %s::vector),
source_chunk_ids = COALESCE(source_chunk_ids, '[]'::jsonb) || (
SELECT COALESCE(jsonb_agg(candidate), '[]'::jsonb)
FROM jsonb_array_elements(%s::jsonb) AS candidate
WHERE NOT COALESCE(source_chunk_ids, '[]'::jsonb)
@> jsonb_build_array(candidate)
)
WHERE id = %s;
""",
(
weight,
description,
fact_embedding,
Jsonb(list(source_chunk_ids or [])),
edge_id,
),
)
return str(edge_id), False
edge_id = str(uuid.uuid4())
cursor.execute(
"""
INSERT INTO graph_edges
(id, source_id, src_node_id, dst_node_id, type, description,
weight, source_chunk_ids)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s);
weight, source_chunk_ids, fact_embedding)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s);
""",
(
edge_id,
@@ -476,9 +553,10 @@ class GraphStore:
description,
weight,
Jsonb(source_chunk_ids or []),
fact_embedding,
),
)
return edge_id
return edge_id, True
def add_edge(
self,
@@ -489,21 +567,28 @@ class GraphStore:
description: Optional[str] = None,
weight: float = 1.0,
source_chunk_ids: Optional[List[str]] = None,
) -> str:
"""Insert an edge and bump the degree of both endpoints. Returns its id."""
fact_embedding: Optional[List[float]] = None,
) -> Optional[str]:
"""Write an edge and bump the degree of both endpoints. Returns its id.
Returns ``None`` for a self-loop, which is not written. A repeat of an
existing pair merges into that row and returns its id, leaving degree
alone — the endpoints gained no new neighbour.
"""
self._ensure_tables_once()
conn = self._get_connection()
cursor = conn.cursor()
try:
edge_id = self._add_edge(
edge_id, created = self._add_edge(
cursor, source_id, src_node_id, dst_node_id, type, description,
weight, source_chunk_ids,
)
cursor.execute(
"UPDATE graph_nodes SET degree = degree + 1 "
"WHERE source_id = %s AND id IN (%s, %s);",
(source_id, src_node_id, dst_node_id),
weight, source_chunk_ids, fact_embedding,
)
if created:
cursor.execute(
"UPDATE graph_nodes SET degree = degree + 1 "
"WHERE source_id = %s AND id IN (%s, %s);",
(source_id, src_node_id, dst_node_id),
)
conn.commit()
return edge_id
except Exception as e:
@@ -608,7 +693,7 @@ class GraphStore:
)
if src_id is None or dst_id is None:
continue
self._add_edge(
_, created = self._add_edge(
cursor,
source_id,
src_id,
@@ -617,8 +702,10 @@ class GraphStore:
description=rel.get("description"),
weight=float(rel.get("weight") or 1.0),
source_chunk_ids=[chunk_id],
fact_embedding=rel.get("fact_embedding"),
)
edges_added += 1
if created:
edges_added += 1
cursor.execute(
"""
@@ -653,7 +740,11 @@ class GraphStore:
clean = str(name).strip()
if not clean:
return None
normalized_name = clean.lower()
from docsgpt.graphrag.naming import normalize_entity_name
normalized_name = normalize_entity_name(clean)
if not normalized_name:
return None
if normalized_name in node_ids:
return node_ids[normalized_name]
node_id = self._upsert_node(
@@ -839,6 +930,7 @@ class GraphStore:
FROM graph_edges
WHERE source_id = %s
AND (src_node_id = ANY(%s) OR dst_node_id = ANY(%s))
ORDER BY weight DESC NULLS LAST
LIMIT %s;
""",
(
@@ -886,6 +978,7 @@ class GraphStore:
FROM graph_edges
WHERE source_id = %s
AND src_node_id = ANY(%s) AND dst_node_id = ANY(%s)
ORDER BY weight DESC NULLS LAST
LIMIT %s;
""",
(source_id, node_id_list, node_id_list, MAX_SUBGRAPH_EDGES),
@@ -999,6 +1092,206 @@ class GraphStore:
cursor.close()
conn.rollback()
def seed_nodes_from_facts(
self,
source_id: str,
query_embedding: List[float],
fact_limit: int = 5,
limit: int = 10,
) -> List[Dict[str, Any]]:
"""Seed nodes drawn from the *relationships* nearest the question.
Name matching asks "which entity is this question about", which a
multi-document question cannot answer: the entity holding the answer is
named in another document, not in the question. A fact string carries
the relation — "Alder streams_to Quill: ..." — so a question about what
a service writes to can match the edge itself and seed the walk on both
of its endpoints, including the one nothing in the question names.
Endpoints are weighted by fact score divided by the entity's
``doc_freq``: an entity appearing in every chunk is a poor seed even
when it sits on a well-matched fact, and dividing by how widely it
occurs prefers the specific endpoint over the hub.
Rows match :meth:`search_nodes_by_embedding`'s shape, so the caller's
seed weighting is unchanged. Returns nothing when the source has no
fact embeddings, which is the signal to fall back to name matching.
"""
if not query_embedding:
return []
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
WITH top_facts AS (
SELECT src_node_id, dst_node_id,
1 - (fact_embedding <=> %s::vector) AS score
FROM graph_edges
WHERE source_id = %s AND fact_embedding IS NOT NULL
ORDER BY fact_embedding <=> %s::vector
LIMIT %s
)
SELECT n.id::text, n.name, n.description,
MAX(f.score / GREATEST(COALESCE(n.doc_freq, 1), 1)) AS weight
FROM top_facts f
JOIN graph_nodes n
ON n.id = f.src_node_id OR n.id = f.dst_node_id
WHERE n.source_id = %s
GROUP BY n.id, n.name, n.description
ORDER BY weight DESC
LIMIT %s;
""",
(
query_embedding,
source_id,
query_embedding,
max(1, int(fact_limit)),
source_id,
max(1, int(limit)),
),
)
return [
{
"id": row[0],
"name": row[1],
"description": row[2],
# The caller reads weight back as ``1 - distance``.
"distance": 1.0 - float(row[3] or 0.0),
}
for row in cursor.fetchall()
]
except Exception as e:
logging.error(f"Error seeding nodes from facts: {e}")
return []
finally:
cursor.close()
conn.rollback()
def entity_relationships(
self, source_id: str, name: str, limit: int = 25
) -> List[Dict[str, Any]]:
"""The relationships an entity takes part in, strongest first.
This is the one thing a caller cannot get from vector search: which
*named* thing an entity is connected to. Matching is on the name rather
than a node id because the caller is an LLM holding a name it read in
the text, not an id.
"""
clean = (name or "").strip()
if not clean:
return []
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
SELECT s.name, e.type, d.name, e.description
FROM graph_edges e
JOIN graph_nodes s ON s.id = e.src_node_id
JOIN graph_nodes d ON d.id = e.dst_node_id
WHERE e.source_id = %s AND (s.name ILIKE %s OR d.name ILIKE %s)
ORDER BY e.weight DESC NULLS LAST
LIMIT %s;
""",
(source_id, f"%{clean}%", f"%{clean}%", max(1, int(limit))),
)
return [
{"source": row[0], "type": row[1], "target": row[2], "description": row[3]}
for row in cursor.fetchall()
]
except Exception as e:
logging.error(f"Error reading relationships for {name!r}: {e}")
return []
finally:
cursor.close()
conn.rollback()
def entity_pages(
self, source_id: str, name: str, limit: int = 4
) -> List[Dict[str, Any]]:
"""Chunks an entity appears in, with the chunk it is *about* first.
A plain substring match answers "Halvard" with pages that merely mention
Halvard, and an unordered ``LIMIT`` then decides which of those the
caller sees. Nodes whose name is the entity (or the entity plus a
qualifier the extractor appended, "Quill" -> "Quill Store") are
preferred, and among those the chunk whose text opens with the name
comes first; a substring match is the fallback so an unusual name still
resolves.
"""
clean = (name or "").strip()
if not clean:
return []
table, text_col, metadata_col, source_col = _pgvector_identifiers()
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
f"""
SELECT d.{metadata_col}, d.{text_col},
(lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject
FROM graph_node_chunks gc
JOIN graph_nodes n ON n.id = gc.node_id
JOIN {table} d ON d.id::text = gc.chunk_id
WHERE gc.source_id = %s AND d.{source_col} = %s
AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s)
GROUP BY d.{metadata_col}, d.{text_col}, is_subject
ORDER BY is_subject DESC, (d.{text_col} ILIKE %s) DESC
LIMIT %s;
""",
(
clean.lower(), f"{clean.lower()} %",
source_id, source_id,
clean.lower(), f"{clean.lower()} %", f"%{clean}%",
f"{clean}%",
max(1, int(limit)),
),
)
return [
{"metadata": row[0] or {}, "text": row[1] or ""}
for row in cursor.fetchall()
]
except Exception as e:
logging.error(f"Error reading pages for {name!r}: {e}")
return []
finally:
cursor.close()
conn.rollback()
def chunk_similarities(
self, source_id: str, chunk_ids: List[str], query_embedding: List[float]
) -> Dict[str, float]:
"""Cosine similarity between the query and specific chunks of a source.
Passage nodes need their own relevance to claim a share of the walk's
restart mass, and that number lives in the co-located pgvector table —
the same one :meth:`get_chunk_texts` reads. Restricted to the chunk ids
the subgraph actually reached, so this never scans the whole source.
"""
if not chunk_ids or not query_embedding:
return {}
table, _text_col, _metadata_col, source_col = _pgvector_identifiers()
vector_col = _pgvector_vector_column()
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
f"""
SELECT id::text, 1 - ({vector_col} <=> %s::vector)
FROM {table}
WHERE {source_col} = %s AND id::text = ANY(%s);
""",
(query_embedding, source_id, [str(c) for c in chunk_ids]),
)
return {row[0]: float(row[1]) for row in cursor.fetchall()}
except Exception as e:
logging.error(f"Error scoring chunks against the query: {e}")
return {}
finally:
cursor.close()
conn.rollback()
def get_chunk_texts(
self,
source_id: str,
+261 -25
View File
@@ -32,6 +32,7 @@ from docsgpt.graphrag.store import GraphStore
from docsgpt.retriever.base import BaseRetriever
from docsgpt.retriever.classic_rag import ClassicRAG
from docsgpt.retriever.labels import labels_from_metadata
from docsgpt.storage.db.source_config import GraphRetrievalConfig
from docsgpt.utils import num_tokens_from_string
from docsgpt.vectorstore.base import get_embeddings
@@ -39,11 +40,27 @@ SEED_NODES = 10
SUBGRAPH_HOPS = 1
PASSAGE_NODE_WEIGHT = 0.05
FACT_SEED_FACTS = 5
RRF_K = 60
# PageRank damping per ranking mode — each is the value that mode was measured
# at. Lower keeps mass nearer the seeds; with passages in the walk 0.5 measured
# better, while entity-only ranking was measured at the conventional 0.85.
DAMPING_WITH_PASSAGES = 0.5
DAMPING_ENTITIES_ONLY = 0.85
def _idf(doc_freq: Any) -> float:
"""Node-specificity weight: rarer entities (low ``doc_freq``) score higher."""
return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0)
def _damping(passage_nodes: bool) -> float:
"""PageRank damping for a ranking mode: the value that mode was measured at."""
return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY
def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]:
"""Normalized restart distribution over ``nodes``.
@@ -210,18 +227,9 @@ class GraphRAGRetriever(BaseRetriever):
After PPR, each node's mass is scaled by ``1/log(2 + doc_freq)`` so a
high-degree hub contributes less than a specific entity at equal mass.
"""
graph = nx.Graph()
for node in subgraph.get("nodes", []):
graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0))
for edge in subgraph.get("edges", []):
src, dst = edge["src_node_id"], edge["dst_node_id"]
if src in graph and dst in graph:
raw_weight = edge.get("weight")
# Same rule the ranker applies: default only a missing or null
# weight. Coercing an explicit 0 to 1.0 here would make "these
# entities are not related" the strongest possible link.
edge_weight = 1.0 if raw_weight is None else float(raw_weight)
graph.add_edge(src, dst, weight=edge_weight)
# Through the class, not ``self``: this method reads no instance state,
# and callers (and tests) rely on being able to invoke it unbound.
graph = GraphRAGRetriever._subgraph_graph(subgraph)
if graph.number_of_nodes() == 0:
return {}
@@ -230,13 +238,33 @@ class GraphRAGRetriever(BaseRetriever):
personalization = None
ranks = _personalized_pagerank(
graph, personalization=personalization, weight="weight"
graph,
personalization=personalization,
weight="weight",
alpha=_damping(passage_nodes=False),
)
return {
node: rank * _idf(graph.nodes[node].get("doc_freq", 0))
for node, rank in ranks.items()
}
@staticmethod
def _subgraph_graph(subgraph) -> "nx.Graph":
"""The fetched subgraph as a weighted undirected graph."""
graph = nx.Graph()
for node in subgraph.get("nodes", []):
graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0))
for edge in subgraph.get("edges", []):
src, dst = edge["src_node_id"], edge["dst_node_id"]
if src in graph and dst in graph:
raw_weight = edge.get("weight")
# Default only a missing or null weight. Coercing an explicit 0
# to 1.0 would make "these entities are not related" the
# strongest possible link.
edge_weight = 1.0 if raw_weight is None else float(raw_weight)
graph.add_edge(src, dst, weight=edge_weight)
return graph
def _rank_chunks(self, store, source_id, node_scores) -> List[str]:
"""Score chunks by summed (PPR mass x IDF) of their linked nodes; top candidates.
@@ -254,6 +282,74 @@ class GraphRAGRetriever(BaseRetriever):
candidates = max(self.chunks * 2, self.chunks + 5)
return ranked[: max(1, candidates)]
def _rank_chunks_with_passages(
self, store, source_id, subgraph, seeds, query_embedding
) -> List[str]:
"""Rank chunks by walking a graph that contains the chunks themselves.
:meth:`_rank_chunks` reads a chunk's score *off* its entities, summing
their PPR mass — so a chunk touching many mid-scoring generic entities
outranks one touching the few entities the question is about. Putting
the chunks in the walk instead, each joined to its own entities and
carrying a small share of the restart mass proportional to its own
vector similarity, makes a chunk reachable both ways: by being about the
question, and by being connected to what is. Graph retrieval then
contains vector retrieval rather than competing with it.
Only an improvement when the seeds are good: measured across five
corpora it helped alongside richer seed embeddings and *hurt* with
bare-name seeds (0.73 -> 0.57 on one corpus). Graphs are now always
built with the richer seed text; one built before that change should be
rebuilt before this is relied on.
"""
node_ids = [node["id"] for node in subgraph.get("nodes", [])]
chunk_links = store.get_chunk_ids_for_nodes(source_id, node_ids)
candidate_ids = sorted({c for chunks in chunk_links.values() for c in chunks})
if not candidate_ids:
return []
graph = self._subgraph_graph(subgraph)
similarities = store.chunk_similarities(
source_id, candidate_ids, query_embedding
)
# Normalised so the passage share is a fixed fraction of the restart
# mass rather than whatever absolute cosine this embedding model emits.
scores = [similarities.get(c, 0.0) for c in candidate_ids]
low, high = (min(scores), max(scores)) if scores else (0.0, 0.0)
spread = high - low
personalization = dict(seeds)
passage_of: Dict[str, str] = {}
for chunk_id in candidate_ids:
linked = [n for n, chunks in chunk_links.items() if chunk_id in chunks]
linked = [n for n in linked if n in graph]
if not linked:
continue
passage_node = f"chunk::{chunk_id}"
passage_of[passage_node] = chunk_id
for node in linked:
graph.add_edge(passage_node, node, weight=1.0)
similarity = similarities.get(chunk_id, 0.0)
normalized = (similarity - low) / spread if spread > 0 else 0.0
personalization[passage_node] = normalized * PASSAGE_NODE_WEIGHT
if graph.number_of_nodes() == 0 or not any(personalization.values()):
return []
ranks = _personalized_pagerank(
graph,
personalization=personalization,
weight="weight",
alpha=_damping(passage_nodes=True),
)
chunk_scores = {
chunk_id: ranks.get(passage_node, 0.0)
for passage_node, chunk_id in passage_of.items()
}
ranked = sorted(chunk_scores, key=lambda c: chunk_scores[c], reverse=True)
candidates = max(self.chunks * 2, self.chunks + 5)
return ranked[: max(1, candidates)]
def _source_top_k(self, source_id) -> int:
"""How many chunks this source may contribute — its own top-k.
@@ -273,6 +369,111 @@ class GraphRAGRetriever(BaseRetriever):
base = self.base_chunks if self.base_chunks is not None else self.chunks
return max(1, base // max(1, len(self.vectorstores)))
def _vector_ranking(self, source_id, query_embedding: List[float]) -> List[tuple]:
"""The source's own vector ranking, as ``(text, metadata)`` in score order.
Used only by the hybrid path. Vector hits carry no row id, so the fused
ranking is keyed on the chunk text itself — the one identifier both
rankings share — and the metadata travels with it so a hit the graph
never surfaced can still be emitted as a document.
"""
from docsgpt.vectorstore.vector_creator import VectorCreator
store = None
try:
store = VectorCreator.create_vectorstore(
settings.VECTOR_STORE, source_id, settings.EMBEDDINGS_KEY
)
hits = store.search(
self._classic._get_rephrased_question(),
k=max(self.chunks * 4, 20),
query_vector=query_embedding,
)
except Exception as e:
logging.error(
"GraphRAG hybrid: vector ranking failed for %s: %s", source_id, e
)
return []
finally:
close = getattr(store, "close", None)
if close is not None:
try:
close()
except Exception as e:
logging.debug("Error closing hybrid vector store: %s", e)
ranked = []
for hit in hits:
text = getattr(hit, "page_content", None)
metadata = getattr(hit, "metadata", None)
if text is None and isinstance(hit, dict):
text = hit.get("text") or hit.get("page_content")
metadata = hit.get("metadata")
if text:
ranked.append((text, metadata or {}))
return ranked
@staticmethod
def _rrf_order(rankings: List[List[str]], k: int) -> Dict[str, float]:
"""Reciprocal rank fusion over ranked lists of the same key type.
Rank-based on purpose: PPR mass and cosine similarity are not on
comparable scales, and normalising either one invents a calibration
that does not exist.
"""
scores: Dict[str, float] = {}
for ranking in rankings:
for position, key in enumerate(ranking):
scores[key] = scores.get(key, 0.0) + 1.0 / (k + position + 1)
return scores
def _graph_options(self, source_id) -> GraphRetrievalConfig:
"""This source's graph retrieval options, or the recommended defaults.
Options travel on the per-source retrieval config the Dispatcher hands
over. A request that carries no per-source detail gets the defaults,
which are the measured-best configuration rather than a neutral one.
"""
cfg = (getattr(self, "per_source_retrieval", None) or {}).get(source_id)
options = cfg.get("graph") if isinstance(cfg, dict) else getattr(cfg, "graph", None)
if isinstance(options, GraphRetrievalConfig):
return options
try:
return GraphRetrievalConfig.model_validate(options or {})
except Exception:
return GraphRetrievalConfig()
def _seed_rows(
self, store, source_id, query_embedding: List[float]
) -> List[Dict[str, Any]]:
"""The nodes the walk restarts from, per the source's ``seed_strategy``.
Seeding decides more than ranking does: a walk that starts on the wrong
nodes cannot be rescued downstream.
``entities``
Cosine NN over entity embeddings, built from each entity's name,
type and description so a whole question has something to match.
The default: best or tied-best on every corpus measured.
``relationships``
Cosine NN over relationship sentences ("A streams_to B: ..."),
seeding both endpoints of the best-matching facts. The only way to
start on an entity the question never names; strongest on
chain-structured content, weaker on ordinary prose.
Relationship seeding falls back to entity matching for a source with no
fact embeddings (one built before they were recorded), so it still
retrieves rather than returning nothing.
"""
if self._graph_options(source_id).seed_strategy == "relationships":
by_fact = store.seed_nodes_from_facts(
source_id, query_embedding, fact_limit=FACT_SEED_FACTS, limit=SEED_NODES
)
if by_fact:
return by_fact
return store.search_nodes_by_embedding(
source_id, query_embedding, k=SEED_NODES
)
def _graph_docs_for_source(
self, store, source_id, query_embedding: List[float]
) -> List[Dict[str, Any]]:
@@ -284,9 +485,7 @@ class GraphRAGRetriever(BaseRetriever):
query_embedding: Embedding of the rephrased question, computed once
by the caller for the whole retrieval.
"""
seed_rows = store.search_nodes_by_embedding(
source_id, query_embedding, k=SEED_NODES
)
seed_rows = self._seed_rows(store, source_id, query_embedding)
if not seed_rows:
return []
@@ -300,26 +499,63 @@ class GraphRAGRetriever(BaseRetriever):
for row in seed_rows
}
options = self._graph_options(source_id)
subgraph = store.get_subgraph(source_id, seed_ids, hops=SUBGRAPH_HOPS)
node_scores = self._ppr_scores(subgraph, seeds)
if not node_scores:
if options.passage_nodes:
chunk_ids = self._rank_chunks_with_passages(
store, source_id, subgraph, seeds, query_embedding
)
else:
node_scores = self._ppr_scores(subgraph, seeds)
if not node_scores:
return []
chunk_ids = self._rank_chunks(store, source_id, node_scores)
if not chunk_ids:
return []
chunk_ids = self._rank_chunks(store, source_id, node_scores)
chunk_data = store.get_chunk_texts(source_id, chunk_ids)
# ``(text, metadata)`` in rank order. Chunk ids stop being the currency
# here: a hit contributed by the vector ranking has no graph chunk id,
# and keying on ids is what made an earlier version of this fusion able
# only to reorder the graph's own candidates.
candidates: List[tuple] = []
for chunk_id in chunk_ids:
chunk = chunk_data.get(chunk_id)
text = chunk.get("text") if chunk else None
if text:
candidates.append((text, chunk.get("metadata")))
if options.blend_vector:
# The graph ranks by how much PPR mass landed on a chunk's
# entities, which says nothing about whether the chunk is about the
# question. Fusing with the source's own vector ranking keeps the
# graph's reach while letting plain relevance back in — including
# chunks the graph never surfaced, which is where most of the value
# is: no reordering can rescue a question whose answer the graph
# missed entirely.
vector_hits = self._vector_ranking(source_id, query_embedding)
if vector_hits:
metadata_by_text = {text: meta for text, meta in candidates}
for text, meta in vector_hits:
metadata_by_text.setdefault(text, meta)
fused = self._rrf_order(
[[t for t, _ in candidates], [t for t, _ in vector_hits]],
RRF_K,
)
candidates = [
(text, metadata_by_text.get(text))
for text in sorted(fused, key=lambda t: fused[t], reverse=True)
]
docs: List[Dict[str, Any]] = []
token_budget = max(int(self.doc_token_limit * 0.9), 100)
cumulative_tokens = 0
source_top_k = self._source_top_k(source_id)
for chunk_id in chunk_ids:
for text, metadata in candidates:
if len(docs) >= source_top_k:
break
chunk = chunk_data.get(chunk_id)
text = chunk.get("text") if chunk else None
if not text:
continue
labels = labels_from_metadata(chunk.get("metadata"), text, source_id)
labels = labels_from_metadata(metadata, text, source_id)
doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}")
if cumulative_tokens + doc_tokens >= token_budget:
break
+23 -1
View File
@@ -12,7 +12,7 @@ reproduces today's chunking byte-for-byte.
from __future__ import annotations
from typing import Optional
from typing import Literal, Optional
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
@@ -90,6 +90,27 @@ class ChunkingConfig(BaseModel):
duplicate_headers: bool = False
class GraphRetrievalConfig(BaseModel):
"""How the graph retriever walks a graphrag source (live; no re-ingest).
The defaults are the configuration that measured best across the corpora
tested rather than a neutral starting point: seed from entity matches, put
the passages in the walk, and blend with the source's own vector ranking.
"""
model_config = ConfigDict(extra="forbid")
# Where the walk starts: entities whose descriptions match the question, or
# relationships ("A streams_to B") that do. Relationships can start the
# walk on an entity the question never names.
seed_strategy: Literal["entities", "relationships"] = "entities"
# Chunks join the walk as nodes, so a passage is reachable both by being
# about the question and by being connected to what is.
passage_nodes: bool = True
# Fuse the graph ranking with plain vector search by reciprocal rank.
blend_vector: bool = True
class RetrievalConfig(BaseModel):
"""Query-time retrieval knobs (live; no re-ingest needed)."""
@@ -102,6 +123,7 @@ class RetrievalConfig(BaseModel):
rephrase_query: bool = True # toggle ClassicRAG._rephrase_query side-call
reranker: Optional[dict] = None # reserved: future cross-encoder/LLM reorder
prescreen: Optional[dict] = None # None = off; else PreScreenConfig dict (D12)
graph: GraphRetrievalConfig = GraphRetrievalConfig() # graphrag retriever only
@field_validator("chunks")
@classmethod
+213
View File
@@ -138,6 +138,129 @@ def _extraction_json(entities, relationships):
return json.dumps({"entities": entities, "relationships": relationships})
class TestFactText:
"""A relationship rendered as the sentence it asserts.
This is what fact seeding matches a question against, so it has to read as
a claim rather than as three fields concatenated.
"""
def test_renders_the_relationship_as_a_sentence(self):
text = extraction_module._fact_text(
{
"source": "Alder",
"target": "Quill",
"type": "streams_to",
"description": "Alder streams audit events to Quill.",
}
)
assert text == "Alder streams_to Quill: Alder streams audit events to Quill."
def test_omits_an_absent_description(self):
text = extraction_module._fact_text(
{"source": "Alder", "target": "Quill", "type": "streams_to"}
)
assert text == "Alder streams_to Quill"
def test_defaults_a_missing_relation(self):
text = extraction_module._fact_text({"source": "Alder", "target": "Quill"})
assert text == "Alder related to Quill"
@pytest.mark.parametrize(
"rel",
[
{"source": "Alder", "target": ""},
{"source": "", "target": "Quill"},
{},
],
)
def test_an_edge_without_both_endpoints_has_no_fact(self, rel):
assert extraction_module._fact_text(rel) == ""
class TestEmbedFacts:
"""Fact embeddings are always recorded, so a source can switch to
relationship seeding at query time without being rebuilt."""
def _relationships(self):
return [{"source": "Alder", "target": "Quill", "type": "streams_to"}]
def test_attaches_one_embedding_per_fact_in_a_single_call(self):
relationships = self._relationships() + [{"source": "", "target": "Nowhere"}]
calls = []
class _Embedding:
def embed_documents(self, texts):
calls.append(texts)
return [[0.5] * 4 for _ in texts]
extraction_module._embed_facts(_Embedding(), relationships)
# One batched call, and the endpoint-less relationship is skipped
# rather than embedded as an empty string.
assert calls == [["Alder streams_to Quill"]]
assert relationships[0]["fact_embedding"] == [0.5] * 4
assert "fact_embedding" not in relationships[1]
def test_survives_an_embedding_failure(self):
"""The graph is still correct without fact embeddings — only
relationship seeding degrades, and it falls back to entities — so a
failure here must not fail the chunk."""
relationships = self._relationships()
class _Embedding:
def embed_documents(self, texts):
raise RuntimeError("embeddings down")
extraction_module._embed_facts(_Embedding(), relationships)
assert "fact_embedding" not in relationships[0]
class TestSeedText:
"""What a node's embedding is computed from.
Retrieval matches a whole question against these embeddings, so what goes
into them decides what the graph walk can start from.
"""
def _entity(self):
return {
"name": "Quill",
"normalized_name": "quill",
"type": "store",
"description": "A write-ahead store.",
}
def test_includes_type_and_description(self):
assert (
extraction_module._seed_text(self._entity())
== "Quill (store): A write-ahead store."
)
def test_falls_back_to_the_name_when_fields_are_missing(self):
assert extraction_module._seed_text({"name": "Quill"}) == "Quill"
def test_embedded_text_is_keyed_by_the_normalized_name(self):
"""The richer text must reach ``embed_documents``, keyed by the same
normalized name the store resolves nodes by — otherwise the embedding
is computed for a node it never reaches."""
captured = {}
class _Embedding:
def embed_documents(self, texts):
captured["texts"] = texts
return [[0.0] * 4 for _ in texts]
result = extraction_module._embed_names(_Embedding(), [self._entity()], [])
assert captured["texts"] == ["Quill (store): A write-ahead store."]
assert set(result) == {"quill"}
@pytest.mark.integration
class TestExtractionLive:
@pytest.fixture
@@ -193,6 +316,96 @@ class TestExtractionLive:
finally:
store.delete_by_source(source_id)
def test_parallel_workers_process_every_chunk_once(
self, store, source_id, monkeypatch, stub_embedding
):
"""Running the model calls concurrently must not change what gets written.
Extraction spends nearly all of a chunk's time waiting on the model, so
the calls run in a pool while every graph write stays on the calling
thread. Six chunks share one entity here: whatever order the pool
finishes in, that entity is upserted once, each chunk is linked, and all
six are marked processed.
"""
from docsgpt.core.settings import settings
try:
payload = _extraction_json(
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
llm = _StubLLM([payload] * 6)
_install_stub_llm(monkeypatch, llm)
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
summary = extract_graph_for_source(
source_id,
user="owner-1",
chunks=[
_chunk(f"c{i}", f"Ada appears here, take {i}.") for i in range(6)
],
config=SourceConfig(),
request_id="req-parallel",
)
assert summary["chunks_processed"] == 6
assert summary["failed_chunks"] == 0
assert summary["nodes"] == 1
assert len(llm.gen_calls) == 6
node = store.get_node_by_normalized(source_id, "ada")
assert node is not None
mapping = store.get_chunk_ids_for_nodes(source_id, [node["id"]])
assert sorted(mapping[node["id"]]) == [f"c{i}" for i in range(6)]
finally:
store.delete_by_source(source_id)
def test_embedding_runs_on_the_calling_thread(
self, store, source_id, monkeypatch, stub_embedding
):
"""Only the LLM call may run in the extraction pool, never embedding.
Inside a Celery worker the embeddings client decides to embed locally
from the task on the *current thread's* stack. A pool thread has none,
so from there it dispatches an embed task to the worker and waits on
it — which Celery refuses inside a task, so every chunk of a graph
build failed.
"""
import threading
from docsgpt.core.settings import settings
caller = threading.current_thread()
seen = []
real_embed_names = extraction_module._embed_names
def _recording_embed_names(*args, **kwargs):
seen.append(threading.current_thread())
return real_embed_names(*args, **kwargs)
monkeypatch.setattr(extraction_module, "_embed_names", _recording_embed_names)
try:
payload = _extraction_json(
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
_install_stub_llm(monkeypatch, _StubLLM([payload] * 4))
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
summary = extract_graph_for_source(
source_id,
user="owner-1",
chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(4)],
config=SourceConfig(),
request_id="req-thread",
)
assert summary["failed_chunks"] == 0
assert len(seen) == 4
assert all(thread is caller for thread in seen)
finally:
store.delete_by_source(source_id)
def test_same_entity_across_chunks_merges(
self, store, source_id, monkeypatch, stub_embedding
):
@@ -0,0 +1,109 @@
"""The graph retriever's default path, end to end through ``_graph_docs_for_source``.
The shipped defaults — seed from entities, walk the passages, blend with vector
search — are the configuration that measured best, so they are what most graph
sources run. This drives that whole path with a store that returns real values,
and checks each per-source option actually switches its stage off.
"""
from __future__ import annotations
from docsgpt.retriever.graph_rag import GraphRAGRetriever
from docsgpt.storage.db.source_config import RetrievalConfig
TEXTS = {
"c-alder": "Alder streams audit events to Quill.",
"c-quill": "Quill is compacted every six hours.",
}
VECTOR_ONLY = "A passage only plain vector search found."
class _Store:
"""A two-entity chain: the question matches Alder, the answer is on Quill."""
def __init__(self):
self.calls: list[str] = []
def search_nodes_by_embedding(self, source_id, query_embedding, k=10):
return [{"id": "alder", "name": "Alder", "distance": 0.1}]
def get_subgraph(self, source_id, node_ids, hops=1):
return {
"nodes": [{"id": "alder", "doc_freq": 1}, {"id": "quill", "doc_freq": 1}],
"edges": [{"src_node_id": "alder", "dst_node_id": "quill", "weight": 1.0}],
}
def get_chunk_ids_for_nodes(self, source_id, node_ids):
return {"alder": ["c-alder"], "quill": ["c-quill"]}
def chunk_similarities(self, source_id, chunk_ids, query_embedding):
self.calls.append("chunk_similarities")
return {"c-alder": 0.9, "c-quill": 0.2}
def get_chunk_texts(self, source_id, chunk_ids):
return {
c: {"text": TEXTS[c], "metadata": {"title": c}}
for c in chunk_ids
if c in TEXTS
}
def _retriever(per_source=None):
"""A retriever without its constructor (which builds a ClassicRAG)."""
retriever = object.__new__(GraphRAGRetriever)
retriever.chunks = 3
retriever.base_chunks = None
retriever.doc_token_limit = 50000
retriever.vectorstores = ["src"]
retriever.per_source_retrieval = per_source or {}
retriever.vector_calls = 0
def _vector_ranking(source_id, query_embedding):
retriever.vector_calls += 1
return [(VECTOR_ONLY, {"title": "vector"})]
retriever._vector_ranking = _vector_ranking
return retriever
def _texts(docs):
return [doc["text"] for doc in docs]
class TestDefaultPath:
def test_walks_passages_and_blends_in_vector_hits(self):
store = _Store()
retriever = _retriever()
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
# The answer sits one edge away from the seed: the walk reached it.
assert TEXTS["c-quill"] in _texts(docs)
# A hit only vector search found is blended in, not lost.
assert VECTOR_ONLY in _texts(docs)
assert store.calls == ["chunk_similarities"]
assert retriever.vector_calls == 1
class TestPerSourceOptions:
def test_passage_walk_can_be_switched_off(self):
store = _Store()
retriever = _retriever(
{"src": RetrievalConfig(chunks=3, graph={"passage_nodes": False})}
)
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
assert "chunk_similarities" not in store.calls
assert TEXTS["c-quill"] in _texts(docs)
def test_vector_blending_can_be_switched_off(self):
store = _Store()
retriever = _retriever(
{"src": RetrievalConfig(chunks=3, graph={"blend_vector": False})}
)
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
assert retriever.vector_calls == 0
assert VECTOR_ONLY not in _texts(docs)
+132
View File
@@ -0,0 +1,132 @@
"""Chunks as nodes in the walk, and the damping that decides how far mass spreads.
``_rank_chunks`` reads a chunk's score off its entities by summing their PPR
mass, which rewards a chunk for touching *many* entities rather than the right
ones. The passage-node path puts the chunks in the graph instead, so a chunk is
reachable both by being about the question and by being connected to what is.
These tests use a stub store: the ranking is graph arithmetic, and pinning it
against a real database would measure Postgres rather than the ranking.
"""
from __future__ import annotations
import pytest
from docsgpt.retriever.graph_rag import GraphRAGRetriever, _damping
class _StubStore:
"""The two reads the passage path makes, and nothing else."""
def __init__(self, chunk_links, similarities):
self._chunk_links = chunk_links
self._similarities = similarities
def get_chunk_ids_for_nodes(self, source_id, node_ids):
return {n: c for n, c in self._chunk_links.items() if n in set(node_ids)}
def chunk_similarities(self, source_id, chunk_ids, query_embedding):
return {c: self._similarities.get(c, 0.0) for c in chunk_ids}
def _retriever(chunks=2):
"""A retriever without its constructor — which builds a ClassicRAG, opens
settings-driven collaborators, and has nothing to do with ranking."""
retriever = object.__new__(GraphRAGRetriever)
retriever.chunks = chunks
return retriever
def _subgraph():
return {
"nodes": [
{"id": "a", "doc_freq": 1},
{"id": "b", "doc_freq": 1},
{"id": "hub", "doc_freq": 40},
],
"edges": [
{"src_node_id": "a", "dst_node_id": "hub", "weight": 1.0},
{"src_node_id": "b", "dst_node_id": "hub", "weight": 1.0},
],
}
class TestDamping:
"""Each ranking mode runs at the damping it was measured at."""
def test_passage_walk_keeps_mass_near_the_seeds(self):
assert _damping(passage_nodes=True) == 0.5
def test_entity_only_ranking_keeps_the_conventional_value(self):
assert _damping(passage_nodes=False) == 0.85
class TestPassageNodes:
def test_ranks_the_chunk_the_question_matches(self, monkeypatch):
"""Two chunks are equally connected; only their own relevance differs,
so the more relevant one must win."""
store = _StubStore(
chunk_links={"a": ["c1"], "b": ["c2"]},
similarities={"c1": 0.1, "c2": 0.9},
)
ranked = _retriever()._rank_chunks_with_passages(
store, "src", _subgraph(), {"a": 1.0, "b": 1.0}, [0.0] * 4
)
assert ranked[0] == "c2"
def test_a_chunk_reached_only_through_the_graph_still_ranks(self, monkeypatch):
"""The point of the walk: a chunk with no similarity of its own is
still reachable through the entity the seeds point at."""
store = _StubStore(
chunk_links={"a": ["c1"], "b": ["c2"]},
similarities={"c1": 0.0, "c2": 0.0},
)
ranked = _retriever()._rank_chunks_with_passages(
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
)
assert set(ranked) == {"c1", "c2"}
def test_no_linked_chunks_returns_nothing(self):
store = _StubStore(chunk_links={}, similarities={})
assert (
_retriever()._rank_chunks_with_passages(
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
)
== []
)
def test_over_fetches_past_the_chunk_budget(self, monkeypatch):
"""Same contract as ``_rank_chunks``: candidates exceed the budget so
chunks with missing text cannot drop the final count below it."""
links = {"a": [f"c{i}" for i in range(10)]}
store = _StubStore(
chunk_links=links,
similarities={f"c{i}": i / 10 for i in range(10)},
)
ranked = _retriever(chunks=2)._rank_chunks_with_passages(
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
)
assert len(ranked) == max(2 * 2, 2 + 5)
class TestChunkSimilaritiesGuard:
"""The store call the passage path depends on short-circuits before it
touches a connection, so an empty subgraph costs no query."""
@pytest.mark.parametrize(
"chunk_ids,embedding", [([], [0.1]), (["c1"], []), ([], [])]
)
def test_empty_inputs_return_empty(self, chunk_ids, embedding):
from docsgpt.graphrag.store import GraphStore
store = object.__new__(GraphStore)
assert store.chunk_similarities("src", chunk_ids, embedding) == {}
+129
View File
@@ -0,0 +1,129 @@
"""Where the graph walk starts, per the source's graph retrieval options.
Seeding decides more than ranking does — a walk that starts on the wrong nodes
cannot be rescued downstream. The options are per source and live (no
re-ingest), carried on the per-source retrieval config the Dispatcher hands the
retriever, so both the dispatch and how the options are resolved are pinned.
The fallback matters most: relationship seeding reads fact embeddings written
at ingest, and a source built before they were recorded has none. It must keep
retrieving through entity matching rather than returning nothing.
"""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from docsgpt.retriever.graph_rag import GraphRAGRetriever
from docsgpt.storage.db.source_config import GraphRetrievalConfig, RetrievalConfig
ENTITY_ROWS = [{"id": "n1", "name": "Quill", "distance": 0.2}]
FACT_ROWS = [
{"id": "f1", "name": "Alder", "distance": 0.1},
{"id": "f2", "name": "Quill", "distance": 0.1},
]
class _StubStore:
def __init__(self, fact_rows=None):
self.fact_rows = list(FACT_ROWS) if fact_rows is None else fact_rows
self.calls: list[str] = []
def seed_nodes_from_facts(self, source_id, query_embedding, fact_limit=5, limit=10):
self.calls.append("facts")
return self.fact_rows
def search_nodes_by_embedding(self, source_id, query_embedding, k=10):
self.calls.append("entities")
return list(ENTITY_ROWS)
def _retriever(per_source=None):
"""A retriever without its constructor, which builds a ClassicRAG."""
retriever = object.__new__(GraphRAGRetriever)
if per_source is not None:
retriever.per_source_retrieval = per_source
return retriever
def _relationships_config():
return RetrievalConfig(graph={"seed_strategy": "relationships"})
class TestDefaults:
def test_measured_best_configuration_is_the_default(self):
options = GraphRetrievalConfig()
assert options.seed_strategy == "entities"
assert options.passage_nodes is True
assert options.blend_vector is True
def test_a_source_with_no_per_source_config_gets_the_defaults(self):
assert _retriever()._graph_options("src") == GraphRetrievalConfig()
def test_existing_retrieval_configs_validate_without_the_new_block(self):
"""Source configs saved before this existed carry no ``graph`` key."""
config = RetrievalConfig.model_validate({"retriever": "graphrag"})
assert config.graph == GraphRetrievalConfig()
@pytest.mark.parametrize(
"bad", [{"seed_strategy": "vector"}, {"seed_strategy": "union"}, {"damping": 0.5}]
)
def test_retired_and_unknown_options_are_rejected(self, bad):
with pytest.raises(ValidationError):
GraphRetrievalConfig.model_validate(bad)
class TestEntitySeeding:
def test_seeds_from_entities_and_never_reads_facts(self):
store = _StubStore()
rows = _retriever()._seed_rows(store, "src", [0.0, 0.1])
assert [r["id"] for r in rows] == ["n1"]
assert store.calls == ["entities"]
class TestRelationshipSeeding:
def test_seeds_from_facts_when_the_source_asks_for_it(self):
store = _StubStore()
retriever = _retriever({"src": _relationships_config()})
rows = retriever._seed_rows(store, "src", [0.0, 0.1])
assert [r["id"] for r in rows] == ["f1", "f2"]
assert store.calls == ["facts"]
def test_reads_the_option_from_a_plain_dict_config_too(self):
store = _StubStore()
retriever = _retriever({"src": {"graph": {"seed_strategy": "relationships"}}})
retriever._seed_rows(store, "src", [0.0, 0.1])
assert store.calls == ["facts"]
def test_falls_back_to_entities_without_fact_embeddings(self):
"""A source built before fact embeddings were recorded still retrieves."""
store = _StubStore(fact_rows=[])
retriever = _retriever({"src": _relationships_config()})
rows = retriever._seed_rows(store, "src", [0.0, 0.1])
assert [r["id"] for r in rows] == ["n1"]
assert store.calls == ["facts", "entities"]
def test_options_are_per_source(self):
store = _StubStore()
retriever = _retriever({"other": _relationships_config()})
retriever._seed_rows(store, "src", [0.0, 0.1])
assert store.calls == ["entities"]
def test_a_malformed_stored_option_falls_back_to_the_defaults(self):
"""A bad value must not take graph retrieval down with it."""
retriever = _retriever({"src": {"graph": {"seed_strategy": "nonsense"}}})
assert retriever._graph_options("src") == GraphRetrievalConfig()
+126 -5
View File
@@ -157,6 +157,122 @@ class TestGraphStoreLive:
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."""
@@ -203,18 +319,23 @@ class TestGraphStoreLive:
store.delete_by_source(source_id)
def test_self_loop_degree_agrees_across_paths(self, store, source_id):
"""``add_edge``'s incremental +1 and ``set_node_degrees`` recompute must
agree on a self-loop (count it once)."""
"""``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")
store.add_edge(source_id, node, node, "self")
assert store.add_edge(source_id, node, node, "self") is None
incremental = store.get_node_by_normalized(source_id, "solo")["degree"]
assert incremental == 1
assert incremental == 0
store.set_node_degrees(source_id)
recomputed = store.get_node_by_normalized(source_id, "solo")["degree"]
assert recomputed == 1
assert recomputed == incremental == 0
finally:
store.delete_by_source(source_id)
+23
View File
@@ -45,6 +45,29 @@ def _patch_embed(monkeypatch):
)
@pytest.fixture(autouse=True)
def _entity_only_ranking(monkeypatch):
"""Pin the ranking path these tests were written for.
Everything here exercises entity-only PPR ranking without vector blending,
driven through ``MagicMock`` stores. The shipped default now walks the
passages and blends with vector search — covered end to end in
``tests/graphrag/test_retriever_default_path.py`` with a store that returns
real values. Pinning keeps each test here asserting what it was written to
assert, rather than whatever a mock happens to return on a path it never set
up.
"""
from docsgpt.storage.db.source_config import GraphRetrievalConfig
monkeypatch.setattr(
GraphRAGRetriever,
"_graph_options",
lambda self, source_id: GraphRetrievalConfig(
passage_nodes=False, blend_vector=False
),
)
# ── Fallback to ClassicRAG ────────────────────────────────────────────────────