diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx
index 32181adb..3221867f 100644
--- a/docs/content/Deploying/Settings-Reference.mdx
+++ b/docs/content/Deploying/Settings-Reference.mdx
@@ -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
diff --git a/docsgpt/agents/agentic_agent.py b/docsgpt/agents/agentic_agent.py
index b83c485d..0b87485c 100644
--- a/docsgpt/agents/agentic_agent.py
+++ b/docsgpt/agents/agentic_agent.py
@@ -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)
diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py
index d731a415..030d76fc 100644
--- a/docsgpt/agents/base.py
+++ b/docsgpt/agents/base.py
@@ -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)
diff --git a/docsgpt/agents/classic_agent.py b/docsgpt/agents/classic_agent.py
index 2bc25130..5956ed3c 100644
--- a/docsgpt/agents/classic_agent.py
+++ b/docsgpt/agents/classic_agent.py
@@ -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)
diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py
index 1d19a7a9..de96cf35 100644
--- a/docsgpt/agents/research_agent.py
+++ b/docsgpt/agents/research_agent.py
@@ -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
diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py
new file mode 100644
index 00000000..b370949e
--- /dev/null
+++ b/docsgpt/agents/tools/graph_search.py
@@ -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}
diff --git a/docsgpt/agents/tools/read_document.py b/docsgpt/agents/tools/read_document.py
index ab9d2294..94e488df 100644
--- a/docsgpt/agents/tools/read_document.py
+++ b/docsgpt/agents/tools/read_document.py
@@ -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:
diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py
index 1381f241..1cfc1b80 100644
--- a/docsgpt/api/user/idempotency.py
+++ b/docsgpt/api/user/idempotency.py
@@ -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(
diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py
index 4272eac1..2eda5a3a 100644
--- a/docsgpt/celery_init.py
+++ b/docsgpt/celery_init.py
@@ -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``.
diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py
index 651d3f3d..53dbf56b 100644
--- a/docsgpt/core/settings/retrieval.py
+++ b/docsgpt/core/settings/retrieval.py
@@ -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
diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py
index beaf3d66..daa8f910 100644
--- a/docsgpt/graphrag/extraction.py
+++ b/docsgpt/graphrag/extraction.py
@@ -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"\n{text}\n"},
@@ -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 {}
diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py
new file mode 100644
index 00000000..2a173eea
--- /dev/null
+++ b/docsgpt/graphrag/naming.py
@@ -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)
diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py
index 818004c0..0b73d031 100644
--- a/docsgpt/graphrag/store.py
+++ b/docsgpt/graphrag/store.py
@@ -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:
diff --git a/docsgpt/retriever/dispatcher.py b/docsgpt/retriever/dispatcher.py
index a14fc02f..31cc75e0 100644
--- a/docsgpt/retriever/dispatcher.py
+++ b/docsgpt/retriever/dispatcher.py
@@ -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
diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py
index c7ac538e..8d191f7b 100644
--- a/docsgpt/retriever/graph_rag.py
+++ b/docsgpt/retriever/graph_rag.py
@@ -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]
diff --git a/docsgpt/storage/db/source_config.py b/docsgpt/storage/db/source_config.py
index 20873b46..8f8d6209 100644
--- a/docsgpt/storage/db/source_config.py
+++ b/docsgpt/storage/db/source_config.py
@@ -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
diff --git a/docsgpt/vectorstore/embeddings_delegated.py b/docsgpt/vectorstore/embeddings_delegated.py
index 75f730b8..bc184f34 100644
--- a/docsgpt/vectorstore/embeddings_delegated.py
+++ b/docsgpt/vectorstore/embeddings_delegated.py
@@ -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
diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json
index e64ed52f..b914fbfd 100644
--- a/frontend/src/locale/de.json
+++ b/frontend/src/locale/de.json
@@ -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.",
diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json
index 03f0f0dd..ba67a1cb 100644
--- a/frontend/src/locale/en.json
+++ b/frontend/src/locale/en.json
@@ -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.",
diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json
index c7b4f794..f30cadd7 100644
--- a/frontend/src/locale/es.json
+++ b/frontend/src/locale/es.json
@@ -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.",
diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json
index dc022fe8..e475f954 100644
--- a/frontend/src/locale/jp.json
+++ b/frontend/src/locale/jp.json
@@ -268,6 +268,19 @@
},
"exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。"
},
+ "graphRetrieval": {
+ "title": "グラフ検索",
+ "tag": "再取り込み不要",
+ "seedStrategy": "探索の開始点",
+ "seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。",
+ "seedEntities": "エンティティ(推奨)",
+ "seedRelationships": "リレーション",
+ "passageNodes": "パッセージを探索に含める",
+ "passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。",
+ "blendVector": "ベクトル検索と組み合わせる",
+ "blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。",
+ "agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。"
+ },
"prescreen": {
"enable": "LLMプリスクリーニングを有効にする",
"warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。",
diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json
index dd7aef1e..ba064d54 100644
--- a/frontend/src/locale/ru.json
+++ b/frontend/src/locale/ru.json
@@ -268,6 +268,19 @@
},
"exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту."
},
+ "graphRetrieval": {
+ "title": "Поиск по графу",
+ "tag": "без повторной загрузки",
+ "seedStrategy": "Начинать обход с",
+ "seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.",
+ "seedEntities": "Сущностей (рекомендуется)",
+ "seedRelationships": "Связей",
+ "passageNodes": "Включать фрагменты в обход",
+ "passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.",
+ "blendVector": "Сочетать с векторным поиском",
+ "blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.",
+ "agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент."
+ },
"prescreen": {
"enable": "Включить предварительный отбор LLM",
"warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.",
diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json
index 16236cec..c8c3e7e1 100644
--- a/frontend/src/locale/zh-TW.json
+++ b/frontend/src/locale/zh-TW.json
@@ -268,6 +268,19 @@
},
"exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。"
},
+ "graphRetrieval": {
+ "title": "圖譜檢索",
+ "tag": "無需重新匯入",
+ "seedStrategy": "走訪起點",
+ "seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。",
+ "seedEntities": "實體(建議)",
+ "seedRelationships": "關係",
+ "passageNodes": "將段落納入走訪",
+ "passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。",
+ "blendVector": "與向量檢索結合",
+ "blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。",
+ "agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。"
+ },
"prescreen": {
"enable": "啟用 LLM 預篩選",
"warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。",
diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json
index 5b4cf5a6..40e8fa0e 100644
--- a/frontend/src/locale/zh.json
+++ b/frontend/src/locale/zh.json
@@ -268,6 +268,19 @@
},
"exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。"
},
+ "graphRetrieval": {
+ "title": "图谱检索",
+ "tag": "无需重新导入",
+ "seedStrategy": "遍历起点",
+ "seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。",
+ "seedEntities": "实体(推荐)",
+ "seedRelationships": "关系",
+ "passageNodes": "将段落纳入遍历",
+ "passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。",
+ "blendVector": "与向量检索结合",
+ "blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。",
+ "agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。"
+ },
"prescreen": {
"enable": "启用 LLM 预筛选",
"warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。",
diff --git a/frontend/src/models/misc.ts b/frontend/src/models/misc.ts
index fc8b0782..c0c5efb2 100644
--- a/frontend/src/models/misc.ts
+++ b/frontend/src/models/misc.ts
@@ -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').
diff --git a/frontend/src/settings/components/RetrievalOptions.test.tsx b/frontend/src/settings/components/RetrievalOptions.test.tsx
index 92142b23..264c1794 100644
--- a/frontend/src/settings/components/RetrievalOptions.test.tsx
+++ b/frontend/src/settings/components/RetrievalOptions.test.tsx
@@ -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,
diff --git a/frontend/src/settings/components/RetrievalOptions.tsx b/frontend/src/settings/components/RetrievalOptions.tsx
index e2723358..b9dec5d6 100644
--- a/frontend/src/settings/components/RetrievalOptions.tsx
+++ b/frontend/src/settings/components/RetrievalOptions.tsx
@@ -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,
+ ) => {
+ setRetrieval({ graph: { ...value.retrieval.graph, ...patch } });
+ };
+
const modelOptions = useMemo(() => {
const builtin: Model[] = [];
const user: Model[] = [];
@@ -611,6 +643,84 @@ export default function RetrievalOptions({
)}
+ {/* Graph retrieval group (graphrag only; live, so shown when testing too) */}
+ {isGraphRAG && (
+