mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
Merge pull request #2804 from arc53/fix/graphrag-extraction-and-retrieval
fix(graphrag): make graph retrieval and extraction work in a default install
This commit is contained in:
43 files changed
+4821
-219
No files matched your search
@@ -447,6 +447,12 @@ Type `int`, default `2000`, must be `>= 0`.
|
||||
|
||||
Hard cap on chunks extracted per source (cost control); 0 extracts nothing.
|
||||
|
||||
### `GRAPHRAG_EXTRACTION_WORKERS`
|
||||
|
||||
Type `int`, default `8`, must be `>= 1` and `<= 32`.
|
||||
|
||||
Concurrent extraction calls during ingest. Model calls run in parallel while graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial.
|
||||
|
||||
|
||||
## Vector stores
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
from typing import Dict, Generator, Optional
|
||||
|
||||
from docsgpt.agents.base import BaseAgent
|
||||
from docsgpt.agents.tools.graph_search import add_graph_search_tool
|
||||
from docsgpt.agents.tools.internal_search import add_internal_search_tool
|
||||
from docsgpt.agents.tools.wiki import add_wiki_tool
|
||||
from docsgpt.logging import LogContext
|
||||
@@ -33,6 +34,7 @@ class AgenticAgent(BaseAgent):
|
||||
) -> Generator[Dict, None, None]:
|
||||
tools_dict = self.tool_executor.get_tools()
|
||||
add_internal_search_tool(tools_dict, self.retriever_config)
|
||||
add_graph_search_tool(tools_dict, self.retriever_config)
|
||||
if self.wiki_config:
|
||||
add_wiki_tool(tools_dict, self.wiki_config)
|
||||
self._prepare_tools(tools_dict)
|
||||
|
||||
+21
-7
@@ -955,16 +955,30 @@ class BaseAgent(ABC):
|
||||
)
|
||||
self.retrieved_docs = scrubbed
|
||||
|
||||
def _collect_internal_sources(self) -> None:
|
||||
"""Merge the cached InternalSearchTool's docs into ``retrieved_docs``,
|
||||
deduped, preserving any pre-fetched docs so a mixed-exposure agent cites
|
||||
both pre-fetched and tool-retrieved sources (not just the tool's)."""
|
||||
def _search_tool_docs(self) -> List[Dict]:
|
||||
"""Documents this run's search tools read: internal search and the graph tool.
|
||||
|
||||
Both record what they surface in ``retrieved_docs``; a page read from
|
||||
the graph carries the answer as much as a search hit does, so both are
|
||||
cited. Tools are looked up the way the executor caches them.
|
||||
"""
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID
|
||||
|
||||
executor = getattr(self, "tool_executor", None)
|
||||
loaded = getattr(executor, "_loaded_tools", None) or {}
|
||||
tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}")
|
||||
if not (tool and getattr(tool, "retrieved_docs", None)):
|
||||
docs: List[Dict] = []
|
||||
for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)):
|
||||
tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}")
|
||||
docs.extend(getattr(tool, "retrieved_docs", None) or [])
|
||||
return docs
|
||||
|
||||
def _collect_internal_sources(self) -> None:
|
||||
"""Merge the search tools' docs into ``retrieved_docs``, deduped,
|
||||
preserving any pre-fetched docs so a mixed-exposure agent cites both
|
||||
pre-fetched and tool-retrieved sources (not just the tools')."""
|
||||
tool_docs = self._search_tool_docs()
|
||||
if not tool_docs:
|
||||
return
|
||||
|
||||
def _key(d):
|
||||
@@ -974,7 +988,7 @@ class BaseAgent(ABC):
|
||||
|
||||
merged = list(self.retrieved_docs or [])
|
||||
seen = {_key(d) for d in merged}
|
||||
for doc in tool.retrieved_docs:
|
||||
for doc in tool_docs:
|
||||
k = _key(doc)
|
||||
if k not in seen:
|
||||
seen.add(k)
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
from typing import Dict, Generator, Optional
|
||||
|
||||
from docsgpt.agents.base import BaseAgent
|
||||
from docsgpt.agents.tools.graph_search import add_graph_search_tool
|
||||
from docsgpt.agents.tools.internal_search import add_internal_search_tool
|
||||
from docsgpt.agents.tools.wiki import add_wiki_tool
|
||||
from docsgpt.logging import LogContext
|
||||
@@ -38,6 +39,7 @@ class ClassicAgent(BaseAgent):
|
||||
tools_dict = self.tool_executor.get_tools()
|
||||
if self.retriever_config:
|
||||
add_internal_search_tool(tools_dict, self.retriever_config)
|
||||
add_graph_search_tool(tools_dict, self.retriever_config)
|
||||
if self.wiki_config:
|
||||
add_wiki_tool(tools_dict, self.wiki_config)
|
||||
self._prepare_tools(tools_dict)
|
||||
|
||||
@@ -6,10 +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
|
||||
|
||||
@@ -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}
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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``.
|
||||
|
||||
|
||||
|
||||
@@ -31,6 +31,15 @@ class RetrievalSettings(SettingsGroup):
|
||||
GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field(
|
||||
default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing."
|
||||
)
|
||||
GRAPHRAG_EXTRACTION_WORKERS: int = Field(
|
||||
default=8,
|
||||
ge=1,
|
||||
le=32,
|
||||
description=(
|
||||
"Concurrent extraction calls during ingest. Model calls run in parallel while "
|
||||
"graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial."
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator("VECTOR_STORE", mode="before")
|
||||
@classmethod
|
||||
|
||||
+261
-38
@@ -23,6 +23,11 @@ import logging
|
||||
import re
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
from docsgpt.core.model_utils import (
|
||||
get_api_key_for_provider,
|
||||
get_provider_from_model_id,
|
||||
)
|
||||
from docsgpt.graphrag.naming import normalize_entity_name
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.llm.llm_creator import LLMCreator
|
||||
from docsgpt.storage.db.source_config import SourceConfig
|
||||
@@ -63,14 +68,37 @@ def _resolve_max_chunks(config: SourceConfig) -> int:
|
||||
return config.graph.max_chunks or settings.GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION
|
||||
|
||||
|
||||
def _resolve_extraction_provider(
|
||||
model_id: Optional[str], user: Optional[str]
|
||||
) -> str:
|
||||
"""The provider that serves ``model_id``, else the deployment default.
|
||||
|
||||
``settings.LLM_PROVIDER`` is only a default (``docsgpt``, the hosted public
|
||||
endpoint, out of the box). Dispatching the resolved extraction model
|
||||
through it sends the request to a provider that does not serve that model:
|
||||
the call is rejected, the shared fallback answers instead, and the graph is
|
||||
built by a different model than the one configured — with nothing in the
|
||||
summary to say so. ``user`` scopes the lookup so a per-user (BYOM) model id
|
||||
resolves as well.
|
||||
"""
|
||||
provider = (
|
||||
get_provider_from_model_id(model_id, user_id=user) if model_id else None
|
||||
)
|
||||
return provider or settings.LLM_PROVIDER
|
||||
|
||||
|
||||
def _build_extraction_llm(
|
||||
model_id: Optional[str], user: Optional[str], request_id: Optional[str]
|
||||
):
|
||||
"""Build the extraction LLM tagged for token-usage attribution to the owner."""
|
||||
decoded_token = {"sub": user} if user else None
|
||||
provider = _resolve_extraction_provider(model_id, user)
|
||||
logger.info(
|
||||
"Graph extraction dispatching model=%s via provider=%s", model_id, provider
|
||||
)
|
||||
llm = LLMCreator.create_llm(
|
||||
settings.LLM_PROVIDER,
|
||||
api_key=settings.API_KEY,
|
||||
provider,
|
||||
api_key=get_api_key_for_provider(provider),
|
||||
user_api_key=None,
|
||||
decoded_token=decoded_token,
|
||||
model_id=model_id,
|
||||
@@ -122,8 +150,15 @@ def _parse_extraction(raw: Any) -> Optional[Dict[str, List[Dict[str, Any]]]]:
|
||||
}
|
||||
|
||||
|
||||
def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]:
|
||||
"""Run exactly one extraction call for a chunk (gleanings off)."""
|
||||
def _extract_chunk(
|
||||
llm, text: str, chunk_id: Optional[str] = None
|
||||
) -> Optional[Dict[str, List[Dict[str, Any]]]]:
|
||||
"""Run exactly one extraction call for a chunk (gleanings off).
|
||||
|
||||
Both failure modes name the chunk: an unparseable response used to return
|
||||
``None`` silently, so a graph could come back short with nothing in the
|
||||
logs to say which chunk was dropped or why.
|
||||
"""
|
||||
messages = [
|
||||
{"role": "system", "content": _SYSTEM_PROMPT},
|
||||
{"role": "user", "content": f"<chunk>\n{text}\n</chunk>"},
|
||||
@@ -134,9 +169,17 @@ def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]:
|
||||
messages=messages,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("Graph extraction call failed, skipping chunk: %s", exc)
|
||||
logger.warning(
|
||||
"Graph extraction call failed for chunk %s: %s", chunk_id, exc
|
||||
)
|
||||
return None
|
||||
return _parse_extraction(response)
|
||||
parsed = _parse_extraction(response)
|
||||
if parsed is None:
|
||||
logger.warning(
|
||||
"Graph extraction returned unparseable output for chunk %s.",
|
||||
chunk_id,
|
||||
)
|
||||
return parsed
|
||||
|
||||
|
||||
def _coerce_weight(value: Any) -> float:
|
||||
@@ -160,8 +203,10 @@ def extract_graph_for_source(
|
||||
Resumable and idempotent: chunks already marked ``done`` are skipped via the
|
||||
``graph_ingest_progress`` checkpoint, so a retry never re-extracts (and never
|
||||
re-bills). Processes at most the resolved chunk cap; excess chunks are
|
||||
reported under ``skipped_over_cap``. A malformed response or an LLM error on
|
||||
a single chunk marks it ``failed`` and continues — the pipeline never crashes.
|
||||
reported under ``skipped_over_cap``. A malformed response, an LLM error or a
|
||||
failed write on a single chunk is retried once after the rest of the build;
|
||||
a chunk that fails again is marked ``failed`` and the run continues — the
|
||||
pipeline never crashes.
|
||||
|
||||
Each chunk is written in a single transaction with one batched embedding
|
||||
call (entity + relationship-endpoint names together).
|
||||
@@ -178,8 +223,13 @@ def extract_graph_for_source(
|
||||
|
||||
Returns:
|
||||
A summary ``{nodes, edges, chunks_processed, skipped_over_cap,
|
||||
failed_chunks}``.
|
||||
failed_chunks}``, where ``nodes`` is how many distinct nodes the
|
||||
source's graph holds after the run — not how many upserts ran, which
|
||||
counts the same entity once per chunk it appears in.
|
||||
"""
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from docsgpt.graphrag.store import GraphStore
|
||||
|
||||
store = GraphStore()
|
||||
@@ -197,11 +247,26 @@ def extract_graph_for_source(
|
||||
|
||||
embedding = get_embeddings()
|
||||
|
||||
llm = _build_extraction_llm(
|
||||
_resolve_extraction_model(config), user, request_id
|
||||
)
|
||||
model_id = _resolve_extraction_model(config)
|
||||
# Built here first so a misconfigured model fails the run before any
|
||||
# chunk is touched; this instance serves the calling thread.
|
||||
thread_llm = threading.local()
|
||||
thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
|
||||
|
||||
nodes = 0
|
||||
def _llm():
|
||||
"""This thread's extraction LLM.
|
||||
|
||||
Provider-reported usage is kept on the LLM instance (``_last_usage``)
|
||||
and claimed by whichever call finishes next, so two calls in flight on
|
||||
one instance can bill each other's tokens. Each pool thread therefore
|
||||
builds its own.
|
||||
"""
|
||||
llm = getattr(thread_llm, "llm", None)
|
||||
if llm is None:
|
||||
llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
|
||||
return llm
|
||||
|
||||
node_upserts = 0
|
||||
edges = 0
|
||||
chunks_processed = 0
|
||||
failed_chunks = 0
|
||||
@@ -215,54 +280,140 @@ def extract_graph_for_source(
|
||||
{
|
||||
"current": chunks_processed + failed_chunks,
|
||||
"total": total,
|
||||
"nodes": nodes,
|
||||
"nodes": node_upserts,
|
||||
"edges": edges,
|
||||
}
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("graph progress callback failed: %s", exc)
|
||||
|
||||
for chunk, chunk_id in to_process:
|
||||
def _prepare(item):
|
||||
"""One chunk's LLM extraction — the only step run concurrently.
|
||||
|
||||
A chunk spends almost all of its time waiting on the model, so that is
|
||||
what runs in the pool. Graph writes and embedding stay on the calling
|
||||
thread, so transactions and the progress checkpoint are exactly what
|
||||
they were serially and the pool never touches the embeddings client.
|
||||
"""
|
||||
chunk, chunk_id = item
|
||||
text = _chunk_text(chunk)
|
||||
if not text:
|
||||
store.mark_chunk(source_id, chunk_id, "done")
|
||||
chunks_processed += 1
|
||||
_report()
|
||||
continue
|
||||
return chunk_id, "empty", None
|
||||
|
||||
extracted = _extract_chunk(llm, text)
|
||||
extracted = _extract_chunk(_llm(), text, chunk_id)
|
||||
if extracted is None:
|
||||
store.mark_chunk(source_id, chunk_id, "failed")
|
||||
failed_chunks += 1
|
||||
_report()
|
||||
continue
|
||||
|
||||
return chunk_id, "failed", None
|
||||
try:
|
||||
entities = _build_entities(extracted["entities"])
|
||||
relationships = _build_relationships(extracted["relationships"])
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Graph extraction failed for chunk %s: %s", chunk_id, exc
|
||||
)
|
||||
return chunk_id, "failed", None
|
||||
return chunk_id, "ok", (entities, relationships)
|
||||
|
||||
def _write(chunk_id, status, payload) -> bool:
|
||||
"""Apply one prepared chunk to the graph; False when it did not land."""
|
||||
nonlocal node_upserts, edges, chunks_processed
|
||||
if status == "empty":
|
||||
store.mark_chunk(source_id, chunk_id, "done")
|
||||
chunks_processed += 1
|
||||
return True
|
||||
if status == "failed":
|
||||
return False
|
||||
|
||||
entities, relationships = payload
|
||||
try:
|
||||
name_embeddings = _embed_names(embedding, entities, relationships)
|
||||
_embed_facts(embedding, relationships)
|
||||
chunk_nodes, chunk_edges = store.apply_chunk(
|
||||
source_id, chunk_id, entities, relationships, name_embeddings
|
||||
)
|
||||
nodes += chunk_nodes
|
||||
edges += chunk_edges
|
||||
store.mark_chunk(source_id, chunk_id, "done")
|
||||
chunks_processed += 1
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Graph extraction write failed for chunk %s, skipping: %s",
|
||||
chunk_id,
|
||||
exc,
|
||||
"Graph extraction embed/write failed for chunk %s: %s", chunk_id, exc
|
||||
)
|
||||
return False
|
||||
# ``apply_chunk`` marks the chunk done inside the transaction that
|
||||
# writes its rows, so the checkpoint cannot disagree with the graph
|
||||
# and a replayed write cannot apply the chunk twice.
|
||||
node_upserts += chunk_nodes
|
||||
edges += chunk_edges
|
||||
chunks_processed += 1
|
||||
return True
|
||||
|
||||
workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1))
|
||||
pool = None
|
||||
if workers > 1 and len(to_process) > 1:
|
||||
pool = ThreadPoolExecutor(max_workers=workers)
|
||||
|
||||
def _pass(items):
|
||||
"""Extract and write ``items``; return the ones that did not land."""
|
||||
if pool is not None:
|
||||
# ``map`` yields in submission order, so chunks are still applied in
|
||||
# the order they were given and a run stays reproducible.
|
||||
prepared = pool.map(_prepare, items)
|
||||
else:
|
||||
prepared = (_prepare(item) for item in items)
|
||||
missed = []
|
||||
for item, (chunk_id, status, payload) in zip(items, prepared):
|
||||
if not _write(chunk_id, status, payload):
|
||||
missed.append(item)
|
||||
_report()
|
||||
return missed
|
||||
|
||||
try:
|
||||
missed = _pass(to_process)
|
||||
if missed:
|
||||
# A failure is usually transient — a provider error, one response
|
||||
# that did not parse — and the checkpoint only picks it up on a
|
||||
# rerun nothing schedules. One more attempt, after the rest of the
|
||||
# build so a burst of rate limiting has passed, and no more: a
|
||||
# chunk that cannot be extracted costs at most two calls.
|
||||
logger.info(
|
||||
"Graph extraction retrying %d failed chunk(s) for source %s",
|
||||
len(missed),
|
||||
source_id,
|
||||
)
|
||||
missed = _pass(missed)
|
||||
for _, chunk_id in missed:
|
||||
store.mark_chunk(source_id, chunk_id, "failed")
|
||||
failed_chunks += 1
|
||||
_report()
|
||||
failed_chunks = len(missed)
|
||||
if missed:
|
||||
logger.warning(
|
||||
"Graph extraction gave up on %d chunk(s) for source %s after a retry: %s",
|
||||
failed_chunks,
|
||||
source_id,
|
||||
", ".join(str(chunk_id) for _, chunk_id in missed),
|
||||
)
|
||||
_report()
|
||||
finally:
|
||||
if pool is not None:
|
||||
pool.shutdown(wait=True)
|
||||
|
||||
try:
|
||||
store.set_node_degrees(source_id)
|
||||
except Exception as exc:
|
||||
logger.warning("set_node_degrees failed for source %s: %s", source_id, exc)
|
||||
|
||||
# Upserts are writes, not nodes: one entity seen in ten chunks is ten
|
||||
# upserts and a single node, so the old count overstated every graph whose
|
||||
# entities recur. Report what the graph holds, falling back to the write
|
||||
# count only if the count query itself fails.
|
||||
# ``strict`` is what makes the fallback below reachable: the default
|
||||
# count swallows query failures and answers 0, which would report a
|
||||
# successful build as an empty graph.
|
||||
nodes = node_upserts
|
||||
try:
|
||||
nodes = store.count_nodes(source_id, strict=True)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"count_nodes failed for source %s; reporting upserts instead: %s",
|
||||
source_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
return {
|
||||
"nodes": nodes,
|
||||
"edges": edges,
|
||||
@@ -273,7 +424,7 @@ def extract_graph_for_source(
|
||||
|
||||
|
||||
def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
|
||||
"""Normalize the LLM's entity dicts (drop nameless ones)."""
|
||||
"""Normalize the LLM's entity dicts (drop the ones with no usable name)."""
|
||||
entities = []
|
||||
for e in raw_entities:
|
||||
if not isinstance(e, dict):
|
||||
@@ -281,10 +432,16 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
|
||||
name = str(e.get("name", "")).strip()
|
||||
if not name:
|
||||
continue
|
||||
normalized_name = normalize_entity_name(name)
|
||||
if not normalized_name:
|
||||
# A punctuation-only name normalizes to nothing, and nodes merge on
|
||||
# that key: keeping it collapses every such entity onto one shared
|
||||
# node. The relationship side already drops them.
|
||||
continue
|
||||
entities.append(
|
||||
{
|
||||
"name": name,
|
||||
"normalized_name": name.lower(),
|
||||
"normalized_name": normalized_name,
|
||||
"type": str(e.get("type") or "") or None,
|
||||
"description": str(e.get("description") or "") or None,
|
||||
}
|
||||
@@ -310,6 +467,70 @@ def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]:
|
||||
return relationships
|
||||
|
||||
|
||||
def _fact_text(rel: Dict[str, Any]) -> str:
|
||||
"""A relationship rendered as the sentence it asserts.
|
||||
|
||||
Embedded and stored on the edge so retrieval can match a question against
|
||||
the *relation* rather than against entity names — the difference between
|
||||
"which entity is this about" and "which fact answers this".
|
||||
"""
|
||||
source = str(rel.get("source") or "").strip()
|
||||
target = str(rel.get("target") or "").strip()
|
||||
if not source or not target:
|
||||
return ""
|
||||
relation = str(rel.get("type") or "related to").strip() or "related to"
|
||||
text = f"{source} {relation} {target}"
|
||||
description = str(rel.get("description") or "").strip()
|
||||
return f"{text}: {description}" if description else text
|
||||
|
||||
|
||||
def _embed_facts(embedding, relationships: List[Dict[str, Any]]) -> None:
|
||||
"""Attach a fact embedding to each relationship, in one batched call.
|
||||
|
||||
Mutates the relationship dicts so the embedding travels with the edge into
|
||||
``apply_chunk`` without a second mapping to keep in step. Always on: it is
|
||||
one extra batched call per chunk against an LLM call that already costs
|
||||
far more, and it lets a source switch to relationship seeding at query time
|
||||
without being rebuilt.
|
||||
"""
|
||||
pending = [(rel, _fact_text(rel)) for rel in relationships]
|
||||
pending = [(rel, text) for rel, text in pending if text]
|
||||
if not pending:
|
||||
return
|
||||
try:
|
||||
vectors = embedding.embed_documents([text for _rel, text in pending])
|
||||
except Exception as exc: # noqa: BLE001
|
||||
# The graph is still correct without them; only fact seeding degrades.
|
||||
logger.warning("Fact embedding failed, continuing without: %s", exc)
|
||||
return
|
||||
for (rel, _text), vector in zip(pending, vectors):
|
||||
rel["fact_embedding"] = vector
|
||||
|
||||
|
||||
def _seed_text(entity: Dict[str, Any]) -> str:
|
||||
"""The text a node's embedding is computed from.
|
||||
|
||||
Retrieval seeds the graph walk by matching a whole question against these
|
||||
embeddings, and a bare entity name is a poor thing to match a question
|
||||
against — a question about what a service writes to shares almost no
|
||||
surface with the name ``Quill``. Including the type and description gives
|
||||
the match something to work with; measured across five corpora it moved
|
||||
recall@4 by +0.07 to +0.50.
|
||||
|
||||
Relationship endpoints keep their bare names: they arrive as strings with
|
||||
no type or description attached.
|
||||
"""
|
||||
name = str(entity.get("name") or "").strip()
|
||||
text = name
|
||||
entity_type = str(entity.get("type") or "").strip()
|
||||
if entity_type:
|
||||
text += f" ({entity_type})"
|
||||
description = str(entity.get("description") or "").strip()
|
||||
if description:
|
||||
text += f": {description}"
|
||||
return text or name
|
||||
|
||||
|
||||
def _embed_names(
|
||||
embedding,
|
||||
entities: List[Dict[str, Any]],
|
||||
@@ -322,14 +543,16 @@ def _embed_names(
|
||||
"""
|
||||
name_by_norm: Dict[str, str] = {}
|
||||
for entity in entities:
|
||||
name_by_norm.setdefault(entity["normalized_name"], entity["name"])
|
||||
name_by_norm.setdefault(entity["normalized_name"], _seed_text(entity))
|
||||
for rel in relationships:
|
||||
for endpoint in (rel.get("source"), rel.get("target")):
|
||||
if endpoint is None:
|
||||
continue
|
||||
clean = str(endpoint).strip()
|
||||
if clean:
|
||||
name_by_norm.setdefault(clean.lower(), clean)
|
||||
# Same key the store resolves endpoints by, or the embedding
|
||||
# computed here never reaches the node it was computed for.
|
||||
name_by_norm.setdefault(normalize_entity_name(clean), clean)
|
||||
|
||||
if not name_by_norm:
|
||||
return {}
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Canonical entity naming for the per-source knowledge graph.
|
||||
|
||||
Nodes are merged on ``normalized_name``, which has been ``name.lower()``. That
|
||||
splits entities a reader would call the same thing: measured on the DocsGPT docs
|
||||
corpus, ``agent``/``agents``, ``VECTOR_STORE``/``Vector store``/``vector stores``,
|
||||
``Celery worker``/``Celery workers`` and ``.env file``/``env_file`` all landed as
|
||||
separate nodes — 58 such collisions across 1,704 entities, with 75% of entities
|
||||
appearing in exactly one chunk as a result.
|
||||
|
||||
:func:`canonical_name` folds the differences that are purely orthographic:
|
||||
case, surrounding punctuation, underscore/hyphen word breaks, and a *cautious*
|
||||
plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``,
|
||||
``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s"
|
||||
would corrupt them into new entities rather than merge anything.
|
||||
|
||||
The result is a merge key, never shown to anyone, so it only has to be the
|
||||
same for a word's singular and plural — not to be a word itself.
|
||||
|
||||
Always on: every graph is built with canonical names.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
_PUNCT = re.compile(r"[^\w\s]+", re.UNICODE)
|
||||
_UNDERSCORE = re.compile(r"[_\-]+")
|
||||
_SPACE = re.compile(r"\s+")
|
||||
|
||||
#: Words that end in "s" without being plural. Singularising these would invent
|
||||
#: entities ("postgre", "kubernete") instead of merging existing ones.
|
||||
_NOT_PLURAL = frozenset(
|
||||
{
|
||||
"postgres", "kubernetes", "https", "aws", "dns", "tls", "cors", "css",
|
||||
"js", "sas", "gas", "ss", "class", "access", "process", "status",
|
||||
"analysis", "basis", "axis", "https", "rss", "less", "express",
|
||||
"redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos",
|
||||
"always", "sometimes", "series", "docs", "ops", "devops", "sse",
|
||||
"alias", "canvas", "atlas", "bias", "pandas",
|
||||
}
|
||||
)
|
||||
|
||||
#: Plural endings that drop ``es``, and the singular endings that meet them.
|
||||
#: ``caches`` cannot say whether it is ``cache`` + "s" or ``cach`` + "es"
|
||||
#: (as ``batches`` is ``batch`` + "es"), so rather than guess, both
|
||||
#: ``caches`` and ``cache`` fold to ``cach`` — as ``databases``/``database``
|
||||
#: fold to ``databas``. Hardly any real word differs from one of these singulars
|
||||
#: by its final "e" alone, so the fold merges next to nothing it should not.
|
||||
_ES_PLURAL = ("ches", "shes", "ses", "zes", "xes")
|
||||
_E_SINGULAR = ("che", "she", "se", "ze", "xe")
|
||||
|
||||
|
||||
def _singular(word: str) -> str:
|
||||
"""Fold one word so its singular and plural share a key, else leave it alone.
|
||||
|
||||
``-ies`` and a singular's ``-ie`` both fold to ``-y`` (``policies``,
|
||||
``cookies``/``cookie``). The ``-es`` endings in :data:`_ES_PLURAL` drop
|
||||
``es`` and the singular endings in :data:`_E_SINGULAR` drop their ``e``, so
|
||||
both sides of an ambiguous plural meet (``caches``/``cache`` -> ``cach``).
|
||||
Otherwise a bare trailing ``s`` is dropped on a word long enough to be
|
||||
safe. Everything in :data:`_NOT_PLURAL`, and anything ending in
|
||||
``ss``/``us``/``is``, is returned unchanged.
|
||||
"""
|
||||
if len(word) < 4 or word in _NOT_PLURAL:
|
||||
return word
|
||||
if word.endswith(("ss", "us", "is")):
|
||||
return word
|
||||
if word.endswith("ies") and len(word) > 4:
|
||||
return word[:-3] + "y"
|
||||
if word.endswith(_ES_PLURAL):
|
||||
return word[:-2]
|
||||
if word.endswith("ie") and len(word) > 4:
|
||||
return word[:-2] + "y"
|
||||
if word.endswith(_E_SINGULAR):
|
||||
return word[:-1]
|
||||
if word.endswith("s"):
|
||||
return word[:-1]
|
||||
return word
|
||||
|
||||
|
||||
def canonical_name(name: str) -> str:
|
||||
"""Merge key for an entity name.
|
||||
|
||||
Args:
|
||||
name: The entity name as the model wrote it.
|
||||
|
||||
Returns:
|
||||
A lowercase, punctuation-free key shared by a name's singular and
|
||||
plural. Returns ``""`` for an empty or punctuation-only name, which
|
||||
callers treat as "no entity".
|
||||
|
||||
Examples:
|
||||
``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``;
|
||||
``.env file`` and ``env_file`` -> ``env file``;
|
||||
``cache`` and ``caches`` -> ``cach``;
|
||||
``postgres`` stays ``postgres``.
|
||||
"""
|
||||
if not name:
|
||||
return ""
|
||||
text = _UNDERSCORE.sub(" ", str(name))
|
||||
text = _PUNCT.sub(" ", text)
|
||||
text = _SPACE.sub(" ", text).strip().lower()
|
||||
if not text:
|
||||
return ""
|
||||
return " ".join(_singular(word) for word in text.split())
|
||||
|
||||
|
||||
def normalize_entity_name(name: str) -> str:
|
||||
"""The key an entity is merged on: its :func:`canonical_name`.
|
||||
|
||||
Every graph the corpora were measured on was built this way, so it is the
|
||||
only mode rather than a flag. A graph built before this used plain
|
||||
``lower()`` keys; re-extracting it merges onto these instead.
|
||||
"""
|
||||
return canonical_name(name)
|
||||
+553
-87
@@ -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:
|
||||
|
||||
@@ -185,12 +185,23 @@ class Dispatcher(BaseRetriever):
|
||||
score_threshold / rephrase_query) plus an opted-in prescreen config; a
|
||||
source left at defaults takes the global path so all-classic retrieval
|
||||
stays byte-identical with zero extra LLM calls.
|
||||
|
||||
A graph source's ``graph`` options count too: they are read from the
|
||||
per-source config this records, so a source that changes only those
|
||||
would otherwise run the defaults and the options would do nothing.
|
||||
They mean nothing to any other retriever, so they only count for
|
||||
``graphrag`` -- an override hands the source its own chunk budget as
|
||||
well, which a classic source must not pick up from a graph setting.
|
||||
"""
|
||||
return (
|
||||
retrieval.chunks != _DEFAULT_RETRIEVAL.chunks
|
||||
or retrieval.score_threshold != _DEFAULT_RETRIEVAL.score_threshold
|
||||
or retrieval.rephrase_query != _DEFAULT_RETRIEVAL.rephrase_query
|
||||
or retrieval.prescreen is not None
|
||||
or (
|
||||
(retrieval.retriever or "").lower() == "graphrag"
|
||||
and retrieval.graph != _DEFAULT_RETRIEVAL.graph
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
+416
-36
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.",
|
||||
|
||||
@@ -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.",
|
||||
|
||||
@@ -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.",
|
||||
|
||||
@@ -268,6 +268,19 @@
|
||||
},
|
||||
"exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。"
|
||||
},
|
||||
"graphRetrieval": {
|
||||
"title": "グラフ検索",
|
||||
"tag": "再取り込み不要",
|
||||
"seedStrategy": "探索の開始点",
|
||||
"seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。",
|
||||
"seedEntities": "エンティティ(推奨)",
|
||||
"seedRelationships": "リレーション",
|
||||
"passageNodes": "パッセージを探索に含める",
|
||||
"passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。",
|
||||
"blendVector": "ベクトル検索と組み合わせる",
|
||||
"blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。",
|
||||
"agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。"
|
||||
},
|
||||
"prescreen": {
|
||||
"enable": "LLMプリスクリーニングを有効にする",
|
||||
"warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。",
|
||||
|
||||
@@ -268,6 +268,19 @@
|
||||
},
|
||||
"exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту."
|
||||
},
|
||||
"graphRetrieval": {
|
||||
"title": "Поиск по графу",
|
||||
"tag": "без повторной загрузки",
|
||||
"seedStrategy": "Начинать обход с",
|
||||
"seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.",
|
||||
"seedEntities": "Сущностей (рекомендуется)",
|
||||
"seedRelationships": "Связей",
|
||||
"passageNodes": "Включать фрагменты в обход",
|
||||
"passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.",
|
||||
"blendVector": "Сочетать с векторным поиском",
|
||||
"blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.",
|
||||
"agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент."
|
||||
},
|
||||
"prescreen": {
|
||||
"enable": "Включить предварительный отбор LLM",
|
||||
"warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.",
|
||||
|
||||
@@ -268,6 +268,19 @@
|
||||
},
|
||||
"exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。"
|
||||
},
|
||||
"graphRetrieval": {
|
||||
"title": "圖譜檢索",
|
||||
"tag": "無需重新匯入",
|
||||
"seedStrategy": "走訪起點",
|
||||
"seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。",
|
||||
"seedEntities": "實體(建議)",
|
||||
"seedRelationships": "關係",
|
||||
"passageNodes": "將段落納入走訪",
|
||||
"passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。",
|
||||
"blendVector": "與向量檢索結合",
|
||||
"blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。",
|
||||
"agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。"
|
||||
},
|
||||
"prescreen": {
|
||||
"enable": "啟用 LLM 預篩選",
|
||||
"warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。",
|
||||
|
||||
@@ -268,6 +268,19 @@
|
||||
},
|
||||
"exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。"
|
||||
},
|
||||
"graphRetrieval": {
|
||||
"title": "图谱检索",
|
||||
"tag": "无需重新导入",
|
||||
"seedStrategy": "遍历起点",
|
||||
"seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。",
|
||||
"seedEntities": "实体(推荐)",
|
||||
"seedRelationships": "关系",
|
||||
"passageNodes": "将段落纳入遍历",
|
||||
"passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。",
|
||||
"blendVector": "与向量检索结合",
|
||||
"blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。",
|
||||
"agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。"
|
||||
},
|
||||
"prescreen": {
|
||||
"enable": "启用 LLM 预筛选",
|
||||
"warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。",
|
||||
|
||||
@@ -32,6 +32,17 @@ export type SourcePrescreenConfig = {
|
||||
max_keep?: number; // default 8, <= candidate_k
|
||||
};
|
||||
|
||||
// Where the graph walk starts: matching entities, or matching relationships
|
||||
// ("A streams_to B"), which can reach an entity the question never names.
|
||||
export type GraphSeedStrategy = 'entities' | 'relationships';
|
||||
|
||||
// Query-time graph retrieval knobs (graphrag only; live, no re-ingest).
|
||||
export type SourceGraphRetrievalConfig = {
|
||||
seed_strategy?: GraphSeedStrategy; // default 'entities'
|
||||
passage_nodes?: boolean; // default true
|
||||
blend_vector?: boolean; // default true
|
||||
};
|
||||
|
||||
// Query-time retrieval knobs (live; no re-ingest needed).
|
||||
export type SourceRetrievalConfig = {
|
||||
retriever?: string; // default 'classic' (only option for now)
|
||||
@@ -40,6 +51,7 @@ export type SourceRetrievalConfig = {
|
||||
score_threshold?: number | null; // default null
|
||||
rephrase_query?: boolean; // default true
|
||||
prescreen?: SourcePrescreenConfig | null; // null = off
|
||||
graph?: SourceGraphRetrievalConfig; // graphrag retriever only
|
||||
};
|
||||
|
||||
// Ingest-time GraphRAG extraction knobs (only used when kind === 'graphrag').
|
||||
|
||||
@@ -146,6 +146,11 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => {
|
||||
batch_size: 5,
|
||||
max_keep: 10,
|
||||
},
|
||||
graph: {
|
||||
seed_strategy: 'relationships',
|
||||
passage_nodes: false,
|
||||
blend_vector: false,
|
||||
},
|
||||
},
|
||||
graph: {
|
||||
extraction_model: null,
|
||||
@@ -157,6 +162,44 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => {
|
||||
});
|
||||
});
|
||||
|
||||
describe('graph retrieval options', () => {
|
||||
it('defaults to the measured-best configuration', () => {
|
||||
expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph).toEqual({
|
||||
seed_strategy: 'entities',
|
||||
passage_nodes: true,
|
||||
blend_vector: true,
|
||||
});
|
||||
});
|
||||
|
||||
it('fills the defaults for a source saved before the options existed', () => {
|
||||
const opts = configToOptions({ retrieval: { retriever: 'graphrag' } });
|
||||
expect(opts.retrieval.graph).toEqual(
|
||||
DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph,
|
||||
);
|
||||
});
|
||||
|
||||
it('honors stored options and fills only the missing ones', () => {
|
||||
const opts = configToOptions({
|
||||
retrieval: { graph: { seed_strategy: 'relationships' } },
|
||||
});
|
||||
expect(opts.retrieval.graph).toEqual({
|
||||
seed_strategy: 'relationships',
|
||||
passage_nodes: true,
|
||||
blend_vector: true,
|
||||
});
|
||||
});
|
||||
|
||||
it('writes the options into the retrieval block', () => {
|
||||
const v = clone(DEFAULT_RETRIEVAL_OPTIONS);
|
||||
v.retrieval.graph.blend_vector = false;
|
||||
expect(optionsToConfig(v).retrieval?.graph).toEqual({
|
||||
seed_strategy: 'entities',
|
||||
passage_nodes: true,
|
||||
blend_vector: false,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe('isPrescreenConfigValid', () => {
|
||||
const withPrescreen = (
|
||||
chunks: number,
|
||||
|
||||
@@ -16,6 +16,7 @@ import {
|
||||
import { Switch } from '../../components/ui/switch';
|
||||
import type {
|
||||
ChunkingStrategy,
|
||||
GraphSeedStrategy,
|
||||
RetrievalExposure,
|
||||
SourceConfig,
|
||||
} from '../../models/misc';
|
||||
@@ -59,6 +60,11 @@ export type RetrievalOptionsValue = {
|
||||
batch_size: number;
|
||||
max_keep: number;
|
||||
};
|
||||
graph: {
|
||||
seed_strategy: GraphSeedStrategy;
|
||||
passage_nodes: boolean;
|
||||
blend_vector: boolean;
|
||||
};
|
||||
};
|
||||
graph: {
|
||||
extraction_model: string | null;
|
||||
@@ -85,6 +91,12 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = {
|
||||
enabled: false,
|
||||
...DEFAULT_PRESCREEN,
|
||||
},
|
||||
// The configuration that measured best across the corpora tested.
|
||||
graph: {
|
||||
seed_strategy: 'entities',
|
||||
passage_nodes: true,
|
||||
blend_vector: true,
|
||||
},
|
||||
},
|
||||
graph: {
|
||||
extraction_model: null,
|
||||
@@ -204,6 +216,7 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue {
|
||||
const chunking = config?.chunking ?? {};
|
||||
const retrieval = config?.retrieval ?? {};
|
||||
const prescreen = retrieval.prescreen ?? null;
|
||||
const retrievalGraph = retrieval.graph ?? {};
|
||||
const graph = config?.graph ?? {};
|
||||
const d = DEFAULT_RETRIEVAL_OPTIONS;
|
||||
return {
|
||||
@@ -228,6 +241,14 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue {
|
||||
batch_size: prescreen?.batch_size ?? DEFAULT_PRESCREEN.batch_size,
|
||||
max_keep: prescreen?.max_keep ?? DEFAULT_PRESCREEN.max_keep,
|
||||
},
|
||||
graph: {
|
||||
seed_strategy:
|
||||
retrievalGraph.seed_strategy ?? d.retrieval.graph.seed_strategy,
|
||||
passage_nodes:
|
||||
retrievalGraph.passage_nodes ?? d.retrieval.graph.passage_nodes,
|
||||
blend_vector:
|
||||
retrievalGraph.blend_vector ?? d.retrieval.graph.blend_vector,
|
||||
},
|
||||
},
|
||||
graph: {
|
||||
extraction_model: graph.extraction_model ?? d.graph.extraction_model,
|
||||
@@ -271,6 +292,11 @@ export function optionsToConfig(value: RetrievalOptionsValue): SourceConfig {
|
||||
max_keep: ps.max_keep,
|
||||
}
|
||||
: null,
|
||||
graph: {
|
||||
seed_strategy: value.retrieval.graph.seed_strategy,
|
||||
passage_nodes: value.retrieval.graph.passage_nodes,
|
||||
blend_vector: value.retrieval.graph.blend_vector,
|
||||
},
|
||||
},
|
||||
graph: {
|
||||
extraction_model: value.graph.extraction_model?.trim()
|
||||
@@ -399,6 +425,12 @@ export default function RetrievalOptions({
|
||||
});
|
||||
};
|
||||
|
||||
const setGraphRetrieval = (
|
||||
patch: Partial<RetrievalOptionsValue['retrieval']['graph']>,
|
||||
) => {
|
||||
setRetrieval({ graph: { ...value.retrieval.graph, ...patch } });
|
||||
};
|
||||
|
||||
const modelOptions = useMemo(() => {
|
||||
const builtin: Model[] = [];
|
||||
const user: Model[] = [];
|
||||
@@ -611,6 +643,84 @@ export default function RetrievalOptions({
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Graph retrieval group (graphrag only; live, so shown when testing too) */}
|
||||
{isGraphRAG && (
|
||||
<div className="flex flex-col gap-3">
|
||||
<GroupHeader
|
||||
title={tr('graphRetrieval.title')}
|
||||
tag={tr('graphRetrieval.tag')}
|
||||
/>
|
||||
<p className="text-muted-foreground text-xs">
|
||||
{tr('graphRetrieval.agentToolHint')}
|
||||
</p>
|
||||
|
||||
<div className="divide-border/50 divide-y">
|
||||
<SettingRow
|
||||
label={tr('graphRetrieval.seedStrategy')}
|
||||
htmlFor="graph-seed-strategy"
|
||||
description={tr('graphRetrieval.seedStrategyHint')}
|
||||
alignStart
|
||||
>
|
||||
<Select
|
||||
value={value.retrieval.graph.seed_strategy}
|
||||
disabled={disabled}
|
||||
onValueChange={(v) =>
|
||||
setGraphRetrieval({ seed_strategy: v as GraphSeedStrategy })
|
||||
}
|
||||
>
|
||||
<SelectTrigger
|
||||
id="graph-seed-strategy"
|
||||
className="w-52 rounded-md"
|
||||
size="lg"
|
||||
>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="entities">
|
||||
{tr('graphRetrieval.seedEntities')}
|
||||
</SelectItem>
|
||||
<SelectItem value="relationships">
|
||||
{tr('graphRetrieval.seedRelationships')}
|
||||
</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</SettingRow>
|
||||
|
||||
<SettingRow
|
||||
label={tr('graphRetrieval.passageNodes')}
|
||||
htmlFor="graph-passage-nodes"
|
||||
description={tr('graphRetrieval.passageNodesHint')}
|
||||
alignStart
|
||||
>
|
||||
<Switch
|
||||
id="graph-passage-nodes"
|
||||
checked={value.retrieval.graph.passage_nodes}
|
||||
disabled={disabled}
|
||||
onCheckedChange={(checked) =>
|
||||
setGraphRetrieval({ passage_nodes: checked })
|
||||
}
|
||||
/>
|
||||
</SettingRow>
|
||||
|
||||
<SettingRow
|
||||
label={tr('graphRetrieval.blendVector')}
|
||||
htmlFor="graph-blend-vector"
|
||||
description={tr('graphRetrieval.blendVectorHint')}
|
||||
alignStart
|
||||
>
|
||||
<Switch
|
||||
id="graph-blend-vector"
|
||||
checked={value.retrieval.graph.blend_vector}
|
||||
disabled={disabled}
|
||||
onCheckedChange={(checked) =>
|
||||
setGraphRetrieval({ blend_vector: checked })
|
||||
}
|
||||
/>
|
||||
</SettingRow>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Graph extraction group (graphrag only; re-ingest required to apply) */}
|
||||
{isGraphRAG && !queryOnly && (
|
||||
<div className="flex flex-col gap-3">
|
||||
|
||||
@@ -313,6 +313,26 @@ class TestClassicAgentSearchExposure:
|
||||
"Tool Doc",
|
||||
]
|
||||
|
||||
def test_collect_internal_sources_includes_graph_pages(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
# Pages the graph tool read carry the answer as much as search hits do,
|
||||
# so they are cited the same way.
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
|
||||
retriever_config = {"source": {"active_docs": ["b"]}}
|
||||
agent = ClassicAgent(retriever_config=retriever_config, **agent_base_params)
|
||||
search = Mock()
|
||||
search.retrieved_docs = [{"text": "Found", "title": "Search Doc", "source": "b"}]
|
||||
graph = Mock()
|
||||
graph.retrieved_docs = [{"text": "Quill is a store.", "title": "quill.md", "source": "b"}]
|
||||
user = agent.user or ""
|
||||
agent.tool_executor._loaded_tools[f"internal_search:{INTERNAL_TOOL_ID}:{user}"] = search
|
||||
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{user}"] = graph
|
||||
|
||||
agent._collect_internal_sources()
|
||||
assert [d["title"] for d in agent.retrieved_docs] == ["Search Doc", "quill.md"]
|
||||
|
||||
def test_collect_internal_sources_dedupes(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
|
||||
@@ -779,6 +779,20 @@ class TestCollectStepSources:
|
||||
|
||||
assert len(agent.citations.citations) == 2
|
||||
|
||||
def test_collects_pages_the_graph_tool_read(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
|
||||
agent = ResearchAgent(**agent_base_params)
|
||||
graph = Mock()
|
||||
graph.retrieved_docs = [{"source": "s3", "title": "quill.md", "text": "Quill"}]
|
||||
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{agent.user or ''}"] = graph
|
||||
|
||||
agent._collect_step_sources()
|
||||
|
||||
assert len(agent.citations.citations) == 1
|
||||
|
||||
def test_no_tool_no_error(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
|
||||
@@ -357,9 +357,9 @@ def test_malformed_json_schema_rejected_before_enqueue(monkeypatch):
|
||||
@pytest.mark.unit
|
||||
def test_dispatch_inline_when_in_worker(monkeypatch):
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
# Inside a worker current_task is truthy -> parse inline, never enqueue (else the
|
||||
# Inside a worker -> parse inline, never enqueue (else the
|
||||
# parsing queue self-deadlocks the worker that also serves it).
|
||||
monkeypatch.setattr(rd, "current_task", object())
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: True)
|
||||
|
||||
import docsgpt.api.user.tasks as tasks
|
||||
monkeypatch.setattr(
|
||||
@@ -387,8 +387,8 @@ def test_dispatch_inline_when_in_worker(monkeypatch):
|
||||
@pytest.mark.unit
|
||||
def test_dispatch_enqueues_when_not_in_worker(monkeypatch):
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
# Web process: current_task falsy -> dispatch to the parsing queue, never inline.
|
||||
monkeypatch.setattr(rd, "current_task", None)
|
||||
# Web process -> dispatch to the parsing queue, never inline.
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: False)
|
||||
captured = _patch_task(monkeypatch, payload={"status": "ok", "content": "queued", "truncated": False})
|
||||
|
||||
import docsgpt.worker as worker
|
||||
@@ -414,9 +414,9 @@ _TIMED_OUT = "document parsing timed out after"
|
||||
|
||||
|
||||
def _inline(monkeypatch, run_parse, *, timeout=0.2) -> ReadDocumentTool:
|
||||
"""Drive the inline branch (current_task truthy) with a patched parse window."""
|
||||
"""Drive the inline (in-worker) branch with a patched parse window."""
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
monkeypatch.setattr(rd, "current_task", object())
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: True)
|
||||
|
||||
import docsgpt.api.user.tasks as tasks
|
||||
monkeypatch.setattr(
|
||||
|
||||
@@ -335,6 +335,79 @@ class TestLiveLeaseDefersConcurrentRun:
|
||||
assert row[2] == "completed"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLeaseDeferralGivesUpQuietly:
|
||||
"""Deferral is bookkeeping, not failure.
|
||||
|
||||
A task that outruns the broker's visibility timeout is redelivered while
|
||||
the first worker is still running it. The lease keeps the duplicate from
|
||||
doing the work, but the duplicate kept re-queueing itself until celery
|
||||
exhausted ``LEASE_RETRY_MAX`` and raised ``MaxRetriesExceededError``, so a
|
||||
healthy long task logged a task failure. The duplicate should stand down
|
||||
instead and leave the run to the worker that holds the lease.
|
||||
"""
|
||||
|
||||
def _hold_lease(self, pg_conn, key):
|
||||
from docsgpt.storage.db.repositories.idempotency import (
|
||||
IdempotencyRepository,
|
||||
)
|
||||
|
||||
IdempotencyRepository(pg_conn).try_claim_lease(
|
||||
key=key, task_name="thing",
|
||||
task_id="t-worker-1", owner_id="worker-1",
|
||||
)
|
||||
|
||||
def test_exhausted_retries_stand_down_without_recording_a_result(self, pg_conn):
|
||||
from celery.exceptions import Ignore, MaxRetriesExceededError
|
||||
|
||||
from docsgpt.api.user.idempotency import with_idempotency
|
||||
|
||||
self._hold_lease(pg_conn, "k-long-run")
|
||||
invocations = {"count": 0}
|
||||
|
||||
@with_idempotency(task_name="thing")
|
||||
def task(self, idempotency_key=None):
|
||||
invocations["count"] += 1
|
||||
return {"ran": True}
|
||||
|
||||
# Celery raises this from ``self.retry`` once max_retries is hit.
|
||||
worker2 = _fake_celery_self("t-worker-2")
|
||||
worker2.retry.side_effect = MaxRetriesExceededError("out of retries")
|
||||
|
||||
# Ignore rather than a return value: a redelivery reuses the original
|
||||
# task id, so returning would mark the id the client is polling
|
||||
# SUCCESS — /api/task_status hands that straight to the UI, which
|
||||
# would announce a finished (empty) build while the holder is still
|
||||
# working. Ignore leaves the id's state to the holder.
|
||||
with _patch_decorator_db(pg_conn), pytest.raises(Ignore):
|
||||
task(worker2, idempotency_key="k-long-run")
|
||||
|
||||
# The lease holder is still running it; the duplicate did not.
|
||||
assert invocations["count"] == 0
|
||||
# The holder's row is untouched — not failed, not completed.
|
||||
row = _row_for(pg_conn, "k-long-run")
|
||||
assert row[2] == "pending"
|
||||
|
||||
def test_a_normal_retry_still_propagates(self, pg_conn):
|
||||
"""Only exhaustion stands down; the first deferrals must re-queue."""
|
||||
from docsgpt.api.user.idempotency import with_idempotency
|
||||
|
||||
self._hold_lease(pg_conn, "k-busy-once")
|
||||
|
||||
@with_idempotency(task_name="thing")
|
||||
def task(self, idempotency_key=None):
|
||||
return {"ran": True}
|
||||
|
||||
class _RetrySignal(Exception):
|
||||
pass
|
||||
|
||||
worker2 = _fake_celery_self("t-worker-2")
|
||||
worker2.retry.side_effect = _RetrySignal("retry scheduled")
|
||||
|
||||
with _patch_decorator_db(pg_conn), pytest.raises(_RetrySignal):
|
||||
task(worker2, idempotency_key="k-busy-once")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestExceptionPathReleasesLease:
|
||||
"""When ``fn`` raises, the lease is dropped so the next attempt
|
||||
|
||||
@@ -183,6 +183,21 @@ def _no_worker_delegation(monkeypatch):
|
||||
monkeypatch.setattr("docsgpt.cache._pubsub_redis_creation_failed", True)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _graphrag_off_by_default(monkeypatch):
|
||||
"""Run with GraphRAG at its shipped default (off), as CI does.
|
||||
|
||||
Every agent that gets a search tool checks its sources for a graph, and
|
||||
that check reads the configured vector database. A dev ``.env`` enabling
|
||||
GraphRAG sent unrelated agent tests to the developer's real database and
|
||||
left a pool to it in ``pgconn._POOLS``, failing a live test that asserts it
|
||||
owns the only pool. Tests that exercise GraphRAG turn it on themselves.
|
||||
"""
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False, raising=False)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm():
|
||||
llm = Mock()
|
||||
|
||||
@@ -88,6 +88,37 @@ class _StubLLM:
|
||||
return response
|
||||
|
||||
|
||||
class _ScriptedLLM:
|
||||
"""Stub LLM answering per chunk, so results do not depend on call order.
|
||||
|
||||
The extraction pool runs calls concurrently, so a stub that hands out
|
||||
responses in call order gives each chunk whichever response its thread
|
||||
happened to grab first. ``script`` maps a chunk's text to the responses
|
||||
for that chunk, consumed one per call.
|
||||
"""
|
||||
|
||||
def __init__(self, script):
|
||||
self._script = {text: list(responses) for text, responses in script.items()}
|
||||
self.model_id = "stub-model"
|
||||
self.calls = []
|
||||
self._token_usage_source = None
|
||||
self._request_id = None
|
||||
|
||||
def gen(self, model=None, messages=None, **kwargs):
|
||||
text = messages[-1]["content"].removeprefix("<chunk>\n").removesuffix("\n</chunk>")
|
||||
self.calls.append(text)
|
||||
responses = self._script.get(text)
|
||||
if not responses:
|
||||
raise AssertionError(f"unexpected extraction call for {text!r}")
|
||||
response = responses.pop(0)
|
||||
if isinstance(response, Exception):
|
||||
raise response
|
||||
return response
|
||||
|
||||
def calls_for(self, text):
|
||||
return self.calls.count(text)
|
||||
|
||||
|
||||
class _StubEmbedding:
|
||||
"""Stub embeddings model producing deterministic fixed-dim vectors."""
|
||||
|
||||
@@ -138,6 +169,129 @@ def _extraction_json(entities, relationships):
|
||||
return json.dumps({"entities": entities, "relationships": relationships})
|
||||
|
||||
|
||||
class TestFactText:
|
||||
"""A relationship rendered as the sentence it asserts.
|
||||
|
||||
This is what fact seeding matches a question against, so it has to read as
|
||||
a claim rather than as three fields concatenated.
|
||||
"""
|
||||
|
||||
def test_renders_the_relationship_as_a_sentence(self):
|
||||
text = extraction_module._fact_text(
|
||||
{
|
||||
"source": "Alder",
|
||||
"target": "Quill",
|
||||
"type": "streams_to",
|
||||
"description": "Alder streams audit events to Quill.",
|
||||
}
|
||||
)
|
||||
|
||||
assert text == "Alder streams_to Quill: Alder streams audit events to Quill."
|
||||
|
||||
def test_omits_an_absent_description(self):
|
||||
text = extraction_module._fact_text(
|
||||
{"source": "Alder", "target": "Quill", "type": "streams_to"}
|
||||
)
|
||||
|
||||
assert text == "Alder streams_to Quill"
|
||||
|
||||
def test_defaults_a_missing_relation(self):
|
||||
text = extraction_module._fact_text({"source": "Alder", "target": "Quill"})
|
||||
|
||||
assert text == "Alder related to Quill"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rel",
|
||||
[
|
||||
{"source": "Alder", "target": ""},
|
||||
{"source": "", "target": "Quill"},
|
||||
{},
|
||||
],
|
||||
)
|
||||
def test_an_edge_without_both_endpoints_has_no_fact(self, rel):
|
||||
assert extraction_module._fact_text(rel) == ""
|
||||
|
||||
|
||||
class TestEmbedFacts:
|
||||
"""Fact embeddings are always recorded, so a source can switch to
|
||||
relationship seeding at query time without being rebuilt."""
|
||||
|
||||
def _relationships(self):
|
||||
return [{"source": "Alder", "target": "Quill", "type": "streams_to"}]
|
||||
|
||||
def test_attaches_one_embedding_per_fact_in_a_single_call(self):
|
||||
relationships = self._relationships() + [{"source": "", "target": "Nowhere"}]
|
||||
calls = []
|
||||
|
||||
class _Embedding:
|
||||
def embed_documents(self, texts):
|
||||
calls.append(texts)
|
||||
return [[0.5] * 4 for _ in texts]
|
||||
|
||||
extraction_module._embed_facts(_Embedding(), relationships)
|
||||
|
||||
# One batched call, and the endpoint-less relationship is skipped
|
||||
# rather than embedded as an empty string.
|
||||
assert calls == [["Alder streams_to Quill"]]
|
||||
assert relationships[0]["fact_embedding"] == [0.5] * 4
|
||||
assert "fact_embedding" not in relationships[1]
|
||||
|
||||
def test_survives_an_embedding_failure(self):
|
||||
"""The graph is still correct without fact embeddings — only
|
||||
relationship seeding degrades, and it falls back to entities — so a
|
||||
failure here must not fail the chunk."""
|
||||
relationships = self._relationships()
|
||||
|
||||
class _Embedding:
|
||||
def embed_documents(self, texts):
|
||||
raise RuntimeError("embeddings down")
|
||||
|
||||
extraction_module._embed_facts(_Embedding(), relationships)
|
||||
|
||||
assert "fact_embedding" not in relationships[0]
|
||||
|
||||
|
||||
class TestSeedText:
|
||||
"""What a node's embedding is computed from.
|
||||
|
||||
Retrieval matches a whole question against these embeddings, so what goes
|
||||
into them decides what the graph walk can start from.
|
||||
"""
|
||||
|
||||
def _entity(self):
|
||||
return {
|
||||
"name": "Quill",
|
||||
"normalized_name": "quill",
|
||||
"type": "store",
|
||||
"description": "A write-ahead store.",
|
||||
}
|
||||
|
||||
def test_includes_type_and_description(self):
|
||||
assert (
|
||||
extraction_module._seed_text(self._entity())
|
||||
== "Quill (store): A write-ahead store."
|
||||
)
|
||||
|
||||
def test_falls_back_to_the_name_when_fields_are_missing(self):
|
||||
assert extraction_module._seed_text({"name": "Quill"}) == "Quill"
|
||||
|
||||
def test_embedded_text_is_keyed_by_the_normalized_name(self):
|
||||
"""The richer text must reach ``embed_documents``, keyed by the same
|
||||
normalized name the store resolves nodes by — otherwise the embedding
|
||||
is computed for a node it never reaches."""
|
||||
captured = {}
|
||||
|
||||
class _Embedding:
|
||||
def embed_documents(self, texts):
|
||||
captured["texts"] = texts
|
||||
return [[0.0] * 4 for _ in texts]
|
||||
|
||||
result = extraction_module._embed_names(_Embedding(), [self._entity()], [])
|
||||
|
||||
assert captured["texts"] == ["Quill (store): A write-ahead store."]
|
||||
assert set(result) == {"quill"}
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestExtractionLive:
|
||||
@pytest.fixture
|
||||
@@ -193,6 +347,97 @@ class TestExtractionLive:
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_parallel_workers_process_every_chunk_once(
|
||||
self, store, source_id, monkeypatch, stub_embedding
|
||||
):
|
||||
"""Running the model calls concurrently must not change what gets written.
|
||||
|
||||
Extraction spends nearly all of a chunk's time waiting on the model, so
|
||||
the calls run in a pool while every graph write stays on the calling
|
||||
thread. Six chunks share one entity here: whatever order the pool
|
||||
finishes in, that entity is upserted once, each chunk is linked, and all
|
||||
six are marked processed.
|
||||
"""
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
try:
|
||||
payload = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
llm = _StubLLM([payload] * 6)
|
||||
_install_stub_llm(monkeypatch, llm)
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
source_id,
|
||||
user="owner-1",
|
||||
chunks=[
|
||||
_chunk(f"c{i}", f"Ada appears here, take {i}.") for i in range(6)
|
||||
],
|
||||
config=SourceConfig(),
|
||||
request_id="req-parallel",
|
||||
)
|
||||
|
||||
assert summary["chunks_processed"] == 6
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert summary["nodes"] == 1
|
||||
assert len(llm.gen_calls) == 6
|
||||
|
||||
node = store.get_node_by_normalized(source_id, "ada")
|
||||
assert node is not None
|
||||
mapping = store.get_chunk_ids_for_nodes(source_id, [node["id"]])
|
||||
assert sorted(mapping[node["id"]]) == [f"c{i}" for i in range(6)]
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_embedding_runs_on_the_calling_thread(
|
||||
self, store, source_id, monkeypatch, stub_embedding
|
||||
):
|
||||
"""Only the LLM call may run in the extraction pool, never embedding.
|
||||
|
||||
Embedding from the pool is what broke every graph build inside a
|
||||
worker: the embeddings client used to decide "embed locally" from the
|
||||
task on the *current thread's* stack, so a pool thread dispatched to
|
||||
the worker instead and Celery refused the wait. ``in_worker`` is
|
||||
process-wide now, but the pool still has no reason to touch the
|
||||
embeddings client — it exists to overlap model latency.
|
||||
"""
|
||||
import threading
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
caller = threading.current_thread()
|
||||
seen = []
|
||||
real_embed_names = extraction_module._embed_names
|
||||
|
||||
def _recording_embed_names(*args, **kwargs):
|
||||
seen.append(threading.current_thread())
|
||||
return real_embed_names(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(extraction_module, "_embed_names", _recording_embed_names)
|
||||
try:
|
||||
payload = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
_install_stub_llm(monkeypatch, _StubLLM([payload] * 4))
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
source_id,
|
||||
user="owner-1",
|
||||
chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(4)],
|
||||
config=SourceConfig(),
|
||||
request_id="req-thread",
|
||||
)
|
||||
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert len(seen) == 4
|
||||
assert all(thread is caller for thread in seen)
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_same_entity_across_chunks_merges(
|
||||
self, store, source_id, monkeypatch, stub_embedding
|
||||
):
|
||||
@@ -292,11 +537,13 @@ class TestExtractionLive:
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
llm = _StubLLM([
|
||||
"not json at all",
|
||||
RuntimeError("model exploded"),
|
||||
good,
|
||||
])
|
||||
# Each failing chunk fails its retry too; one that recovers on
|
||||
# retry is covered in ``TestFailedChunksAreRetried``.
|
||||
llm = _ScriptedLLM({
|
||||
"garbage": ["not json at all", "still not json"],
|
||||
"boom": [RuntimeError("model exploded"), RuntimeError("model exploded again")],
|
||||
"Ada.": [good],
|
||||
})
|
||||
_install_stub_llm(monkeypatch, llm)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
@@ -359,6 +606,61 @@ class TestExtractionTokenUsage:
|
||||
assert built._request_id == "req-99"
|
||||
assert captured["model_id"] == "stub-model"
|
||||
|
||||
def test_concurrent_extraction_calls_never_share_an_llm(self, monkeypatch, stub_embedding):
|
||||
"""Provider usage is recorded on the LLM instance (``_last_usage``) and
|
||||
claimed by whichever call finishes next, so two calls in flight on one
|
||||
instance can bill each other's tokens. Each extraction thread needs its
|
||||
own instance."""
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
payload = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
|
||||
class _ThreadRecordingLLM:
|
||||
model_id = "stub-model"
|
||||
|
||||
def __init__(self):
|
||||
self.threads = set()
|
||||
|
||||
def gen(self, model=None, messages=None, **kwargs):
|
||||
self.threads.add(threading.get_ident())
|
||||
time.sleep(0.01) # keep calls overlapping
|
||||
return payload
|
||||
|
||||
built = []
|
||||
|
||||
def _create(*args, **kwargs):
|
||||
llm = _ThreadRecordingLLM()
|
||||
built.append(llm)
|
||||
return llm
|
||||
|
||||
monkeypatch.setattr(extraction_module.LLMCreator, "create_llm", staticmethod(_create))
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
|
||||
store = MagicMock(name="GraphStore")
|
||||
store.pending_chunks.return_value = [f"c{i}" for i in range(8)]
|
||||
store.apply_chunk.return_value = (1, 0)
|
||||
store.count_nodes.return_value = 1
|
||||
monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", lambda *a, **k: store)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(8)],
|
||||
config=SourceConfig(),
|
||||
request_id="req-threads",
|
||||
)
|
||||
|
||||
assert summary["chunks_processed"] == 8
|
||||
used = [llm for llm in built if llm.threads]
|
||||
assert len(used) > 1, "calls did not run concurrently"
|
||||
assert all(len(llm.threads) == 1 for llm in used)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestModelResolution:
|
||||
@@ -403,6 +705,407 @@ class TestModelResolution:
|
||||
assert extraction_module._resolve_max_chunks(config) == 5
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestExtractionProviderResolution:
|
||||
"""The extraction model decides the provider, not ``LLM_PROVIDER``.
|
||||
|
||||
``settings.LLM_PROVIDER`` is the deployment default (``docsgpt`` out of the
|
||||
box, i.e. the hosted public endpoint). Dispatching the resolved extraction
|
||||
model through it sends the call to a provider that never serves that model:
|
||||
the request is rejected, the shared fallback answers instead, and the graph
|
||||
is quietly built by a different model than the one configured.
|
||||
"""
|
||||
|
||||
def _capture_create_llm(self, monkeypatch, llm=None):
|
||||
captured = {}
|
||||
|
||||
def _create(provider, *args, **kwargs):
|
||||
captured["provider"] = provider
|
||||
captured["args"] = args
|
||||
captured["kwargs"] = kwargs
|
||||
return llm or _StubLLM([])
|
||||
|
||||
monkeypatch.setattr(
|
||||
extraction_module.LLMCreator, "create_llm", staticmethod(_create)
|
||||
)
|
||||
return captured
|
||||
|
||||
def test_provider_comes_from_the_model_registry(self, monkeypatch):
|
||||
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_provider_from_model_id", lambda *a, **k: "openai"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_api_key_for_provider", lambda provider: "sk-openai"
|
||||
)
|
||||
captured = self._capture_create_llm(monkeypatch)
|
||||
|
||||
extraction_module._build_extraction_llm("gpt-4o-mini", "owner-1", "req-1")
|
||||
|
||||
assert captured["provider"] == "openai"
|
||||
assert captured["kwargs"]["api_key"] == "sk-openai"
|
||||
assert captured["kwargs"]["model_id"] == "gpt-4o-mini"
|
||||
|
||||
def test_owner_scopes_the_registry_lookup(self, monkeypatch):
|
||||
"""A per-user (BYOM) model only resolves when the owner is passed."""
|
||||
seen = {}
|
||||
|
||||
def _resolve(model_id, user_id=None):
|
||||
seen["model_id"] = model_id
|
||||
seen["user_id"] = user_id
|
||||
return "anthropic"
|
||||
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_provider_from_model_id", _resolve
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_api_key_for_provider", lambda provider: "k"
|
||||
)
|
||||
self._capture_create_llm(monkeypatch)
|
||||
|
||||
extraction_module._build_extraction_llm("byom-uuid", "owner-7", "req-1")
|
||||
|
||||
assert seen == {"model_id": "byom-uuid", "user_id": "owner-7"}
|
||||
|
||||
def test_unknown_model_falls_back_to_the_configured_provider(self, monkeypatch):
|
||||
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_provider_from_model_id", lambda *a, **k: None
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_api_key_for_provider", lambda provider: "fallback-key"
|
||||
)
|
||||
captured = self._capture_create_llm(monkeypatch)
|
||||
|
||||
extraction_module._build_extraction_llm("mystery-model", "owner-1", "req-1")
|
||||
|
||||
assert captured["provider"] == "docsgpt"
|
||||
assert captured["kwargs"]["api_key"] == "fallback-key"
|
||||
|
||||
def test_no_model_id_skips_the_lookup(self, monkeypatch):
|
||||
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "openai")
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
extraction_module,
|
||||
"get_provider_from_model_id",
|
||||
lambda *a, **k: calls.append(a) or "anthropic",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_api_key_for_provider", lambda provider: "k"
|
||||
)
|
||||
captured = self._capture_create_llm(monkeypatch)
|
||||
|
||||
extraction_module._build_extraction_llm(None, "owner-1", "req-1")
|
||||
|
||||
assert calls == []
|
||||
assert captured["provider"] == "openai"
|
||||
|
||||
def test_api_key_follows_the_resolved_provider(self, monkeypatch):
|
||||
"""The key must match the provider actually dispatched to."""
|
||||
monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt")
|
||||
monkeypatch.setattr(extraction_module.settings, "API_KEY", "generic-key")
|
||||
monkeypatch.setattr(
|
||||
extraction_module, "get_provider_from_model_id", lambda *a, **k: "anthropic"
|
||||
)
|
||||
keyed_for = {}
|
||||
|
||||
def _key(provider):
|
||||
keyed_for["provider"] = provider
|
||||
return "sk-anthropic"
|
||||
|
||||
monkeypatch.setattr(extraction_module, "get_api_key_for_provider", _key)
|
||||
captured = self._capture_create_llm(monkeypatch)
|
||||
|
||||
extraction_module._build_extraction_llm("claude-x", "owner-1", "req-1")
|
||||
|
||||
assert keyed_for["provider"] == "anthropic"
|
||||
assert captured["kwargs"]["api_key"] == "sk-anthropic"
|
||||
assert captured["kwargs"]["api_key"] != "generic-key"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFailedChunksAreReported:
|
||||
"""Every dropped chunk has to leave a trace.
|
||||
|
||||
A chunk whose extraction cannot be parsed is marked ``failed`` and skipped.
|
||||
That path logged nothing at all, so a graph could come back short with the
|
||||
summary's ``failed_chunks`` count as the only hint and no way to tell which
|
||||
chunk, or why, from the logs.
|
||||
"""
|
||||
|
||||
def _fake_store(self, monkeypatch, chunk_ids):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
store = MagicMock(name="GraphStore")
|
||||
store.pending_chunks.return_value = list(chunk_ids)
|
||||
store.apply_chunk.return_value = (1, 0)
|
||||
store.count_nodes.return_value = 1
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
|
||||
)
|
||||
return store
|
||||
|
||||
def test_unparseable_output_is_logged_with_the_chunk_id(
|
||||
self, monkeypatch, caplog, stub_embedding
|
||||
):
|
||||
import logging
|
||||
|
||||
store = self._fake_store(monkeypatch, ["c1"])
|
||||
_install_stub_llm(monkeypatch, _StubLLM(["not json at all", "still not json"]))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"):
|
||||
summary = extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=[_chunk("c1", "some text")],
|
||||
config=SourceConfig(),
|
||||
request_id="req-1",
|
||||
)
|
||||
|
||||
assert summary["failed_chunks"] == 1
|
||||
store.mark_chunk.assert_called_once()
|
||||
assert store.mark_chunk.call_args.args[2] == "failed"
|
||||
messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING]
|
||||
assert any("c1" in message for message in messages), messages
|
||||
|
||||
def test_llm_errors_still_name_the_chunk(
|
||||
self, monkeypatch, caplog, stub_embedding
|
||||
):
|
||||
import logging
|
||||
|
||||
self._fake_store(monkeypatch, ["c7"])
|
||||
_install_stub_llm(
|
||||
monkeypatch,
|
||||
_StubLLM([RuntimeError("model exploded"), RuntimeError("model exploded again")]),
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"):
|
||||
extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=[_chunk("c7", "some text")],
|
||||
config=SourceConfig(),
|
||||
request_id="req-1",
|
||||
)
|
||||
|
||||
messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING]
|
||||
assert any("c7" in message for message in messages), messages
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFailedChunksAreRetried:
|
||||
"""A chunk that fails once gets one more attempt before the build ends.
|
||||
|
||||
Failures are recorded as ``failed`` and the checkpoint treats them as
|
||||
pending, but nothing ever ran the build again, so a single transient error
|
||||
— a provider hiccup, one response that did not parse — left a permanent
|
||||
hole in the graph until someone rebuilt the whole source. Retries run
|
||||
after the rest of the build, which gives a burst of rate limiting time to
|
||||
pass, and are bounded at one per chunk so a chunk that can never be
|
||||
extracted costs at most two calls.
|
||||
"""
|
||||
|
||||
GOOD = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
|
||||
def _fake_store(self, monkeypatch, chunk_ids):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
store = MagicMock(name="GraphStore")
|
||||
store.pending_chunks.return_value = list(chunk_ids)
|
||||
store.apply_chunk.return_value = (1, 0)
|
||||
store.count_nodes.return_value = 1
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
|
||||
)
|
||||
return store
|
||||
|
||||
def _run(self, chunks, progress=None):
|
||||
return extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=chunks,
|
||||
config=SourceConfig(),
|
||||
request_id="req-retry",
|
||||
progress_cb=progress,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _marked_failed(store):
|
||||
return [c.args[1] for c in store.mark_chunk.call_args_list if c.args[2] == "failed"]
|
||||
|
||||
def test_a_transient_failure_is_retried_and_written(self, monkeypatch, stub_embedding):
|
||||
store = self._fake_store(monkeypatch, ["c1"])
|
||||
llm = _ScriptedLLM({"flaky": [RuntimeError("rate limited"), self.GOOD]})
|
||||
_install_stub_llm(monkeypatch, llm)
|
||||
|
||||
summary = self._run([_chunk("c1", "flaky")])
|
||||
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert summary["chunks_processed"] == 1
|
||||
assert store.apply_chunk.call_args.args[1] == "c1"
|
||||
assert self._marked_failed(store) == []
|
||||
|
||||
def test_an_unparseable_response_is_retried(self, monkeypatch, stub_embedding):
|
||||
store = self._fake_store(monkeypatch, ["c1"])
|
||||
_install_stub_llm(monkeypatch, _ScriptedLLM({"odd": ["not json", self.GOOD]}))
|
||||
|
||||
summary = self._run([_chunk("c1", "odd")])
|
||||
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert self._marked_failed(store) == []
|
||||
|
||||
def test_a_chunk_that_fails_again_is_marked_failed_once(self, monkeypatch, stub_embedding):
|
||||
store = self._fake_store(monkeypatch, ["c1"])
|
||||
llm = _ScriptedLLM({"broken": ["not json", "still not json"]})
|
||||
_install_stub_llm(monkeypatch, llm)
|
||||
|
||||
summary = self._run([_chunk("c1", "broken")])
|
||||
|
||||
assert summary["failed_chunks"] == 1
|
||||
assert summary["chunks_processed"] == 0
|
||||
assert llm.calls_for("broken") == 2
|
||||
assert self._marked_failed(store) == ["c1"]
|
||||
|
||||
def test_only_failed_chunks_are_retried(self, monkeypatch, stub_embedding):
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
|
||||
self._fake_store(monkeypatch, ["c1", "c2", "c3"])
|
||||
llm = _ScriptedLLM({
|
||||
"one": [self.GOOD],
|
||||
"two": [RuntimeError("timeout"), self.GOOD],
|
||||
"three": [self.GOOD],
|
||||
})
|
||||
_install_stub_llm(monkeypatch, llm)
|
||||
|
||||
summary = self._run([_chunk("c1", "one"), _chunk("c2", "two"), _chunk("c3", "three")])
|
||||
|
||||
assert summary["chunks_processed"] == 3
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert (llm.calls_for("one"), llm.calls_for("two"), llm.calls_for("three")) == (1, 2, 1)
|
||||
|
||||
def test_a_failed_write_is_retried(self, monkeypatch, stub_embedding):
|
||||
store = self._fake_store(monkeypatch, ["c1"])
|
||||
store.apply_chunk.side_effect = [RuntimeError("write failed"), (1, 0)]
|
||||
_install_stub_llm(monkeypatch, _ScriptedLLM({"text": [self.GOOD, self.GOOD]}))
|
||||
|
||||
summary = self._run([_chunk("c1", "text")])
|
||||
|
||||
assert summary["failed_chunks"] == 0
|
||||
assert summary["chunks_processed"] == 1
|
||||
assert self._marked_failed(store) == []
|
||||
|
||||
def test_progress_ends_at_the_total(self, monkeypatch, stub_embedding):
|
||||
self._fake_store(monkeypatch, ["c1", "c2"])
|
||||
_install_stub_llm(monkeypatch, _ScriptedLLM({
|
||||
"fine": [self.GOOD],
|
||||
"broken": ["not json", "still not json"],
|
||||
}))
|
||||
events = []
|
||||
|
||||
self._run([_chunk("c1", "fine"), _chunk("c2", "broken")], progress=events.append)
|
||||
|
||||
assert all(e["current"] <= e["total"] for e in events)
|
||||
assert events[-1]["current"] == events[-1]["total"] == 2
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestSummaryNodeCount:
|
||||
"""``nodes`` must describe the graph, not the number of upserts."""
|
||||
|
||||
@pytest.fixture
|
||||
def store(self, monkeypatch, postgresql):
|
||||
store = _live_store(monkeypatch, postgresql.info)
|
||||
yield store
|
||||
store.close()
|
||||
|
||||
def test_repeated_entity_counts_once(
|
||||
self, store, monkeypatch, stub_embedding
|
||||
):
|
||||
source_id = str(uuid.uuid4())
|
||||
try:
|
||||
payload = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
_install_stub_llm(monkeypatch, _StubLLM([payload, payload]))
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
source_id,
|
||||
user="owner-1",
|
||||
chunks=[_chunk("c1", "Ada one."), _chunk("c2", "Ada two.")],
|
||||
config=SourceConfig(),
|
||||
request_id="req-1",
|
||||
)
|
||||
|
||||
# Two chunks upserted the same entity: one node in the graph.
|
||||
assert store.count_nodes(source_id) == 1
|
||||
assert summary["nodes"] == 1
|
||||
assert summary["chunks_processed"] == 2
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestSummaryCountFailure:
|
||||
"""A broken count query must not be reported as an empty graph."""
|
||||
|
||||
def test_a_failed_count_reports_the_write_count(
|
||||
self, monkeypatch, stub_embedding
|
||||
):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
store = MagicMock(name="GraphStore")
|
||||
store.pending_chunks.return_value = ["c1"]
|
||||
store.apply_chunk.return_value = (2, 1)
|
||||
store.count_nodes.side_effect = RuntimeError("count query failed")
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.graphrag.store.GraphStore", lambda *a, **k: store
|
||||
)
|
||||
_install_stub_llm(
|
||||
monkeypatch,
|
||||
_StubLLM([_extraction_json([{"name": "Ada"}], [])]),
|
||||
)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=[_chunk("c1", "Ada.")],
|
||||
config=SourceConfig(),
|
||||
request_id="req-1",
|
||||
)
|
||||
|
||||
# Falls back to what was actually written, not to zero.
|
||||
assert summary["nodes"] == 2
|
||||
# And it asked for a count that raises rather than one that returns 0,
|
||||
# or the fallback above could never run.
|
||||
assert store.count_nodes.call_args.kwargs.get("strict") is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEntityNormalization:
|
||||
"""An entity whose name normalizes to nothing is not an entity.
|
||||
|
||||
``canonical_name`` answers "" for a punctuation-only name, and nodes are
|
||||
merged on that key, so keeping them collapses every such entity onto one
|
||||
shared node. ``_resolve_endpoint`` already drops them on the relationship
|
||||
side.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize("name", ["!!!", "--", "?", " *** "])
|
||||
def test_a_name_that_normalizes_to_nothing_is_dropped(self, name):
|
||||
assert extraction_module._build_entities([{"name": name}]) == []
|
||||
|
||||
def test_real_names_survive(self):
|
||||
built = extraction_module._build_entities(
|
||||
[{"name": "Quill Store"}, {"name": "!!!"}, {"name": "Alder"}]
|
||||
)
|
||||
assert [e["normalized_name"] for e in built] == ["quill store", "alder"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestParsing:
|
||||
def test_parses_embedded_json(self):
|
||||
|
||||
@@ -0,0 +1,361 @@
|
||||
"""The graph exposed to an agent as callable tools.
|
||||
|
||||
Ranking with a graph never beat vector search in measurement; letting a model
|
||||
*follow* an edge did, on content where the answer is two documents away. These
|
||||
tests cover the contract that makes that possible — the tool must return the
|
||||
relationships verbatim enough for the model to read a name out of them, and
|
||||
must refuse clearly rather than silently when it has nothing to offer.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.agents.tools.graph_search import (
|
||||
GRAPH_TOOL_ID,
|
||||
GraphSearchTool,
|
||||
add_graph_search_tool,
|
||||
build_graph_tool_entry,
|
||||
)
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
SOURCE = {"active_docs": ["src-1"]}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _graph_store_is_pgvector(monkeypatch):
|
||||
"""The graph only exists under pgvector, and CI's default is faiss.
|
||||
|
||||
Set here rather than per test so a case that forgets it fails for its own
|
||||
reason instead of the gate; the cases about a different vector store
|
||||
override it in the test body.
|
||||
"""
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
|
||||
|
||||
|
||||
class _StubStore:
|
||||
def __init__(self, nodes=None, relationships=None, pages=None):
|
||||
self._nodes = nodes or []
|
||||
self._relationships = relationships or []
|
||||
self._pages = pages or []
|
||||
|
||||
def search_nodes_by_embedding(self, source_id, embedding, k=10):
|
||||
return self._nodes[:k]
|
||||
|
||||
def entity_relationships(self, source_id, name, limit=25):
|
||||
return self._relationships
|
||||
|
||||
def entity_pages(self, source_id, name, limit=4):
|
||||
return self._pages
|
||||
|
||||
|
||||
def _tool(monkeypatch, store, enabled=True):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
|
||||
tool = GraphSearchTool({"source": SOURCE})
|
||||
tool._store = store
|
||||
monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1])
|
||||
return tool
|
||||
|
||||
|
||||
class TestGating:
|
||||
def test_reports_when_graphs_are_disabled(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore(), enabled=False)
|
||||
|
||||
assert "not enabled" in tool.execute_action("search_entities", query="x")
|
||||
|
||||
def test_reports_when_no_sources_are_configured(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
|
||||
tool = GraphSearchTool({"source": {"active_docs": []}})
|
||||
|
||||
assert "No graph-backed sources" in tool.execute_action(
|
||||
"search_entities", query="x"
|
||||
)
|
||||
|
||||
def test_unknown_action_is_named(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore())
|
||||
|
||||
assert "Unknown action" in tool.execute_action("wander")
|
||||
|
||||
|
||||
class TestActions:
|
||||
def test_search_entities_lists_names_with_match_strength(self, monkeypatch):
|
||||
tool = _tool(
|
||||
monkeypatch,
|
||||
_StubStore(nodes=[{"name": "Quill", "distance": 0.2, "description": "A store."}]),
|
||||
)
|
||||
|
||||
result = tool.execute_action("search_entities", query="quill")
|
||||
|
||||
assert "Quill" in result
|
||||
assert "0.80" in result
|
||||
|
||||
def test_search_entities_requires_a_query(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore())
|
||||
|
||||
assert "required" in tool.execute_action("search_entities", query=" ")
|
||||
|
||||
def test_relationships_are_rendered_as_triples(self, monkeypatch):
|
||||
"""The model reads the *target* out of this line to take its next step,
|
||||
so the target name has to survive rendering intact."""
|
||||
tool = _tool(
|
||||
monkeypatch,
|
||||
_StubStore(
|
||||
relationships=[
|
||||
{"source": "Alder", "type": "streams_to", "target": "Quill", "description": ""}
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
result = tool.execute_action("get_relationships", entity="Alder")
|
||||
|
||||
assert "Alder --streams_to--> Quill" in result
|
||||
|
||||
def test_missing_relationships_suggest_the_next_step(self, monkeypatch):
|
||||
"""A dead end should point at search_entities rather than stop the agent."""
|
||||
tool = _tool(monkeypatch, _StubStore(relationships=[]))
|
||||
|
||||
result = tool.execute_action("get_relationships", entity="Nope")
|
||||
|
||||
assert "search_entities" in result
|
||||
|
||||
def test_pages_are_titled_truncated_and_recorded(self, monkeypatch):
|
||||
tool = _tool(
|
||||
monkeypatch,
|
||||
_StubStore(
|
||||
pages=[
|
||||
{
|
||||
"metadata": {"title": "quill-store.md", "source": "quill-store.md"},
|
||||
"text": "x" * 5000,
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
result = tool.execute_action("read_entity_pages", entity="Quill")
|
||||
|
||||
assert "--- quill-store.md ---" in result
|
||||
# Only what the model reads is truncated.
|
||||
assert len(result) < 3000
|
||||
# Accumulated so the answer can cite what the walk actually read.
|
||||
assert tool.retrieved_docs[0]["title"] == "quill-store.md"
|
||||
assert len(tool.retrieved_docs[0]["text"]) == 5000
|
||||
|
||||
def test_page_labels_match_what_the_retrievers_record(self, monkeypatch):
|
||||
"""A page read here and the same chunk retrieved by internal_search are
|
||||
one document. The citation manager keys on (source, title), so labels
|
||||
derived differently give the same document two citation numbers."""
|
||||
from docsgpt.retriever.labels import labels_from_metadata
|
||||
|
||||
metadata = {"title": "Quill Store", "source": "quill-store.md"}
|
||||
text = "Quill is a write-ahead store."
|
||||
tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": metadata, "text": text}]))
|
||||
|
||||
tool.execute_action("read_entity_pages", entity="Quill")
|
||||
|
||||
expected = labels_from_metadata(metadata, text, "src-1")
|
||||
doc = tool.retrieved_docs[0]
|
||||
assert {k: doc[k] for k in ("title", "source", "filename")} == expected
|
||||
# The full chunk text, so the doc dedupes against the retriever's copy;
|
||||
# only what the model reads is truncated.
|
||||
assert doc["text"] == text
|
||||
|
||||
def test_a_page_with_no_metadata_falls_back_to_its_source_id(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": {}, "text": "body"}]))
|
||||
|
||||
tool.execute_action("read_entity_pages", entity="Quill")
|
||||
|
||||
assert tool.retrieved_docs[0]["source"] == "src-1"
|
||||
|
||||
def test_pages_absent_is_stated_plainly(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore(pages=[]))
|
||||
|
||||
assert "No documents" in tool.execute_action("read_entity_pages", entity="Quill")
|
||||
|
||||
|
||||
class TestWiring:
|
||||
def test_entry_exposes_every_action(self):
|
||||
entry = build_graph_tool_entry()
|
||||
|
||||
assert {a["name"] for a in entry["actions"]} == {
|
||||
"search_entities",
|
||||
"get_relationships",
|
||||
"read_entity_pages",
|
||||
}
|
||||
assert all(action["active"] for action in entry["actions"])
|
||||
|
||||
def test_not_added_when_graphs_are_disabled(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False)
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True
|
||||
)
|
||||
tools: dict = {}
|
||||
|
||||
add_graph_search_tool(tools, {"source": SOURCE})
|
||||
|
||||
assert tools == {}
|
||||
|
||||
def test_not_added_when_the_sources_have_no_graph(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False
|
||||
)
|
||||
tools: dict = {}
|
||||
|
||||
add_graph_search_tool(tools, {"source": SOURCE})
|
||||
|
||||
assert tools == {}
|
||||
|
||||
def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector")
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True
|
||||
)
|
||||
tools: dict = {}
|
||||
|
||||
add_graph_search_tool(tools, {"source": SOURCE})
|
||||
|
||||
assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID
|
||||
assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE
|
||||
|
||||
|
||||
class TestPlumbing:
|
||||
def test_a_single_source_id_and_empty_entries_are_accepted(self):
|
||||
tool = GraphSearchTool({"source": {"active_docs": "src-1"}})
|
||||
assert tool._sources() == ["src-1"]
|
||||
|
||||
tool = GraphSearchTool({"source": {"active_docs": ["src-1", "", None]}})
|
||||
assert tool._sources() == ["src-1"]
|
||||
|
||||
def test_the_store_is_built_once_and_reused(self, monkeypatch):
|
||||
built = []
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.graphrag.store.GraphStore", lambda: built.append(object()) or built[-1]
|
||||
)
|
||||
tool = GraphSearchTool({"source": SOURCE})
|
||||
|
||||
assert tool._get_store() is tool._get_store()
|
||||
assert len(built) == 1
|
||||
|
||||
def test_an_embedding_failure_makes_entity_search_unavailable(self, monkeypatch):
|
||||
def _broken_embeddings():
|
||||
raise RuntimeError("no model")
|
||||
|
||||
monkeypatch.setattr("docsgpt.vectorstore.base.get_embeddings", _broken_embeddings)
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
tool = GraphSearchTool({"source": SOURCE})
|
||||
tool._store = _StubStore()
|
||||
|
||||
assert tool.execute_action("search_entities", query="quill") == "Entity search is unavailable."
|
||||
|
||||
def test_no_matching_entities_is_stated_plainly(self, monkeypatch):
|
||||
tool = _tool(monkeypatch, _StubStore())
|
||||
|
||||
assert "No entities found" in tool.execute_action("search_entities", query="quill")
|
||||
|
||||
def test_a_failing_store_is_reported_not_raised(self, monkeypatch):
|
||||
class _BrokenStore(_StubStore):
|
||||
def entity_relationships(self, source_id, name, limit=25):
|
||||
raise RuntimeError("connection lost")
|
||||
|
||||
tool = _tool(monkeypatch, _BrokenStore())
|
||||
|
||||
assert tool.execute_action("get_relationships", entity="Quill") == "The graph lookup failed."
|
||||
|
||||
|
||||
class TestSourcesHaveGraph:
|
||||
"""Whether to offer the tool at all: only when some source has a graph."""
|
||||
|
||||
def _patch_counts(self, monkeypatch, counts=None, error=None):
|
||||
class _Store:
|
||||
def count_nodes_many(self, source_ids):
|
||||
if error:
|
||||
raise error
|
||||
return {s: counts.get(s, 0) for s in source_ids}
|
||||
|
||||
monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", _Store)
|
||||
|
||||
def test_true_when_any_source_has_nodes(self, monkeypatch):
|
||||
from docsgpt.agents.tools.graph_search import sources_have_graph
|
||||
|
||||
self._patch_counts(monkeypatch, {"b": 12})
|
||||
assert sources_have_graph({"active_docs": ["a", "b"]}) is True
|
||||
|
||||
def test_false_when_no_source_has_nodes(self, monkeypatch):
|
||||
from docsgpt.agents.tools.graph_search import sources_have_graph
|
||||
|
||||
self._patch_counts(monkeypatch, {})
|
||||
assert sources_have_graph({"active_docs": "a"}) is False
|
||||
|
||||
def test_false_without_sources_or_when_the_check_fails(self, monkeypatch):
|
||||
from docsgpt.agents.tools.graph_search import sources_have_graph
|
||||
|
||||
assert sources_have_graph({"active_docs": []}) is False
|
||||
self._patch_counts(monkeypatch, error=RuntimeError("no pgvector"))
|
||||
assert sources_have_graph({"active_docs": ["a"]}) is False
|
||||
|
||||
|
||||
class TestGraphsMustBeAvailable:
|
||||
"""The graph lives in the pgvector store, so the flag alone is not enough.
|
||||
|
||||
With another vector store configured the graph tables are not the ones the
|
||||
sources were ingested into; everything else in the app asks
|
||||
``graphrag_available()``, which requires both.
|
||||
"""
|
||||
|
||||
def test_the_tool_is_not_offered_without_pgvector(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "faiss")
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.agents.tools.graph_search.sources_have_graph",
|
||||
lambda source: pytest.fail("must not reach the database"),
|
||||
)
|
||||
tools = {}
|
||||
add_graph_search_tool(tools, {"source": SOURCE})
|
||||
assert tools == {}
|
||||
|
||||
def test_actions_report_it_rather_than_querying(self, monkeypatch):
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True)
|
||||
monkeypatch.setattr(settings, "VECTOR_STORE", "faiss")
|
||||
tool = GraphSearchTool({"source": SOURCE})
|
||||
tool._store = _StubStore(nodes=[{"name": "Quill", "distance": 0.1}])
|
||||
|
||||
assert "not enabled" in tool.execute_action("search_entities", query="quill")
|
||||
|
||||
|
||||
class TestPooledConnection:
|
||||
"""The tool is cached for the whole agent run; its connection must not be."""
|
||||
|
||||
class _ClosingStore(_StubStore):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.closed = 0
|
||||
|
||||
def close(self):
|
||||
self.closed += 1
|
||||
|
||||
def test_the_connection_goes_back_after_each_action(self, monkeypatch):
|
||||
store = self._ClosingStore(relationships=[{"source": "A", "target": "B", "type": "r"}])
|
||||
tool = _tool(monkeypatch, store)
|
||||
|
||||
tool.execute_action("get_relationships", entity="A")
|
||||
|
||||
# Held open, one pooled connection would be pinned across every LLM
|
||||
# round trip of the run.
|
||||
assert store.closed == 1
|
||||
assert tool._store is None
|
||||
|
||||
def test_a_failing_action_still_releases_it(self, monkeypatch):
|
||||
class _Broken(self._ClosingStore):
|
||||
def entity_relationships(self, source_id, name, limit=25):
|
||||
raise RuntimeError("connection lost")
|
||||
|
||||
store = _Broken()
|
||||
tool = _tool(monkeypatch, store)
|
||||
|
||||
assert tool.execute_action("get_relationships", entity="A") == "The graph lookup failed."
|
||||
assert store.closed == 1
|
||||
assert tool._store is None
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Tests for canonical entity naming (the key graph nodes are merged on).
|
||||
|
||||
Two failure directions matter. Too little folding splits one entity across
|
||||
nodes ("agent" / "agents", "VECTOR_STORE" / "vector stores"), so the walk never
|
||||
connects what the text connects. Too much folding invents entities: stripping
|
||||
the "s" off ``postgres`` or ``redis`` would merge nothing and create a node no
|
||||
chunk ever named.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.graphrag.naming import canonical_name, normalize_entity_name
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCanonicalName:
|
||||
@pytest.mark.parametrize(
|
||||
"variants, key",
|
||||
[
|
||||
(["VECTOR_STORE", "Vector store", "vector stores", "vector-stores"], "vector store"),
|
||||
([".env file", "env_file", "ENV FILE"], "env file"),
|
||||
(["Celery worker", "Celery workers"], "celery worker"),
|
||||
(["agent", "Agents", "agents!"], "agent"),
|
||||
],
|
||||
)
|
||||
def test_orthographic_variants_share_one_key(self, variants, key):
|
||||
assert {canonical_name(v) for v in variants} == {key}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"singular, plural",
|
||||
[
|
||||
("policy", "policies"),
|
||||
("index", "indexes"),
|
||||
("batch", "batches"),
|
||||
("hash", "hashes"),
|
||||
("class", "classes"),
|
||||
("process", "processes"),
|
||||
("status", "statuses"),
|
||||
("bus", "buses"),
|
||||
("alias", "aliases"),
|
||||
("document", "documents"),
|
||||
("service", "services"),
|
||||
# Singulars ending in "e" whose plural also ends in "-es": the
|
||||
# plural alone cannot say whether to drop "s" or "es".
|
||||
("cache", "caches"),
|
||||
("database", "databases"),
|
||||
("response", "responses"),
|
||||
("release", "releases"),
|
||||
("case", "cases"),
|
||||
("size", "sizes"),
|
||||
("cookie", "cookies"),
|
||||
],
|
||||
)
|
||||
def test_singular_and_plural_share_one_key(self, singular, plural):
|
||||
assert canonical_name(singular) == canonical_name(plural)
|
||||
|
||||
def test_the_key_need_not_be_a_word(self):
|
||||
# It is a merge key, never shown: "cache" and "caches" meet at the
|
||||
# stem an "-es" plural cannot see past, rather than guessing a form.
|
||||
assert canonical_name("caches") == "cach"
|
||||
assert canonical_name("batches") == "batch"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"word",
|
||||
["postgres", "kubernetes", "redis", "https", "status", "analysis", "access", "docs", "series"],
|
||||
)
|
||||
def test_words_that_only_look_plural_are_left_alone(self, word):
|
||||
assert canonical_name(word) == word
|
||||
|
||||
@pytest.mark.parametrize("word", ["class", "corpus", "thesis"])
|
||||
def test_ss_us_is_endings_are_never_stripped(self, word):
|
||||
assert canonical_name(word) == word
|
||||
|
||||
@pytest.mark.parametrize("word", ["aws", "ids", "ops"])
|
||||
def test_short_words_are_left_alone(self, word):
|
||||
assert canonical_name(word) == word
|
||||
|
||||
def test_each_word_of_a_phrase_is_folded(self):
|
||||
assert canonical_name("Postgres Replicas") == "postgres replica"
|
||||
|
||||
@pytest.mark.parametrize("name", [None, "", " ", "!!!", "--_--"])
|
||||
def test_a_name_with_nothing_left_is_no_entity(self, name):
|
||||
assert canonical_name(name) == ""
|
||||
|
||||
def test_normalize_entity_name_is_the_canonical_key(self):
|
||||
assert normalize_entity_name("Vector Stores") == canonical_name("Vector Stores")
|
||||
@@ -0,0 +1,173 @@
|
||||
"""The graph retriever's default path, end to end through ``_graph_docs_for_source``.
|
||||
|
||||
The shipped defaults — seed from entities, walk the passages, blend with vector
|
||||
search — are the configuration that measured best, so they are what most graph
|
||||
sources run. This drives that whole path with a store that returns real values,
|
||||
and checks each per-source option actually switches its stage off.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from docsgpt.retriever.graph_rag import GraphRAGRetriever
|
||||
from docsgpt.storage.db.source_config import RetrievalConfig
|
||||
|
||||
TEXTS = {
|
||||
"c-alder": "Alder streams audit events to Quill.",
|
||||
"c-quill": "Quill is compacted every six hours.",
|
||||
}
|
||||
VECTOR_ONLY = "A passage only plain vector search found."
|
||||
|
||||
|
||||
class _Store:
|
||||
"""A two-entity chain: the question matches Alder, the answer is on Quill."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls: list[str] = []
|
||||
|
||||
def search_nodes_by_embedding(self, source_id, query_embedding, k=10):
|
||||
return [{"id": "alder", "name": "Alder", "distance": 0.1}]
|
||||
|
||||
def get_subgraph(self, source_id, node_ids, hops=1):
|
||||
return {
|
||||
"nodes": [{"id": "alder", "doc_freq": 1}, {"id": "quill", "doc_freq": 1}],
|
||||
"edges": [{"src_node_id": "alder", "dst_node_id": "quill", "weight": 1.0}],
|
||||
}
|
||||
|
||||
def get_chunk_ids_for_nodes(self, source_id, node_ids):
|
||||
return {"alder": ["c-alder"], "quill": ["c-quill"]}
|
||||
|
||||
def chunk_similarities(self, source_id, chunk_ids, query_embedding):
|
||||
self.calls.append("chunk_similarities")
|
||||
return {"c-alder": 0.9, "c-quill": 0.2}
|
||||
|
||||
def get_chunk_texts(self, source_id, chunk_ids):
|
||||
return {
|
||||
c: {"text": TEXTS[c], "metadata": {"title": c}}
|
||||
for c in chunk_ids
|
||||
if c in TEXTS
|
||||
}
|
||||
|
||||
|
||||
def _retriever(per_source=None):
|
||||
"""A retriever without its constructor (which builds a ClassicRAG)."""
|
||||
retriever = object.__new__(GraphRAGRetriever)
|
||||
retriever.chunks = 3
|
||||
retriever.base_chunks = None
|
||||
retriever.doc_token_limit = 50000
|
||||
retriever.vectorstores = ["src"]
|
||||
retriever.per_source_retrieval = per_source or {}
|
||||
retriever.vector_calls = 0
|
||||
|
||||
def _vector_ranking(source_id, query_embedding):
|
||||
retriever.vector_calls += 1
|
||||
return [(VECTOR_ONLY, {"title": "vector"})]
|
||||
|
||||
retriever._vector_ranking = _vector_ranking
|
||||
return retriever
|
||||
|
||||
|
||||
def _texts(docs):
|
||||
return [doc["text"] for doc in docs]
|
||||
|
||||
|
||||
class TestDefaultPath:
|
||||
def test_walks_passages_and_blends_in_vector_hits(self):
|
||||
store = _Store()
|
||||
retriever = _retriever()
|
||||
|
||||
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
||||
|
||||
# The answer sits one edge away from the seed: the walk reached it.
|
||||
assert TEXTS["c-quill"] in _texts(docs)
|
||||
# A hit only vector search found is blended in, not lost.
|
||||
assert VECTOR_ONLY in _texts(docs)
|
||||
assert store.calls == ["chunk_similarities"]
|
||||
assert retriever.vector_calls == 1
|
||||
|
||||
|
||||
class TestPerSourceOptions:
|
||||
def test_passage_walk_can_be_switched_off(self):
|
||||
store = _Store()
|
||||
retriever = _retriever(
|
||||
{"src": RetrievalConfig(chunks=3, graph={"passage_nodes": False})}
|
||||
)
|
||||
|
||||
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
||||
|
||||
assert "chunk_similarities" not in store.calls
|
||||
assert TEXTS["c-quill"] in _texts(docs)
|
||||
|
||||
def test_vector_blending_can_be_switched_off(self):
|
||||
store = _Store()
|
||||
retriever = _retriever(
|
||||
{"src": RetrievalConfig(chunks=3, graph={"blend_vector": False})}
|
||||
)
|
||||
|
||||
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
||||
|
||||
assert retriever.vector_calls == 0
|
||||
assert VECTOR_ONLY not in _texts(docs)
|
||||
|
||||
|
||||
class TestVectorRanking:
|
||||
"""The vector half of the blend, keyed on chunk text since hits carry no id."""
|
||||
|
||||
class _VectorStore:
|
||||
def __init__(self, hits=None, error=None):
|
||||
self.hits = hits or []
|
||||
self.error = error
|
||||
self.searched = None
|
||||
self.closed = False
|
||||
|
||||
def search(self, question, k, query_vector=None):
|
||||
self.searched = (question, k, query_vector)
|
||||
if self.error:
|
||||
raise self.error
|
||||
return self.hits
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
@staticmethod
|
||||
def _real_retriever(monkeypatch, store):
|
||||
from types import SimpleNamespace
|
||||
|
||||
retriever = object.__new__(GraphRAGRetriever)
|
||||
retriever.chunks = 3
|
||||
retriever._classic = SimpleNamespace(_get_rephrased_question=lambda: "where does Alder stream?")
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.vectorstore.vector_creator.VectorCreator.create_vectorstore",
|
||||
lambda *args, **kwargs: store,
|
||||
)
|
||||
return retriever
|
||||
|
||||
def test_object_and_dict_hits_become_text_and_metadata(self, monkeypatch):
|
||||
from types import SimpleNamespace
|
||||
|
||||
store = self._VectorStore(
|
||||
hits=[
|
||||
SimpleNamespace(page_content="Alder streams to Quill.", metadata={"title": "alder.md"}),
|
||||
{"text": "Quill is compacted every six hours.", "metadata": {"title": "quill.md"}},
|
||||
{"page_content": "A passage without metadata."},
|
||||
{"metadata": {"title": "no text"}},
|
||||
]
|
||||
)
|
||||
retriever = self._real_retriever(monkeypatch, store)
|
||||
|
||||
ranked = retriever._vector_ranking("src", [0.1, 0.2])
|
||||
|
||||
assert ranked == [
|
||||
("Alder streams to Quill.", {"title": "alder.md"}),
|
||||
("Quill is compacted every six hours.", {"title": "quill.md"}),
|
||||
("A passage without metadata.", {}),
|
||||
]
|
||||
# The rephrased question and the shared query vector, with room to fuse.
|
||||
assert store.searched == ("where does Alder stream?", 20, [0.1, 0.2])
|
||||
assert store.closed
|
||||
|
||||
def test_a_failed_search_ranks_nothing_and_still_closes_the_store(self, monkeypatch):
|
||||
store = self._VectorStore(error=RuntimeError("pgvector down"))
|
||||
retriever = self._real_retriever(monkeypatch, store)
|
||||
|
||||
assert retriever._vector_ranking("src", [0.1]) == []
|
||||
assert store.closed
|
||||
@@ -0,0 +1,151 @@
|
||||
"""Chunks as nodes in the walk, and the damping that decides how far mass spreads.
|
||||
|
||||
``_rank_chunks`` reads a chunk's score off its entities by summing their PPR
|
||||
mass, which rewards a chunk for touching *many* entities rather than the right
|
||||
ones. The passage-node path puts the chunks in the graph instead, so a chunk is
|
||||
reachable both by being about the question and by being connected to what is.
|
||||
|
||||
These tests use a stub store: the ranking is graph arithmetic, and pinning it
|
||||
against a real database would measure Postgres rather than the ranking.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.retriever.graph_rag import GraphRAGRetriever, _damping
|
||||
|
||||
|
||||
class _StubStore:
|
||||
"""The two reads the passage path makes, and nothing else."""
|
||||
|
||||
def __init__(self, chunk_links, similarities):
|
||||
self._chunk_links = chunk_links
|
||||
self._similarities = similarities
|
||||
|
||||
def get_chunk_ids_for_nodes(self, source_id, node_ids):
|
||||
return {n: c for n, c in self._chunk_links.items() if n in set(node_ids)}
|
||||
|
||||
def chunk_similarities(self, source_id, chunk_ids, query_embedding):
|
||||
return {c: self._similarities.get(c, 0.0) for c in chunk_ids}
|
||||
|
||||
|
||||
def _retriever(chunks=2):
|
||||
"""A retriever without its constructor — which builds a ClassicRAG, opens
|
||||
settings-driven collaborators, and has nothing to do with ranking."""
|
||||
retriever = object.__new__(GraphRAGRetriever)
|
||||
retriever.chunks = chunks
|
||||
return retriever
|
||||
|
||||
|
||||
def _subgraph():
|
||||
return {
|
||||
"nodes": [
|
||||
{"id": "a", "doc_freq": 1},
|
||||
{"id": "b", "doc_freq": 1},
|
||||
{"id": "hub", "doc_freq": 40},
|
||||
],
|
||||
"edges": [
|
||||
{"src_node_id": "a", "dst_node_id": "hub", "weight": 1.0},
|
||||
{"src_node_id": "b", "dst_node_id": "hub", "weight": 1.0},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class TestDamping:
|
||||
"""Each ranking mode runs at the damping it was measured at."""
|
||||
|
||||
def test_passage_walk_keeps_mass_near_the_seeds(self):
|
||||
assert _damping(passage_nodes=True) == 0.5
|
||||
|
||||
def test_entity_only_ranking_keeps_the_conventional_value(self):
|
||||
assert _damping(passage_nodes=False) == 0.85
|
||||
|
||||
|
||||
class TestPassageNodes:
|
||||
def test_ranks_the_chunk_the_question_matches(self, monkeypatch):
|
||||
"""Two chunks are equally connected; only their own relevance differs,
|
||||
so the more relevant one must win."""
|
||||
store = _StubStore(
|
||||
chunk_links={"a": ["c1"], "b": ["c2"]},
|
||||
similarities={"c1": 0.1, "c2": 0.9},
|
||||
)
|
||||
|
||||
ranked = _retriever()._rank_chunks_with_passages(
|
||||
store, "src", _subgraph(), {"a": 1.0, "b": 1.0}, [0.0] * 4
|
||||
)
|
||||
|
||||
assert ranked[0] == "c2"
|
||||
|
||||
def test_a_chunk_reached_only_through_the_graph_still_ranks(self, monkeypatch):
|
||||
"""The point of the walk: a chunk with no similarity of its own is
|
||||
still reachable through the entity the seeds point at."""
|
||||
store = _StubStore(
|
||||
chunk_links={"a": ["c1"], "b": ["c2"]},
|
||||
similarities={"c1": 0.0, "c2": 0.0},
|
||||
)
|
||||
|
||||
ranked = _retriever()._rank_chunks_with_passages(
|
||||
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
|
||||
)
|
||||
|
||||
assert set(ranked) == {"c1", "c2"}
|
||||
|
||||
def test_no_linked_chunks_returns_nothing(self):
|
||||
store = _StubStore(chunk_links={}, similarities={})
|
||||
|
||||
assert (
|
||||
_retriever()._rank_chunks_with_passages(
|
||||
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
|
||||
)
|
||||
== []
|
||||
)
|
||||
|
||||
def test_over_fetches_past_the_chunk_budget(self, monkeypatch):
|
||||
"""Same contract as ``_rank_chunks``: candidates exceed the budget so
|
||||
chunks with missing text cannot drop the final count below it."""
|
||||
links = {"a": [f"c{i}" for i in range(10)]}
|
||||
store = _StubStore(
|
||||
chunk_links=links,
|
||||
similarities={f"c{i}": i / 10 for i in range(10)},
|
||||
)
|
||||
|
||||
ranked = _retriever(chunks=2)._rank_chunks_with_passages(
|
||||
store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4
|
||||
)
|
||||
|
||||
assert len(ranked) == max(2 * 2, 2 + 5)
|
||||
|
||||
|
||||
class TestChunkSimilaritiesGuard:
|
||||
"""The store call the passage path depends on short-circuits before it
|
||||
touches a connection, so an empty subgraph costs no query."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"chunk_ids,embedding", [([], [0.1]), (["c1"], []), ([], [])]
|
||||
)
|
||||
def test_empty_inputs_return_empty(self, chunk_ids, embedding):
|
||||
from docsgpt.graphrag.store import GraphStore
|
||||
|
||||
store = object.__new__(GraphStore)
|
||||
|
||||
assert store.chunk_similarities("src", chunk_ids, embedding) == {}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestNodesByChunk:
|
||||
"""The passage stage inverts node->chunks once instead of rescanning."""
|
||||
|
||||
def test_inverts_and_keeps_node_order(self):
|
||||
from docsgpt.retriever.graph_rag import _nodes_by_chunk
|
||||
|
||||
assert _nodes_by_chunk({"n1": ["c1", "c2"], "n2": ["c2"], "n3": []}) == {
|
||||
"c1": ["n1"],
|
||||
"c2": ["n1", "n2"],
|
||||
}
|
||||
|
||||
def test_no_links_invert_to_nothing(self):
|
||||
from docsgpt.retriever.graph_rag import _nodes_by_chunk
|
||||
|
||||
assert _nodes_by_chunk({}) == {}
|
||||
assert _nodes_by_chunk({"n1": None}) == {}
|
||||
@@ -0,0 +1,129 @@
|
||||
"""Where the graph walk starts, per the source's graph retrieval options.
|
||||
|
||||
Seeding decides more than ranking does — a walk that starts on the wrong nodes
|
||||
cannot be rescued downstream. The options are per source and live (no
|
||||
re-ingest), carried on the per-source retrieval config the Dispatcher hands the
|
||||
retriever, so both the dispatch and how the options are resolved are pinned.
|
||||
|
||||
The fallback matters most: relationship seeding reads fact embeddings written
|
||||
at ingest, and a source built before they were recorded has none. It must keep
|
||||
retrieving through entity matching rather than returning nothing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from docsgpt.retriever.graph_rag import GraphRAGRetriever
|
||||
from docsgpt.storage.db.source_config import GraphRetrievalConfig, RetrievalConfig
|
||||
|
||||
ENTITY_ROWS = [{"id": "n1", "name": "Quill", "distance": 0.2}]
|
||||
FACT_ROWS = [
|
||||
{"id": "f1", "name": "Alder", "distance": 0.1},
|
||||
{"id": "f2", "name": "Quill", "distance": 0.1},
|
||||
]
|
||||
|
||||
|
||||
class _StubStore:
|
||||
def __init__(self, fact_rows=None):
|
||||
self.fact_rows = list(FACT_ROWS) if fact_rows is None else fact_rows
|
||||
self.calls: list[str] = []
|
||||
|
||||
def seed_nodes_from_facts(self, source_id, query_embedding, fact_limit=5, limit=10):
|
||||
self.calls.append("facts")
|
||||
return self.fact_rows
|
||||
|
||||
def search_nodes_by_embedding(self, source_id, query_embedding, k=10):
|
||||
self.calls.append("entities")
|
||||
return list(ENTITY_ROWS)
|
||||
|
||||
|
||||
def _retriever(per_source=None):
|
||||
"""A retriever without its constructor, which builds a ClassicRAG."""
|
||||
retriever = object.__new__(GraphRAGRetriever)
|
||||
if per_source is not None:
|
||||
retriever.per_source_retrieval = per_source
|
||||
return retriever
|
||||
|
||||
|
||||
def _relationships_config():
|
||||
return RetrievalConfig(graph={"seed_strategy": "relationships"})
|
||||
|
||||
|
||||
class TestDefaults:
|
||||
def test_measured_best_configuration_is_the_default(self):
|
||||
options = GraphRetrievalConfig()
|
||||
|
||||
assert options.seed_strategy == "entities"
|
||||
assert options.passage_nodes is True
|
||||
assert options.blend_vector is True
|
||||
|
||||
def test_a_source_with_no_per_source_config_gets_the_defaults(self):
|
||||
assert _retriever()._graph_options("src") == GraphRetrievalConfig()
|
||||
|
||||
def test_existing_retrieval_configs_validate_without_the_new_block(self):
|
||||
"""Source configs saved before this existed carry no ``graph`` key."""
|
||||
config = RetrievalConfig.model_validate({"retriever": "graphrag"})
|
||||
|
||||
assert config.graph == GraphRetrievalConfig()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad", [{"seed_strategy": "vector"}, {"seed_strategy": "union"}, {"damping": 0.5}]
|
||||
)
|
||||
def test_retired_and_unknown_options_are_rejected(self, bad):
|
||||
with pytest.raises(ValidationError):
|
||||
GraphRetrievalConfig.model_validate(bad)
|
||||
|
||||
|
||||
class TestEntitySeeding:
|
||||
def test_seeds_from_entities_and_never_reads_facts(self):
|
||||
store = _StubStore()
|
||||
|
||||
rows = _retriever()._seed_rows(store, "src", [0.0, 0.1])
|
||||
|
||||
assert [r["id"] for r in rows] == ["n1"]
|
||||
assert store.calls == ["entities"]
|
||||
|
||||
|
||||
class TestRelationshipSeeding:
|
||||
def test_seeds_from_facts_when_the_source_asks_for_it(self):
|
||||
store = _StubStore()
|
||||
retriever = _retriever({"src": _relationships_config()})
|
||||
|
||||
rows = retriever._seed_rows(store, "src", [0.0, 0.1])
|
||||
|
||||
assert [r["id"] for r in rows] == ["f1", "f2"]
|
||||
assert store.calls == ["facts"]
|
||||
|
||||
def test_reads_the_option_from_a_plain_dict_config_too(self):
|
||||
store = _StubStore()
|
||||
retriever = _retriever({"src": {"graph": {"seed_strategy": "relationships"}}})
|
||||
|
||||
retriever._seed_rows(store, "src", [0.0, 0.1])
|
||||
|
||||
assert store.calls == ["facts"]
|
||||
|
||||
def test_falls_back_to_entities_without_fact_embeddings(self):
|
||||
"""A source built before fact embeddings were recorded still retrieves."""
|
||||
store = _StubStore(fact_rows=[])
|
||||
retriever = _retriever({"src": _relationships_config()})
|
||||
|
||||
rows = retriever._seed_rows(store, "src", [0.0, 0.1])
|
||||
|
||||
assert [r["id"] for r in rows] == ["n1"]
|
||||
assert store.calls == ["facts", "entities"]
|
||||
|
||||
def test_options_are_per_source(self):
|
||||
store = _StubStore()
|
||||
retriever = _retriever({"other": _relationships_config()})
|
||||
|
||||
retriever._seed_rows(store, "src", [0.0, 0.1])
|
||||
|
||||
assert store.calls == ["entities"]
|
||||
|
||||
def test_a_malformed_stored_option_falls_back_to_the_defaults(self):
|
||||
"""A bad value must not take graph retrieval down with it."""
|
||||
retriever = _retriever({"src": {"graph": {"seed_strategy": "nonsense"}}})
|
||||
|
||||
assert retriever._graph_options("src") == GraphRetrievalConfig()
|
||||
@@ -19,6 +19,7 @@ import uuid
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from psycopg.types.json import Jsonb
|
||||
|
||||
import docsgpt.graphrag.store as store_module
|
||||
from docsgpt.vectorstore import pgconn
|
||||
@@ -157,6 +158,122 @@ class TestGraphStoreLive:
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_seed_nodes_from_facts_returns_both_endpoints_of_the_match(
|
||||
self, store, source_id
|
||||
):
|
||||
"""Fact seeding's whole point: the question matches the *relationship*,
|
||||
and both of its endpoints become seeds — including the one the question
|
||||
never names."""
|
||||
try:
|
||||
alder = store.upsert_node(source_id, "Alder", "alder", "service", "d")
|
||||
quill = store.upsert_node(source_id, "Quill", "quill", "store", "d")
|
||||
birch = store.upsert_node(source_id, "Birch", "birch", "service", "d")
|
||||
ridge = store.upsert_node(source_id, "Ridge", "ridge", "store", "d")
|
||||
store.add_edge(
|
||||
source_id, alder, quill, "streams_to", "Alder streams to Quill",
|
||||
1.0, ["c1"], fact_embedding=_embedding(1.0),
|
||||
)
|
||||
store.add_edge(
|
||||
source_id, birch, ridge, "streams_to", "Birch streams to Ridge",
|
||||
1.0, ["c2"], fact_embedding=_embedding(-1.0),
|
||||
)
|
||||
|
||||
rows = store.seed_nodes_from_facts(
|
||||
source_id, _embedding(1.0), fact_limit=1, limit=10
|
||||
)
|
||||
|
||||
assert {row["name"] for row in rows} == {"Alder", "Quill"}
|
||||
assert all(row["distance"] <= 1.0 for row in rows)
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_seed_nodes_from_facts_is_empty_without_fact_embeddings(
|
||||
self, store, source_id
|
||||
):
|
||||
"""A source ingested before fact embeddings existed returns nothing,
|
||||
which is the signal the retriever falls back to name matching on."""
|
||||
try:
|
||||
a = store.upsert_node(source_id, "A", "a")
|
||||
b = store.upsert_node(source_id, "B", "b")
|
||||
store.add_edge(source_id, a, b, "rel")
|
||||
|
||||
assert store.seed_nodes_from_facts(source_id, _embedding(1.0)) == []
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_add_edge_skips_self_loops(self, store, source_id):
|
||||
"""A relationship whose endpoints resolve to one node is noise.
|
||||
|
||||
A self-loop feeds a node's PageRank mass straight back to itself, and a
|
||||
real extraction produced 121 of them on a 98-page corpus.
|
||||
"""
|
||||
try:
|
||||
a = store.upsert_node(source_id, "A", "a", "thing", "desc a")
|
||||
assert (
|
||||
store.add_edge(source_id, a, a, "related", "a relates to a", 1.0, ["c1"])
|
||||
is None
|
||||
)
|
||||
assert store.get_subgraph(source_id, [a], hops=1)["edges"] == []
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_add_edge_merges_a_repeated_pair(self, store, source_id):
|
||||
"""The same relationship seen in many chunks is one edge, not many rows.
|
||||
|
||||
``graph_edges`` carries no uniqueness constraint, so re-extracting a
|
||||
relationship used to insert a row per chunk — 19.9% of a real corpus's
|
||||
edges — inflating traversal weight and wasting the subgraph fetch
|
||||
budget. The surviving row keeps the strongest weight and both chunk ids.
|
||||
"""
|
||||
try:
|
||||
a = store.upsert_node(source_id, "A", "a", "thing", "desc a")
|
||||
b = store.upsert_node(source_id, "B", "b", "thing", "desc b")
|
||||
first = store.add_edge(source_id, a, b, "related", "d", 2.0, ["chunk-1"])
|
||||
second = store.add_edge(source_id, a, b, "related", "d", 5.0, ["chunk-2"])
|
||||
|
||||
assert second == first
|
||||
edges = store.get_subgraph(source_id, [a, b], hops=1)["edges"]
|
||||
assert len(edges) == 1
|
||||
assert float(edges[0]["weight"]) == 5.0
|
||||
|
||||
# Both chunks are still recorded as evidence for the merged edge.
|
||||
conn = store._get_connection()
|
||||
cursor = conn.cursor()
|
||||
try:
|
||||
cursor.execute(
|
||||
"SELECT source_chunk_ids FROM graph_edges WHERE id = %s;", (first,)
|
||||
)
|
||||
chunk_ids = cursor.fetchone()[0]
|
||||
finally:
|
||||
cursor.close()
|
||||
conn.rollback()
|
||||
assert sorted(chunk_ids) == ["chunk-1", "chunk-2"]
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_get_subgraph_keeps_the_heaviest_edges_when_capped(
|
||||
self, store, source_id, monkeypatch
|
||||
):
|
||||
"""A capped fetch must drop the weakest edges, not an arbitrary subset.
|
||||
|
||||
The cap is applied with ``LIMIT``; without an ordering Postgres is free
|
||||
to return any rows at all, so a dense graph silently retrieves a random
|
||||
neighbourhood.
|
||||
"""
|
||||
try:
|
||||
a = store.upsert_node(source_id, "A", "a", "thing", "d")
|
||||
b = store.upsert_node(source_id, "B", "b", "thing", "d")
|
||||
c = store.upsert_node(source_id, "C", "c", "thing", "d")
|
||||
store.add_edge(source_id, a, b, "light", "d", 1.0, ["c1"])
|
||||
store.add_edge(source_id, a, c, "heavy", "d", 9.0, ["c1"])
|
||||
|
||||
monkeypatch.setattr(store_module, "MAX_SUBGRAPH_EDGES", 1)
|
||||
edges = store.get_subgraph(source_id, [a], hops=1)["edges"]
|
||||
|
||||
assert [e["type"] for e in edges] == ["heavy"]
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_apply_chunk_writes_nodes_links_and_edges(self, store, source_id):
|
||||
"""One transactional write: entities linked to the chunk, edges added,
|
||||
and a bare relationship endpoint upserted but not chunk-linked."""
|
||||
@@ -203,18 +320,23 @@ class TestGraphStoreLive:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_self_loop_degree_agrees_across_paths(self, store, source_id):
|
||||
"""``add_edge``'s incremental +1 and ``set_node_degrees`` recompute must
|
||||
agree on a self-loop (count it once)."""
|
||||
"""``add_edge``'s incremental bump and ``set_node_degrees`` recompute must
|
||||
agree on a self-loop.
|
||||
|
||||
They now agree on zero rather than one: the self-loop is rejected at
|
||||
write time, so neither path has an edge to count. The property under
|
||||
test is that the two paths agree, not the number they agree on.
|
||||
"""
|
||||
try:
|
||||
node = store.upsert_node(source_id, "Solo", "solo")
|
||||
store.add_edge(source_id, node, node, "self")
|
||||
assert store.add_edge(source_id, node, node, "self") is None
|
||||
|
||||
incremental = store.get_node_by_normalized(source_id, "solo")["degree"]
|
||||
assert incremental == 1
|
||||
assert incremental == 0
|
||||
|
||||
store.set_node_degrees(source_id)
|
||||
recomputed = store.get_node_by_normalized(source_id, "solo")["degree"]
|
||||
assert recomputed == 1
|
||||
assert recomputed == incremental
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
@@ -436,17 +558,63 @@ class TestGraphStoreParameterization:
|
||||
store._tables_ensured = True
|
||||
return store, cursor
|
||||
|
||||
def test_graph_writes_for_a_source_are_serialized(self):
|
||||
# A chunk write and a reset each take the source's transaction-scoped
|
||||
# advisory lock before touching a row, so overlapping builds of one
|
||||
# source cannot interleave inside a chunk.
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
cursor.fetchone.return_value = None
|
||||
sid = str(uuid.uuid4())
|
||||
|
||||
store.apply_chunk(sid, "c1", [], [], {})
|
||||
first_sql, first_params = cursor.execute.call_args_list[0].args
|
||||
assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql
|
||||
assert first_params == (f"graphrag:source:{sid}",)
|
||||
|
||||
cursor.execute.reset_mock()
|
||||
store.delete_by_source(sid)
|
||||
first_sql, first_params = cursor.execute.call_args_list[0].args
|
||||
assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql
|
||||
assert first_params == (f"graphrag:source:{sid}",)
|
||||
|
||||
def test_apply_chunk_keeps_an_explicit_zero_weight(self, monkeypatch):
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
cursor.fetchone.side_effect = [None, ["n1"], ["n2"]]
|
||||
weights = []
|
||||
|
||||
def _capture(cursor, source_id, src, dst, type=None, description=None, weight=1.0, **kwargs):
|
||||
weights.append(weight)
|
||||
return "e1", True
|
||||
|
||||
monkeypatch.setattr(store, "_add_edge", _capture)
|
||||
store.apply_chunk(
|
||||
"sid", "c1", [],
|
||||
[{"source": "A", "target": "B", "weight": 0}, {"source": "A", "target": "B"}],
|
||||
{},
|
||||
)
|
||||
# Zero is a real weight; only a missing one defaults.
|
||||
assert weights == [0.0, 1.0]
|
||||
|
||||
def test_delete_by_source_binds_source_id(self):
|
||||
from psycopg import sql as pgsql
|
||||
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
sid = str(uuid.uuid4())
|
||||
store.delete_by_source(sid)
|
||||
|
||||
for call in cursor.execute.call_args_list:
|
||||
sql = call.args[0]
|
||||
tables = []
|
||||
lock, *deletes = cursor.execute.call_args_list
|
||||
assert "pg_advisory_xact_lock" in lock.args[0]
|
||||
for call in deletes:
|
||||
query = call.args[0]
|
||||
params = call.args[1] if len(call.args) > 1 else None
|
||||
assert isinstance(query, pgsql.Composable)
|
||||
sql = query.as_string()
|
||||
assert "WHERE source_id = %s" in sql
|
||||
assert sid not in sql
|
||||
assert params == (sid,)
|
||||
tables.append(sql.split('"')[1])
|
||||
assert tables == ["graph_node_chunks", "graph_edges", "graph_nodes", "graph_ingest_progress"]
|
||||
|
||||
def test_search_binds_embedding_and_source(self):
|
||||
store, cursor = self._store_with_mock_conn()
|
||||
@@ -494,6 +662,217 @@ class TestGraphStoreParameterization:
|
||||
assert params[-1] == embedding
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestEntityPagesLive:
|
||||
"""``entity_pages`` against a real pgvector-shaped table.
|
||||
|
||||
The graph tables alone cannot answer it: the rows it returns live in the
|
||||
documents table the sources were ingested into, so the test creates a
|
||||
minimal one with the same column names ``PGVectorStore`` uses.
|
||||
"""
|
||||
|
||||
@pytest.fixture
|
||||
def store(self, postgresql):
|
||||
store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info))
|
||||
try:
|
||||
store._ensure_tables()
|
||||
except Exception as exc:
|
||||
pytest.skip(f"pgvector extension unavailable: {exc}")
|
||||
conn = store._get_connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS documents (
|
||||
id SERIAL PRIMARY KEY,
|
||||
text TEXT,
|
||||
metadata JSONB,
|
||||
source_id TEXT
|
||||
);
|
||||
"""
|
||||
)
|
||||
conn.commit()
|
||||
cursor.close()
|
||||
yield store
|
||||
store.close()
|
||||
|
||||
def test_a_page_linked_by_two_entities_is_returned_once(self, store):
|
||||
"""One chunk, two nodes whose names both match: an exact hit and a
|
||||
mention. They differ only in whether the page is *about* the entity, so
|
||||
grouping on that flag returned the same page twice and spent a quarter
|
||||
of the page budget on it."""
|
||||
source_id = str(uuid.uuid4())
|
||||
conn = store._get_connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.execute(
|
||||
"INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;",
|
||||
("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id),
|
||||
)
|
||||
chunk_id = str(cursor.fetchone()[0])
|
||||
conn.commit()
|
||||
cursor.close()
|
||||
try:
|
||||
subject = store.upsert_node(source_id, "Quill", "quill")
|
||||
mention = store.upsert_node(source_id, "Legacy Quill", "legacy quill")
|
||||
store.link_node_chunk(source_id, subject, chunk_id)
|
||||
store.link_node_chunk(source_id, mention, chunk_id)
|
||||
|
||||
pages = store.entity_pages(source_id, "Quill", limit=4)
|
||||
|
||||
assert [page["text"] for page in pages] == ["Quill is a write-ahead store."]
|
||||
assert pages[0]["metadata"] == {"title": "quill.md"}
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
|
||||
def test_two_chunks_with_identical_text_collapse_into_one_page(self, store):
|
||||
"""Deliberate: the caller gets at most four pages to hand a model, and a
|
||||
crawl that ingested the same text twice would spend two of them saying
|
||||
the same thing. The rows differ only by an id the model never sees."""
|
||||
source_id = str(uuid.uuid4())
|
||||
conn = store._get_connection()
|
||||
cursor = conn.cursor()
|
||||
chunk_ids = []
|
||||
for _ in range(2):
|
||||
cursor.execute(
|
||||
"INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;",
|
||||
("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id),
|
||||
)
|
||||
chunk_ids.append(str(cursor.fetchone()[0]))
|
||||
conn.commit()
|
||||
cursor.close()
|
||||
try:
|
||||
node = store.upsert_node(source_id, "Quill", "quill")
|
||||
for chunk_id in chunk_ids:
|
||||
store.link_node_chunk(source_id, node, chunk_id)
|
||||
|
||||
pages = store.entity_pages(source_id, "Quill", limit=4)
|
||||
|
||||
assert [page["text"] for page in pages] == ["Quill is a write-ahead store."]
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGraphReadQueries:
|
||||
"""The reads behind fact seeding and the agent's graph tool, without a DB.
|
||||
|
||||
The live class covers what these return from real rows; these pin the
|
||||
contract that holds without one. The entity name reaching
|
||||
``entity_relationships``/``entity_pages`` comes from an LLM tool call, so
|
||||
it must only ever travel as a bound parameter.
|
||||
"""
|
||||
|
||||
def _store(self, rows=(), fail=False):
|
||||
store = GraphStore.__new__(GraphStore)
|
||||
cursor = MagicMock()
|
||||
cursor.fetchall.return_value = list(rows)
|
||||
if fail:
|
||||
cursor.execute.side_effect = RuntimeError("relation does not exist")
|
||||
conn = MagicMock()
|
||||
conn.cursor.return_value = cursor
|
||||
store._get_connection = lambda: conn
|
||||
return store, cursor, conn
|
||||
|
||||
def test_identifiers_are_quoted_as_postgres_folds_them_unquoted(self):
|
||||
# PGVectorStore writes these names unquoted, which Postgres folds to
|
||||
# lower case; quoting keeps case, so the fold happens first or the two
|
||||
# stores would address different tables.
|
||||
assert store_module._identifier("Documents").as_string() == '"documents"'
|
||||
with pytest.raises(ValueError):
|
||||
store_module._identifier('documents"; DROP TABLE graph_nodes; --')
|
||||
|
||||
def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self):
|
||||
store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)])
|
||||
sid = str(uuid.uuid4())
|
||||
embedding = _embedding(0.3)
|
||||
|
||||
rows = store.seed_nodes_from_facts(sid, embedding, fact_limit=0, limit=3)
|
||||
|
||||
sql, params = cursor.execute.call_args.args
|
||||
assert sid not in sql and str(embedding) not in sql
|
||||
# Limits are clamped to at least one before binding.
|
||||
assert params == (embedding, sid, embedding, 1, sid, 3)
|
||||
assert rows[0] == {"id": "n1", "name": "Quill", "description": "a store", "distance": pytest.approx(0.2)}
|
||||
assert rows[1]["distance"] == 1.0
|
||||
|
||||
def test_fact_seeds_need_a_query_vector(self):
|
||||
store, cursor, _ = self._store()
|
||||
assert store.seed_nodes_from_facts(str(uuid.uuid4()), []) == []
|
||||
cursor.execute.assert_not_called()
|
||||
|
||||
def test_relationships_bind_the_name_as_a_pattern(self):
|
||||
store, cursor, _ = self._store(rows=[("Alder", "streams_to", "Quill", "audit events")])
|
||||
sid = str(uuid.uuid4())
|
||||
name = "Quill'; DROP TABLE graph_nodes; --"
|
||||
|
||||
rows = store.entity_relationships(sid, f" {name} ", limit=500)
|
||||
|
||||
sql, params = cursor.execute.call_args.args
|
||||
assert name not in sql
|
||||
assert params == (sid, f"%{name}%", f"%{name}%", 500)
|
||||
assert rows == [
|
||||
{"source": "Alder", "type": "streams_to", "target": "Quill", "description": "audit events"}
|
||||
]
|
||||
|
||||
def test_pages_prefer_the_entity_itself_over_a_mention(self):
|
||||
store, cursor, _ = self._store(rows=[({"title": "quill.md"}, "Quill is a store."), (None, None)])
|
||||
sid = str(uuid.uuid4())
|
||||
|
||||
pages = store.entity_pages(sid, "Quill", limit=0)
|
||||
|
||||
query, params = cursor.execute.call_args.args
|
||||
sql = query.as_string()
|
||||
assert 'JOIN "documents" d' in sql and 'd."source_id" = %s' in sql
|
||||
assert "Quill" not in sql
|
||||
# Exact name, name plus a qualifier ("Quill Store"), substring fallback,
|
||||
# text-opens-with ordering, then the clamped limit.
|
||||
assert params == ("quill", "quill %", sid, sid, "quill", "quill %", "%Quill%", "Quill%", 1)
|
||||
assert pages == [{"metadata": {"title": "quill.md"}, "text": "Quill is a store."}, {"metadata": {}, "text": ""}]
|
||||
|
||||
def test_chunk_similarities_are_restricted_to_the_reached_chunks(self):
|
||||
store, cursor, _ = self._store(rows=[("11", 0.75)])
|
||||
sid = str(uuid.uuid4())
|
||||
embedding = _embedding(0.9)
|
||||
|
||||
scores = store.chunk_similarities(sid, [11, "12"], embedding)
|
||||
|
||||
query, params = cursor.execute.call_args.args
|
||||
sql = query.as_string()
|
||||
assert '1 - ("embedding" <=> %s::vector)' in sql and 'FROM "documents"' in sql
|
||||
assert "= ANY(%s)" in sql and sid not in sql
|
||||
assert params == (embedding, sid, ["11", "12"])
|
||||
assert scores == {"11": 0.75}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call",
|
||||
[
|
||||
lambda s: s.entity_relationships("sid", " "),
|
||||
lambda s: s.entity_pages("sid", ""),
|
||||
lambda s: s.chunk_similarities("sid", [], [0.1]),
|
||||
lambda s: s.chunk_similarities("sid", ["1"], []),
|
||||
],
|
||||
)
|
||||
def test_empty_input_runs_no_query(self, call):
|
||||
store, cursor, _ = self._store()
|
||||
assert not call(store)
|
||||
cursor.execute.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call",
|
||||
[
|
||||
lambda s: s.seed_nodes_from_facts("sid", [0.1]),
|
||||
lambda s: s.entity_relationships("sid", "Quill"),
|
||||
lambda s: s.entity_pages("sid", "Quill"),
|
||||
lambda s: s.chunk_similarities("sid", ["1"], [0.1]),
|
||||
],
|
||||
)
|
||||
def test_a_failed_query_returns_nothing_and_releases_the_connection(self, call):
|
||||
store, cursor, conn = self._store(fail=True)
|
||||
assert not call(store)
|
||||
cursor.close.assert_called_once()
|
||||
conn.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEmbeddingDim:
|
||||
"""The graph table dimension is derived from the configured model (FIX 1)."""
|
||||
@@ -890,3 +1269,280 @@ class TestCountNodesMany:
|
||||
store, _, _ = self._store_with_mock_conn([(source_id.lower(), 3)])
|
||||
|
||||
assert store.count_nodes_many([source_id]) == {source_id: 3}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestWritesSurviveALostConnection:
|
||||
"""A graph build holds one pooled connection across its LLM calls.
|
||||
|
||||
Extraction spends minutes per chunk waiting on a model, so the connection
|
||||
sits idle between writes and the server (or a pooler) can drop it. The pool
|
||||
only validates a connection at checkout, and this one was checked out once
|
||||
at the start of the build, so the next write raises and the chunk is marked
|
||||
``failed`` — silently losing it from the graph. The write reconnects and
|
||||
retries once instead; the statements are idempotent upserts, so a retry
|
||||
cannot double-write.
|
||||
"""
|
||||
|
||||
def _store_with_connections(self, conns):
|
||||
"""Store that hands out ``conns`` in order, one per (re)connect."""
|
||||
store = GraphStore.__new__(GraphStore)
|
||||
store._tables_ensured = True
|
||||
store._connection = None
|
||||
handed = []
|
||||
closed = []
|
||||
|
||||
def _get_connection():
|
||||
if store._connection is None:
|
||||
store._connection = conns[len(handed)]
|
||||
handed.append(store._connection)
|
||||
return store._connection
|
||||
|
||||
def _close():
|
||||
if store._connection is not None:
|
||||
closed.append(store._connection)
|
||||
store._connection = None
|
||||
|
||||
store._get_connection = _get_connection
|
||||
store.close = _close
|
||||
return store, handed, closed
|
||||
|
||||
@staticmethod
|
||||
def _conn(execute_error=None):
|
||||
cursor = MagicMock()
|
||||
cursor.fetchone.return_value = [str(uuid.uuid4())]
|
||||
cursor.fetchall.return_value = []
|
||||
if execute_error is not None:
|
||||
cursor.execute.side_effect = execute_error
|
||||
conn = MagicMock()
|
||||
conn.cursor.return_value = cursor
|
||||
return conn
|
||||
|
||||
def test_mark_chunk_retries_on_a_dropped_connection(self):
|
||||
import psycopg
|
||||
|
||||
dead = self._conn(psycopg.OperationalError("the connection is lost"))
|
||||
alive = self._conn()
|
||||
store, handed, closed = self._store_with_connections([dead, alive])
|
||||
|
||||
store.mark_chunk(str(uuid.uuid4()), "c1", "done")
|
||||
|
||||
assert handed == [dead, alive]
|
||||
assert closed == [dead]
|
||||
alive.commit.assert_called_once()
|
||||
|
||||
def test_apply_chunk_retries_on_a_dropped_connection(self):
|
||||
import psycopg
|
||||
|
||||
dead = self._conn(psycopg.OperationalError("the connection is lost"))
|
||||
alive = self._conn()
|
||||
store, handed, closed = self._store_with_connections([dead, alive])
|
||||
entities = [
|
||||
{
|
||||
"name": "Ada",
|
||||
"normalized_name": "ada",
|
||||
"type": "person",
|
||||
"description": "d",
|
||||
}
|
||||
]
|
||||
|
||||
nodes, edges = store.apply_chunk(
|
||||
str(uuid.uuid4()), "c1", entities, [], {"ada": _embedding(0.5)}
|
||||
)
|
||||
|
||||
assert (nodes, edges) == (1, 0)
|
||||
assert handed == [dead, alive]
|
||||
assert closed == [dead]
|
||||
alive.commit.assert_called_once()
|
||||
|
||||
def test_a_second_connection_failure_is_not_retried_again(self):
|
||||
"""One retry, not a loop: a genuinely unreachable DB still fails."""
|
||||
import psycopg
|
||||
|
||||
dead = self._conn(psycopg.OperationalError("the connection is lost"))
|
||||
also_dead = self._conn(psycopg.OperationalError("the connection is lost"))
|
||||
store, handed, _ = self._store_with_connections([dead, also_dead])
|
||||
|
||||
with pytest.raises(psycopg.OperationalError):
|
||||
store.mark_chunk(str(uuid.uuid4()), "c1", "done")
|
||||
|
||||
assert handed == [dead, also_dead]
|
||||
|
||||
def test_a_query_error_is_not_retried(self):
|
||||
"""Only connection loss is retryable; a bad statement must surface."""
|
||||
import psycopg
|
||||
|
||||
broken = self._conn(psycopg.ProgrammingError("syntax error"))
|
||||
spare = self._conn()
|
||||
store, handed, _ = self._store_with_connections([broken, spare])
|
||||
|
||||
with pytest.raises(psycopg.ProgrammingError):
|
||||
store.mark_chunk(str(uuid.uuid4()), "c1", "done")
|
||||
|
||||
assert handed == [broken]
|
||||
broken.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCountNodesFailureModes:
|
||||
"""Retrieval wants a swallowed count; extraction wants to hear about it."""
|
||||
|
||||
def _store_with_failing_cursor(self):
|
||||
store = GraphStore.__new__(GraphStore)
|
||||
store._tables_ensured = True
|
||||
cursor = MagicMock()
|
||||
cursor.execute.side_effect = RuntimeError("relation does not exist")
|
||||
conn = MagicMock()
|
||||
conn.cursor.return_value = cursor
|
||||
store._connection = conn
|
||||
store._get_connection = lambda: conn
|
||||
return store
|
||||
|
||||
def test_default_reports_zero_to_drive_the_classic_fallback(self):
|
||||
store = self._store_with_failing_cursor()
|
||||
|
||||
assert store.count_nodes(str(uuid.uuid4())) == 0
|
||||
|
||||
def test_strict_surfaces_the_query_failure(self):
|
||||
"""A caller reporting graph size must not read a broken query as empty."""
|
||||
store = self._store_with_failing_cursor()
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
store.count_nodes(str(uuid.uuid4()), strict=True)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestApplyChunkIsReplaySafe:
|
||||
"""A retry after an ambiguous commit must not apply a chunk twice.
|
||||
|
||||
``_write_with_reconnect`` replays the write when the connection dies, and
|
||||
``commit()`` itself can raise connection loss *after* the server committed.
|
||||
Replaying then bumps ``doc_freq`` a second time and inserts a second
|
||||
logical edge (``graph_edges`` has no uniqueness constraint), so the chunk's
|
||||
own progress row is written in the same transaction and short-circuits it.
|
||||
"""
|
||||
|
||||
@pytest.fixture
|
||||
def store(self, postgresql):
|
||||
store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info))
|
||||
try:
|
||||
store._ensure_tables()
|
||||
except Exception as exc:
|
||||
pytest.skip(f"pgvector extension unavailable: {exc}")
|
||||
yield store
|
||||
store.close()
|
||||
|
||||
def test_a_replayed_chunk_is_not_applied_twice(self, store):
|
||||
source_id = str(uuid.uuid4())
|
||||
entities = [
|
||||
{
|
||||
"name": "Ada",
|
||||
"normalized_name": "ada",
|
||||
"type": "person",
|
||||
"description": "d",
|
||||
}
|
||||
]
|
||||
relationships = [
|
||||
{
|
||||
"source": "Ada",
|
||||
"target": "Engine",
|
||||
"type": "worked_on",
|
||||
"description": "x",
|
||||
"weight": 2.0,
|
||||
}
|
||||
]
|
||||
embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)}
|
||||
try:
|
||||
first = store.apply_chunk(
|
||||
source_id, "c1", entities, relationships, embeddings
|
||||
)
|
||||
replay = store.apply_chunk(
|
||||
source_id, "c1", entities, relationships, embeddings
|
||||
)
|
||||
|
||||
assert first == (1, 1)
|
||||
assert replay == (0, 0)
|
||||
node = store.get_node_by_normalized(source_id, "ada")
|
||||
assert node["doc_freq"] == 1
|
||||
overview = store.get_graph_overview(source_id)
|
||||
assert len(overview["edges"]) == 1
|
||||
# The write records its own progress, so the caller's checkpoint
|
||||
# and the rows it describes commit together.
|
||||
assert store.get_progress(source_id)["c1"] == "done"
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_overlapping_applies_of_one_chunk_write_it_once(self, store, postgresql, monkeypatch):
|
||||
"""Two builds of one source can overlap: a rebuild dispatched while the
|
||||
last one runs gets a new lease key. Both may reach the same chunk at
|
||||
once, and the second must wait for the first to commit instead of
|
||||
passing the done check while the first is still in flight."""
|
||||
import threading
|
||||
import time
|
||||
|
||||
source_id = str(uuid.uuid4())
|
||||
entities = [{"name": "Ada", "normalized_name": "ada", "type": "person", "description": "d"}]
|
||||
relationships = [
|
||||
{"source": "Ada", "target": "Engine", "type": "worked_on", "description": "x", "weight": 2.0}
|
||||
]
|
||||
embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)}
|
||||
writers = [GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) for _ in range(2)]
|
||||
real_upsert = GraphStore._upsert_node
|
||||
|
||||
def _slow_upsert(self, *args, **kwargs):
|
||||
time.sleep(0.3) # hold the first writer inside its transaction
|
||||
return real_upsert(self, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(GraphStore, "_upsert_node", _slow_upsert)
|
||||
results = []
|
||||
|
||||
def _apply(writer):
|
||||
results.append(writer.apply_chunk(source_id, "c1", entities, relationships, embeddings))
|
||||
|
||||
try:
|
||||
threads = [threading.Thread(target=_apply, args=(w,)) for w in writers]
|
||||
threads[0].start()
|
||||
time.sleep(0.05)
|
||||
threads[1].start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
assert sorted(results) == [(0, 0), (1, 1)]
|
||||
assert store.get_node_by_normalized(source_id, "ada")["doc_freq"] == 1
|
||||
assert len(store.get_graph_overview(source_id)["edges"]) == 1
|
||||
finally:
|
||||
for writer in writers:
|
||||
writer.close()
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_a_zero_weight_relationship_stays_zero(self, store):
|
||||
source_id = str(uuid.uuid4())
|
||||
relationships = [{"source": "Ada", "target": "Engine", "type": "mentions", "weight": 0.0}]
|
||||
try:
|
||||
store.apply_chunk(source_id, "c1", [], relationships, {})
|
||||
edges = store.get_graph_overview(source_id)["edges"]
|
||||
assert [edge["weight"] for edge in edges] == [0.0]
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
|
||||
def test_a_different_chunk_still_applies(self, store):
|
||||
"""The guard is per chunk, not a blanket 'already saw this source'."""
|
||||
source_id = str(uuid.uuid4())
|
||||
entities = [
|
||||
{
|
||||
"name": "Ada",
|
||||
"normalized_name": "ada",
|
||||
"type": "person",
|
||||
"description": "d",
|
||||
}
|
||||
]
|
||||
embeddings = {"ada": _embedding(0.1)}
|
||||
try:
|
||||
store.apply_chunk(source_id, "c1", entities, [], embeddings)
|
||||
second = store.apply_chunk(source_id, "c2", entities, [], embeddings)
|
||||
|
||||
assert second == (1, 0)
|
||||
node = store.get_node_by_normalized(source_id, "ada")
|
||||
assert node["doc_freq"] == 2
|
||||
finally:
|
||||
store.delete_by_source(source_id)
|
||||
@@ -45,6 +45,29 @@ def _patch_embed(monkeypatch):
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _entity_only_ranking(monkeypatch):
|
||||
"""Pin the ranking path these tests were written for.
|
||||
|
||||
Everything here exercises entity-only PPR ranking without vector blending,
|
||||
driven through ``MagicMock`` stores. The shipped default now walks the
|
||||
passages and blends with vector search — covered end to end in
|
||||
``tests/graphrag/test_retriever_default_path.py`` with a store that returns
|
||||
real values. Pinning keeps each test here asserting what it was written to
|
||||
assert, rather than whatever a mock happens to return on a path it never set
|
||||
up.
|
||||
"""
|
||||
from docsgpt.storage.db.source_config import GraphRetrievalConfig
|
||||
|
||||
monkeypatch.setattr(
|
||||
GraphRAGRetriever,
|
||||
"_graph_options",
|
||||
lambda self, source_id: GraphRetrievalConfig(
|
||||
passage_nodes=False, blend_vector=False
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# ── Fallback to ClassicRAG ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -157,7 +180,9 @@ class TestGraphRAGPoolDiscipline:
|
||||
mock_store_cls.return_value = store
|
||||
|
||||
rag = _make_retriever()
|
||||
with patch.object(rag, "_graph_docs_for_source", return_value=[]):
|
||||
# A real result: an empty one now falls back like a failure does.
|
||||
graph_docs = [{"title": "g", "text": "graph text", "source": "src1", "filename": "g"}]
|
||||
with patch.object(rag, "_graph_docs_for_source", return_value=graph_docs):
|
||||
with patch.object(rag, "_classic_for_sources") as classic:
|
||||
rag._get_data()
|
||||
|
||||
@@ -332,15 +357,24 @@ class TestGraphRAGHappyPath:
|
||||
@patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10)
|
||||
@patch("docsgpt.retriever.graph_rag.GraphStore")
|
||||
@patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True)
|
||||
def test_no_seeds_returns_empty(
|
||||
def test_a_graph_that_answers_nothing_falls_back_to_classic(
|
||||
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed
|
||||
):
|
||||
"""Empty is not an answer. Every graph read swallows its own errors and
|
||||
returns nothing, so "no rows" covers a broken query as much as a walk
|
||||
that found nothing — and the source would contribute nothing at all,
|
||||
with no fallback, because only a raise routes one to ClassicRAG."""
|
||||
store = _store_with_graph([], [], {}, {}, [])
|
||||
store.count_nodes_many.side_effect = lambda ids: {s: 5 for s in ids}
|
||||
mock_store_cls.return_value = store
|
||||
|
||||
rag = _make_retriever()
|
||||
assert rag._get_data() == []
|
||||
seen = _recording_classic(rag, [_CLASSIC_DOC])
|
||||
|
||||
docs = rag._get_data()
|
||||
|
||||
assert seen == [["src1"]]
|
||||
assert [doc["text"] for doc in docs] == ["classic"]
|
||||
|
||||
|
||||
# ── IDF down-weighting ────────────────────────────────────────────────────────
|
||||
@@ -449,11 +483,16 @@ class TestGetChunkTexts:
|
||||
sid = str(uuid.uuid4())
|
||||
store.get_chunk_texts(sid, ["1", "2"])
|
||||
|
||||
sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
|
||||
assert f"FROM {table}" in sql
|
||||
assert text_col in sql
|
||||
assert metadata_col in sql
|
||||
assert f"{source_col} = %s" in sql
|
||||
from psycopg import sql as pgsql
|
||||
|
||||
query, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1]
|
||||
# Identifiers are composed and quoted by psycopg, never formatted in.
|
||||
assert isinstance(query, pgsql.Composable)
|
||||
sql = query.as_string()
|
||||
assert f'FROM "{table}"' in sql
|
||||
assert f'"{text_col}"' in sql
|
||||
assert f'"{metadata_col}"' in sql
|
||||
assert f'"{source_col}" = %s' in sql
|
||||
assert "id::text = ANY(%s)" in sql
|
||||
assert sid not in sql
|
||||
assert params == (sid, ["1", "2"])
|
||||
@@ -800,3 +839,185 @@ class TestGraphRAGBatching:
|
||||
assert rag._get_data() == []
|
||||
mock_store_cls.assert_not_called()
|
||||
rag._classic._get_data.assert_not_called()
|
||||
|
||||
|
||||
# ── Personalized PageRank without scipy ──────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _no_scipy(monkeypatch):
|
||||
"""Make ``import scipy`` fail, as it does in a default install.
|
||||
|
||||
``scipy`` is not a DocsGPT dependency — it only reaches this test env
|
||||
through the optional docling extra. ``networkx.pagerank`` delegates to its
|
||||
scipy implementation, so ranking must not go through it.
|
||||
"""
|
||||
import sys
|
||||
|
||||
for name in [m for m in list(sys.modules) if m == "scipy" or m.startswith("scipy.")]:
|
||||
monkeypatch.delitem(sys.modules, name)
|
||||
monkeypatch.setitem(sys.modules, "scipy", None)
|
||||
|
||||
|
||||
def _chain_graph():
|
||||
"""Weighted chain a-b-c-d plus a heavier shortcut a-d."""
|
||||
import networkx as nx
|
||||
|
||||
graph = nx.Graph()
|
||||
graph.add_weighted_edges_from(
|
||||
[("a", "b", 1.0), ("b", "c", 2.0), ("c", "d", 1.0), ("a", "d", 0.5)]
|
||||
)
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPersonalizedPageRankWithoutScipy:
|
||||
def test_ranking_runs_when_scipy_is_missing(self, _no_scipy):
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
ranks = _personalized_pagerank(
|
||||
graph, personalization={"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
|
||||
)
|
||||
|
||||
assert set(ranks) == {"a", "b", "c", "d"}
|
||||
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
|
||||
assert all(rank > 0 for rank in ranks.values())
|
||||
# Pinned from the parity test below, which runs the library
|
||||
# implementation over the same graph while scipy is installed here.
|
||||
assert sorted(ranks, key=ranks.get, reverse=True) == ["b", "a", "c", "d"]
|
||||
# The seed outranks the node furthest from it along the heavy path.
|
||||
assert ranks["a"] > ranks["d"]
|
||||
|
||||
def test_matches_networkx_within_tolerance(self):
|
||||
"""Parity with the library implementation, while it is installed here."""
|
||||
import networkx as nx
|
||||
|
||||
pytest.importorskip("scipy")
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
personalization = {"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
|
||||
|
||||
ours = _personalized_pagerank(graph, personalization=personalization)
|
||||
theirs = nx.pagerank(graph, personalization=personalization, weight="weight")
|
||||
|
||||
for node in theirs:
|
||||
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
|
||||
|
||||
def test_uniform_personalization_when_none(self):
|
||||
pytest.importorskip("scipy")
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
ours = _personalized_pagerank(graph, personalization=None)
|
||||
theirs = nx.pagerank(graph, personalization=None, weight="weight")
|
||||
|
||||
for node in theirs:
|
||||
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
|
||||
|
||||
def test_isolated_node_still_gets_mass(self):
|
||||
"""A node with no edges is dangling; its mass must not vanish."""
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = nx.Graph()
|
||||
graph.add_edge("a", "b", weight=1.0)
|
||||
graph.add_node("lonely")
|
||||
|
||||
ranks = _personalized_pagerank(graph, personalization=None)
|
||||
|
||||
assert ranks["lonely"] > 0
|
||||
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_empty_graph_returns_empty(self):
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
assert _personalized_pagerank(nx.Graph(), personalization=None) == {}
|
||||
|
||||
def test_a_zero_weight_edge_is_not_traversable(self):
|
||||
"""Zero means "not related", not "use the default weight"."""
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = nx.Graph()
|
||||
graph.add_edge("seed", "zero", weight=0.0)
|
||||
graph.add_edge("seed", "real", weight=1.0)
|
||||
|
||||
ranks = _personalized_pagerank(
|
||||
graph, personalization={"seed": 1.0, "zero": 0.0, "real": 0.0}
|
||||
)
|
||||
|
||||
# ``zero`` is reachable only across the zero-weight edge, so no mass
|
||||
# walks to it; ``real`` is on a live edge and must outrank it.
|
||||
assert ranks["real"] > ranks["zero"]
|
||||
assert ranks["zero"] == pytest.approx(0.0, abs=1e-9)
|
||||
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_stored_zero_weights_reach_the_ranker_intact(self):
|
||||
"""The subgraph builder must not coerce a stored 0 into a real edge.
|
||||
|
||||
Without this the ranker's zero-weight rule is unreachable in
|
||||
production: every 0 from ``graph_edges`` arrives as 1.0.
|
||||
"""
|
||||
subgraph = {
|
||||
"nodes": [
|
||||
{"id": "seed", "doc_freq": 1},
|
||||
{"id": "zero", "doc_freq": 1},
|
||||
{"id": "real", "doc_freq": 1},
|
||||
],
|
||||
"edges": [
|
||||
{"src_node_id": "seed", "dst_node_id": "zero", "weight": 0},
|
||||
{"src_node_id": "seed", "dst_node_id": "real", "weight": 1.0},
|
||||
],
|
||||
}
|
||||
# Called unbound with ``None`` for self: _ppr_scores reads no state.
|
||||
scores = GraphRAGRetriever._ppr_scores(None, subgraph, {"seed": 1.0})
|
||||
|
||||
assert scores["real"] > scores["zero"]
|
||||
assert scores["zero"] == pytest.approx(0.0, abs=1e-9)
|
||||
|
||||
def test_missing_and_null_weights_default_to_one(self):
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
absent = nx.Graph()
|
||||
absent.add_edge("a", "b") # no weight attribute at all
|
||||
null = nx.Graph()
|
||||
null.add_edge("a", "b", weight=None)
|
||||
|
||||
personalization = {"a": 1.0, "b": 0.0}
|
||||
from_absent = _personalized_pagerank(absent, personalization=personalization)
|
||||
from_null = _personalized_pagerank(null, personalization=personalization)
|
||||
|
||||
assert from_absent["b"] == pytest.approx(from_null["b"], abs=1e-9)
|
||||
assert from_absent["b"] > 0
|
||||
|
||||
@patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10)
|
||||
@patch("docsgpt.retriever.graph_rag.GraphStore")
|
||||
@patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True)
|
||||
def test_graph_retrieval_does_not_fall_back_without_scipy(
|
||||
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed, _no_scipy
|
||||
):
|
||||
"""The whole PPR path runs with scipy absent — no ClassicRAG fallback."""
|
||||
nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}]
|
||||
edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}]
|
||||
node_chunks = {"n1": ["c1"], "n2": ["c2"]}
|
||||
chunk_texts = {"c1": "near", "c2": "far"}
|
||||
seed_rows = [{"id": "n1", "distance": 0.0}]
|
||||
store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows)
|
||||
mock_store_cls.return_value = store
|
||||
|
||||
rag = _make_retriever(chunks=2)
|
||||
rag._classic_for_sources = Mock(side_effect=AssertionError("fell back"))
|
||||
|
||||
docs = rag._get_data()
|
||||
|
||||
assert [doc["text"] for doc in docs] == ["near", "far"]
|
||||
@@ -276,3 +276,77 @@ class TestReclaimIsSkippedForEmbeds:
|
||||
from docsgpt.vectorstore.embeddings_delegated import EMBED_TASK
|
||||
|
||||
assert EMBED_TASK in _NO_RECLAIM_TASKS
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestInWorker:
|
||||
"""Whether code runs inside a worker must not depend on which thread asks.
|
||||
|
||||
Celery records the executing task on the thread that runs it, so a thread
|
||||
that task starts sees no task at all. Code deciding "am I in the worker?"
|
||||
from that alone takes the web-process branch there: it dispatches to the
|
||||
worker it is running in and blocks on the result, which Celery refuses
|
||||
("Never call result.get() within a task!") or, where joins are allowed,
|
||||
waits on a queue only this busy process serves.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _ask_from_a_new_thread():
|
||||
import threading
|
||||
|
||||
from docsgpt.celery_init import in_worker
|
||||
|
||||
seen = []
|
||||
thread = threading.Thread(target=lambda: seen.append(in_worker()))
|
||||
thread.start()
|
||||
thread.join()
|
||||
return seen[0]
|
||||
|
||||
def test_false_outside_a_worker(self):
|
||||
from docsgpt.celery_init import in_worker
|
||||
|
||||
assert in_worker() is False
|
||||
assert self._ask_from_a_new_thread() is False
|
||||
|
||||
def test_true_on_a_thread_started_inside_a_worker(self):
|
||||
# Blocking pools (prefork, solo, threads) mark the whole process as one
|
||||
# where joining a task would block; ``denied_join_result`` sets exactly
|
||||
# that flag.
|
||||
from celery.result import denied_join_result
|
||||
|
||||
with denied_join_result():
|
||||
assert self._ask_from_a_new_thread() is True
|
||||
|
||||
def test_true_on_any_thread_of_a_non_blocking_pool_worker(self, monkeypatch):
|
||||
# eventlet/gevent leave the join flag unset and scope the current task
|
||||
# to one greenlet, so only the worker's own startup can say this
|
||||
# process is a worker. The lifecycle signal records that for every
|
||||
# thread and greenlet in it.
|
||||
from celery.signals import worker_init
|
||||
|
||||
monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False)
|
||||
assert self._ask_from_a_new_thread() is False
|
||||
|
||||
worker_init.send(sender=None)
|
||||
|
||||
assert self._ask_from_a_new_thread() is True
|
||||
|
||||
def test_prefork_children_record_it_on_their_own_start(self, monkeypatch):
|
||||
from celery.signals import worker_process_init
|
||||
|
||||
monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False)
|
||||
worker_process_init.send(sender=None)
|
||||
|
||||
assert self._ask_from_a_new_thread() is True
|
||||
|
||||
def test_true_on_the_task_thread_of_a_non_blocking_pool(self):
|
||||
# eventlet/gevent pools leave the process flag unset; the thread
|
||||
# running the task still knows it is in one.
|
||||
from unittest.mock import PropertyMock
|
||||
|
||||
from docsgpt.celery_init import celery, in_worker
|
||||
|
||||
with patch.object(
|
||||
type(celery), "current_worker_task", new_callable=PropertyMock, return_value=object()
|
||||
):
|
||||
assert in_worker() is True
|
||||
@@ -74,6 +74,35 @@ class TestDispatcherGrouping:
|
||||
assert "b" not in retrievals
|
||||
|
||||
|
||||
def test_graph_options_count_as_an_override(self, _patch_llm_creator):
|
||||
"""A graph source that changes only its graph options still needs its
|
||||
config carried over: those options live on the per-source retrieval the
|
||||
Dispatcher hands the retriever, so without this the UI toggles are
|
||||
no-ops and every source runs the defaults."""
|
||||
sources = [
|
||||
{
|
||||
"id": "a",
|
||||
"retrieval": RetrievalConfig(
|
||||
retriever="graphrag", graph={"seed_strategy": "relationships"}
|
||||
),
|
||||
},
|
||||
{"id": "b", "retrieval": RetrievalConfig(retriever="graphrag")},
|
||||
]
|
||||
d = Dispatcher(source={"question": "q", "active_docs": ["a", "b"]}, sources=sources)
|
||||
retrievals = d._groups[0]["retrievals"]
|
||||
assert "a" in retrievals
|
||||
assert retrievals["a"].graph.seed_strategy == "relationships"
|
||||
# A source on the defaults still takes the shared path.
|
||||
assert "b" not in retrievals
|
||||
|
||||
def test_graph_options_on_a_classic_source_are_not_an_override(self, _patch_llm_creator):
|
||||
# They only mean anything to the graph retriever; treating them as an
|
||||
# override would hand a classic source its own chunk budget.
|
||||
sources = [{"id": "a", "retrieval": RetrievalConfig(graph={"blend_vector": False})}]
|
||||
d = Dispatcher(source={"question": "q", "active_docs": ["a"]}, sources=sources)
|
||||
assert d._groups[0]["retrievals"] == {}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestDispatcherSharedBudget:
|
||||
def test_single_group_full_budget(self, _patch_llm_creator):
|
||||
|
||||
@@ -72,6 +72,32 @@ class TestInsideAWorker:
|
||||
assert vector == [1.0, 2.0]
|
||||
celery.send_task.assert_not_called()
|
||||
|
||||
def test_a_thread_started_inside_the_worker_embeds_locally(self):
|
||||
"""The task's own thread is not the only one in a worker.
|
||||
|
||||
Graph extraction and per-source retrieval both fan out to thread pools
|
||||
inside tasks. The check used to read the task off the current thread
|
||||
only, so from those threads it dispatched to the worker it was running
|
||||
in -- and Celery refuses that ``get()`` inside a worker, failing the
|
||||
call and latching the 30s dispatch cooldown for every caller after it.
|
||||
"""
|
||||
from celery.result import denied_join_result
|
||||
|
||||
from docsgpt.celery_init import celery
|
||||
|
||||
local = MagicMock()
|
||||
local.embed_documents.return_value = [[1.0, 2.0]]
|
||||
client = DelegatedEmbeddings("some/model")
|
||||
vectors = []
|
||||
with denied_join_result():
|
||||
with patch("docsgpt.vectorstore.base.build_local_embeddings", return_value=local):
|
||||
with patch.object(celery, "send_task") as send_task:
|
||||
thread = threading.Thread(target=lambda: vectors.append(client.embed_query("hi")))
|
||||
thread.start()
|
||||
thread.join()
|
||||
assert vectors == [[1.0, 2.0]]
|
||||
send_task.assert_not_called()
|
||||
|
||||
def test_the_local_model_is_built_once(self):
|
||||
local = MagicMock()
|
||||
local.embed_documents.return_value = [[1.0]]
|
||||
|
||||
Reference in new issue
Block a user