Merge pull request #2804 from arc53/fix/graphrag-extraction-and-retrieval

fix(graphrag): make graph retrieval and extraction work in a default install
This commit is contained in:
Alex authored and GitHub committed 2026-09-20 11:20:54 +01:00
commit 9ad09039e4
43 files changed
+4821 -219

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
+2
View File
@@ -2,6 +2,7 @@ import logging
from typing import Dict, Generator, Optional
from docsgpt.agents.base import BaseAgent
from docsgpt.agents.tools.graph_search import add_graph_search_tool
from docsgpt.agents.tools.internal_search import add_internal_search_tool
from docsgpt.agents.tools.wiki import add_wiki_tool
from docsgpt.logging import LogContext
@@ -33,6 +34,7 @@ class AgenticAgent(BaseAgent):
) -> Generator[Dict, None, None]:
tools_dict = self.tool_executor.get_tools()
add_internal_search_tool(tools_dict, self.retriever_config)
add_graph_search_tool(tools_dict, self.retriever_config)
if self.wiki_config:
add_wiki_tool(tools_dict, self.wiki_config)
self._prepare_tools(tools_dict)
+21 -7
View File
@@ -955,16 +955,30 @@ class BaseAgent(ABC):
)
self.retrieved_docs = scrubbed
def _collect_internal_sources(self) -> None:
"""Merge the cached InternalSearchTool's docs into ``retrieved_docs``,
deduped, preserving any pre-fetched docs so a mixed-exposure agent cites
both pre-fetched and tool-retrieved sources (not just the tool's)."""
def _search_tool_docs(self) -> List[Dict]:
"""Documents this run's search tools read: internal search and the graph tool.
Both record what they surface in ``retrieved_docs``; a page read from
the graph carries the answer as much as a search hit does, so both are
cited. Tools are looked up the way the executor caches them.
"""
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID
executor = getattr(self, "tool_executor", None)
loaded = getattr(executor, "_loaded_tools", None) or {}
tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}")
if not (tool and getattr(tool, "retrieved_docs", None)):
docs: List[Dict] = []
for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)):
tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}")
docs.extend(getattr(tool, "retrieved_docs", None) or [])
return docs
def _collect_internal_sources(self) -> None:
"""Merge the search tools' docs into ``retrieved_docs``, deduped,
preserving any pre-fetched docs so a mixed-exposure agent cites both
pre-fetched and tool-retrieved sources (not just the tools')."""
tool_docs = self._search_tool_docs()
if not tool_docs:
return
def _key(d):
@@ -974,7 +988,7 @@ class BaseAgent(ABC):
merged = list(self.retrieved_docs or [])
seen = {_key(d) for d in merged}
for doc in tool.retrieved_docs:
for doc in tool_docs:
k = _key(doc)
if k not in seen:
seen.add(k)
+2
View File
@@ -2,6 +2,7 @@ import logging
from typing import Dict, Generator, Optional
from docsgpt.agents.base import BaseAgent
from docsgpt.agents.tools.graph_search import add_graph_search_tool
from docsgpt.agents.tools.internal_search import add_internal_search_tool
from docsgpt.agents.tools.wiki import add_wiki_tool
from docsgpt.logging import LogContext
@@ -38,6 +39,7 @@ class ClassicAgent(BaseAgent):
tools_dict = self.tool_executor.get_tools()
if self.retriever_config:
add_internal_search_tool(tools_dict, self.retriever_config)
add_graph_search_tool(tools_dict, self.retriever_config)
if self.wiki_config:
add_wiki_tool(tools_dict, self.wiki_config)
self._prepare_tools(tools_dict)
+6 -10
View File
@@ -6,10 +6,8 @@ from typing import Dict, Generator, List, Optional
from docsgpt.agents.base import BaseAgent
from docsgpt.agents.tool_executor import ToolExecutor
from docsgpt.agents.tools.internal_search import (
INTERNAL_TOOL_ID,
add_internal_search_tool,
)
from docsgpt.agents.tools.graph_search import add_graph_search_tool
from docsgpt.agents.tools.internal_search import add_internal_search_tool
from docsgpt.agents.tools.wiki import add_wiki_tool
from docsgpt.agents.tools.think import THINK_TOOL_ENTRY, THINK_TOOL_ID
from docsgpt.logging import LogContext
@@ -277,6 +275,7 @@ class ResearchAgent(BaseAgent):
tools_dict = self.tool_executor.get_tools()
add_internal_search_tool(tools_dict, self.retriever_config)
add_graph_search_tool(tools_dict, self.retriever_config)
if self.wiki_config:
add_wiki_tool(tools_dict, self.wiki_config)
@@ -620,12 +619,9 @@ class ResearchAgent(BaseAgent):
return messages, search_returned_empty
def _collect_step_sources(self):
"""Collect sources from InternalSearchTool and register with CitationManager."""
cache_key = f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}"
tool = self.tool_executor._loaded_tools.get(cache_key)
if tool and hasattr(tool, "retrieved_docs"):
for doc in tool.retrieved_docs:
self.citations.add(doc)
"""Register the search tools' docs (internal search and graph pages) with CitationManager."""
for doc in self._search_tool_docs():
self.citations.add(doc)
# ------------------------------------------------------------------
# Phase 3: Synthesis
+299
View File
@@ -0,0 +1,299 @@
"""Let the model search the knowledge graph itself, one edge at a time.
Graph retrieval normally runs as a ranker: seed a walk from the question,
diffuse mass over a subgraph, hand back the highest-scoring chunks. Measured
across five corpora that never beat plain vector search, because a question
whose answer lives two documents away has nothing in it for the seeding step to
match — the bridging entity is named in the *first* document, not the question.
Exposing the graph as tools removes the guess. The model can look up the
service, read which store it names, then fetch that store's page: the chain
followed deliberately rather than approximated by a diffusion. On a corpus built
so that vector search cannot shortcut the chain, this took two-hop answers from
1/8 to 8/8, against 0.40 for vector and 0.47 for one-shot graph retrieval.
It is not a general win, and is deliberately not a default. On ordinary prose
documentation it *lost* to plain vector search (0.50 against 0.90): it answers
well when a question names an entity and wanders when the question is a task
description. It also costs several model round-trips per answer instead of one.
So it is offered only where a source owner has already chosen search over
prefetch — the per-source exposure setting, or an agentic/research agent — and
suits content that is genuinely chain-structured: runbooks, service catalogues,
infrastructure inventories.
"""
from __future__ import annotations
import logging
from typing import Any, Dict, List, Optional
from docsgpt.agents.tools.base import Tool
from docsgpt.graphrag import graphrag_available
from docsgpt.retriever.labels import labels_from_metadata
logger = logging.getLogger(__name__)
GRAPH_TOOL_ID = "graph_search"
MAX_PAGE_CHARS = 1500
class GraphSearchTool(Tool):
"""Entity lookup, relationship traversal and page reads over a source's graph."""
internal = True
def __init__(self, config: Dict):
self.config = config or {}
self._store = None
self.retrieved_docs: List[Dict] = []
# -- plumbing ------------------------------------------------------------
def _sources(self) -> List[str]:
source = self.config.get("source") or {}
active = source.get("active_docs") or []
if isinstance(active, str):
active = [active]
return [str(s) for s in active if s]
def _get_store(self):
if self._store is None:
from docsgpt.graphrag.store import GraphStore
self._store = GraphStore()
return self._store
def _release_store(self) -> None:
"""Hand the pooled connection back at the end of an action.
The executor caches this tool for the whole agent run, so a store kept
between actions pins one connection of the shared pgvector pool across
every LLM round trip of that run -- minutes at a time, and enough
concurrent runs exhaust the pool. ``GraphRAGRetriever`` releases its
store before falling back for the same reason. Checking one back out
costs a pool acquire.
"""
store, self._store = self._store, None
if store is None:
return
try:
store.close()
except Exception as exc: # noqa: BLE001 -- releasing must not fail an action
logger.debug(f"Graph tool could not release its store: {exc}")
def _embed(self, text: str) -> Optional[List[float]]:
try:
from docsgpt.vectorstore.base import get_embeddings
return get_embeddings().embed_query(text)
except Exception as e: # noqa: BLE001
logger.error(f"Graph tool could not embed the query: {e}")
return None
# -- actions -------------------------------------------------------------
def execute_action(self, action_name: str, **kwargs):
# The graph lives in the pgvector store, so the flag alone is not
# enough: under another vector store there is no graph to read.
if not graphrag_available():
return "The knowledge graph is not enabled for this deployment."
if not self._sources():
return "No graph-backed sources are configured."
try:
if action_name == "search_entities":
return self._search_entities(**kwargs)
if action_name == "get_relationships":
return self._get_relationships(**kwargs)
if action_name == "read_entity_pages":
return self._read_entity_pages(**kwargs)
except Exception as e: # noqa: BLE001
logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True)
return "The graph lookup failed."
finally:
self._release_store()
return f"Unknown action: {action_name}"
def _search_entities(self, **kwargs) -> str:
query = str(kwargs.get("query") or "").strip()
if not query:
return "Error: 'query' parameter is required."
limit = max(1, min(int(kwargs.get("k") or 8), 25))
embedding = self._embed(query)
if embedding is None:
return "Entity search is unavailable."
store = self._get_store()
lines: List[str] = []
for source_id in self._sources():
for row in store.search_nodes_by_embedding(source_id, embedding, k=limit):
similarity = 1.0 - float(row.get("distance") or 0.0)
description = (row.get("description") or "").strip()
suffix = f" — {description[:160]}" if description else ""
lines.append(f"- {row['name']} (match {similarity:.2f}){suffix}")
if not lines:
return f"No entities found for {query!r}."
return "Entities:\n" + "\n".join(lines[:limit])
def _get_relationships(self, **kwargs) -> str:
entity = str(kwargs.get("entity") or "").strip()
if not entity:
return "Error: 'entity' parameter is required."
store = self._get_store()
lines: List[str] = []
for source_id in self._sources():
for edge in store.entity_relationships(source_id, entity):
relation = edge.get("type") or "related to"
lines.append(f"- {edge['source']} --{relation}--> {edge['target']}")
if not lines:
return (
f"No relationships found for {entity!r}. Try search_entities first "
"to get the exact name used in the graph."
)
return f"Relationships for {entity!r}:\n" + "\n".join(lines)
def _read_entity_pages(self, **kwargs) -> str:
entity = str(kwargs.get("entity") or "").strip()
if not entity:
return "Error: 'entity' parameter is required."
store = self._get_store()
parts: List[str] = []
for source_id in self._sources():
for page in store.entity_pages(source_id, entity):
text = page.get("text") or ""
# The retrievers' own labelling: a page read here and the same
# chunk retrieved by internal_search are one document, and
# citations key on (source, title). Labelling it differently
# gives that document two citation numbers.
labels = labels_from_metadata(page.get("metadata"), text, source_id)
doc = {**labels, "text": text}
if doc not in self.retrieved_docs:
self.retrieved_docs.append(doc)
header = labels["filename"] or labels["title"]
parts.append(f"--- {header} ---\n{text[:MAX_PAGE_CHARS]}")
if not parts:
return f"No documents mention {entity!r}."
return "\n\n".join(parts)
# -- metadata ------------------------------------------------------------
def get_actions_metadata(self):
return [
{
"name": "search_entities",
"description": (
"Find named things in the knowledge graph — services, components, "
"settings, people — whose names resemble a query. Use this first to "
"learn the exact name the graph uses before asking for its "
"relationships."
),
"parameters": {
"properties": {
"query": {
"type": "string",
"description": "What to look for, e.g. a service or component name.",
"filled_by_llm": True,
"required": True,
},
"k": {
"type": "integer",
"description": "How many entities to return (default 8).",
"filled_by_llm": True,
"required": False,
},
}
},
},
{
"name": "get_relationships",
"description": (
"List what an entity is connected to, as 'source --relation--> target'. "
"This is how you answer a question about something the question does not "
"name: look up what it points at, then read that thing's pages."
),
"parameters": {
"properties": {
"entity": {
"type": "string",
"description": "Exact entity name, as returned by search_entities.",
"filled_by_llm": True,
"required": True,
}
}
},
},
{
"name": "read_entity_pages",
"description": (
"Read the documentation an entity appears in, the page it is about "
"first. Use this once you know which entity holds the answer."
),
"parameters": {
"properties": {
"entity": {
"type": "string",
"description": "Exact entity name, as returned by search_entities.",
"filled_by_llm": True,
"required": True,
}
}
},
},
]
def get_config_requirements(self):
return {}
def build_graph_tool_entry() -> Dict:
"""The synthetic ``tools_dict`` entry for the graph tool."""
tool = GraphSearchTool({})
actions = []
for action in tool.get_actions_metadata():
entry = dict(action)
entry["active"] = True
actions.append(entry)
return {"name": "graph_search", "actions": actions}
def sources_have_graph(source: Dict) -> bool:
"""Whether any active source actually has a graph to search."""
active = source.get("active_docs") or []
if isinstance(active, str):
active = [active]
if not active:
return False
try:
from docsgpt.graphrag.store import GraphStore
counts = GraphStore().count_nodes_many([str(a) for a in active])
return any(count > 0 for count in counts.values())
except Exception as e: # noqa: BLE001
logger.debug(f"Could not check for graphs: {e}")
return False
def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None:
"""Add the graph tool when the agent's search-tool sources include a graph.
No setting of its own: ``retriever_config`` already carries exactly the
sources the agent may *search* — the ones a source owner exposed as a
search tool, or every source for an agentic/research agent — so the graph
tool follows that same per-source exposure choice. A graph source left at
``prefetch`` in a classic agent is used for ranking only.
"""
if not graphrag_available():
return
source = retriever_config.get("source") or {}
if not source.get("active_docs") or not sources_have_graph(source):
return
entry = build_graph_tool_entry()
# The executor resolves tools by ``id``; this one is synthetic (no DB row).
entry["id"] = GRAPH_TOOL_ID
entry["config"] = {"source": source}
tools_dict[GRAPH_TOOL_ID] = entry
def build_graph_tool_config(source: Dict, **_ignored: Any) -> Dict:
"""Config for :class:`GraphSearchTool` — it only needs the source ids."""
return {"source": source}
+4 -5
View File
@@ -17,8 +17,6 @@ import signal
import threading
from typing import Any, Callable, Dict, List, Optional
from celery import current_task
from docsgpt.agents.tools.artifact_ref import resolve_artifact_id
from docsgpt.agents.tools.attachment_bridge import (
AttachmentBridgeError,
@@ -26,6 +24,7 @@ from docsgpt.agents.tools.attachment_bridge import (
match_attachment,
)
from docsgpt.agents.tools.base import Tool
from docsgpt.celery_init import in_worker
from docsgpt.core.json_schema_utils import (
JsonSchemaValidationError,
normalize_json_schema_payload,
@@ -229,9 +228,9 @@ class ReadDocumentTool(Tool):
# (floored at DOCUMENT_PARSE_TIMEOUT).
timeout = parse_timeout_for_size(self._input_size)
# ``current_task`` is a Celery proxy: truthy only while this runs inside a worker task,
# falsy in the web process (the bare proxy is NOT identity-None, so test truthiness).
if current_task:
# Process-wide, not the thread-local ``current_task``: a thread a task starts has no
# task of its own, and dispatching from there is the self-deadlock described above.
if in_worker():
from docsgpt.worker import run_parse_document
try:
+26 -4
View File
@@ -9,6 +9,8 @@ import threading
import uuid
from typing import Any, Callable, Optional
from celery.exceptions import Ignore, MaxRetriesExceededError
from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository
from docsgpt.storage.db.session import db_readonly, db_session
@@ -81,10 +83,30 @@ def with_idempotency(
"idempotency: live lease held; deferring task=%s key=%s",
task_name, key,
)
raise self.retry(
countdown=LEASE_TTL_SECONDS,
max_retries=LEASE_RETRY_MAX,
)
try:
raise self.retry(
countdown=LEASE_TTL_SECONDS,
max_retries=LEASE_RETRY_MAX,
)
except MaxRetriesExceededError:
# The holder is simply slower than LEASE_RETRY_MAX
# deferrals: a task that outruns the broker's visibility
# timeout is redelivered while its first run is still
# going. Letting the exhaustion propagate would report a
# failure for a task that is running normally — but so
# would returning a value, only less visibly. A redelivery
# reuses the original task id (``Context`` carries
# ``task_id`` into the retry), so a return marks the very
# id the client polls SUCCESS, and ``/api/task_status``
# hands that to the UI as a finished build. ``Ignore``
# records no state at all, leaving the outcome to the run
# that actually holds the lease.
logger.info(
"idempotency: lease still held after %s deferrals; "
"leaving task=%s key=%s to its holder",
LEASE_RETRY_MAX, task_name, key,
)
raise Ignore() from None
if attempt > MAX_TASK_ATTEMPTS:
logger.error(
+46
View File
@@ -13,6 +13,7 @@ from celery.signals import (
setup_logging,
task_postrun,
task_prerun,
worker_init,
worker_process_init,
worker_ready,
)
@@ -172,6 +173,51 @@ def _run_version_check(*args, **kwargs):
celery = make_celery()
celery.config_from_object("docsgpt.celeryconfig")
#: Set once this process starts as a worker; see :func:`_mark_worker_process`.
_IS_WORKER_PROCESS = False
@worker_init.connect
@worker_process_init.connect
def _mark_worker_process(*args, **kwargs):
"""Record that this process runs tasks, for :func:`in_worker`.
``worker_init`` fires in every worker's main process before its pool
starts: that is where solo, threads, eventlet and gevent run tasks, and
what prefork children fork from. ``worker_process_init`` covers prefork
children however they were started.
"""
global _IS_WORKER_PROCESS
_IS_WORKER_PROCESS = True
def in_worker() -> bool:
"""True anywhere in a Celery worker process, on any thread or greenlet.
``current_worker_task`` alone is not enough: Celery records the executing
task on the thread (or greenlet) that runs it, so one the task starts sees
none and would take the web-process branch — dispatching to the worker it
is running in and blocking on the result. Celery refuses that ``get()``
("Never call result.get() within a task!"), or, where joins are allowed,
it waits on a queue that only this busy process may be able to serve.
The worker's own startup (:func:`_mark_worker_process`) answers for every
pool. ``task_join_will_block`` — process-wide, set for every blocking pool
— and the task's own ``current_worker_task`` still count for a process
that runs tasks without having gone through that startup.
Returns:
bool: Whether this call is running inside a worker process.
"""
from celery.result import task_join_will_block
return (
_IS_WORKER_PROCESS
or task_join_will_block()
or celery.current_worker_task is not None
)
#: Task-name prefix the package carried before the rename to ``docsgpt``.
+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
+261 -38
View File
@@ -23,6 +23,11 @@ import logging
import re
from typing import Any, Callable, Dict, List, Optional
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
@@ -63,14 +68,37 @@ def _resolve_max_chunks(config: SourceConfig) -> int:
return config.graph.max_chunks or settings.GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION
def _resolve_extraction_provider(
model_id: Optional[str], user: Optional[str]
) -> str:
"""The provider that serves ``model_id``, else the deployment default.
``settings.LLM_PROVIDER`` is only a default (``docsgpt``, the hosted public
endpoint, out of the box). Dispatching the resolved extraction model
through it sends the request to a provider that does not serve that model:
the call is rejected, the shared fallback answers instead, and the graph is
built by a different model than the one configured — with nothing in the
summary to say so. ``user`` scopes the lookup so a per-user (BYOM) model id
resolves as well.
"""
provider = (
get_provider_from_model_id(model_id, user_id=user) if model_id else None
)
return provider or settings.LLM_PROVIDER
def _build_extraction_llm(
model_id: Optional[str], user: Optional[str], request_id: Optional[str]
):
"""Build the extraction LLM tagged for token-usage attribution to the owner."""
decoded_token = {"sub": user} if user else None
provider = _resolve_extraction_provider(model_id, user)
logger.info(
"Graph extraction dispatching model=%s via provider=%s", model_id, provider
)
llm = LLMCreator.create_llm(
settings.LLM_PROVIDER,
api_key=settings.API_KEY,
provider,
api_key=get_api_key_for_provider(provider),
user_api_key=None,
decoded_token=decoded_token,
model_id=model_id,
@@ -122,8 +150,15 @@ def _parse_extraction(raw: Any) -> Optional[Dict[str, List[Dict[str, Any]]]]:
}
def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]:
"""Run exactly one extraction call for a chunk (gleanings off)."""
def _extract_chunk(
llm, text: str, chunk_id: Optional[str] = None
) -> Optional[Dict[str, List[Dict[str, Any]]]]:
"""Run exactly one extraction call for a chunk (gleanings off).
Both failure modes name the chunk: an unparseable response used to return
``None`` silently, so a graph could come back short with nothing in the
logs to say which chunk was dropped or why.
"""
messages = [
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": f"<chunk>\n{text}\n</chunk>"},
@@ -134,9 +169,17 @@ def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]:
messages=messages,
)
except Exception as exc:
logger.warning("Graph extraction call failed, skipping chunk: %s", exc)
logger.warning(
"Graph extraction call failed for chunk %s: %s", chunk_id, exc
)
return None
return _parse_extraction(response)
parsed = _parse_extraction(response)
if parsed is None:
logger.warning(
"Graph extraction returned unparseable output for chunk %s.",
chunk_id,
)
return parsed
def _coerce_weight(value: Any) -> float:
@@ -160,8 +203,10 @@ def extract_graph_for_source(
Resumable and idempotent: chunks already marked ``done`` are skipped via the
``graph_ingest_progress`` checkpoint, so a retry never re-extracts (and never
re-bills). Processes at most the resolved chunk cap; excess chunks are
reported under ``skipped_over_cap``. A malformed response or an LLM error on
a single chunk marks it ``failed`` and continues — the pipeline never crashes.
reported under ``skipped_over_cap``. A malformed response, an LLM error or a
failed write on a single chunk is retried once after the rest of the build;
a chunk that fails again is marked ``failed`` and the run continues — the
pipeline never crashes.
Each chunk is written in a single transaction with one batched embedding
call (entity + relationship-endpoint names together).
@@ -178,8 +223,13 @@ def extract_graph_for_source(
Returns:
A summary ``{nodes, edges, chunks_processed, skipped_over_cap,
failed_chunks}``.
failed_chunks}``, where ``nodes`` is how many distinct nodes the
source's graph holds after the run — not how many upserts ran, which
counts the same entity once per chunk it appears in.
"""
import threading
from concurrent.futures import ThreadPoolExecutor
from docsgpt.graphrag.store import GraphStore
store = GraphStore()
@@ -197,11 +247,26 @@ def extract_graph_for_source(
embedding = get_embeddings()
llm = _build_extraction_llm(
_resolve_extraction_model(config), user, request_id
)
model_id = _resolve_extraction_model(config)
# Built here first so a misconfigured model fails the run before any
# chunk is touched; this instance serves the calling thread.
thread_llm = threading.local()
thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
nodes = 0
def _llm():
"""This thread's extraction LLM.
Provider-reported usage is kept on the LLM instance (``_last_usage``)
and claimed by whichever call finishes next, so two calls in flight on
one instance can bill each other's tokens. Each pool thread therefore
builds its own.
"""
llm = getattr(thread_llm, "llm", None)
if llm is None:
llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
return llm
node_upserts = 0
edges = 0
chunks_processed = 0
failed_chunks = 0
@@ -215,54 +280,140 @@ def extract_graph_for_source(
{
"current": chunks_processed + failed_chunks,
"total": total,
"nodes": nodes,
"nodes": node_upserts,
"edges": edges,
}
)
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. Graph writes and embedding stay on the calling
thread, so transactions and the progress checkpoint are exactly what
they were serially and the pool never touches the embeddings client.
"""
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)
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"])
except Exception as exc:
logger.warning(
"Graph extraction failed for chunk %s: %s", chunk_id, exc
)
return chunk_id, "failed", None
return chunk_id, "ok", (entities, relationships)
def _write(chunk_id, status, payload) -> bool:
"""Apply one prepared chunk to the graph; False when it did not land."""
nonlocal node_upserts, edges, chunks_processed
if status == "empty":
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
return True
if status == "failed":
return False
entities, relationships = payload
try:
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
)
nodes += chunk_nodes
edges += chunk_edges
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
except Exception as exc:
logger.warning(
"Graph extraction write failed for chunk %s, skipping: %s",
chunk_id,
exc,
"Graph extraction embed/write failed for chunk %s: %s", chunk_id, exc
)
return False
# ``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.
node_upserts += chunk_nodes
edges += chunk_edges
chunks_processed += 1
return True
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)
def _pass(items):
"""Extract and write ``items``; return the ones that did not land."""
if pool is not None:
# ``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, items)
else:
prepared = (_prepare(item) for item in items)
missed = []
for item, (chunk_id, status, payload) in zip(items, prepared):
if not _write(chunk_id, status, payload):
missed.append(item)
_report()
return missed
try:
missed = _pass(to_process)
if missed:
# A failure is usually transient — a provider error, one response
# that did not parse — and the checkpoint only picks it up on a
# rerun nothing schedules. One more attempt, after the rest of the
# build so a burst of rate limiting has passed, and no more: a
# chunk that cannot be extracted costs at most two calls.
logger.info(
"Graph extraction retrying %d failed chunk(s) for source %s",
len(missed),
source_id,
)
missed = _pass(missed)
for _, chunk_id in missed:
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
failed_chunks = len(missed)
if missed:
logger.warning(
"Graph extraction gave up on %d chunk(s) for source %s after a retry: %s",
failed_chunks,
source_id,
", ".join(str(chunk_id) for _, chunk_id in missed),
)
_report()
finally:
if pool is not None:
pool.shutdown(wait=True)
try:
store.set_node_degrees(source_id)
except Exception as exc:
logger.warning("set_node_degrees failed for source %s: %s", source_id, exc)
# Upserts are writes, not nodes: one entity seen in ten chunks is ten
# upserts and a single node, so the old count overstated every graph whose
# entities recur. Report what the graph holds, falling back to the write
# count only if the count query itself fails.
# ``strict`` is what makes the fallback below reachable: the default
# count swallows query failures and answers 0, which would report a
# successful build as an empty graph.
nodes = node_upserts
try:
nodes = store.count_nodes(source_id, strict=True)
except Exception as exc:
logger.warning(
"count_nodes failed for source %s; reporting upserts instead: %s",
source_id,
exc,
)
return {
"nodes": nodes,
"edges": edges,
@@ -273,7 +424,7 @@ def extract_graph_for_source(
def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
"""Normalize the LLM's entity dicts (drop nameless ones)."""
"""Normalize the LLM's entity dicts (drop the ones with no usable name)."""
entities = []
for e in raw_entities:
if not isinstance(e, dict):
@@ -281,10 +432,16 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
name = str(e.get("name", "")).strip()
if not name:
continue
normalized_name = normalize_entity_name(name)
if not normalized_name:
# A punctuation-only name normalizes to nothing, and nodes merge on
# that key: keeping it collapses every such entity onto one shared
# node. The relationship side already drops them.
continue
entities.append(
{
"name": name,
"normalized_name": name.lower(),
"normalized_name": normalized_name,
"type": str(e.get("type") or "") or None,
"description": str(e.get("description") or "") or None,
}
@@ -310,6 +467,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]],
@@ -322,14 +543,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 {}
+115
View File
@@ -0,0 +1,115 @@
"""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.
The result is a merge key, never shown to anyone, so it only has to be the
same for a word's singular and plural — not to be a word itself.
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",
"alias", "canvas", "atlas", "bias", "pandas",
}
)
#: Plural endings that drop ``es``, and the singular endings that meet them.
#: ``caches`` cannot say whether it is ``cache`` + "s" or ``cach`` + "es"
#: (as ``batches`` is ``batch`` + "es"), so rather than guess, both
#: ``caches`` and ``cache`` fold to ``cach`` — as ``databases``/``database``
#: fold to ``databas``. Hardly any real word differs from one of these singulars
#: by its final "e" alone, so the fold merges next to nothing it should not.
_ES_PLURAL = ("ches", "shes", "ses", "zes", "xes")
_E_SINGULAR = ("che", "she", "se", "ze", "xe")
def _singular(word: str) -> str:
"""Fold one word so its singular and plural share a key, else leave it alone.
``-ies`` and a singular's ``-ie`` both fold to ``-y`` (``policies``,
``cookies``/``cookie``). The ``-es`` endings in :data:`_ES_PLURAL` drop
``es`` and the singular endings in :data:`_E_SINGULAR` drop their ``e``, so
both sides of an ambiguous plural meet (``caches``/``cache`` -> ``cach``).
Otherwise a bare trailing ``s`` is dropped 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(_ES_PLURAL):
return word[:-2]
if word.endswith("ie") and len(word) > 4:
return word[:-2] + "y"
if word.endswith(_E_SINGULAR):
return word[:-1]
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 key shared by a name's singular and
plural. 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``;
``cache`` and ``caches`` -> ``cach``;
``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)
+553 -87
View File
@@ -17,6 +17,8 @@ import logging
import uuid
from typing import Any, Dict, List, Optional
import psycopg
from psycopg import sql
from psycopg.types.json import Jsonb
from docsgpt.core.settings import settings
@@ -51,6 +53,18 @@ def _safe_identifier(name: str) -> str:
return name
def _identifier(name: str) -> sql.Identifier:
"""``name`` as a quoted identifier, folded the way Postgres folds it unquoted.
Composing identifiers through psycopg keeps every query a fixed statement
with bound values: nothing is formatted into the SQL string. The fold
matters because ``PGVectorStore`` writes these names unquoted, which
Postgres lower-cases, while a quoted identifier keeps its case; folding
first keeps both stores addressing the same table.
"""
return sql.Identifier(_safe_identifier(name).lower())
def _pgvector_identifiers() -> tuple[str, str, str, str]:
"""Resolve ``(table, text_col, metadata_col, source_col)`` from ``PGVectorStore``.
@@ -73,6 +87,51 @@ 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.
psycopg raises ``OperationalError`` ("the connection is lost") when the
socket dies under a statement and ``InterfaceError`` when the connection
object is already closed. Everything else — a bad statement, a constraint
violation — is a real failure that a retry would only repeat.
"""
return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError))
def _lock_source(cursor, source_id: str) -> None:
"""Serialize graph writes for one source until this transaction ends.
Writes within one build are already serial, but two builds of the same
source can overlap: a rebuild dispatched while the last one is still
running gets a new idempotency key, so its lease does not stop it. Without
this, both could pass a chunk's "done" check before either commits and
apply it twice. A transaction-scoped advisory lock keyed by the source
makes them take turns chunk by chunk; the lock is released on commit or
rollback, and a hash collision only makes two sources take turns.
"""
cursor.execute(
"SELECT pg_advisory_xact_lock(hashtext(%s));", (f"graphrag:source:{source_id}",)
)
def _safe_rollback(conn) -> None:
"""Roll back, tolerating a connection too broken to roll back."""
try:
conn.rollback()
except Exception as exc:
logging.debug("Rollback on a broken connection failed: %s", exc)
class GraphStore:
"""Stores and queries a per-source knowledge graph in the pgvector DB."""
@@ -138,6 +197,42 @@ class GraphStore:
self._pooled = False
return self._connection
def _write_with_reconnect(self, operation):
"""Run ``operation(conn)``, once more on a fresh connection if it was dead.
A graph build holds one checked-out connection for the length of the
whole extraction and spends minutes per chunk waiting on the model, so
the connection idles long enough for the server (or a pooler) to drop
it. The pool only validates a connection when it is handed out, and
this one was handed out at the start of the build, so the next write
raises and its chunk is lost from the graph. Every statement here is an
idempotent upsert, so replaying one on a new connection cannot
double-write.
Args:
operation: Callable taking the connection and doing one write.
Returns:
Whatever ``operation`` returns.
Raises:
Exception: Anything ``operation`` raises that is not connection
loss, and anything the single retry raises.
"""
try:
return operation(self._get_connection())
except Exception as exc:
if not _is_connection_lost(exc):
raise
logging.warning(
"Graph write lost its connection (%s); reconnecting and retrying once.",
exc,
)
self.close()
# Second and final attempt, on a connection freshly checked out by
# ``_get_connection``. A failure here belongs to the caller.
return operation(self._get_connection())
def _register_pgvector_types(self, conn) -> None:
"""Register pgvector's adapters, tolerating a not-yet-created extension.
@@ -197,7 +292,7 @@ class GraphStore:
)
cursor.execute(
"""
f"""
CREATE TABLE IF NOT EXISTS graph_edges (
id UUID PRIMARY KEY,
source_id UUID NOT NULL,
@@ -206,10 +301,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(
"""
@@ -397,19 +500,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,
@@ -420,9 +582,10 @@ class GraphStore:
description,
weight,
Jsonb(source_chunk_ids or []),
fact_embedding,
),
)
return edge_id
return edge_id, True
def add_edge(
self,
@@ -433,21 +596,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:
@@ -499,56 +669,96 @@ class GraphStore:
(not linked to the chunk), mirroring the per-call path.
``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees
are not bumped here — the caller runs ``set_node_degrees`` once at the
end. Returns ``(nodes_upserted, edges_added)``.
end. Reconnects and retries once if the connection died while the
extraction was waiting on the model.
The chunk's ``graph_ingest_progress`` row is written in this same
transaction, so the checkpoint and the rows it describes commit
together and a replay of an already-applied chunk returns ``(0, 0)``
without touching the graph. Returns ``(nodes_upserted, edges_added)``.
"""
self._ensure_tables_once()
conn = self._get_connection()
cursor = conn.cursor()
node_ids: Dict[str, str] = {}
edges_added = 0
try:
for entity in entities:
normalized_name = entity["normalized_name"]
node_id = self._upsert_node(
cursor,
source_id,
entity["name"],
normalized_name,
entity.get("type"),
entity.get("description"),
name_embeddings.get(normalized_name),
)
node_ids[normalized_name] = node_id
self._link_node_chunk(cursor, source_id, node_id, chunk_id)
for rel in relationships:
src_id = self._resolve_endpoint(
cursor, source_id, rel.get("source"), node_ids, name_embeddings
def _write(conn):
cursor = conn.cursor()
node_ids: Dict[str, str] = {}
edges_added = 0
try:
# ``commit()`` can report connection loss *after* the server
# committed, and the retry then replays this write: doc_freq
# would be bumped twice and a second logical edge inserted
# (graph_edges has no uniqueness constraint). The progress row
# below is written in this transaction, so a replay sees it —
# and so does an overlapping build, once the source lock makes
# it wait for this one to commit.
_lock_source(cursor, source_id)
cursor.execute(
"SELECT status FROM graph_ingest_progress "
"WHERE source_id = %s AND chunk_id = %s;",
(source_id, str(chunk_id)),
)
dst_id = self._resolve_endpoint(
cursor, source_id, rel.get("target"), node_ids, name_embeddings
)
if src_id is None or dst_id is None:
continue
self._add_edge(
cursor,
source_id,
src_id,
dst_id,
type=rel.get("type"),
description=rel.get("description"),
weight=float(rel.get("weight") or 1.0),
source_chunk_ids=[chunk_id],
)
edges_added += 1
applied = cursor.fetchone()
if applied is not None and applied[0] == "done":
conn.rollback()
return 0, 0
conn.commit()
return len(entities), edges_added
except Exception:
conn.rollback()
raise
finally:
cursor.close()
for entity in entities:
normalized_name = entity["normalized_name"]
node_id = self._upsert_node(
cursor,
source_id,
entity["name"],
normalized_name,
entity.get("type"),
entity.get("description"),
name_embeddings.get(normalized_name),
)
node_ids[normalized_name] = node_id
self._link_node_chunk(cursor, source_id, node_id, chunk_id)
for rel in relationships:
src_id = self._resolve_endpoint(
cursor, source_id, rel.get("source"), node_ids, name_embeddings
)
dst_id = self._resolve_endpoint(
cursor, source_id, rel.get("target"), node_ids, name_embeddings
)
if src_id is None or dst_id is None:
continue
_, created = self._add_edge(
cursor,
source_id,
src_id,
dst_id,
type=rel.get("type"),
description=rel.get("description"),
# Only a missing weight defaults: 0 is a real one,
# and the ranker drops non-positive edges.
weight=1.0 if rel.get("weight") is None else float(rel["weight"]),
source_chunk_ids=[chunk_id],
fact_embedding=rel.get("fact_embedding"),
)
if created:
edges_added += 1
cursor.execute(
"""
INSERT INTO graph_ingest_progress (source_id, chunk_id, status)
VALUES (%s, %s, 'done')
ON CONFLICT (source_id, chunk_id)
DO UPDATE SET status = EXCLUDED.status;
""",
(source_id, str(chunk_id)),
)
conn.commit()
return len(entities), edges_added
except Exception:
_safe_rollback(conn)
raise
finally:
cursor.close()
return self._write_with_reconnect(_write)
def _resolve_endpoint(
self,
@@ -564,7 +774,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(
@@ -610,8 +824,22 @@ class GraphStore:
cursor.close()
conn.rollback()
def count_nodes(self, source_id: str) -> int:
"""Number of nodes for a source. Zero drives the ClassicRAG fallback."""
def count_nodes(self, source_id: str, strict: bool = False) -> int:
"""Number of nodes for a source. Zero drives the ClassicRAG fallback.
Args:
source_id: Source whose nodes to count.
strict: Re-raise a query failure instead of reporting ``0``.
Retrieval wants the swallow — a broken count there just routes
the source to ClassicRAG — but a caller reporting how big a
graph is must not read a failed query as "the graph is empty".
Returns:
int: The node count, or ``0`` when a query failure is swallowed.
Raises:
Exception: The underlying query failure, when ``strict`` is set.
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
@@ -622,6 +850,8 @@ class GraphStore:
return int(cursor.fetchone()[0])
except Exception as e:
logging.error(f"Error counting nodes: {e}")
if strict:
raise
return 0
finally:
cursor.close()
@@ -734,6 +964,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;
""",
(
@@ -781,6 +1012,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),
@@ -894,6 +1126,223 @@ 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(
sql.SQL(
"""
SELECT d.{metadata}, d.{text},
bool_or(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} = %s
AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s)
-- One page per chunk. Grouping on the subject flag as well
-- split a chunk two entities link -- one naming it, one
-- merely mentioned -- into two identical pages, spending
-- the caller's page budget twice on the same text.
GROUP BY d.{metadata}, d.{text}
ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC
LIMIT %s;
"""
).format(
metadata=_identifier(metadata_col),
text=_identifier(text_col),
table=_identifier(table),
source=_identifier(source_col),
),
(
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(
sql.SQL(
"""
SELECT id::text, 1 - ({vector} <=> %s::vector)
FROM {table}
WHERE {source} = %s AND id::text = ANY(%s);
"""
).format(
vector=_identifier(vector_col),
table=_identifier(table),
source=_identifier(source_col),
),
(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,
@@ -914,10 +1363,17 @@ class GraphStore:
cursor = conn.cursor()
try:
cursor.execute(
f"""
SELECT id, {text_col}, {metadata_col} FROM {table}
WHERE {source_col} = %s AND id::text = ANY(%s);
""",
sql.SQL(
"""
SELECT id, {text}, {metadata} FROM {table}
WHERE {source} = %s AND id::text = ANY(%s);
"""
).format(
text=_identifier(text_col),
metadata=_identifier(metadata_col),
table=_identifier(table),
source=_identifier(source_col),
),
(source_id, [str(c) for c in chunk_ids]),
)
return {
@@ -1026,25 +1482,29 @@ class GraphStore:
cursor.close()
def mark_chunk(self, source_id: str, chunk_id: str, status: str):
"""Record a chunk's extraction status, reconnecting once if the connection died."""
self._ensure_tables_once()
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO graph_ingest_progress (source_id, chunk_id, status)
VALUES (%s, %s, %s)
ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status;
""",
(source_id, str(chunk_id), status),
)
conn.commit()
except Exception as e:
conn.rollback()
logging.error(f"Error marking chunk: {e}")
raise
finally:
cursor.close()
def _write(conn):
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO graph_ingest_progress (source_id, chunk_id, status)
VALUES (%s, %s, %s)
ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status;
""",
(source_id, str(chunk_id), status),
)
conn.commit()
except Exception as e:
_safe_rollback(conn)
logging.error(f"Error marking chunk: {e}")
raise
finally:
cursor.close()
return self._write_with_reconnect(_write)
def pending_chunks(self, source_id: str, all_chunk_ids: List[str]) -> List[str]:
"""Chunk ids from ``all_chunk_ids`` not yet marked ``done`` for the source."""
@@ -1092,6 +1552,9 @@ class GraphStore:
conn = self._get_connection()
cursor = conn.cursor()
try:
# A reset while a build is still writing must not land in the
# middle of one of its chunks.
_lock_source(cursor, source_id)
for table in (
"graph_node_chunks",
"graph_edges",
@@ -1099,7 +1562,10 @@ class GraphStore:
"graph_ingest_progress",
):
cursor.execute(
f"DELETE FROM {table} WHERE source_id = %s;", (source_id,)
sql.SQL("DELETE FROM {} WHERE source_id = %s;").format(
sql.Identifier(table)
),
(source_id,),
)
conn.commit()
except Exception as e:
+11
View File
@@ -185,12 +185,23 @@ class Dispatcher(BaseRetriever):
score_threshold / rephrase_query) plus an opted-in prescreen config; a
source left at defaults takes the global path so all-classic retrieval
stays byte-identical with zero extra LLM calls.
A graph source's ``graph`` options count too: they are read from the
per-source config this records, so a source that changes only those
would otherwise run the defaults and the options would do nothing.
They mean nothing to any other retriever, so they only count for
``graphrag`` -- an override hands the source its own chunk budget as
well, which a classic source must not pick up from a graph setting.
"""
return (
retrieval.chunks != _DEFAULT_RETRIEVAL.chunks
or retrieval.score_threshold != _DEFAULT_RETRIEVAL.score_threshold
or retrieval.rephrase_query != _DEFAULT_RETRIEVAL.rephrase_query
or retrieval.prescreen is not None
or (
(retrieval.retriever or "").lower() == "graphrag"
and retrieval.graph != _DEFAULT_RETRIEVAL.graph
)
)
@staticmethod
+416 -36
View File
@@ -1,9 +1,13 @@
"""GraphRAG local retriever — Personalized PageRank over a per-source graph.
Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx
Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR
mass -> shared token budget. No LLM call at query time beyond the (optional,
reused) rephrase.
Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> Personalized
PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR mass -> shared
token budget. No LLM call at query time beyond the (optional, reused) rephrase.
``networkx`` supplies the graph structure, but the ranking is the local power
iteration in :func:`_personalized_pagerank`: ``nx.pagerank`` delegates to scipy,
which DocsGPT does not depend on, so calling it turned every graph retrieval
into a silent ClassicRAG fallback.
Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the
``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the
@@ -28,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
@@ -35,11 +40,144 @@ 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 _nodes_by_chunk(chunk_links: Dict[str, List[str]]) -> Dict[str, List[str]]:
"""Invert ``node -> chunk ids`` into ``chunk id -> node ids``.
Node order within a chunk follows the node order of ``chunk_links``, so the
passage edges are added in the same order as before.
"""
inverted: Dict[str, List[str]] = {}
for node, chunks in chunk_links.items():
for chunk_id in chunks or ():
inverted.setdefault(chunk_id, []).append(node)
return inverted
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``.
Weights are clamped at zero (a cosine distance above 1 yields a negative
seed weight) and normalized across the nodes actually present in the graph,
so restart mass can never leak to a node the subgraph does not contain. An
absent or all-zero personalization collapses to a uniform restart.
"""
if personalization:
weights = {
node: max(float(personalization.get(node, 0.0) or 0.0), 0.0)
for node in nodes
}
total = sum(weights.values())
if total > 0:
return {node: weight / total for node, weight in weights.items()}
uniform = 1.0 / len(nodes)
return {node: uniform for node in nodes}
def _personalized_pagerank(
graph: nx.Graph,
personalization: Dict[Any, float] | None = None,
*,
weight: str = "weight",
alpha: float = 0.85,
max_iter: int = 100,
tol: float = 1.0e-6,
) -> Dict[Any, float]:
"""Personalized PageRank by power iteration — no scipy.
``networkx.pagerank`` delegates to a scipy implementation, and scipy is not
a DocsGPT dependency: in a default install the import raises and every graph
retrieval silently degrades to the ClassicRAG fallback. This is the same
algorithm over the same row-normalized transition matrix, so the ranking is
unchanged where scipy happens to be installed.
Args:
graph: Undirected graph whose edges may carry a ``weight`` attribute.
personalization: Node -> restart weight; ``None`` means uniform.
weight: Edge attribute holding the weight.
alpha: Damping factor.
max_iter: Iteration cap. The last iterate is returned if it is hit —
retrieval degrades to a slightly less converged ranking rather than
raising, which is what the library does.
tol: Convergence tolerance; iteration stops below ``len(graph) * tol``.
Returns:
Node -> PageRank mass, summing to ~1.0. Empty dict for an empty graph.
"""
nodes = list(graph.nodes)
node_count = len(nodes)
if node_count == 0:
return {}
restart = _restart_vector(nodes, personalization)
# Row-normalized transitions. An undirected edge is traversable from both
# endpoints, so each node normalizes over its own incident weights.
transitions: Dict[Any, List[tuple[Any, float]]] = {}
for node in nodes:
neighbors = []
total = 0.0
for neighbor, data in graph[node].items():
raw_weight = data.get(weight, 1.0)
# Default only a missing or null weight. ``or 1.0`` would also
# rewrite an explicit 0 — "these entities are not related" — into a
# full-strength transition, which changes the ranking.
edge_weight = 1.0 if raw_weight is None else float(raw_weight)
if edge_weight <= 0:
continue
neighbors.append((neighbor, edge_weight))
total += edge_weight
transitions[node] = (
[(n, w / total) for n, w in neighbors] if total > 0 else []
)
# A node with no usable edge is dangling: its mass would vanish each pass,
# so it is redistributed along the restart vector instead.
dangling = [node for node in nodes if not transitions[node]]
ranks = {node: 1.0 / node_count for node in nodes}
for _ in range(max_iter):
previous = ranks
ranks = dict.fromkeys(nodes, 0.0)
leaked = alpha * sum(previous[node] for node in dangling)
for node in nodes:
share = alpha * previous[node]
for neighbor, transition in transitions[node]:
ranks[neighbor] += share * transition
for node in nodes:
ranks[node] += (leaked + 1.0 - alpha) * restart[node]
if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol:
break
else:
logging.debug(
"Personalized PageRank hit its %s-iteration cap on a %s-node "
"subgraph; ranking with the last iterate.",
max_iter,
node_count,
)
return ranks
class GraphRAGRetriever(BaseRetriever):
"""Per-source PPR retriever; falls back to ClassicRAG when a source has no graph."""
@@ -102,14 +240,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:
weight = float(edge.get("weight") or 1.0)
graph.add_edge(src, dst, weight=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 {}
@@ -117,12 +250,34 @@ class GraphRAGRetriever(BaseRetriever):
if not any(personalization.values()):
personalization = None
ranks = nx.pagerank(graph, personalization=personalization, weight="weight")
ranks = _personalized_pagerank(
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.
@@ -140,6 +295,77 @@ 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] = {}
# Inverted once: scanning every node's chunk list per candidate is
# quadratic, and at the candidate cap it cost more than the walk it
# feeds (76 ms against 0.5 ms measured).
nodes_by_chunk = _nodes_by_chunk(chunk_links)
for chunk_id in candidate_ids:
linked = [n for n in nodes_by_chunk.get(chunk_id, ()) 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.
@@ -159,6 +385,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]]:
@@ -170,9 +501,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 []
@@ -186,26 +515,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
@@ -252,8 +618,9 @@ class GraphRAGRetriever(BaseRetriever):
Graph sources keep their own slot in source order; every graphless
source collapses into a single ClassicRAG run that occupies the slot of
the first graphless source. Sources whose PPR retrieval raises are
collected and retried as one more classic batch, appended at the end.
the first graphless source. Sources whose PPR retrieval raises, or
answers nothing, are collected and retried as one more classic batch,
appended at the end.
"""
try:
counts = store.count_nodes_many(sources)
@@ -279,7 +646,7 @@ class GraphRAGRetriever(BaseRetriever):
segments.append([])
graphless.append(source_id)
failed: List[str] = []
fallback: List[str] = []
query_embedding = None
if graphed:
# Embedded once for the whole retrieval, not once per graph source.
@@ -292,26 +659,39 @@ class GraphRAGRetriever(BaseRetriever):
f"GraphRAG query embedding failed, falling back: {e}",
exc_info=True,
)
failed, graphed = list(graphed), []
fallback, graphed = list(graphed), []
for source_id in graphed:
try:
segments[graph_slots[source_id]] = self._graph_docs_for_source(
store, source_id, query_embedding
)
docs = self._graph_docs_for_source(store, source_id, query_embedding)
except Exception as e:
logging.error(
f"GraphRAG retrieval failed for {source_id}, falling back: {e}",
exc_info=True,
)
failed.append(source_id)
fallback.append(source_id)
continue
if not docs:
# Empty is not an answer. Every graph read reports its own
# failure and returns nothing, so "no rows" covers a query that
# broke or a half-built graph as much as a walk that found
# nothing — and only a raise reaches the fallback, so the
# source would otherwise contribute nothing at all. Searching
# it classically is what a source with no graph already gets.
logging.info(
"GraphRAG retrieval returned nothing for %s, falling back",
source_id,
)
fallback.append(source_id)
continue
segments[graph_slots[source_id]] = docs
# Every remaining segment is a ClassicRAG fan-out, and each of its legs
# checks out of the *same* per-DSN pool this store is holding. Hand the
# graph connection back first, or concurrent GraphRAG retrievals occupy
# every slot and then block on their own fallbacks until PoolTimeout.
# ``close()`` nulls the connection, so ``_get_data``'s finally stays correct.
if graphless or failed:
if graphless or fallback:
try:
store.close()
except Exception as e:
@@ -319,8 +699,8 @@ class GraphRAGRetriever(BaseRetriever):
if graphless:
segments[classic_slot] = self._classic_for_sources(graphless)
if failed:
segments.append(self._classic_for_sources(failed))
if fallback:
segments.append(self._classic_for_sources(fallback))
return [doc for segment in segments for doc in segment]
+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
+6 -5
View File
@@ -10,8 +10,9 @@ Celery and the vector comes back. The API pays a broker round trip per query
and no resident model.
Inside a worker there is nothing to delegate to -- dispatching would queue work
behind the task already running and wait on itself -- so a call made while a
task is executing runs locally, on a model this process loads once and caches.
behind the task already running and wait on itself -- so a call made anywhere in
a worker process, including from a thread a task started, runs locally, on a
model this process loads once and caches.
``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side.
Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding
@@ -79,11 +80,11 @@ def _forget(result) -> None:
def _in_worker() -> bool:
"""True when a Celery task is executing in this process."""
"""True anywhere in a Celery worker process -- on any thread, not only the task's."""
try:
from docsgpt.celery_init import celery
from docsgpt.celery_init import in_worker
return celery.current_worker_task is not None
return in_worker()
except Exception:
return False
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "Lade diese Quelle vorab in den Prompt oder lass den Agenten sie bei Bedarf als Werkzeug durchsuchen."
},
"graphRetrieval": {
"title": "Graph-Abruf",
"tag": "ohne Neuimport",
"seedStrategy": "Suche beginnt bei",
"seedStrategyHint": "Entitäten eignen sich für die meisten Dokumente. Beziehungen erreichen auch Entitäten, die in der Frage nicht vorkommen – ideal für Inhalte, die beschreiben, wie Dinge zusammenhängen.",
"seedEntities": "Entitäten (empfohlen)",
"seedRelationships": "Beziehungen",
"passageNodes": "Textabschnitte in die Suche einbeziehen",
"passageNodesHint": "Ein Abschnitt wird gefunden, wenn er zur Frage passt oder mit etwas Passendem verbunden ist. Am besten bei Graphen, die mit dieser Version erstellt wurden.",
"blendVector": "Mit Vektorsuche kombinieren",
"blendVectorHint": "Ergänzt Ergebnisse der Vektorsuche, damit kein Abschnitt verloren geht, den der Graph übersieht.",
"agentToolHint": "Agenten können diesen Beziehungen auch selbst folgen, wenn die Bereitstellung dieser Quelle „Suchwerkzeug auf Abruf“ ist oder ein agentischer Agent sie nutzt."
},
"prescreen": {
"enable": "LLM-Vorfilterung aktivieren",
"warning": "Ruft eine größere Kandidatenmenge ab und filtert sie mit einem LLM. Das erhöht Latenz und Kosten pro Anfrage.",
+13
View File
@@ -272,6 +272,19 @@
},
"exposureHint": "Pre-fetch this source into the prompt, or let the agent search it on demand as a tool."
},
"graphRetrieval": {
"title": "Graph retrieval",
"tag": "no re-ingest",
"seedStrategy": "Start the walk from",
"seedStrategyHint": "Entities suit most documents. Relationships can reach an entity the question never names, and suit content that describes how things connect.",
"seedEntities": "Entities (recommended)",
"seedRelationships": "Relationships",
"passageNodes": "Include passages in the walk",
"passageNodesHint": "Lets a passage be found both by matching the question and by being connected to what does. Works best on graphs built with this version.",
"blendVector": "Blend with vector search",
"blendVectorHint": "Adds plain vector search results, so a passage the graph misses is not lost.",
"agentToolHint": "Agents can also follow these relationships themselves when this source's exposure is “On-demand search tool”, or when an agentic agent uses it."
},
"prescreen": {
"enable": "Enable LLM prescreen",
"warning": "Fetches a larger candidate set and uses an LLM to filter it. This adds query-time latency and cost.",
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "Precarga esta fuente en el prompt, o deja que el agente la busque bajo demanda como herramienta."
},
"graphRetrieval": {
"title": "Recuperación por grafo",
"tag": "sin reingesta",
"seedStrategy": "Iniciar el recorrido desde",
"seedStrategyHint": "Las entidades funcionan para la mayoría de documentos. Las relaciones pueden llegar a una entidad que la pregunta no menciona; son ideales para contenido que describe cómo se conectan las cosas.",
"seedEntities": "Entidades (recomendado)",
"seedRelationships": "Relaciones",
"passageNodes": "Incluir fragmentos en el recorrido",
"passageNodesHint": "Un fragmento puede encontrarse por coincidir con la pregunta o por estar conectado con lo que coincide. Funciona mejor en grafos creados con esta versión.",
"blendVector": "Combinar con búsqueda vectorial",
"blendVectorHint": "Añade resultados de la búsqueda vectorial para no perder fragmentos que el grafo pase por alto.",
"agentToolHint": "Los agentes también pueden seguir estas relaciones por sí mismos cuando la exposición de esta fuente es «Herramienta de búsqueda bajo demanda» o cuando la usa un agente agéntico."
},
"prescreen": {
"enable": "Habilitar preselección con LLM",
"warning": "Obtiene un conjunto de candidatos más grande y usa un LLM para filtrarlo. Esto añade latencia y costo por consulta.",
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。"
},
"graphRetrieval": {
"title": "グラフ検索",
"tag": "再取り込み不要",
"seedStrategy": "探索の開始点",
"seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。",
"seedEntities": "エンティティ(推奨)",
"seedRelationships": "リレーション",
"passageNodes": "パッセージを探索に含める",
"passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。",
"blendVector": "ベクトル検索と組み合わせる",
"blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。",
"agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。"
},
"prescreen": {
"enable": "LLMプリスクリーニングを有効にする",
"warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。",
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту."
},
"graphRetrieval": {
"title": "Поиск по графу",
"tag": "без повторной загрузки",
"seedStrategy": "Начинать обход с",
"seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.",
"seedEntities": "Сущностей (рекомендуется)",
"seedRelationships": "Связей",
"passageNodes": "Включать фрагменты в обход",
"passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.",
"blendVector": "Сочетать с векторным поиском",
"blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.",
"agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент."
},
"prescreen": {
"enable": "Включить предварительный отбор LLM",
"warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.",
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。"
},
"graphRetrieval": {
"title": "圖譜檢索",
"tag": "無需重新匯入",
"seedStrategy": "走訪起點",
"seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。",
"seedEntities": "實體(建議)",
"seedRelationships": "關係",
"passageNodes": "將段落納入走訪",
"passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。",
"blendVector": "與向量檢索結合",
"blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。",
"agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。"
},
"prescreen": {
"enable": "啟用 LLM 預篩選",
"warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。",
+13
View File
@@ -268,6 +268,19 @@
},
"exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。"
},
"graphRetrieval": {
"title": "图谱检索",
"tag": "无需重新导入",
"seedStrategy": "遍历起点",
"seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。",
"seedEntities": "实体(推荐)",
"seedRelationships": "关系",
"passageNodes": "将段落纳入遍历",
"passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。",
"blendVector": "与向量检索结合",
"blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。",
"agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。"
},
"prescreen": {
"enable": "启用 LLM 预筛选",
"warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。",
+12
View File
@@ -32,6 +32,17 @@ export type SourcePrescreenConfig = {
max_keep?: number; // default 8, <= candidate_k
};
// Where the graph walk starts: matching entities, or matching relationships
// ("A streams_to B"), which can reach an entity the question never names.
export type GraphSeedStrategy = 'entities' | 'relationships';
// Query-time graph retrieval knobs (graphrag only; live, no re-ingest).
export type SourceGraphRetrievalConfig = {
seed_strategy?: GraphSeedStrategy; // default 'entities'
passage_nodes?: boolean; // default true
blend_vector?: boolean; // default true
};
// Query-time retrieval knobs (live; no re-ingest needed).
export type SourceRetrievalConfig = {
retriever?: string; // default 'classic' (only option for now)
@@ -40,6 +51,7 @@ export type SourceRetrievalConfig = {
score_threshold?: number | null; // default null
rephrase_query?: boolean; // default true
prescreen?: SourcePrescreenConfig | null; // null = off
graph?: SourceGraphRetrievalConfig; // graphrag retriever only
};
// Ingest-time GraphRAG extraction knobs (only used when kind === 'graphrag').
@@ -146,6 +146,11 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => {
batch_size: 5,
max_keep: 10,
},
graph: {
seed_strategy: 'relationships',
passage_nodes: false,
blend_vector: false,
},
},
graph: {
extraction_model: null,
@@ -157,6 +162,44 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => {
});
});
describe('graph retrieval options', () => {
it('defaults to the measured-best configuration', () => {
expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph).toEqual({
seed_strategy: 'entities',
passage_nodes: true,
blend_vector: true,
});
});
it('fills the defaults for a source saved before the options existed', () => {
const opts = configToOptions({ retrieval: { retriever: 'graphrag' } });
expect(opts.retrieval.graph).toEqual(
DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph,
);
});
it('honors stored options and fills only the missing ones', () => {
const opts = configToOptions({
retrieval: { graph: { seed_strategy: 'relationships' } },
});
expect(opts.retrieval.graph).toEqual({
seed_strategy: 'relationships',
passage_nodes: true,
blend_vector: true,
});
});
it('writes the options into the retrieval block', () => {
const v = clone(DEFAULT_RETRIEVAL_OPTIONS);
v.retrieval.graph.blend_vector = false;
expect(optionsToConfig(v).retrieval?.graph).toEqual({
seed_strategy: 'entities',
passage_nodes: true,
blend_vector: false,
});
});
});
describe('isPrescreenConfigValid', () => {
const withPrescreen = (
chunks: number,
@@ -16,6 +16,7 @@ import {
import { Switch } from '../../components/ui/switch';
import type {
ChunkingStrategy,
GraphSeedStrategy,
RetrievalExposure,
SourceConfig,
} from '../../models/misc';
@@ -59,6 +60,11 @@ export type RetrievalOptionsValue = {
batch_size: number;
max_keep: number;
};
graph: {
seed_strategy: GraphSeedStrategy;
passage_nodes: boolean;
blend_vector: boolean;
};
};
graph: {
extraction_model: string | null;
@@ -85,6 +91,12 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = {
enabled: false,
...DEFAULT_PRESCREEN,
},
// The configuration that measured best across the corpora tested.
graph: {
seed_strategy: 'entities',
passage_nodes: true,
blend_vector: true,
},
},
graph: {
extraction_model: null,
@@ -204,6 +216,7 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue {
const chunking = config?.chunking ?? {};
const retrieval = config?.retrieval ?? {};
const prescreen = retrieval.prescreen ?? null;
const retrievalGraph = retrieval.graph ?? {};
const graph = config?.graph ?? {};
const d = DEFAULT_RETRIEVAL_OPTIONS;
return {
@@ -228,6 +241,14 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue {
batch_size: prescreen?.batch_size ?? DEFAULT_PRESCREEN.batch_size,
max_keep: prescreen?.max_keep ?? DEFAULT_PRESCREEN.max_keep,
},
graph: {
seed_strategy:
retrievalGraph.seed_strategy ?? d.retrieval.graph.seed_strategy,
passage_nodes:
retrievalGraph.passage_nodes ?? d.retrieval.graph.passage_nodes,
blend_vector:
retrievalGraph.blend_vector ?? d.retrieval.graph.blend_vector,
},
},
graph: {
extraction_model: graph.extraction_model ?? d.graph.extraction_model,
@@ -271,6 +292,11 @@ export function optionsToConfig(value: RetrievalOptionsValue): SourceConfig {
max_keep: ps.max_keep,
}
: null,
graph: {
seed_strategy: value.retrieval.graph.seed_strategy,
passage_nodes: value.retrieval.graph.passage_nodes,
blend_vector: value.retrieval.graph.blend_vector,
},
},
graph: {
extraction_model: value.graph.extraction_model?.trim()
@@ -399,6 +425,12 @@ export default function RetrievalOptions({
});
};
const setGraphRetrieval = (
patch: Partial<RetrievalOptionsValue['retrieval']['graph']>,
) => {
setRetrieval({ graph: { ...value.retrieval.graph, ...patch } });
};
const modelOptions = useMemo(() => {
const builtin: Model[] = [];
const user: Model[] = [];
@@ -611,6 +643,84 @@ export default function RetrievalOptions({
)}
</div>
{/* Graph retrieval group (graphrag only; live, so shown when testing too) */}
{isGraphRAG && (
<div className="flex flex-col gap-3">
<GroupHeader
title={tr('graphRetrieval.title')}
tag={tr('graphRetrieval.tag')}
/>
<p className="text-muted-foreground text-xs">
{tr('graphRetrieval.agentToolHint')}
</p>
<div className="divide-border/50 divide-y">
<SettingRow
label={tr('graphRetrieval.seedStrategy')}
htmlFor="graph-seed-strategy"
description={tr('graphRetrieval.seedStrategyHint')}
alignStart
>
<Select
value={value.retrieval.graph.seed_strategy}
disabled={disabled}
onValueChange={(v) =>
setGraphRetrieval({ seed_strategy: v as GraphSeedStrategy })
}
>
<SelectTrigger
id="graph-seed-strategy"
className="w-52 rounded-md"
size="lg"
>
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="entities">
{tr('graphRetrieval.seedEntities')}
</SelectItem>
<SelectItem value="relationships">
{tr('graphRetrieval.seedRelationships')}
</SelectItem>
</SelectContent>
</Select>
</SettingRow>
<SettingRow
label={tr('graphRetrieval.passageNodes')}
htmlFor="graph-passage-nodes"
description={tr('graphRetrieval.passageNodesHint')}
alignStart
>
<Switch
id="graph-passage-nodes"
checked={value.retrieval.graph.passage_nodes}
disabled={disabled}
onCheckedChange={(checked) =>
setGraphRetrieval({ passage_nodes: checked })
}
/>
</SettingRow>
<SettingRow
label={tr('graphRetrieval.blendVector')}
htmlFor="graph-blend-vector"
description={tr('graphRetrieval.blendVectorHint')}
alignStart
>
<Switch
id="graph-blend-vector"
checked={value.retrieval.graph.blend_vector}
disabled={disabled}
onCheckedChange={(checked) =>
setGraphRetrieval({ blend_vector: checked })
}
/>
</SettingRow>
</div>
</div>
)}
{/* Graph extraction group (graphrag only; re-ingest required to apply) */}
{isGraphRAG && !queryOnly && (
<div className="flex flex-col gap-3">
+20
View File
@@ -313,6 +313,26 @@ class TestClassicAgentSearchExposure:
"Tool Doc",
]
def test_collect_internal_sources_includes_graph_pages(
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
):
# Pages the graph tool read carry the answer as much as search hits do,
# so they are cited the same way.
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
retriever_config = {"source": {"active_docs": ["b"]}}
agent = ClassicAgent(retriever_config=retriever_config, **agent_base_params)
search = Mock()
search.retrieved_docs = [{"text": "Found", "title": "Search Doc", "source": "b"}]
graph = Mock()
graph.retrieved_docs = [{"text": "Quill is a store.", "title": "quill.md", "source": "b"}]
user = agent.user or ""
agent.tool_executor._loaded_tools[f"internal_search:{INTERNAL_TOOL_ID}:{user}"] = search
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{user}"] = graph
agent._collect_internal_sources()
assert [d["title"] for d in agent.retrieved_docs] == ["Search Doc", "quill.md"]
def test_collect_internal_sources_dedupes(
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
):
+14
View File
@@ -779,6 +779,20 @@ class TestCollectStepSources:
assert len(agent.citations.citations) == 2
def test_collects_pages_the_graph_tool_read(
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
):
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
agent = ResearchAgent(**agent_base_params)
graph = Mock()
graph.retrieved_docs = [{"source": "s3", "title": "quill.md", "text": "Quill"}]
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{agent.user or ''}"] = graph
agent._collect_step_sources()
assert len(agent.citations.citations) == 1
def test_no_tool_no_error(
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
):
@@ -357,9 +357,9 @@ def test_malformed_json_schema_rejected_before_enqueue(monkeypatch):
@pytest.mark.unit
def test_dispatch_inline_when_in_worker(monkeypatch):
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
# Inside a worker current_task is truthy -> parse inline, never enqueue (else the
# Inside a worker -> parse inline, never enqueue (else the
# parsing queue self-deadlocks the worker that also serves it).
monkeypatch.setattr(rd, "current_task", object())
monkeypatch.setattr(rd, "in_worker", lambda: True)
import docsgpt.api.user.tasks as tasks
monkeypatch.setattr(
@@ -387,8 +387,8 @@ def test_dispatch_inline_when_in_worker(monkeypatch):
@pytest.mark.unit
def test_dispatch_enqueues_when_not_in_worker(monkeypatch):
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
# Web process: current_task falsy -> dispatch to the parsing queue, never inline.
monkeypatch.setattr(rd, "current_task", None)
# Web process -> dispatch to the parsing queue, never inline.
monkeypatch.setattr(rd, "in_worker", lambda: False)
captured = _patch_task(monkeypatch, payload={"status": "ok", "content": "queued", "truncated": False})
import docsgpt.worker as worker
@@ -414,9 +414,9 @@ _TIMED_OUT = "document parsing timed out after"
def _inline(monkeypatch, run_parse, *, timeout=0.2) -> ReadDocumentTool:
"""Drive the inline branch (current_task truthy) with a patched parse window."""
"""Drive the inline (in-worker) branch with a patched parse window."""
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
monkeypatch.setattr(rd, "current_task", object())
monkeypatch.setattr(rd, "in_worker", lambda: True)
import docsgpt.api.user.tasks as tasks
monkeypatch.setattr(
@@ -335,6 +335,79 @@ class TestLiveLeaseDefersConcurrentRun:
assert row[2] == "completed"
@pytest.mark.unit
class TestLeaseDeferralGivesUpQuietly:
"""Deferral is bookkeeping, not failure.
A task that outruns the broker's visibility timeout is redelivered while
the first worker is still running it. The lease keeps the duplicate from
doing the work, but the duplicate kept re-queueing itself until celery
exhausted ``LEASE_RETRY_MAX`` and raised ``MaxRetriesExceededError``, so a
healthy long task logged a task failure. The duplicate should stand down
instead and leave the run to the worker that holds the lease.
"""
def _hold_lease(self, pg_conn, key):
from docsgpt.storage.db.repositories.idempotency import (
IdempotencyRepository,
)
IdempotencyRepository(pg_conn).try_claim_lease(
key=key, task_name="thing",
task_id="t-worker-1", owner_id="worker-1",
)
def test_exhausted_retries_stand_down_without_recording_a_result(self, pg_conn):
from celery.exceptions import Ignore, MaxRetriesExceededError
from docsgpt.api.user.idempotency import with_idempotency
self._hold_lease(pg_conn, "k-long-run")
invocations = {"count": 0}
@with_idempotency(task_name="thing")
def task(self, idempotency_key=None):
invocations["count"] += 1
return {"ran": True}
# Celery raises this from ``self.retry`` once max_retries is hit.
worker2 = _fake_celery_self("t-worker-2")
worker2.retry.side_effect = MaxRetriesExceededError("out of retries")
# Ignore rather than a return value: a redelivery reuses the original
# task id, so returning would mark the id the client is polling
# SUCCESS — /api/task_status hands that straight to the UI, which
# would announce a finished (empty) build while the holder is still
# working. Ignore leaves the id's state to the holder.
with _patch_decorator_db(pg_conn), pytest.raises(Ignore):
task(worker2, idempotency_key="k-long-run")
# The lease holder is still running it; the duplicate did not.
assert invocations["count"] == 0
# The holder's row is untouched — not failed, not completed.
row = _row_for(pg_conn, "k-long-run")
assert row[2] == "pending"
def test_a_normal_retry_still_propagates(self, pg_conn):
"""Only exhaustion stands down; the first deferrals must re-queue."""
from docsgpt.api.user.idempotency import with_idempotency
self._hold_lease(pg_conn, "k-busy-once")
@with_idempotency(task_name="thing")
def task(self, idempotency_key=None):
return {"ran": True}
class _RetrySignal(Exception):
pass
worker2 = _fake_celery_self("t-worker-2")
worker2.retry.side_effect = _RetrySignal("retry scheduled")
with _patch_decorator_db(pg_conn), pytest.raises(_RetrySignal):
task(worker2, idempotency_key="k-busy-once")
@pytest.mark.unit
class TestExceptionPathReleasesLease:
"""When ``fn`` raises, the lease is dropped so the next attempt
+15
View File
@@ -183,6 +183,21 @@ def _no_worker_delegation(monkeypatch):
monkeypatch.setattr("docsgpt.cache._pubsub_redis_creation_failed", True)
@pytest.fixture(autouse=True)
def _graphrag_off_by_default(monkeypatch):
"""Run with GraphRAG at its shipped default (off), as CI does.
Every agent that gets a search tool checks its sources for a graph, and
that check reads the configured vector database. A dev ``.env`` enabling
GraphRAG sent unrelated agent tests to the developer's real database and
left a pool to it in ``pgconn._POOLS``, failing a live test that asserts it
owns the only pool. Tests that exercise GraphRAG turn it on themselves.
"""
from docsgpt.core.settings import settings
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False, raising=False)
@pytest.fixture
def mock_llm():
llm = Mock()
+708 -5
View File
@@ -88,6 +88,37 @@ class _StubLLM:
return response
class _ScriptedLLM:
"""Stub LLM answering per chunk, so results do not depend on call order.
The extraction pool runs calls concurrently, so a stub that hands out
responses in call order gives each chunk whichever response its thread
happened to grab first. ``script`` maps a chunk's text to the responses
for that chunk, consumed one per call.
"""
def __init__(self, script):
self._script = {text: list(responses) for text, responses in script.items()}
self.model_id = "stub-model"
self.calls = []
self._token_usage_source = None
self._request_id = None
def gen(self, model=None, messages=None, **kwargs):
text = messages[-1]["content"].removeprefix("<chunk>\n").removesuffix("\n</chunk>")
self.calls.append(text)
responses = self._script.get(text)
if not responses:
raise AssertionError(f"unexpected extraction call for {text!r}")
response = responses.pop(0)
if isinstance(response, Exception):
raise response
return response
def calls_for(self, text):
return self.calls.count(text)
class _StubEmbedding:
"""Stub embeddings model producing deterministic fixed-dim vectors."""
@@ -138,6 +169,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 +347,97 @@ 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.
Embedding from the pool is what broke every graph build inside a
worker: the embeddings client used to decide "embed locally" from the
task on the *current thread's* stack, so a pool thread dispatched to
the worker instead and Celery refused the wait. ``in_worker`` is
process-wide now, but the pool still has no reason to touch the
embeddings client — it exists to overlap model latency.
"""
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
):
@@ -292,11 +537,13 @@ class TestExtractionLive:
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
llm = _StubLLM([
"not json at all",
RuntimeError("model exploded"),
good,
])
# Each failing chunk fails its retry too; one that recovers on
# retry is covered in ``TestFailedChunksAreRetried``.
llm = _ScriptedLLM({
"garbage": ["not json at all", "still not json"],
"boom": [RuntimeError("model exploded"), RuntimeError("model exploded again")],
"Ada.": [good],
})
_install_stub_llm(monkeypatch, llm)
summary = extract_graph_for_source(
@@ -359,6 +606,61 @@ class TestExtractionTokenUsage:
assert built._request_id == "req-99"
assert captured["model_id"] == "stub-model"
def test_concurrent_extraction_calls_never_share_an_llm(self, monkeypatch, stub_embedding):
"""Provider usage is recorded on the LLM instance (``_last_usage``) and
claimed by whichever call finishes next, so two calls in flight on one
instance can bill each other's tokens. Each extraction thread needs its
own instance."""
import threading
import time
from unittest.mock import MagicMock
from docsgpt.core.settings import settings
payload = _extraction_json(
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
class _ThreadRecordingLLM:
model_id = "stub-model"
def __init__(self):
self.threads = set()
def gen(self, model=None, messages=None, **kwargs):
self.threads.add(threading.get_ident())
time.sleep(0.01) # keep calls overlapping
return payload
built = []
def _create(*args, **kwargs):
llm = _ThreadRecordingLLM()
built.append(llm)
return llm
monkeypatch.setattr(extraction_module.LLMCreator, "create_llm", staticmethod(_create))
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
store = MagicMock(name="GraphStore")
store.pending_chunks.return_value = [f"c{i}" for i in range(8)]
store.apply_chunk.return_value = (1, 0)
store.count_nodes.return_value = 1
monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", lambda *a, **k: store)
summary = extract_graph_for_source(
str(uuid.uuid4()),
user="owner-1",
chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(8)],
config=SourceConfig(),
request_id="req-threads",
)
assert summary["chunks_processed"] == 8
used = [llm for llm in built if llm.threads]
assert len(used) > 1, "calls did not run concurrently"
assert all(len(llm.threads) == 1 for llm in used)
@pytest.mark.unit
class TestModelResolution:
@@ -403,6 +705,407 @@ class TestModelResolution:
assert extraction_module._resolve_max_chunks(config) == 5
@pytest.mark.unit
class TestExtractionProviderResolution:
"""The extraction model decides the provider, not ``LLM_PROVIDER``.
``settings.LLM_PROVIDER`` is the deployment default (``docsgpt`` out of the
box, i.e. the hosted public endpoint). Dispatching the resolved extraction
model through it sends the call to a provider that never serves that model:
the request is rejected, the shared fallback answers instead, and the graph
is quietly built by a different model than the one configured.
"""
def _capture_create_llm(self, monkeypatch, llm=None):
captured = {}
def _create(provider, *args, **kwargs):
captured["provider"] = provider
captured["args"] = args
captured["kwargs"] = kwargs
return llm or _StubLLM([])
monkeypatch.setattr(
extraction_module.LLMCreator, "create_llm", staticmethod(_create)
)
return captured
def test_provider_comes_from_the_model_registry(self, monkeypatch):
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
monkeypatch.setattr(
extraction_module, "get_provider_from_model_id", lambda *a, **k: "openai"
)
monkeypatch.setattr(
extraction_module, "get_api_key_for_provider", lambda provider: "sk-openai"
)
captured = self._capture_create_llm(monkeypatch)
extraction_module._build_extraction_llm("gpt-4o-mini", "owner-1", "req-1")
assert captured["provider"] == "openai"
assert captured["kwargs"]["api_key"] == "sk-openai"
assert captured["kwargs"]["model_id"] == "gpt-4o-mini"
def test_owner_scopes_the_registry_lookup(self, monkeypatch):
"""A per-user (BYOM) model only resolves when the owner is passed."""
seen = {}
def _resolve(model_id, user_id=None):
seen["model_id"] = model_id
seen["user_id"] = user_id
return "anthropic"
monkeypatch.setattr(
extraction_module, "get_provider_from_model_id", _resolve
)
monkeypatch.setattr(
extraction_module, "get_api_key_for_provider", lambda provider: "k"
)
self._capture_create_llm(monkeypatch)
extraction_module._build_extraction_llm("byom-uuid", "owner-7", "req-1")
assert seen == {"model_id": "byom-uuid", "user_id": "owner-7"}
def test_unknown_model_falls_back_to_the_configured_provider(self, monkeypatch):
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
monkeypatch.setattr(
extraction_module, "get_provider_from_model_id", lambda *a, **k: None
)
monkeypatch.setattr(
extraction_module, "get_api_key_for_provider", lambda provider: "fallback-key"
)
captured = self._capture_create_llm(monkeypatch)
extraction_module._build_extraction_llm("mystery-model", "owner-1", "req-1")
assert captured["provider"] == "docsgpt"
assert captured["kwargs"]["api_key"] == "fallback-key"
def test_no_model_id_skips_the_lookup(self, monkeypatch):
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "openai")
calls = []
monkeypatch.setattr(
extraction_module,
"get_provider_from_model_id",
lambda *a, **k: calls.append(a) or "anthropic",
)
monkeypatch.setattr(
extraction_module, "get_api_key_for_provider", lambda provider: "k"
)
captured = self._capture_create_llm(monkeypatch)
extraction_module._build_extraction_llm(None, "owner-1", "req-1")
assert calls == []
assert captured["provider"] == "openai"
def test_api_key_follows_the_resolved_provider(self, monkeypatch):
"""The key must match the provider actually dispatched to."""
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
monkeypatch.setattr(extraction_module.settings, "API_KEY", "generic-key")
monkeypatch.setattr(
extraction_module, "get_provider_from_model_id", lambda *a, **k: "anthropic"
)
keyed_for = {}
def _key(provider):
keyed_for["provider"] = provider
return "sk-anthropic"
monkeypatch.setattr(extraction_module, "get_api_key_for_provider", _key)
captured = self._capture_create_llm(monkeypatch)
extraction_module._build_extraction_llm("claude-x", "owner-1", "req-1")
assert keyed_for["provider"] == "anthropic"
assert captured["kwargs"]["api_key"] == "sk-anthropic"
assert captured["kwargs"]["api_key"] != "generic-key"
@pytest.mark.unit
class TestFailedChunksAreReported:
"""Every dropped chunk has to leave a trace.
A chunk whose extraction cannot be parsed is marked ``failed`` and skipped.
That path logged nothing at all, so a graph could come back short with the
summary's ``failed_chunks`` count as the only hint and no way to tell which
chunk, or why, from the logs.
"""
def _fake_store(self, monkeypatch, chunk_ids):
from unittest.mock import MagicMock
store = MagicMock(name="GraphStore")
store.pending_chunks.return_value = list(chunk_ids)
store.apply_chunk.return_value = (1, 0)
store.count_nodes.return_value = 1
monkeypatch.setattr(
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
)
return store
def test_unparseable_output_is_logged_with_the_chunk_id(
self, monkeypatch, caplog, stub_embedding
):
import logging
store = self._fake_store(monkeypatch, ["c1"])
_install_stub_llm(monkeypatch, _StubLLM(["not json at all", "still not json"]))
with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"):
summary = extract_graph_for_source(
str(uuid.uuid4()),
user="owner-1",
chunks=[_chunk("c1", "some text")],
config=SourceConfig(),
request_id="req-1",
)
assert summary["failed_chunks"] == 1
store.mark_chunk.assert_called_once()
assert store.mark_chunk.call_args.args[2] == "failed"
messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING]
assert any("c1" in message for message in messages), messages
def test_llm_errors_still_name_the_chunk(
self, monkeypatch, caplog, stub_embedding
):
import logging
self._fake_store(monkeypatch, ["c7"])
_install_stub_llm(
monkeypatch,
_StubLLM([RuntimeError("model exploded"), RuntimeError("model exploded again")]),
)
with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"):
extract_graph_for_source(
str(uuid.uuid4()),
user="owner-1",
chunks=[_chunk("c7", "some text")],
config=SourceConfig(),
request_id="req-1",
)
messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING]
assert any("c7" in message for message in messages), messages
@pytest.mark.unit
class TestFailedChunksAreRetried:
"""A chunk that fails once gets one more attempt before the build ends.
Failures are recorded as ``failed`` and the checkpoint treats them as
pending, but nothing ever ran the build again, so a single transient error
— a provider hiccup, one response that did not parse — left a permanent
hole in the graph until someone rebuilt the whole source. Retries run
after the rest of the build, which gives a burst of rate limiting time to
pass, and are bounded at one per chunk so a chunk that can never be
extracted costs at most two calls.
"""
GOOD = _extraction_json(
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
def _fake_store(self, monkeypatch, chunk_ids):
from unittest.mock import MagicMock
store = MagicMock(name="GraphStore")
store.pending_chunks.return_value = list(chunk_ids)
store.apply_chunk.return_value = (1, 0)
store.count_nodes.return_value = 1
monkeypatch.setattr(
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
)
return store
def _run(self, chunks, progress=None):
return extract_graph_for_source(
str(uuid.uuid4()),
user="owner-1",
chunks=chunks,
config=SourceConfig(),
request_id="req-retry",
progress_cb=progress,
)
@staticmethod
def _marked_failed(store):
return [c.args[1] for c in store.mark_chunk.call_args_list if c.args[2] == "failed"]
def test_a_transient_failure_is_retried_and_written(self, monkeypatch, stub_embedding):
store = self._fake_store(monkeypatch, ["c1"])
llm = _ScriptedLLM({"flaky": [RuntimeError("rate limited"), self.GOOD]})
_install_stub_llm(monkeypatch, llm)
summary = self._run([_chunk("c1", "flaky")])
assert summary["failed_chunks"] == 0
assert summary["chunks_processed"] == 1
assert store.apply_chunk.call_args.args[1] == "c1"
assert self._marked_failed(store) == []
def test_an_unparseable_response_is_retried(self, monkeypatch, stub_embedding):
store = self._fake_store(monkeypatch, ["c1"])
_install_stub_llm(monkeypatch, _ScriptedLLM({"odd": ["not json", self.GOOD]}))
summary = self._run([_chunk("c1", "odd")])
assert summary["failed_chunks"] == 0
assert self._marked_failed(store) == []
def test_a_chunk_that_fails_again_is_marked_failed_once(self, monkeypatch, stub_embedding):
store = self._fake_store(monkeypatch, ["c1"])
llm = _ScriptedLLM({"broken": ["not json", "still not json"]})
_install_stub_llm(monkeypatch, llm)
summary = self._run([_chunk("c1", "broken")])
assert summary["failed_chunks"] == 1
assert summary["chunks_processed"] == 0
assert llm.calls_for("broken") == 2
assert self._marked_failed(store) == ["c1"]
def test_only_failed_chunks_are_retried(self, monkeypatch, stub_embedding):
from docsgpt.core.settings import settings
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
self._fake_store(monkeypatch, ["c1", "c2", "c3"])
llm = _ScriptedLLM({
"one": [self.GOOD],
"two": [RuntimeError("timeout"), self.GOOD],
"three": [self.GOOD],
})
_install_stub_llm(monkeypatch, llm)
summary = self._run([_chunk("c1", "one"), _chunk("c2", "two"), _chunk("c3", "three")])
assert summary["chunks_processed"] == 3
assert summary["failed_chunks"] == 0
assert (llm.calls_for("one"), llm.calls_for("two"), llm.calls_for("three")) == (1, 2, 1)
def test_a_failed_write_is_retried(self, monkeypatch, stub_embedding):
store = self._fake_store(monkeypatch, ["c1"])
store.apply_chunk.side_effect = [RuntimeError("write failed"), (1, 0)]
_install_stub_llm(monkeypatch, _ScriptedLLM({"text": [self.GOOD, self.GOOD]}))
summary = self._run([_chunk("c1", "text")])
assert summary["failed_chunks"] == 0
assert summary["chunks_processed"] == 1
assert self._marked_failed(store) == []
def test_progress_ends_at_the_total(self, monkeypatch, stub_embedding):
self._fake_store(monkeypatch, ["c1", "c2"])
_install_stub_llm(monkeypatch, _ScriptedLLM({
"fine": [self.GOOD],
"broken": ["not json", "still not json"],
}))
events = []
self._run([_chunk("c1", "fine"), _chunk("c2", "broken")], progress=events.append)
assert all(e["current"] <= e["total"] for e in events)
assert events[-1]["current"] == events[-1]["total"] == 2
@pytest.mark.integration
class TestSummaryNodeCount:
"""``nodes`` must describe the graph, not the number of upserts."""
@pytest.fixture
def store(self, monkeypatch, postgresql):
store = _live_store(monkeypatch, postgresql.info)
yield store
store.close()
def test_repeated_entity_counts_once(
self, store, monkeypatch, stub_embedding
):
source_id = str(uuid.uuid4())
try:
payload = _extraction_json(
entities=[{"name": "Ada", "type": "person", "description": "d"}],
relationships=[],
)
_install_stub_llm(monkeypatch, _StubLLM([payload, payload]))
summary = extract_graph_for_source(
source_id,
user="owner-1",
chunks=[_chunk("c1", "Ada one."), _chunk("c2", "Ada two.")],
config=SourceConfig(),
request_id="req-1",
)
# Two chunks upserted the same entity: one node in the graph.
assert store.count_nodes(source_id) == 1
assert summary["nodes"] == 1
assert summary["chunks_processed"] == 2
finally:
store.delete_by_source(source_id)
@pytest.mark.unit
class TestSummaryCountFailure:
"""A broken count query must not be reported as an empty graph."""
def test_a_failed_count_reports_the_write_count(
self, monkeypatch, stub_embedding
):
from unittest.mock import MagicMock
store = MagicMock(name="GraphStore")
store.pending_chunks.return_value = ["c1"]
store.apply_chunk.return_value = (2, 1)
store.count_nodes.side_effect = RuntimeError("count query failed")
monkeypatch.setattr(
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
)
_install_stub_llm(
monkeypatch,
_StubLLM([_extraction_json([{"name": "Ada"}], [])]),
)
summary = extract_graph_for_source(
str(uuid.uuid4()),
user="owner-1",
chunks=[_chunk("c1", "Ada.")],
config=SourceConfig(),
request_id="req-1",
)
# Falls back to what was actually written, not to zero.
assert summary["nodes"] == 2
# And it asked for a count that raises rather than one that returns 0,
# or the fallback above could never run.
assert store.count_nodes.call_args.kwargs.get("strict") is True
@pytest.mark.unit
class TestEntityNormalization:
"""An entity whose name normalizes to nothing is not an entity.
``canonical_name`` answers "" for a punctuation-only name, and nodes are
merged on that key, so keeping them collapses every such entity onto one
shared node. ``_resolve_endpoint`` already drops them on the relationship
side.
"""
@pytest.mark.parametrize("name", ["!!!", "--", "?", " *** "])
def test_a_name_that_normalizes_to_nothing_is_dropped(self, name):
assert extraction_module._build_entities([{"name": name}]) == []
def test_real_names_survive(self):
built = extraction_module._build_entities(
[{"name": "Quill Store"}, {"name": "!!!"}, {"name": "Alder"}]
)
assert [e["normalized_name"] for e in built] == ["quill store", "alder"]
@pytest.mark.unit
class TestParsing:
def test_parses_embedded_json(self):
+361
View File
@@ -0,0 +1,361 @@
"""The graph exposed to an agent as callable tools.
Ranking with a graph never beat vector search in measurement; letting a model
*follow* an edge did, on content where the answer is two documents away. These
tests cover the contract that makes that possible — the tool must return the
relationships verbatim enough for the model to read a name out of them, and
must refuse clearly rather than silently when it has nothing to offer.
"""
from __future__ import annotations
import pytest
from docsgpt.agents.tools.graph_search import (
GRAPH_TOOL_ID,
GraphSearchTool,
add_graph_search_tool,
build_graph_tool_entry,
)
from docsgpt.core.settings import settings
SOURCE = {"active_docs": ["src-1"]}
@pytest.fixture(autouse=True)
def _graph_store_is_pgvector(monkeypatch):
"""The graph only exists under pgvector, and CI's default is faiss.
Set here rather than per test so a case that forgets it fails for its own
reason instead of the gate; the cases about a different vector store
override it in the test body.
"""
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
class _StubStore:
def __init__(self, nodes=None, relationships=None, pages=None):
self._nodes = nodes or []
self._relationships = relationships or []
self._pages = pages or []
def search_nodes_by_embedding(self, source_id, embedding, k=10):
return self._nodes[:k]
def entity_relationships(self, source_id, name, limit=25):
return self._relationships
def entity_pages(self, source_id, name, limit=4):
return self._pages
def _tool(monkeypatch, store, enabled=True):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled)
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
tool = GraphSearchTool({"source": SOURCE})
tool._store = store
monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1])
return tool
class TestGating:
def test_reports_when_graphs_are_disabled(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore(), enabled=False)
assert "not enabled" in tool.execute_action("search_entities", query="x")
def test_reports_when_no_sources_are_configured(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
tool = GraphSearchTool({"source": {"active_docs": []}})
assert "No graph-backed sources" in tool.execute_action(
"search_entities", query="x"
)
def test_unknown_action_is_named(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore())
assert "Unknown action" in tool.execute_action("wander")
class TestActions:
def test_search_entities_lists_names_with_match_strength(self, monkeypatch):
tool = _tool(
monkeypatch,
_StubStore(nodes=[{"name": "Quill", "distance": 0.2, "description": "A store."}]),
)
result = tool.execute_action("search_entities", query="quill")
assert "Quill" in result
assert "0.80" in result
def test_search_entities_requires_a_query(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore())
assert "required" in tool.execute_action("search_entities", query=" ")
def test_relationships_are_rendered_as_triples(self, monkeypatch):
"""The model reads the *target* out of this line to take its next step,
so the target name has to survive rendering intact."""
tool = _tool(
monkeypatch,
_StubStore(
relationships=[
{"source": "Alder", "type": "streams_to", "target": "Quill", "description": ""}
]
),
)
result = tool.execute_action("get_relationships", entity="Alder")
assert "Alder --streams_to--> Quill" in result
def test_missing_relationships_suggest_the_next_step(self, monkeypatch):
"""A dead end should point at search_entities rather than stop the agent."""
tool = _tool(monkeypatch, _StubStore(relationships=[]))
result = tool.execute_action("get_relationships", entity="Nope")
assert "search_entities" in result
def test_pages_are_titled_truncated_and_recorded(self, monkeypatch):
tool = _tool(
monkeypatch,
_StubStore(
pages=[
{
"metadata": {"title": "quill-store.md", "source": "quill-store.md"},
"text": "x" * 5000,
}
]
),
)
result = tool.execute_action("read_entity_pages", entity="Quill")
assert "--- quill-store.md ---" in result
# Only what the model reads is truncated.
assert len(result) < 3000
# Accumulated so the answer can cite what the walk actually read.
assert tool.retrieved_docs[0]["title"] == "quill-store.md"
assert len(tool.retrieved_docs[0]["text"]) == 5000
def test_page_labels_match_what_the_retrievers_record(self, monkeypatch):
"""A page read here and the same chunk retrieved by internal_search are
one document. The citation manager keys on (source, title), so labels
derived differently give the same document two citation numbers."""
from docsgpt.retriever.labels import labels_from_metadata
metadata = {"title": "Quill Store", "source": "quill-store.md"}
text = "Quill is a write-ahead store."
tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": metadata, "text": text}]))
tool.execute_action("read_entity_pages", entity="Quill")
expected = labels_from_metadata(metadata, text, "src-1")
doc = tool.retrieved_docs[0]
assert {k: doc[k] for k in ("title", "source", "filename")} == expected
# The full chunk text, so the doc dedupes against the retriever's copy;
# only what the model reads is truncated.
assert doc["text"] == text
def test_a_page_with_no_metadata_falls_back_to_its_source_id(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": {}, "text": "body"}]))
tool.execute_action("read_entity_pages", entity="Quill")
assert tool.retrieved_docs[0]["source"] == "src-1"
def test_pages_absent_is_stated_plainly(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore(pages=[]))
assert "No documents" in tool.execute_action("read_entity_pages", entity="Quill")
class TestWiring:
def test_entry_exposes_every_action(self):
entry = build_graph_tool_entry()
assert {a["name"] for a in entry["actions"]} == {
"search_entities",
"get_relationships",
"read_entity_pages",
}
assert all(action["active"] for action in entry["actions"])
def test_not_added_when_graphs_are_disabled(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False)
monkeypatch.setattr(
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True
)
tools: dict = {}
add_graph_search_tool(tools, {"source": SOURCE})
assert tools == {}
def test_not_added_when_the_sources_have_no_graph(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
monkeypatch.setattr(
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False
)
tools: dict = {}
add_graph_search_tool(tools, {"source": SOURCE})
assert tools == {}
def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
monkeypatch.setattr(
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True
)
tools: dict = {}
add_graph_search_tool(tools, {"source": SOURCE})
assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID
assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE
class TestPlumbing:
def test_a_single_source_id_and_empty_entries_are_accepted(self):
tool = GraphSearchTool({"source": {"active_docs": "src-1"}})
assert tool._sources() == ["src-1"]
tool = GraphSearchTool({"source": {"active_docs": ["src-1", "", None]}})
assert tool._sources() == ["src-1"]
def test_the_store_is_built_once_and_reused(self, monkeypatch):
built = []
monkeypatch.setattr(
"docsgpt.graphrag.store.GraphStore", lambda: built.append(object()) or built[-1]
)
tool = GraphSearchTool({"source": SOURCE})
assert tool._get_store() is tool._get_store()
assert len(built) == 1
def test_an_embedding_failure_makes_entity_search_unavailable(self, monkeypatch):
def _broken_embeddings():
raise RuntimeError("no model")
monkeypatch.setattr("docsgpt.vectorstore.base.get_embeddings", _broken_embeddings)
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
tool = GraphSearchTool({"source": SOURCE})
tool._store = _StubStore()
assert tool.execute_action("search_entities", query="quill") == "Entity search is unavailable."
def test_no_matching_entities_is_stated_plainly(self, monkeypatch):
tool = _tool(monkeypatch, _StubStore())
assert "No entities found" in tool.execute_action("search_entities", query="quill")
def test_a_failing_store_is_reported_not_raised(self, monkeypatch):
class _BrokenStore(_StubStore):
def entity_relationships(self, source_id, name, limit=25):
raise RuntimeError("connection lost")
tool = _tool(monkeypatch, _BrokenStore())
assert tool.execute_action("get_relationships", entity="Quill") == "The graph lookup failed."
class TestSourcesHaveGraph:
"""Whether to offer the tool at all: only when some source has a graph."""
def _patch_counts(self, monkeypatch, counts=None, error=None):
class _Store:
def count_nodes_many(self, source_ids):
if error:
raise error
return {s: counts.get(s, 0) for s in source_ids}
monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", _Store)
def test_true_when_any_source_has_nodes(self, monkeypatch):
from docsgpt.agents.tools.graph_search import sources_have_graph
self._patch_counts(monkeypatch, {"b": 12})
assert sources_have_graph({"active_docs": ["a", "b"]}) is True
def test_false_when_no_source_has_nodes(self, monkeypatch):
from docsgpt.agents.tools.graph_search import sources_have_graph
self._patch_counts(monkeypatch, {})
assert sources_have_graph({"active_docs": "a"}) is False
def test_false_without_sources_or_when_the_check_fails(self, monkeypatch):
from docsgpt.agents.tools.graph_search import sources_have_graph
assert sources_have_graph({"active_docs": []}) is False
self._patch_counts(monkeypatch, error=RuntimeError("no pgvector"))
assert sources_have_graph({"active_docs": ["a"]}) is False
class TestGraphsMustBeAvailable:
"""The graph lives in the pgvector store, so the flag alone is not enough.
With another vector store configured the graph tables are not the ones the
sources were ingested into; everything else in the app asks
``graphrag_available()``, which requires both.
"""
def test_the_tool_is_not_offered_without_pgvector(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
monkeypatch.setattr(settings, "VECTOR_STORE", "faiss")
monkeypatch.setattr(
"docsgpt.agents.tools.graph_search.sources_have_graph",
lambda source: pytest.fail("must not reach the database"),
)
tools = {}
add_graph_search_tool(tools, {"source": SOURCE})
assert tools == {}
def test_actions_report_it_rather_than_querying(self, monkeypatch):
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
monkeypatch.setattr(settings, "VECTOR_STORE", "faiss")
tool = GraphSearchTool({"source": SOURCE})
tool._store = _StubStore(nodes=[{"name": "Quill", "distance": 0.1}])
assert "not enabled" in tool.execute_action("search_entities", query="quill")
class TestPooledConnection:
"""The tool is cached for the whole agent run; its connection must not be."""
class _ClosingStore(_StubStore):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.closed = 0
def close(self):
self.closed += 1
def test_the_connection_goes_back_after_each_action(self, monkeypatch):
store = self._ClosingStore(relationships=[{"source": "A", "target": "B", "type": "r"}])
tool = _tool(monkeypatch, store)
tool.execute_action("get_relationships", entity="A")
# Held open, one pooled connection would be pinned across every LLM
# round trip of the run.
assert store.closed == 1
assert tool._store is None
def test_a_failing_action_still_releases_it(self, monkeypatch):
class _Broken(self._ClosingStore):
def entity_relationships(self, source_id, name, limit=25):
raise RuntimeError("connection lost")
store = _Broken()
tool = _tool(monkeypatch, store)
assert tool.execute_action("get_relationships", entity="A") == "The graph lookup failed."
assert store.closed == 1
assert tool._store is None
+88
View File
@@ -0,0 +1,88 @@
"""Tests for canonical entity naming (the key graph nodes are merged on).
Two failure directions matter. Too little folding splits one entity across
nodes ("agent" / "agents", "VECTOR_STORE" / "vector stores"), so the walk never
connects what the text connects. Too much folding invents entities: stripping
the "s" off ``postgres`` or ``redis`` would merge nothing and create a node no
chunk ever named.
"""
from __future__ import annotations
import pytest
from docsgpt.graphrag.naming import canonical_name, normalize_entity_name
@pytest.mark.unit
class TestCanonicalName:
@pytest.mark.parametrize(
"variants, key",
[
(["VECTOR_STORE", "Vector store", "vector stores", "vector-stores"], "vector store"),
([".env file", "env_file", "ENV FILE"], "env file"),
(["Celery worker", "Celery workers"], "celery worker"),
(["agent", "Agents", "agents!"], "agent"),
],
)
def test_orthographic_variants_share_one_key(self, variants, key):
assert {canonical_name(v) for v in variants} == {key}
@pytest.mark.parametrize(
"singular, plural",
[
("policy", "policies"),
("index", "indexes"),
("batch", "batches"),
("hash", "hashes"),
("class", "classes"),
("process", "processes"),
("status", "statuses"),
("bus", "buses"),
("alias", "aliases"),
("document", "documents"),
("service", "services"),
# Singulars ending in "e" whose plural also ends in "-es": the
# plural alone cannot say whether to drop "s" or "es".
("cache", "caches"),
("database", "databases"),
("response", "responses"),
("release", "releases"),
("case", "cases"),
("size", "sizes"),
("cookie", "cookies"),
],
)
def test_singular_and_plural_share_one_key(self, singular, plural):
assert canonical_name(singular) == canonical_name(plural)
def test_the_key_need_not_be_a_word(self):
# It is a merge key, never shown: "cache" and "caches" meet at the
# stem an "-es" plural cannot see past, rather than guessing a form.
assert canonical_name("caches") == "cach"
assert canonical_name("batches") == "batch"
@pytest.mark.parametrize(
"word",
["postgres", "kubernetes", "redis", "https", "status", "analysis", "access", "docs", "series"],
)
def test_words_that_only_look_plural_are_left_alone(self, word):
assert canonical_name(word) == word
@pytest.mark.parametrize("word", ["class", "corpus", "thesis"])
def test_ss_us_is_endings_are_never_stripped(self, word):
assert canonical_name(word) == word
@pytest.mark.parametrize("word", ["aws", "ids", "ops"])
def test_short_words_are_left_alone(self, word):
assert canonical_name(word) == word
def test_each_word_of_a_phrase_is_folded(self):
assert canonical_name("Postgres Replicas") == "postgres replica"
@pytest.mark.parametrize("name", [None, "", " ", "!!!", "--_--"])
def test_a_name_with_nothing_left_is_no_entity(self, name):
assert canonical_name(name) == ""
def test_normalize_entity_name_is_the_canonical_key(self):
assert normalize_entity_name("Vector Stores") == canonical_name("Vector Stores")
@@ -0,0 +1,173 @@
"""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)
class TestVectorRanking:
"""The vector half of the blend, keyed on chunk text since hits carry no id."""
class _VectorStore:
def __init__(self, hits=None, error=None):
self.hits = hits or []
self.error = error
self.searched = None
self.closed = False
def search(self, question, k, query_vector=None):
self.searched = (question, k, query_vector)
if self.error:
raise self.error
return self.hits
def close(self):
self.closed = True
@staticmethod
def _real_retriever(monkeypatch, store):
from types import SimpleNamespace
retriever = object.__new__(GraphRAGRetriever)
retriever.chunks = 3
retriever._classic = SimpleNamespace(_get_rephrased_question=lambda: "where does Alder stream?")
monkeypatch.setattr(
"docsgpt.vectorstore.vector_creator.VectorCreator.create_vectorstore",
lambda *args, **kwargs: store,
)
return retriever
def test_object_and_dict_hits_become_text_and_metadata(self, monkeypatch):
from types import SimpleNamespace
store = self._VectorStore(
hits=[
SimpleNamespace(page_content="Alder streams to Quill.", metadata={"title": "alder.md"}),
{"text": "Quill is compacted every six hours.", "metadata": {"title": "quill.md"}},
{"page_content": "A passage without metadata."},
{"metadata": {"title": "no text"}},
]
)
retriever = self._real_retriever(monkeypatch, store)
ranked = retriever._vector_ranking("src", [0.1, 0.2])
assert ranked == [
("Alder streams to Quill.", {"title": "alder.md"}),
("Quill is compacted every six hours.", {"title": "quill.md"}),
("A passage without metadata.", {}),
]
# The rephrased question and the shared query vector, with room to fuse.
assert store.searched == ("where does Alder stream?", 20, [0.1, 0.2])
assert store.closed
def test_a_failed_search_ranks_nothing_and_still_closes_the_store(self, monkeypatch):
store = self._VectorStore(error=RuntimeError("pgvector down"))
retriever = self._real_retriever(monkeypatch, store)
assert retriever._vector_ranking("src", [0.1]) == []
assert store.closed
+151
View File
@@ -0,0 +1,151 @@
"""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) == {}
@pytest.mark.unit
class TestNodesByChunk:
"""The passage stage inverts node->chunks once instead of rescanning."""
def test_inverts_and_keeps_node_order(self):
from docsgpt.retriever.graph_rag import _nodes_by_chunk
assert _nodes_by_chunk({"n1": ["c1", "c2"], "n2": ["c2"], "n3": []}) == {
"c1": ["n1"],
"c2": ["n1", "n2"],
}
def test_no_links_invert_to_nothing(self):
from docsgpt.retriever.graph_rag import _nodes_by_chunk
assert _nodes_by_chunk({}) == {}
assert _nodes_by_chunk({"n1": None}) == {}
+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()
+663 -7
View File
@@ -19,6 +19,7 @@ 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
@@ -157,6 +158,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 +320,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
finally:
store.delete_by_source(source_id)
@@ -436,17 +558,63 @@ class TestGraphStoreParameterization:
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)
for call in cursor.execute.call_args_list:
sql = call.args[0]
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()
@@ -494,6 +662,217 @@ class TestGraphStoreParameterization:
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)."""
@@ -890,3 +1269,280 @@ class TestCountNodesMany:
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)
+229 -8
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 ────────────────────────────────────────────────────
@@ -157,7 +180,9 @@ class TestGraphRAGPoolDiscipline:
mock_store_cls.return_value = store
rag = _make_retriever()
with patch.object(rag, "_graph_docs_for_source", return_value=[]):
# A real result: an empty one now falls back like a failure does.
graph_docs = [{"title": "g", "text": "graph text", "source": "src1", "filename": "g"}]
with patch.object(rag, "_graph_docs_for_source", return_value=graph_docs):
with patch.object(rag, "_classic_for_sources") as classic:
rag._get_data()
@@ -332,15 +357,24 @@ class TestGraphRAGHappyPath:
@patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10)
@patch("docsgpt.retriever.graph_rag.GraphStore")
@patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True)
def test_no_seeds_returns_empty(
def test_a_graph_that_answers_nothing_falls_back_to_classic(
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed
):
"""Empty is not an answer. Every graph read swallows its own errors and
returns nothing, so "no rows" covers a broken query as much as a walk
that found nothing — and the source would contribute nothing at all,
with no fallback, because only a raise routes one to ClassicRAG."""
store = _store_with_graph([], [], {}, {}, [])
store.count_nodes_many.side_effect = lambda ids: {s: 5 for s in ids}
mock_store_cls.return_value = store
rag = _make_retriever()
assert rag._get_data() == []
seen = _recording_classic(rag, [_CLASSIC_DOC])
docs = rag._get_data()
assert seen == [["src1"]]
assert [doc["text"] for doc in docs] == ["classic"]
# ── IDF down-weighting ────────────────────────────────────────────────────────
@@ -449,11 +483,16 @@ class TestGetChunkTexts:
sid = str(uuid.uuid4())
store.get_chunk_texts(sid, ["1", "2"])
sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
assert f"FROM {table}" in sql
assert text_col in sql
assert metadata_col in sql
assert f"{source_col} = %s" in sql
from psycopg import sql as pgsql
query, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
# Identifiers are composed and quoted by psycopg, never formatted in.
assert isinstance(query, pgsql.Composable)
sql = query.as_string()
assert f'FROM "{table}"' in sql
assert f'"{text_col}"' in sql
assert f'"{metadata_col}"' in sql
assert f'"{source_col}" = %s' in sql
assert "id::text = ANY(%s)" in sql
assert sid not in sql
assert params == (sid, ["1", "2"])
@@ -800,3 +839,185 @@ class TestGraphRAGBatching:
assert rag._get_data() == []
mock_store_cls.assert_not_called()
rag._classic._get_data.assert_not_called()
# ── Personalized PageRank without scipy ──────────────────────────────────────
@pytest.fixture
def _no_scipy(monkeypatch):
"""Make ``import scipy`` fail, as it does in a default install.
``scipy`` is not a DocsGPT dependency — it only reaches this test env
through the optional docling extra. ``networkx.pagerank`` delegates to its
scipy implementation, so ranking must not go through it.
"""
import sys
for name in [m for m in list(sys.modules) if m == "scipy" or m.startswith("scipy.")]:
monkeypatch.delitem(sys.modules, name)
monkeypatch.setitem(sys.modules, "scipy", None)
def _chain_graph():
"""Weighted chain a-b-c-d plus a heavier shortcut a-d."""
import networkx as nx
graph = nx.Graph()
graph.add_weighted_edges_from(
[("a", "b", 1.0), ("b", "c", 2.0), ("c", "d", 1.0), ("a", "d", 0.5)]
)
return graph
@pytest.mark.unit
class TestPersonalizedPageRankWithoutScipy:
def test_ranking_runs_when_scipy_is_missing(self, _no_scipy):
from docsgpt.retriever.graph_rag import _personalized_pagerank
graph = _chain_graph()
ranks = _personalized_pagerank(
graph, personalization={"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
)
assert set(ranks) == {"a", "b", "c", "d"}
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
assert all(rank > 0 for rank in ranks.values())
# Pinned from the parity test below, which runs the library
# implementation over the same graph while scipy is installed here.
assert sorted(ranks, key=ranks.get, reverse=True) == ["b", "a", "c", "d"]
# The seed outranks the node furthest from it along the heavy path.
assert ranks["a"] > ranks["d"]
def test_matches_networkx_within_tolerance(self):
"""Parity with the library implementation, while it is installed here."""
import networkx as nx
pytest.importorskip("scipy")
from docsgpt.retriever.graph_rag import _personalized_pagerank
graph = _chain_graph()
personalization = {"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
ours = _personalized_pagerank(graph, personalization=personalization)
theirs = nx.pagerank(graph, personalization=personalization, weight="weight")
for node in theirs:
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
def test_uniform_personalization_when_none(self):
pytest.importorskip("scipy")
import networkx as nx
from docsgpt.retriever.graph_rag import _personalized_pagerank
graph = _chain_graph()
ours = _personalized_pagerank(graph, personalization=None)
theirs = nx.pagerank(graph, personalization=None, weight="weight")
for node in theirs:
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
def test_isolated_node_still_gets_mass(self):
"""A node with no edges is dangling; its mass must not vanish."""
import networkx as nx
from docsgpt.retriever.graph_rag import _personalized_pagerank
graph = nx.Graph()
graph.add_edge("a", "b", weight=1.0)
graph.add_node("lonely")
ranks = _personalized_pagerank(graph, personalization=None)
assert ranks["lonely"] > 0
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
def test_empty_graph_returns_empty(self):
import networkx as nx
from docsgpt.retriever.graph_rag import _personalized_pagerank
assert _personalized_pagerank(nx.Graph(), personalization=None) == {}
def test_a_zero_weight_edge_is_not_traversable(self):
"""Zero means "not related", not "use the default weight"."""
import networkx as nx
from docsgpt.retriever.graph_rag import _personalized_pagerank
graph = nx.Graph()
graph.add_edge("seed", "zero", weight=0.0)
graph.add_edge("seed", "real", weight=1.0)
ranks = _personalized_pagerank(
graph, personalization={"seed": 1.0, "zero": 0.0, "real": 0.0}
)
# ``zero`` is reachable only across the zero-weight edge, so no mass
# walks to it; ``real`` is on a live edge and must outrank it.
assert ranks["real"] > ranks["zero"]
assert ranks["zero"] == pytest.approx(0.0, abs=1e-9)
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
def test_stored_zero_weights_reach_the_ranker_intact(self):
"""The subgraph builder must not coerce a stored 0 into a real edge.
Without this the ranker's zero-weight rule is unreachable in
production: every 0 from ``graph_edges`` arrives as 1.0.
"""
subgraph = {
"nodes": [
{"id": "seed", "doc_freq": 1},
{"id": "zero", "doc_freq": 1},
{"id": "real", "doc_freq": 1},
],
"edges": [
{"src_node_id": "seed", "dst_node_id": "zero", "weight": 0},
{"src_node_id": "seed", "dst_node_id": "real", "weight": 1.0},
],
}
# Called unbound with ``None`` for self: _ppr_scores reads no state.
scores = GraphRAGRetriever._ppr_scores(None, subgraph, {"seed": 1.0})
assert scores["real"] > scores["zero"]
assert scores["zero"] == pytest.approx(0.0, abs=1e-9)
def test_missing_and_null_weights_default_to_one(self):
import networkx as nx
from docsgpt.retriever.graph_rag import _personalized_pagerank
absent = nx.Graph()
absent.add_edge("a", "b") # no weight attribute at all
null = nx.Graph()
null.add_edge("a", "b", weight=None)
personalization = {"a": 1.0, "b": 0.0}
from_absent = _personalized_pagerank(absent, personalization=personalization)
from_null = _personalized_pagerank(null, personalization=personalization)
assert from_absent["b"] == pytest.approx(from_null["b"], abs=1e-9)
assert from_absent["b"] > 0
@patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10)
@patch("docsgpt.retriever.graph_rag.GraphStore")
@patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True)
def test_graph_retrieval_does_not_fall_back_without_scipy(
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed, _no_scipy
):
"""The whole PPR path runs with scipy absent — no ClassicRAG fallback."""
nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}]
edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}]
node_chunks = {"n1": ["c1"], "n2": ["c2"]}
chunk_texts = {"c1": "near", "c2": "far"}
seed_rows = [{"id": "n1", "distance": 0.0}]
store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows)
mock_store_cls.return_value = store
rag = _make_retriever(chunks=2)
rag._classic_for_sources = Mock(side_effect=AssertionError("fell back"))
docs = rag._get_data()
assert [doc["text"] for doc in docs] == ["near", "far"]
+74
View File
@@ -276,3 +276,77 @@ class TestReclaimIsSkippedForEmbeds:
from docsgpt.vectorstore.embeddings_delegated import EMBED_TASK
assert EMBED_TASK in _NO_RECLAIM_TASKS
@pytest.mark.unit
class TestInWorker:
"""Whether code runs inside a worker must not depend on which thread asks.
Celery records the executing task on the thread that runs it, so a thread
that task starts sees no task at all. Code deciding "am I in the worker?"
from that alone takes the web-process branch there: it dispatches to the
worker it is running in and blocks on the result, which Celery refuses
("Never call result.get() within a task!") or, where joins are allowed,
waits on a queue only this busy process serves.
"""
@staticmethod
def _ask_from_a_new_thread():
import threading
from docsgpt.celery_init import in_worker
seen = []
thread = threading.Thread(target=lambda: seen.append(in_worker()))
thread.start()
thread.join()
return seen[0]
def test_false_outside_a_worker(self):
from docsgpt.celery_init import in_worker
assert in_worker() is False
assert self._ask_from_a_new_thread() is False
def test_true_on_a_thread_started_inside_a_worker(self):
# Blocking pools (prefork, solo, threads) mark the whole process as one
# where joining a task would block; ``denied_join_result`` sets exactly
# that flag.
from celery.result import denied_join_result
with denied_join_result():
assert self._ask_from_a_new_thread() is True
def test_true_on_any_thread_of_a_non_blocking_pool_worker(self, monkeypatch):
# eventlet/gevent leave the join flag unset and scope the current task
# to one greenlet, so only the worker's own startup can say this
# process is a worker. The lifecycle signal records that for every
# thread and greenlet in it.
from celery.signals import worker_init
monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False)
assert self._ask_from_a_new_thread() is False
worker_init.send(sender=None)
assert self._ask_from_a_new_thread() is True
def test_prefork_children_record_it_on_their_own_start(self, monkeypatch):
from celery.signals import worker_process_init
monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False)
worker_process_init.send(sender=None)
assert self._ask_from_a_new_thread() is True
def test_true_on_the_task_thread_of_a_non_blocking_pool(self):
# eventlet/gevent pools leave the process flag unset; the thread
# running the task still knows it is in one.
from unittest.mock import PropertyMock
from docsgpt.celery_init import celery, in_worker
with patch.object(
type(celery), "current_worker_task", new_callable=PropertyMock, return_value=object()
):
assert in_worker() is True
+29
View File
@@ -74,6 +74,35 @@ class TestDispatcherGrouping:
assert "b" not in retrievals
def test_graph_options_count_as_an_override(self, _patch_llm_creator):
"""A graph source that changes only its graph options still needs its
config carried over: those options live on the per-source retrieval the
Dispatcher hands the retriever, so without this the UI toggles are
no-ops and every source runs the defaults."""
sources = [
{
"id": "a",
"retrieval": RetrievalConfig(
retriever="graphrag", graph={"seed_strategy": "relationships"}
),
},
{"id": "b", "retrieval": RetrievalConfig(retriever="graphrag")},
]
d = Dispatcher(source={"question": "q", "active_docs": ["a", "b"]}, sources=sources)
retrievals = d._groups[0]["retrievals"]
assert "a" in retrievals
assert retrievals["a"].graph.seed_strategy == "relationships"
# A source on the defaults still takes the shared path.
assert "b" not in retrievals
def test_graph_options_on_a_classic_source_are_not_an_override(self, _patch_llm_creator):
# They only mean anything to the graph retriever; treating them as an
# override would hand a classic source its own chunk budget.
sources = [{"id": "a", "retrieval": RetrievalConfig(graph={"blend_vector": False})}]
d = Dispatcher(source={"question": "q", "active_docs": ["a"]}, sources=sources)
assert d._groups[0]["retrievals"] == {}
@pytest.mark.unit
class TestDispatcherSharedBudget:
def test_single_group_full_budget(self, _patch_llm_creator):
@@ -72,6 +72,32 @@ class TestInsideAWorker:
assert vector == [1.0, 2.0]
celery.send_task.assert_not_called()
def test_a_thread_started_inside_the_worker_embeds_locally(self):
"""The task's own thread is not the only one in a worker.
Graph extraction and per-source retrieval both fan out to thread pools
inside tasks. The check used to read the task off the current thread
only, so from those threads it dispatched to the worker it was running
in -- and Celery refuses that ``get()`` inside a worker, failing the
call and latching the 30s dispatch cooldown for every caller after it.
"""
from celery.result import denied_join_result
from docsgpt.celery_init import celery
local = MagicMock()
local.embed_documents.return_value = [[1.0, 2.0]]
client = DelegatedEmbeddings("some/model")
vectors = []
with denied_join_result():
with patch("docsgpt.vectorstore.base.build_local_embeddings", return_value=local):
with patch.object(celery, "send_task") as send_task:
thread = threading.Thread(target=lambda: vectors.append(client.embed_query("hi")))
thread.start()
thread.join()
assert vectors == [[1.0, 2.0]]
send_task.assert_not_called()
def test_the_local_model_is_built_once(self):
local = MagicMock()
local.embed_documents.return_value = [[1.0]]