diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 32181adb..3221867f 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -447,6 +447,12 @@ Type `int`, default `2000`, must be `>= 0`. Hard cap on chunks extracted per source (cost control); 0 extracts nothing. +### `GRAPHRAG_EXTRACTION_WORKERS` + +Type `int`, default `8`, must be `>= 1` and `<= 32`. + +Concurrent extraction calls during ingest. Model calls run in parallel while graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial. + ## Vector stores diff --git a/docsgpt/agents/agentic_agent.py b/docsgpt/agents/agentic_agent.py index b83c485d..0b87485c 100644 --- a/docsgpt/agents/agentic_agent.py +++ b/docsgpt/agents/agentic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -33,6 +34,7 @@ class AgenticAgent(BaseAgent): ) -> Generator[Dict, None, None]: tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py index d731a415..030d76fc 100644 --- a/docsgpt/agents/base.py +++ b/docsgpt/agents/base.py @@ -955,16 +955,30 @@ class BaseAgent(ABC): ) self.retrieved_docs = scrubbed - def _collect_internal_sources(self) -> None: - """Merge the cached InternalSearchTool's docs into ``retrieved_docs``, - deduped, preserving any pre-fetched docs so a mixed-exposure agent cites - both pre-fetched and tool-retrieved sources (not just the tool's).""" + def _search_tool_docs(self) -> List[Dict]: + """Documents this run's search tools read: internal search and the graph tool. + + Both record what they surface in ``retrieved_docs``; a page read from + the graph carries the answer as much as a search hit does, so both are + cited. Tools are looked up the way the executor caches them. + """ + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID executor = getattr(self, "tool_executor", None) loaded = getattr(executor, "_loaded_tools", None) or {} - tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}") - if not (tool and getattr(tool, "retrieved_docs", None)): + docs: List[Dict] = [] + for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)): + tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}") + docs.extend(getattr(tool, "retrieved_docs", None) or []) + return docs + + def _collect_internal_sources(self) -> None: + """Merge the search tools' docs into ``retrieved_docs``, deduped, + preserving any pre-fetched docs so a mixed-exposure agent cites both + pre-fetched and tool-retrieved sources (not just the tools').""" + tool_docs = self._search_tool_docs() + if not tool_docs: return def _key(d): @@ -974,7 +988,7 @@ class BaseAgent(ABC): merged = list(self.retrieved_docs or []) seen = {_key(d) for d in merged} - for doc in tool.retrieved_docs: + for doc in tool_docs: k = _key(doc) if k not in seen: seen.add(k) diff --git a/docsgpt/agents/classic_agent.py b/docsgpt/agents/classic_agent.py index 2bc25130..5956ed3c 100644 --- a/docsgpt/agents/classic_agent.py +++ b/docsgpt/agents/classic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -38,6 +39,7 @@ class ClassicAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() if self.retriever_config: add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py index 1d19a7a9..de96cf35 100644 --- a/docsgpt/agents/research_agent.py +++ b/docsgpt/agents/research_agent.py @@ -6,10 +6,8 @@ from typing import Dict, Generator, List, Optional from docsgpt.agents.base import BaseAgent from docsgpt.agents.tool_executor import ToolExecutor -from docsgpt.agents.tools.internal_search import ( - INTERNAL_TOOL_ID, - add_internal_search_tool, -) +from docsgpt.agents.tools.graph_search import add_graph_search_tool +from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.agents.tools.think import THINK_TOOL_ENTRY, THINK_TOOL_ID from docsgpt.logging import LogContext @@ -277,6 +275,7 @@ class ResearchAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) @@ -620,12 +619,9 @@ class ResearchAgent(BaseAgent): return messages, search_returned_empty def _collect_step_sources(self): - """Collect sources from InternalSearchTool and register with CitationManager.""" - cache_key = f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}" - tool = self.tool_executor._loaded_tools.get(cache_key) - if tool and hasattr(tool, "retrieved_docs"): - for doc in tool.retrieved_docs: - self.citations.add(doc) + """Register the search tools' docs (internal search and graph pages) with CitationManager.""" + for doc in self._search_tool_docs(): + self.citations.add(doc) # ------------------------------------------------------------------ # Phase 3: Synthesis diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py new file mode 100644 index 00000000..b370949e --- /dev/null +++ b/docsgpt/agents/tools/graph_search.py @@ -0,0 +1,299 @@ +"""Let the model search the knowledge graph itself, one edge at a time. + +Graph retrieval normally runs as a ranker: seed a walk from the question, +diffuse mass over a subgraph, hand back the highest-scoring chunks. Measured +across five corpora that never beat plain vector search, because a question +whose answer lives two documents away has nothing in it for the seeding step to +match — the bridging entity is named in the *first* document, not the question. + +Exposing the graph as tools removes the guess. The model can look up the +service, read which store it names, then fetch that store's page: the chain +followed deliberately rather than approximated by a diffusion. On a corpus built +so that vector search cannot shortcut the chain, this took two-hop answers from +1/8 to 8/8, against 0.40 for vector and 0.47 for one-shot graph retrieval. + +It is not a general win, and is deliberately not a default. On ordinary prose +documentation it *lost* to plain vector search (0.50 against 0.90): it answers +well when a question names an entity and wanders when the question is a task +description. It also costs several model round-trips per answer instead of one. +So it is offered only where a source owner has already chosen search over +prefetch — the per-source exposure setting, or an agentic/research agent — and +suits content that is genuinely chain-structured: runbooks, service catalogues, +infrastructure inventories. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, List, Optional + +from docsgpt.agents.tools.base import Tool +from docsgpt.graphrag import graphrag_available +from docsgpt.retriever.labels import labels_from_metadata + +logger = logging.getLogger(__name__) + +GRAPH_TOOL_ID = "graph_search" +MAX_PAGE_CHARS = 1500 + + +class GraphSearchTool(Tool): + """Entity lookup, relationship traversal and page reads over a source's graph.""" + + internal = True + + def __init__(self, config: Dict): + self.config = config or {} + self._store = None + self.retrieved_docs: List[Dict] = [] + + # -- plumbing ------------------------------------------------------------ + def _sources(self) -> List[str]: + source = self.config.get("source") or {} + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + return [str(s) for s in active if s] + + def _get_store(self): + if self._store is None: + from docsgpt.graphrag.store import GraphStore + + self._store = GraphStore() + return self._store + + def _release_store(self) -> None: + """Hand the pooled connection back at the end of an action. + + The executor caches this tool for the whole agent run, so a store kept + between actions pins one connection of the shared pgvector pool across + every LLM round trip of that run -- minutes at a time, and enough + concurrent runs exhaust the pool. ``GraphRAGRetriever`` releases its + store before falling back for the same reason. Checking one back out + costs a pool acquire. + """ + store, self._store = self._store, None + if store is None: + return + try: + store.close() + except Exception as exc: # noqa: BLE001 -- releasing must not fail an action + logger.debug(f"Graph tool could not release its store: {exc}") + + def _embed(self, text: str) -> Optional[List[float]]: + try: + from docsgpt.vectorstore.base import get_embeddings + + return get_embeddings().embed_query(text) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool could not embed the query: {e}") + return None + + # -- actions ------------------------------------------------------------- + def execute_action(self, action_name: str, **kwargs): + # The graph lives in the pgvector store, so the flag alone is not + # enough: under another vector store there is no graph to read. + if not graphrag_available(): + return "The knowledge graph is not enabled for this deployment." + if not self._sources(): + return "No graph-backed sources are configured." + try: + if action_name == "search_entities": + return self._search_entities(**kwargs) + if action_name == "get_relationships": + return self._get_relationships(**kwargs) + if action_name == "read_entity_pages": + return self._read_entity_pages(**kwargs) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True) + return "The graph lookup failed." + finally: + self._release_store() + return f"Unknown action: {action_name}" + + def _search_entities(self, **kwargs) -> str: + query = str(kwargs.get("query") or "").strip() + if not query: + return "Error: 'query' parameter is required." + limit = max(1, min(int(kwargs.get("k") or 8), 25)) + + embedding = self._embed(query) + if embedding is None: + return "Entity search is unavailable." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for row in store.search_nodes_by_embedding(source_id, embedding, k=limit): + similarity = 1.0 - float(row.get("distance") or 0.0) + description = (row.get("description") or "").strip() + suffix = f" — {description[:160]}" if description else "" + lines.append(f"- {row['name']} (match {similarity:.2f}){suffix}") + if not lines: + return f"No entities found for {query!r}." + return "Entities:\n" + "\n".join(lines[:limit]) + + def _get_relationships(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for edge in store.entity_relationships(source_id, entity): + relation = edge.get("type") or "related to" + lines.append(f"- {edge['source']} --{relation}--> {edge['target']}") + if not lines: + return ( + f"No relationships found for {entity!r}. Try search_entities first " + "to get the exact name used in the graph." + ) + return f"Relationships for {entity!r}:\n" + "\n".join(lines) + + def _read_entity_pages(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + parts: List[str] = [] + for source_id in self._sources(): + for page in store.entity_pages(source_id, entity): + text = page.get("text") or "" + # The retrievers' own labelling: a page read here and the same + # chunk retrieved by internal_search are one document, and + # citations key on (source, title). Labelling it differently + # gives that document two citation numbers. + labels = labels_from_metadata(page.get("metadata"), text, source_id) + doc = {**labels, "text": text} + if doc not in self.retrieved_docs: + self.retrieved_docs.append(doc) + header = labels["filename"] or labels["title"] + parts.append(f"--- {header} ---\n{text[:MAX_PAGE_CHARS]}") + if not parts: + return f"No documents mention {entity!r}." + return "\n\n".join(parts) + + # -- metadata ------------------------------------------------------------ + def get_actions_metadata(self): + return [ + { + "name": "search_entities", + "description": ( + "Find named things in the knowledge graph — services, components, " + "settings, people — whose names resemble a query. Use this first to " + "learn the exact name the graph uses before asking for its " + "relationships." + ), + "parameters": { + "properties": { + "query": { + "type": "string", + "description": "What to look for, e.g. a service or component name.", + "filled_by_llm": True, + "required": True, + }, + "k": { + "type": "integer", + "description": "How many entities to return (default 8).", + "filled_by_llm": True, + "required": False, + }, + } + }, + }, + { + "name": "get_relationships", + "description": ( + "List what an entity is connected to, as 'source --relation--> target'. " + "This is how you answer a question about something the question does not " + "name: look up what it points at, then read that thing's pages." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + { + "name": "read_entity_pages", + "description": ( + "Read the documentation an entity appears in, the page it is about " + "first. Use this once you know which entity holds the answer." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + ] + + def get_config_requirements(self): + return {} + + +def build_graph_tool_entry() -> Dict: + """The synthetic ``tools_dict`` entry for the graph tool.""" + tool = GraphSearchTool({}) + actions = [] + for action in tool.get_actions_metadata(): + entry = dict(action) + entry["active"] = True + actions.append(entry) + return {"name": "graph_search", "actions": actions} + + +def sources_have_graph(source: Dict) -> bool: + """Whether any active source actually has a graph to search.""" + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + if not active: + return False + try: + from docsgpt.graphrag.store import GraphStore + + counts = GraphStore().count_nodes_many([str(a) for a in active]) + return any(count > 0 for count in counts.values()) + except Exception as e: # noqa: BLE001 + logger.debug(f"Could not check for graphs: {e}") + return False + + +def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None: + """Add the graph tool when the agent's search-tool sources include a graph. + + No setting of its own: ``retriever_config`` already carries exactly the + sources the agent may *search* — the ones a source owner exposed as a + search tool, or every source for an agentic/research agent — so the graph + tool follows that same per-source exposure choice. A graph source left at + ``prefetch`` in a classic agent is used for ranking only. + """ + if not graphrag_available(): + return + source = retriever_config.get("source") or {} + if not source.get("active_docs") or not sources_have_graph(source): + return + + entry = build_graph_tool_entry() + # The executor resolves tools by ``id``; this one is synthetic (no DB row). + entry["id"] = GRAPH_TOOL_ID + entry["config"] = {"source": source} + tools_dict[GRAPH_TOOL_ID] = entry + + +def build_graph_tool_config(source: Dict, **_ignored: Any) -> Dict: + """Config for :class:`GraphSearchTool` — it only needs the source ids.""" + return {"source": source} diff --git a/docsgpt/agents/tools/read_document.py b/docsgpt/agents/tools/read_document.py index ab9d2294..94e488df 100644 --- a/docsgpt/agents/tools/read_document.py +++ b/docsgpt/agents/tools/read_document.py @@ -17,8 +17,6 @@ import signal import threading from typing import Any, Callable, Dict, List, Optional -from celery import current_task - from docsgpt.agents.tools.artifact_ref import resolve_artifact_id from docsgpt.agents.tools.attachment_bridge import ( AttachmentBridgeError, @@ -26,6 +24,7 @@ from docsgpt.agents.tools.attachment_bridge import ( match_attachment, ) from docsgpt.agents.tools.base import Tool +from docsgpt.celery_init import in_worker from docsgpt.core.json_schema_utils import ( JsonSchemaValidationError, normalize_json_schema_payload, @@ -229,9 +228,9 @@ class ReadDocumentTool(Tool): # (floored at DOCUMENT_PARSE_TIMEOUT). timeout = parse_timeout_for_size(self._input_size) - # ``current_task`` is a Celery proxy: truthy only while this runs inside a worker task, - # falsy in the web process (the bare proxy is NOT identity-None, so test truthiness). - if current_task: + # Process-wide, not the thread-local ``current_task``: a thread a task starts has no + # task of its own, and dispatching from there is the self-deadlock described above. + if in_worker(): from docsgpt.worker import run_parse_document try: diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py index 1381f241..1cfc1b80 100644 --- a/docsgpt/api/user/idempotency.py +++ b/docsgpt/api/user/idempotency.py @@ -9,6 +9,8 @@ import threading import uuid from typing import Any, Callable, Optional +from celery.exceptions import Ignore, MaxRetriesExceededError + from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository from docsgpt.storage.db.session import db_readonly, db_session @@ -81,10 +83,30 @@ def with_idempotency( "idempotency: live lease held; deferring task=%s key=%s", task_name, key, ) - raise self.retry( - countdown=LEASE_TTL_SECONDS, - max_retries=LEASE_RETRY_MAX, - ) + try: + raise self.retry( + countdown=LEASE_TTL_SECONDS, + max_retries=LEASE_RETRY_MAX, + ) + except MaxRetriesExceededError: + # The holder is simply slower than LEASE_RETRY_MAX + # deferrals: a task that outruns the broker's visibility + # timeout is redelivered while its first run is still + # going. Letting the exhaustion propagate would report a + # failure for a task that is running normally — but so + # would returning a value, only less visibly. A redelivery + # reuses the original task id (``Context`` carries + # ``task_id`` into the retry), so a return marks the very + # id the client polls SUCCESS, and ``/api/task_status`` + # hands that to the UI as a finished build. ``Ignore`` + # records no state at all, leaving the outcome to the run + # that actually holds the lease. + logger.info( + "idempotency: lease still held after %s deferrals; " + "leaving task=%s key=%s to its holder", + LEASE_RETRY_MAX, task_name, key, + ) + raise Ignore() from None if attempt > MAX_TASK_ATTEMPTS: logger.error( diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py index 4272eac1..2eda5a3a 100644 --- a/docsgpt/celery_init.py +++ b/docsgpt/celery_init.py @@ -13,6 +13,7 @@ from celery.signals import ( setup_logging, task_postrun, task_prerun, + worker_init, worker_process_init, worker_ready, ) @@ -172,6 +173,51 @@ def _run_version_check(*args, **kwargs): celery = make_celery() celery.config_from_object("docsgpt.celeryconfig") + +#: Set once this process starts as a worker; see :func:`_mark_worker_process`. +_IS_WORKER_PROCESS = False + + +@worker_init.connect +@worker_process_init.connect +def _mark_worker_process(*args, **kwargs): + """Record that this process runs tasks, for :func:`in_worker`. + + ``worker_init`` fires in every worker's main process before its pool + starts: that is where solo, threads, eventlet and gevent run tasks, and + what prefork children fork from. ``worker_process_init`` covers prefork + children however they were started. + """ + global _IS_WORKER_PROCESS + _IS_WORKER_PROCESS = True + + +def in_worker() -> bool: + """True anywhere in a Celery worker process, on any thread or greenlet. + + ``current_worker_task`` alone is not enough: Celery records the executing + task on the thread (or greenlet) that runs it, so one the task starts sees + none and would take the web-process branch — dispatching to the worker it + is running in and blocking on the result. Celery refuses that ``get()`` + ("Never call result.get() within a task!"), or, where joins are allowed, + it waits on a queue that only this busy process may be able to serve. + + The worker's own startup (:func:`_mark_worker_process`) answers for every + pool. ``task_join_will_block`` — process-wide, set for every blocking pool + — and the task's own ``current_worker_task`` still count for a process + that runs tasks without having gone through that startup. + + Returns: + bool: Whether this call is running inside a worker process. + """ + from celery.result import task_join_will_block + + return ( + _IS_WORKER_PROCESS + or task_join_will_block() + or celery.current_worker_task is not None + ) + #: Task-name prefix the package carried before the rename to ``docsgpt``. diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index 651d3f3d..53dbf56b 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -31,6 +31,15 @@ class RetrievalSettings(SettingsGroup): GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing." ) + GRAPHRAG_EXTRACTION_WORKERS: int = Field( + default=8, + ge=1, + le=32, + description=( + "Concurrent extraction calls during ingest. Model calls run in parallel while " + "graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial." + ), + ) @field_validator("VECTOR_STORE", mode="before") @classmethod diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index beaf3d66..daa8f910 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -23,6 +23,11 @@ import logging import re from typing import Any, Callable, Dict, List, Optional +from docsgpt.core.model_utils import ( + get_api_key_for_provider, + get_provider_from_model_id, +) +from docsgpt.graphrag.naming import normalize_entity_name from docsgpt.core.settings import settings from docsgpt.llm.llm_creator import LLMCreator from docsgpt.storage.db.source_config import SourceConfig @@ -63,14 +68,37 @@ def _resolve_max_chunks(config: SourceConfig) -> int: return config.graph.max_chunks or settings.GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION +def _resolve_extraction_provider( + model_id: Optional[str], user: Optional[str] +) -> str: + """The provider that serves ``model_id``, else the deployment default. + + ``settings.LLM_PROVIDER`` is only a default (``docsgpt``, the hosted public + endpoint, out of the box). Dispatching the resolved extraction model + through it sends the request to a provider that does not serve that model: + the call is rejected, the shared fallback answers instead, and the graph is + built by a different model than the one configured — with nothing in the + summary to say so. ``user`` scopes the lookup so a per-user (BYOM) model id + resolves as well. + """ + provider = ( + get_provider_from_model_id(model_id, user_id=user) if model_id else None + ) + return provider or settings.LLM_PROVIDER + + def _build_extraction_llm( model_id: Optional[str], user: Optional[str], request_id: Optional[str] ): """Build the extraction LLM tagged for token-usage attribution to the owner.""" decoded_token = {"sub": user} if user else None + provider = _resolve_extraction_provider(model_id, user) + logger.info( + "Graph extraction dispatching model=%s via provider=%s", model_id, provider + ) llm = LLMCreator.create_llm( - settings.LLM_PROVIDER, - api_key=settings.API_KEY, + provider, + api_key=get_api_key_for_provider(provider), user_api_key=None, decoded_token=decoded_token, model_id=model_id, @@ -122,8 +150,15 @@ def _parse_extraction(raw: Any) -> Optional[Dict[str, List[Dict[str, Any]]]]: } -def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: - """Run exactly one extraction call for a chunk (gleanings off).""" +def _extract_chunk( + llm, text: str, chunk_id: Optional[str] = None +) -> Optional[Dict[str, List[Dict[str, Any]]]]: + """Run exactly one extraction call for a chunk (gleanings off). + + Both failure modes name the chunk: an unparseable response used to return + ``None`` silently, so a graph could come back short with nothing in the + logs to say which chunk was dropped or why. + """ messages = [ {"role": "system", "content": _SYSTEM_PROMPT}, {"role": "user", "content": f"\n{text}\n"}, @@ -134,9 +169,17 @@ def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: messages=messages, ) except Exception as exc: - logger.warning("Graph extraction call failed, skipping chunk: %s", exc) + logger.warning( + "Graph extraction call failed for chunk %s: %s", chunk_id, exc + ) return None - return _parse_extraction(response) + parsed = _parse_extraction(response) + if parsed is None: + logger.warning( + "Graph extraction returned unparseable output for chunk %s.", + chunk_id, + ) + return parsed def _coerce_weight(value: Any) -> float: @@ -160,8 +203,10 @@ def extract_graph_for_source( Resumable and idempotent: chunks already marked ``done`` are skipped via the ``graph_ingest_progress`` checkpoint, so a retry never re-extracts (and never re-bills). Processes at most the resolved chunk cap; excess chunks are - reported under ``skipped_over_cap``. A malformed response or an LLM error on - a single chunk marks it ``failed`` and continues — the pipeline never crashes. + reported under ``skipped_over_cap``. A malformed response, an LLM error or a + failed write on a single chunk is retried once after the rest of the build; + a chunk that fails again is marked ``failed`` and the run continues — the + pipeline never crashes. Each chunk is written in a single transaction with one batched embedding call (entity + relationship-endpoint names together). @@ -178,8 +223,13 @@ def extract_graph_for_source( Returns: A summary ``{nodes, edges, chunks_processed, skipped_over_cap, - failed_chunks}``. + failed_chunks}``, where ``nodes`` is how many distinct nodes the + source's graph holds after the run — not how many upserts ran, which + counts the same entity once per chunk it appears in. """ + import threading + from concurrent.futures import ThreadPoolExecutor + from docsgpt.graphrag.store import GraphStore store = GraphStore() @@ -197,11 +247,26 @@ def extract_graph_for_source( embedding = get_embeddings() - llm = _build_extraction_llm( - _resolve_extraction_model(config), user, request_id - ) + model_id = _resolve_extraction_model(config) + # Built here first so a misconfigured model fails the run before any + # chunk is touched; this instance serves the calling thread. + thread_llm = threading.local() + thread_llm.llm = _build_extraction_llm(model_id, user, request_id) - nodes = 0 + def _llm(): + """This thread's extraction LLM. + + Provider-reported usage is kept on the LLM instance (``_last_usage``) + and claimed by whichever call finishes next, so two calls in flight on + one instance can bill each other's tokens. Each pool thread therefore + builds its own. + """ + llm = getattr(thread_llm, "llm", None) + if llm is None: + llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id) + return llm + + node_upserts = 0 edges = 0 chunks_processed = 0 failed_chunks = 0 @@ -215,54 +280,140 @@ def extract_graph_for_source( { "current": chunks_processed + failed_chunks, "total": total, - "nodes": nodes, + "nodes": node_upserts, "edges": edges, } ) except Exception as exc: logger.debug("graph progress callback failed: %s", exc) - for chunk, chunk_id in to_process: + def _prepare(item): + """One chunk's LLM extraction — the only step run concurrently. + + A chunk spends almost all of its time waiting on the model, so that is + what runs in the pool. Graph writes and embedding stay on the calling + thread, so transactions and the progress checkpoint are exactly what + they were serially and the pool never touches the embeddings client. + """ + chunk, chunk_id = item text = _chunk_text(chunk) if not text: - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 - _report() - continue + return chunk_id, "empty", None - extracted = _extract_chunk(llm, text) + extracted = _extract_chunk(_llm(), text, chunk_id) if extracted is None: - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() - continue - + return chunk_id, "failed", None try: entities = _build_entities(extracted["entities"]) relationships = _build_relationships(extracted["relationships"]) + except Exception as exc: + logger.warning( + "Graph extraction failed for chunk %s: %s", chunk_id, exc + ) + return chunk_id, "failed", None + return chunk_id, "ok", (entities, relationships) + + def _write(chunk_id, status, payload) -> bool: + """Apply one prepared chunk to the graph; False when it did not land.""" + nonlocal node_upserts, edges, chunks_processed + if status == "empty": + store.mark_chunk(source_id, chunk_id, "done") + chunks_processed += 1 + return True + if status == "failed": + return False + + entities, relationships = payload + try: name_embeddings = _embed_names(embedding, entities, relationships) + _embed_facts(embedding, relationships) chunk_nodes, chunk_edges = store.apply_chunk( source_id, chunk_id, entities, relationships, name_embeddings ) - nodes += chunk_nodes - edges += chunk_edges - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 except Exception as exc: logger.warning( - "Graph extraction write failed for chunk %s, skipping: %s", - chunk_id, - exc, + "Graph extraction embed/write failed for chunk %s: %s", chunk_id, exc ) + return False + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. + node_upserts += chunk_nodes + edges += chunk_edges + chunks_processed += 1 + return True + + workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1)) + pool = None + if workers > 1 and len(to_process) > 1: + pool = ThreadPoolExecutor(max_workers=workers) + + def _pass(items): + """Extract and write ``items``; return the ones that did not land.""" + if pool is not None: + # ``map`` yields in submission order, so chunks are still applied in + # the order they were given and a run stays reproducible. + prepared = pool.map(_prepare, items) + else: + prepared = (_prepare(item) for item in items) + missed = [] + for item, (chunk_id, status, payload) in zip(items, prepared): + if not _write(chunk_id, status, payload): + missed.append(item) + _report() + return missed + + try: + missed = _pass(to_process) + if missed: + # A failure is usually transient — a provider error, one response + # that did not parse — and the checkpoint only picks it up on a + # rerun nothing schedules. One more attempt, after the rest of the + # build so a burst of rate limiting has passed, and no more: a + # chunk that cannot be extracted costs at most two calls. + logger.info( + "Graph extraction retrying %d failed chunk(s) for source %s", + len(missed), + source_id, + ) + missed = _pass(missed) + for _, chunk_id in missed: store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() + failed_chunks = len(missed) + if missed: + logger.warning( + "Graph extraction gave up on %d chunk(s) for source %s after a retry: %s", + failed_chunks, + source_id, + ", ".join(str(chunk_id) for _, chunk_id in missed), + ) + _report() + finally: + if pool is not None: + pool.shutdown(wait=True) try: store.set_node_degrees(source_id) except Exception as exc: logger.warning("set_node_degrees failed for source %s: %s", source_id, exc) + # Upserts are writes, not nodes: one entity seen in ten chunks is ten + # upserts and a single node, so the old count overstated every graph whose + # entities recur. Report what the graph holds, falling back to the write + # count only if the count query itself fails. + # ``strict`` is what makes the fallback below reachable: the default + # count swallows query failures and answers 0, which would report a + # successful build as an empty graph. + nodes = node_upserts + try: + nodes = store.count_nodes(source_id, strict=True) + except Exception as exc: + logger.warning( + "count_nodes failed for source %s; reporting upserts instead: %s", + source_id, + exc, + ) + return { "nodes": nodes, "edges": edges, @@ -273,7 +424,7 @@ def extract_graph_for_source( def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: - """Normalize the LLM's entity dicts (drop nameless ones).""" + """Normalize the LLM's entity dicts (drop the ones with no usable name).""" entities = [] for e in raw_entities: if not isinstance(e, dict): @@ -281,10 +432,16 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: name = str(e.get("name", "")).strip() if not name: continue + normalized_name = normalize_entity_name(name) + if not normalized_name: + # A punctuation-only name normalizes to nothing, and nodes merge on + # that key: keeping it collapses every such entity onto one shared + # node. The relationship side already drops them. + continue entities.append( { "name": name, - "normalized_name": name.lower(), + "normalized_name": normalized_name, "type": str(e.get("type") or "") or None, "description": str(e.get("description") or "") or None, } @@ -310,6 +467,70 @@ def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]: return relationships +def _fact_text(rel: Dict[str, Any]) -> str: + """A relationship rendered as the sentence it asserts. + + Embedded and stored on the edge so retrieval can match a question against + the *relation* rather than against entity names — the difference between + "which entity is this about" and "which fact answers this". + """ + source = str(rel.get("source") or "").strip() + target = str(rel.get("target") or "").strip() + if not source or not target: + return "" + relation = str(rel.get("type") or "related to").strip() or "related to" + text = f"{source} {relation} {target}" + description = str(rel.get("description") or "").strip() + return f"{text}: {description}" if description else text + + +def _embed_facts(embedding, relationships: List[Dict[str, Any]]) -> None: + """Attach a fact embedding to each relationship, in one batched call. + + Mutates the relationship dicts so the embedding travels with the edge into + ``apply_chunk`` without a second mapping to keep in step. Always on: it is + one extra batched call per chunk against an LLM call that already costs + far more, and it lets a source switch to relationship seeding at query time + without being rebuilt. + """ + pending = [(rel, _fact_text(rel)) for rel in relationships] + pending = [(rel, text) for rel, text in pending if text] + if not pending: + return + try: + vectors = embedding.embed_documents([text for _rel, text in pending]) + except Exception as exc: # noqa: BLE001 + # The graph is still correct without them; only fact seeding degrades. + logger.warning("Fact embedding failed, continuing without: %s", exc) + return + for (rel, _text), vector in zip(pending, vectors): + rel["fact_embedding"] = vector + + +def _seed_text(entity: Dict[str, Any]) -> str: + """The text a node's embedding is computed from. + + Retrieval seeds the graph walk by matching a whole question against these + embeddings, and a bare entity name is a poor thing to match a question + against — a question about what a service writes to shares almost no + surface with the name ``Quill``. Including the type and description gives + the match something to work with; measured across five corpora it moved + recall@4 by +0.07 to +0.50. + + Relationship endpoints keep their bare names: they arrive as strings with + no type or description attached. + """ + name = str(entity.get("name") or "").strip() + text = name + entity_type = str(entity.get("type") or "").strip() + if entity_type: + text += f" ({entity_type})" + description = str(entity.get("description") or "").strip() + if description: + text += f": {description}" + return text or name + + def _embed_names( embedding, entities: List[Dict[str, Any]], @@ -322,14 +543,16 @@ def _embed_names( """ name_by_norm: Dict[str, str] = {} for entity in entities: - name_by_norm.setdefault(entity["normalized_name"], entity["name"]) + name_by_norm.setdefault(entity["normalized_name"], _seed_text(entity)) for rel in relationships: for endpoint in (rel.get("source"), rel.get("target")): if endpoint is None: continue clean = str(endpoint).strip() if clean: - name_by_norm.setdefault(clean.lower(), clean) + # Same key the store resolves endpoints by, or the embedding + # computed here never reaches the node it was computed for. + name_by_norm.setdefault(normalize_entity_name(clean), clean) if not name_by_norm: return {} diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py new file mode 100644 index 00000000..2a173eea --- /dev/null +++ b/docsgpt/graphrag/naming.py @@ -0,0 +1,115 @@ +"""Canonical entity naming for the per-source knowledge graph. + +Nodes are merged on ``normalized_name``, which has been ``name.lower()``. That +splits entities a reader would call the same thing: measured on the DocsGPT docs +corpus, ``agent``/``agents``, ``VECTOR_STORE``/``Vector store``/``vector stores``, +``Celery worker``/``Celery workers`` and ``.env file``/``env_file`` all landed as +separate nodes — 58 such collisions across 1,704 entities, with 75% of entities +appearing in exactly one chunk as a result. + +:func:`canonical_name` folds the differences that are purely orthographic: +case, surrounding punctuation, underscore/hyphen word breaks, and a *cautious* +plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``, +``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s" +would corrupt them into new entities rather than merge anything. + +The result is a merge key, never shown to anyone, so it only has to be the +same for a word's singular and plural — not to be a word itself. + +Always on: every graph is built with canonical names. +""" + +from __future__ import annotations + +import re + +_PUNCT = re.compile(r"[^\w\s]+", re.UNICODE) +_UNDERSCORE = re.compile(r"[_\-]+") +_SPACE = re.compile(r"\s+") + +#: Words that end in "s" without being plural. Singularising these would invent +#: entities ("postgre", "kubernete") instead of merging existing ones. +_NOT_PLURAL = frozenset( + { + "postgres", "kubernetes", "https", "aws", "dns", "tls", "cors", "css", + "js", "sas", "gas", "ss", "class", "access", "process", "status", + "analysis", "basis", "axis", "https", "rss", "less", "express", + "redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos", + "always", "sometimes", "series", "docs", "ops", "devops", "sse", + "alias", "canvas", "atlas", "bias", "pandas", + } +) + +#: Plural endings that drop ``es``, and the singular endings that meet them. +#: ``caches`` cannot say whether it is ``cache`` + "s" or ``cach`` + "es" +#: (as ``batches`` is ``batch`` + "es"), so rather than guess, both +#: ``caches`` and ``cache`` fold to ``cach`` — as ``databases``/``database`` +#: fold to ``databas``. Hardly any real word differs from one of these singulars +#: by its final "e" alone, so the fold merges next to nothing it should not. +_ES_PLURAL = ("ches", "shes", "ses", "zes", "xes") +_E_SINGULAR = ("che", "she", "se", "ze", "xe") + + +def _singular(word: str) -> str: + """Fold one word so its singular and plural share a key, else leave it alone. + + ``-ies`` and a singular's ``-ie`` both fold to ``-y`` (``policies``, + ``cookies``/``cookie``). The ``-es`` endings in :data:`_ES_PLURAL` drop + ``es`` and the singular endings in :data:`_E_SINGULAR` drop their ``e``, so + both sides of an ambiguous plural meet (``caches``/``cache`` -> ``cach``). + Otherwise a bare trailing ``s`` is dropped on a word long enough to be + safe. Everything in :data:`_NOT_PLURAL`, and anything ending in + ``ss``/``us``/``is``, is returned unchanged. + """ + if len(word) < 4 or word in _NOT_PLURAL: + return word + if word.endswith(("ss", "us", "is")): + return word + if word.endswith("ies") and len(word) > 4: + return word[:-3] + "y" + if word.endswith(_ES_PLURAL): + return word[:-2] + if word.endswith("ie") and len(word) > 4: + return word[:-2] + "y" + if word.endswith(_E_SINGULAR): + return word[:-1] + if word.endswith("s"): + return word[:-1] + return word + + +def canonical_name(name: str) -> str: + """Merge key for an entity name. + + Args: + name: The entity name as the model wrote it. + + Returns: + A lowercase, punctuation-free key shared by a name's singular and + plural. Returns ``""`` for an empty or punctuation-only name, which + callers treat as "no entity". + + Examples: + ``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``; + ``.env file`` and ``env_file`` -> ``env file``; + ``cache`` and ``caches`` -> ``cach``; + ``postgres`` stays ``postgres``. + """ + if not name: + return "" + text = _UNDERSCORE.sub(" ", str(name)) + text = _PUNCT.sub(" ", text) + text = _SPACE.sub(" ", text).strip().lower() + if not text: + return "" + return " ".join(_singular(word) for word in text.split()) + + +def normalize_entity_name(name: str) -> str: + """The key an entity is merged on: its :func:`canonical_name`. + + Every graph the corpora were measured on was built this way, so it is the + only mode rather than a flag. A graph built before this used plain + ``lower()`` keys; re-extracting it merges onto these instead. + """ + return canonical_name(name) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 818004c0..0b73d031 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -17,6 +17,8 @@ import logging import uuid from typing import Any, Dict, List, Optional +import psycopg +from psycopg import sql from psycopg.types.json import Jsonb from docsgpt.core.settings import settings @@ -51,6 +53,18 @@ def _safe_identifier(name: str) -> str: return name +def _identifier(name: str) -> sql.Identifier: + """``name`` as a quoted identifier, folded the way Postgres folds it unquoted. + + Composing identifiers through psycopg keeps every query a fixed statement + with bound values: nothing is formatted into the SQL string. The fold + matters because ``PGVectorStore`` writes these names unquoted, which + Postgres lower-cases, while a quoted identifier keeps its case; folding + first keeps both stores addressing the same table. + """ + return sql.Identifier(_safe_identifier(name).lower()) + + def _pgvector_identifiers() -> tuple[str, str, str, str]: """Resolve ``(table, text_col, metadata_col, source_col)`` from ``PGVectorStore``. @@ -73,6 +87,51 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]: ) +def _pgvector_vector_column() -> str: + """Resolve the embedding column name from the same ``PGVectorStore`` defaults.""" + import inspect + + from docsgpt.vectorstore.pgvector import PGVectorStore + + params = inspect.signature(PGVectorStore.__init__).parameters + return _safe_identifier(params["vector_column"].default) + + +def _is_connection_lost(exc: BaseException) -> bool: + """True when ``exc`` says the server connection went away, not that the SQL was bad. + + psycopg raises ``OperationalError`` ("the connection is lost") when the + socket dies under a statement and ``InterfaceError`` when the connection + object is already closed. Everything else — a bad statement, a constraint + violation — is a real failure that a retry would only repeat. + """ + return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError)) + + +def _lock_source(cursor, source_id: str) -> None: + """Serialize graph writes for one source until this transaction ends. + + Writes within one build are already serial, but two builds of the same + source can overlap: a rebuild dispatched while the last one is still + running gets a new idempotency key, so its lease does not stop it. Without + this, both could pass a chunk's "done" check before either commits and + apply it twice. A transaction-scoped advisory lock keyed by the source + makes them take turns chunk by chunk; the lock is released on commit or + rollback, and a hash collision only makes two sources take turns. + """ + cursor.execute( + "SELECT pg_advisory_xact_lock(hashtext(%s));", (f"graphrag:source:{source_id}",) + ) + + +def _safe_rollback(conn) -> None: + """Roll back, tolerating a connection too broken to roll back.""" + try: + conn.rollback() + except Exception as exc: + logging.debug("Rollback on a broken connection failed: %s", exc) + + class GraphStore: """Stores and queries a per-source knowledge graph in the pgvector DB.""" @@ -138,6 +197,42 @@ class GraphStore: self._pooled = False return self._connection + def _write_with_reconnect(self, operation): + """Run ``operation(conn)``, once more on a fresh connection if it was dead. + + A graph build holds one checked-out connection for the length of the + whole extraction and spends minutes per chunk waiting on the model, so + the connection idles long enough for the server (or a pooler) to drop + it. The pool only validates a connection when it is handed out, and + this one was handed out at the start of the build, so the next write + raises and its chunk is lost from the graph. Every statement here is an + idempotent upsert, so replaying one on a new connection cannot + double-write. + + Args: + operation: Callable taking the connection and doing one write. + + Returns: + Whatever ``operation`` returns. + + Raises: + Exception: Anything ``operation`` raises that is not connection + loss, and anything the single retry raises. + """ + try: + return operation(self._get_connection()) + except Exception as exc: + if not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + # Second and final attempt, on a connection freshly checked out by + # ``_get_connection``. A failure here belongs to the caller. + return operation(self._get_connection()) + def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. @@ -197,7 +292,7 @@ class GraphStore: ) cursor.execute( - """ + f""" CREATE TABLE IF NOT EXISTS graph_edges ( id UUID PRIMARY KEY, source_id UUID NOT NULL, @@ -206,10 +301,18 @@ class GraphStore: type TEXT, description TEXT, weight REAL DEFAULT 1.0, - source_chunk_ids JSONB + source_chunk_ids JSONB, + fact_embedding vector({dimension}) ); """ ) + # ``CREATE TABLE IF NOT EXISTS`` is a no-op on a database that + # already has the table, so a column added after the fact needs its + # own statement or every existing deployment silently lacks it. + cursor.execute( + f"ALTER TABLE graph_edges " + f"ADD COLUMN IF NOT EXISTS fact_embedding vector({dimension});" + ) cursor.execute( """ @@ -397,19 +500,78 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge on an open cursor (no commit, no degree bump). + fact_embedding: Optional[List[float]] = None, + ) -> tuple[Optional[str], bool]: + """Write an edge on an open cursor (no commit, no degree bump). + + Returns ``(edge_id, created)``. Two shapes of noise are rejected here + rather than at read time, because once written neither is visible: + + * A self-loop feeds a node's PageRank mass straight back to itself. It + is dropped, reported as ``(None, False)``. + * A pair already related by the same type is *merged* rather than + inserted again. ``graph_edges`` carries no uniqueness constraint, so + re-extracting one relationship across many chunks otherwise writes a + row per chunk — a fifth of a real corpus's edges — inflating that + pair's traversal weight and spending the bounded subgraph fetch on + duplicates. The surviving row keeps the strongest weight seen and + every contributing chunk id. Callers that batch many edges run ``set_node_degrees`` once afterwards instead of bumping degree per edge. """ + if str(src_node_id) == str(dst_node_id): + return None, False + + cursor.execute( + """ + SELECT id + FROM graph_edges + WHERE source_id = %s AND src_node_id = %s AND dst_node_id = %s + AND type IS NOT DISTINCT FROM %s + LIMIT 1; + """, + (source_id, src_node_id, dst_node_id, type), + ) + existing = cursor.fetchone() + if existing: + edge_id = existing[0] + # The chunk ids are merged in SQL, against the row's own current + # value, rather than read here and written back: a read-modify-write + # would drop whatever a concurrent writer appended in between. + cursor.execute( + """ + UPDATE graph_edges + SET weight = GREATEST(COALESCE(weight, 0), %s), + description = COALESCE(description, %s), + -- Backfills the fact embedding for an edge first written + -- before fact embeddings were switched on. + fact_embedding = COALESCE(fact_embedding, %s::vector), + source_chunk_ids = COALESCE(source_chunk_ids, '[]'::jsonb) || ( + SELECT COALESCE(jsonb_agg(candidate), '[]'::jsonb) + FROM jsonb_array_elements(%s::jsonb) AS candidate + WHERE NOT COALESCE(source_chunk_ids, '[]'::jsonb) + @> jsonb_build_array(candidate) + ) + WHERE id = %s; + """, + ( + weight, + description, + fact_embedding, + Jsonb(list(source_chunk_ids or [])), + edge_id, + ), + ) + return str(edge_id), False + edge_id = str(uuid.uuid4()) cursor.execute( """ INSERT INTO graph_edges (id, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids) - VALUES (%s, %s, %s, %s, %s, %s, %s, %s); + weight, source_chunk_ids, fact_embedding) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s); """, ( edge_id, @@ -420,9 +582,10 @@ class GraphStore: description, weight, Jsonb(source_chunk_ids or []), + fact_embedding, ), ) - return edge_id + return edge_id, True def add_edge( self, @@ -433,21 +596,28 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge and bump the degree of both endpoints. Returns its id.""" + fact_embedding: Optional[List[float]] = None, + ) -> Optional[str]: + """Write an edge and bump the degree of both endpoints. Returns its id. + + Returns ``None`` for a self-loop, which is not written. A repeat of an + existing pair merges into that row and returns its id, leaving degree + alone — the endpoints gained no new neighbour. + """ self._ensure_tables_once() conn = self._get_connection() cursor = conn.cursor() try: - edge_id = self._add_edge( + edge_id, created = self._add_edge( cursor, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids, - ) - cursor.execute( - "UPDATE graph_nodes SET degree = degree + 1 " - "WHERE source_id = %s AND id IN (%s, %s);", - (source_id, src_node_id, dst_node_id), + weight, source_chunk_ids, fact_embedding, ) + if created: + cursor.execute( + "UPDATE graph_nodes SET degree = degree + 1 " + "WHERE source_id = %s AND id IN (%s, %s);", + (source_id, src_node_id, dst_node_id), + ) conn.commit() return edge_id except Exception as e: @@ -499,56 +669,96 @@ class GraphStore: (not linked to the chunk), mirroring the per-call path. ``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees are not bumped here — the caller runs ``set_node_degrees`` once at the - end. Returns ``(nodes_upserted, edges_added)``. + end. Reconnects and retries once if the connection died while the + extraction was waiting on the model. + + The chunk's ``graph_ingest_progress`` row is written in this same + transaction, so the checkpoint and the rows it describes commit + together and a replay of an already-applied chunk returns ``(0, 0)`` + without touching the graph. Returns ``(nodes_upserted, edges_added)``. """ self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - node_ids: Dict[str, str] = {} - edges_added = 0 - try: - for entity in entities: - normalized_name = entity["normalized_name"] - node_id = self._upsert_node( - cursor, - source_id, - entity["name"], - normalized_name, - entity.get("type"), - entity.get("description"), - name_embeddings.get(normalized_name), - ) - node_ids[normalized_name] = node_id - self._link_node_chunk(cursor, source_id, node_id, chunk_id) - for rel in relationships: - src_id = self._resolve_endpoint( - cursor, source_id, rel.get("source"), node_ids, name_embeddings + def _write(conn): + cursor = conn.cursor() + node_ids: Dict[str, str] = {} + edges_added = 0 + try: + # ``commit()`` can report connection loss *after* the server + # committed, and the retry then replays this write: doc_freq + # would be bumped twice and a second logical edge inserted + # (graph_edges has no uniqueness constraint). The progress row + # below is written in this transaction, so a replay sees it — + # and so does an overlapping build, once the source lock makes + # it wait for this one to commit. + _lock_source(cursor, source_id) + cursor.execute( + "SELECT status FROM graph_ingest_progress " + "WHERE source_id = %s AND chunk_id = %s;", + (source_id, str(chunk_id)), ) - dst_id = self._resolve_endpoint( - cursor, source_id, rel.get("target"), node_ids, name_embeddings - ) - if src_id is None or dst_id is None: - continue - self._add_edge( - cursor, - source_id, - src_id, - dst_id, - type=rel.get("type"), - description=rel.get("description"), - weight=float(rel.get("weight") or 1.0), - source_chunk_ids=[chunk_id], - ) - edges_added += 1 + applied = cursor.fetchone() + if applied is not None and applied[0] == "done": + conn.rollback() + return 0, 0 - conn.commit() - return len(entities), edges_added - except Exception: - conn.rollback() - raise - finally: - cursor.close() + for entity in entities: + normalized_name = entity["normalized_name"] + node_id = self._upsert_node( + cursor, + source_id, + entity["name"], + normalized_name, + entity.get("type"), + entity.get("description"), + name_embeddings.get(normalized_name), + ) + node_ids[normalized_name] = node_id + self._link_node_chunk(cursor, source_id, node_id, chunk_id) + + for rel in relationships: + src_id = self._resolve_endpoint( + cursor, source_id, rel.get("source"), node_ids, name_embeddings + ) + dst_id = self._resolve_endpoint( + cursor, source_id, rel.get("target"), node_ids, name_embeddings + ) + if src_id is None or dst_id is None: + continue + _, created = self._add_edge( + cursor, + source_id, + src_id, + dst_id, + type=rel.get("type"), + description=rel.get("description"), + # Only a missing weight defaults: 0 is a real one, + # and the ranker drops non-positive edges. + weight=1.0 if rel.get("weight") is None else float(rel["weight"]), + source_chunk_ids=[chunk_id], + fact_embedding=rel.get("fact_embedding"), + ) + if created: + edges_added += 1 + + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, 'done') + ON CONFLICT (source_id, chunk_id) + DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id)), + ) + conn.commit() + return len(entities), edges_added + except Exception: + _safe_rollback(conn) + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def _resolve_endpoint( self, @@ -564,7 +774,11 @@ class GraphStore: clean = str(name).strip() if not clean: return None - normalized_name = clean.lower() + from docsgpt.graphrag.naming import normalize_entity_name + + normalized_name = normalize_entity_name(clean) + if not normalized_name: + return None if normalized_name in node_ids: return node_ids[normalized_name] node_id = self._upsert_node( @@ -610,8 +824,22 @@ class GraphStore: cursor.close() conn.rollback() - def count_nodes(self, source_id: str) -> int: - """Number of nodes for a source. Zero drives the ClassicRAG fallback.""" + def count_nodes(self, source_id: str, strict: bool = False) -> int: + """Number of nodes for a source. Zero drives the ClassicRAG fallback. + + Args: + source_id: Source whose nodes to count. + strict: Re-raise a query failure instead of reporting ``0``. + Retrieval wants the swallow — a broken count there just routes + the source to ClassicRAG — but a caller reporting how big a + graph is must not read a failed query as "the graph is empty". + + Returns: + int: The node count, or ``0`` when a query failure is swallowed. + + Raises: + Exception: The underlying query failure, when ``strict`` is set. + """ conn = self._get_connection() cursor = conn.cursor() try: @@ -622,6 +850,8 @@ class GraphStore: return int(cursor.fetchone()[0]) except Exception as e: logging.error(f"Error counting nodes: {e}") + if strict: + raise return 0 finally: cursor.close() @@ -734,6 +964,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND (src_node_id = ANY(%s) OR dst_node_id = ANY(%s)) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, ( @@ -781,6 +1012,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND src_node_id = ANY(%s) AND dst_node_id = ANY(%s) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, (source_id, node_id_list, node_id_list, MAX_SUBGRAPH_EDGES), @@ -894,6 +1126,223 @@ class GraphStore: cursor.close() conn.rollback() + def seed_nodes_from_facts( + self, + source_id: str, + query_embedding: List[float], + fact_limit: int = 5, + limit: int = 10, + ) -> List[Dict[str, Any]]: + """Seed nodes drawn from the *relationships* nearest the question. + + Name matching asks "which entity is this question about", which a + multi-document question cannot answer: the entity holding the answer is + named in another document, not in the question. A fact string carries + the relation — "Alder streams_to Quill: ..." — so a question about what + a service writes to can match the edge itself and seed the walk on both + of its endpoints, including the one nothing in the question names. + + Endpoints are weighted by fact score divided by the entity's + ``doc_freq``: an entity appearing in every chunk is a poor seed even + when it sits on a well-matched fact, and dividing by how widely it + occurs prefers the specific endpoint over the hub. + + Rows match :meth:`search_nodes_by_embedding`'s shape, so the caller's + seed weighting is unchanged. Returns nothing when the source has no + fact embeddings, which is the signal to fall back to name matching. + """ + if not query_embedding: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + WITH top_facts AS ( + SELECT src_node_id, dst_node_id, + 1 - (fact_embedding <=> %s::vector) AS score + FROM graph_edges + WHERE source_id = %s AND fact_embedding IS NOT NULL + ORDER BY fact_embedding <=> %s::vector + LIMIT %s + ) + SELECT n.id::text, n.name, n.description, + MAX(f.score / GREATEST(COALESCE(n.doc_freq, 1), 1)) AS weight + FROM top_facts f + JOIN graph_nodes n + ON n.id = f.src_node_id OR n.id = f.dst_node_id + WHERE n.source_id = %s + GROUP BY n.id, n.name, n.description + ORDER BY weight DESC + LIMIT %s; + """, + ( + query_embedding, + source_id, + query_embedding, + max(1, int(fact_limit)), + source_id, + max(1, int(limit)), + ), + ) + return [ + { + "id": row[0], + "name": row[1], + "description": row[2], + # The caller reads weight back as ``1 - distance``. + "distance": 1.0 - float(row[3] or 0.0), + } + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error seeding nodes from facts: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_relationships( + self, source_id: str, name: str, limit: int = 25 + ) -> List[Dict[str, Any]]: + """The relationships an entity takes part in, strongest first. + + This is the one thing a caller cannot get from vector search: which + *named* thing an entity is connected to. Matching is on the name rather + than a node id because the caller is an LLM holding a name it read in + the text, not an id. + """ + clean = (name or "").strip() + if not clean: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + SELECT s.name, e.type, d.name, e.description + FROM graph_edges e + JOIN graph_nodes s ON s.id = e.src_node_id + JOIN graph_nodes d ON d.id = e.dst_node_id + WHERE e.source_id = %s AND (s.name ILIKE %s OR d.name ILIKE %s) + ORDER BY e.weight DESC NULLS LAST + LIMIT %s; + """, + (source_id, f"%{clean}%", f"%{clean}%", max(1, int(limit))), + ) + return [ + {"source": row[0], "type": row[1], "target": row[2], "description": row[3]} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading relationships for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_pages( + self, source_id: str, name: str, limit: int = 4 + ) -> List[Dict[str, Any]]: + """Chunks an entity appears in, with the chunk it is *about* first. + + A plain substring match answers "Halvard" with pages that merely mention + Halvard, and an unordered ``LIMIT`` then decides which of those the + caller sees. Nodes whose name is the entity (or the entity plus a + qualifier the extractor appended, "Quill" -> "Quill Store") are + preferred, and among those the chunk whose text opens with the name + comes first; a substring match is the fallback so an unusual name still + resolves. + """ + clean = (name or "").strip() + if not clean: + return [] + table, text_col, metadata_col, source_col = _pgvector_identifiers() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + sql.SQL( + """ + SELECT d.{metadata}, d.{text}, + bool_or(lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + FROM graph_node_chunks gc + JOIN graph_nodes n ON n.id = gc.node_id + JOIN {table} d ON d.id::text = gc.chunk_id + WHERE gc.source_id = %s AND d.{source} = %s + AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) + -- One page per chunk. Grouping on the subject flag as well + -- split a chunk two entities link -- one naming it, one + -- merely mentioned -- into two identical pages, spending + -- the caller's page budget twice on the same text. + GROUP BY d.{metadata}, d.{text} + ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC + LIMIT %s; + """ + ).format( + metadata=_identifier(metadata_col), + text=_identifier(text_col), + table=_identifier(table), + source=_identifier(source_col), + ), + ( + clean.lower(), f"{clean.lower()} %", + source_id, source_id, + clean.lower(), f"{clean.lower()} %", f"%{clean}%", + f"{clean}%", + max(1, int(limit)), + ), + ) + return [ + {"metadata": row[0] or {}, "text": row[1] or ""} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading pages for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def chunk_similarities( + self, source_id: str, chunk_ids: List[str], query_embedding: List[float] + ) -> Dict[str, float]: + """Cosine similarity between the query and specific chunks of a source. + + Passage nodes need their own relevance to claim a share of the walk's + restart mass, and that number lives in the co-located pgvector table — + the same one :meth:`get_chunk_texts` reads. Restricted to the chunk ids + the subgraph actually reached, so this never scans the whole source. + """ + if not chunk_ids or not query_embedding: + return {} + table, _text_col, _metadata_col, source_col = _pgvector_identifiers() + vector_col = _pgvector_vector_column() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + sql.SQL( + """ + SELECT id::text, 1 - ({vector} <=> %s::vector) + FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + vector=_identifier(vector_col), + table=_identifier(table), + source=_identifier(source_col), + ), + (query_embedding, source_id, [str(c) for c in chunk_ids]), + ) + return {row[0]: float(row[1]) for row in cursor.fetchall()} + except Exception as e: + logging.error(f"Error scoring chunks against the query: {e}") + return {} + finally: + cursor.close() + conn.rollback() + def get_chunk_texts( self, source_id: str, @@ -914,10 +1363,17 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT id, {text_col}, {metadata_col} FROM {table} - WHERE {source_col} = %s AND id::text = ANY(%s); - """, + sql.SQL( + """ + SELECT id, {text}, {metadata} FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + text=_identifier(text_col), + metadata=_identifier(metadata_col), + table=_identifier(table), + source=_identifier(source_col), + ), (source_id, [str(c) for c in chunk_ids]), ) return { @@ -1026,25 +1482,29 @@ class GraphStore: cursor.close() def mark_chunk(self, source_id: str, chunk_id: str, status: str): + """Record a chunk's extraction status, reconnecting once if the connection died.""" self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - try: - cursor.execute( - """ - INSERT INTO graph_ingest_progress (source_id, chunk_id, status) - VALUES (%s, %s, %s) - ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; - """, - (source_id, str(chunk_id), status), - ) - conn.commit() - except Exception as e: - conn.rollback() - logging.error(f"Error marking chunk: {e}") - raise - finally: - cursor.close() + + def _write(conn): + cursor = conn.cursor() + try: + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, %s) + ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id), status), + ) + conn.commit() + except Exception as e: + _safe_rollback(conn) + logging.error(f"Error marking chunk: {e}") + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def pending_chunks(self, source_id: str, all_chunk_ids: List[str]) -> List[str]: """Chunk ids from ``all_chunk_ids`` not yet marked ``done`` for the source.""" @@ -1092,6 +1552,9 @@ class GraphStore: conn = self._get_connection() cursor = conn.cursor() try: + # A reset while a build is still writing must not land in the + # middle of one of its chunks. + _lock_source(cursor, source_id) for table in ( "graph_node_chunks", "graph_edges", @@ -1099,7 +1562,10 @@ class GraphStore: "graph_ingest_progress", ): cursor.execute( - f"DELETE FROM {table} WHERE source_id = %s;", (source_id,) + sql.SQL("DELETE FROM {} WHERE source_id = %s;").format( + sql.Identifier(table) + ), + (source_id,), ) conn.commit() except Exception as e: diff --git a/docsgpt/retriever/dispatcher.py b/docsgpt/retriever/dispatcher.py index a14fc02f..31cc75e0 100644 --- a/docsgpt/retriever/dispatcher.py +++ b/docsgpt/retriever/dispatcher.py @@ -185,12 +185,23 @@ class Dispatcher(BaseRetriever): score_threshold / rephrase_query) plus an opted-in prescreen config; a source left at defaults takes the global path so all-classic retrieval stays byte-identical with zero extra LLM calls. + + A graph source's ``graph`` options count too: they are read from the + per-source config this records, so a source that changes only those + would otherwise run the defaults and the options would do nothing. + They mean nothing to any other retriever, so they only count for + ``graphrag`` -- an override hands the source its own chunk budget as + well, which a classic source must not pick up from a graph setting. """ return ( retrieval.chunks != _DEFAULT_RETRIEVAL.chunks or retrieval.score_threshold != _DEFAULT_RETRIEVAL.score_threshold or retrieval.rephrase_query != _DEFAULT_RETRIEVAL.rephrase_query or retrieval.prescreen is not None + or ( + (retrieval.retriever or "").lower() == "graphrag" + and retrieval.graph != _DEFAULT_RETRIEVAL.graph + ) ) @staticmethod diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index c7ac538e..8d191f7b 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -1,9 +1,13 @@ """GraphRAG local retriever — Personalized PageRank over a per-source graph. -Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx -Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR -mass -> shared token budget. No LLM call at query time beyond the (optional, -reused) rephrase. +Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> Personalized +PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR mass -> shared +token budget. No LLM call at query time beyond the (optional, reused) rephrase. + +``networkx`` supplies the graph structure, but the ranking is the local power +iteration in :func:`_personalized_pagerank`: ``nx.pagerank`` delegates to scipy, +which DocsGPT does not depend on, so calling it turned every graph retrieval +into a silent ClassicRAG fallback. Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the ``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the @@ -28,6 +32,7 @@ from docsgpt.graphrag.store import GraphStore from docsgpt.retriever.base import BaseRetriever from docsgpt.retriever.classic_rag import ClassicRAG from docsgpt.retriever.labels import labels_from_metadata +from docsgpt.storage.db.source_config import GraphRetrievalConfig from docsgpt.utils import num_tokens_from_string from docsgpt.vectorstore.base import get_embeddings @@ -35,11 +40,144 @@ SEED_NODES = 10 SUBGRAPH_HOPS = 1 +PASSAGE_NODE_WEIGHT = 0.05 +FACT_SEED_FACTS = 5 +RRF_K = 60 + +# PageRank damping per ranking mode — each is the value that mode was measured +# at. Lower keeps mass nearer the seeds; with passages in the walk 0.5 measured +# better, while entity-only ranking was measured at the conventional 0.85. +DAMPING_WITH_PASSAGES = 0.5 +DAMPING_ENTITIES_ONLY = 0.85 + + def _idf(doc_freq: Any) -> float: """Node-specificity weight: rarer entities (low ``doc_freq``) score higher.""" return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _nodes_by_chunk(chunk_links: Dict[str, List[str]]) -> Dict[str, List[str]]: + """Invert ``node -> chunk ids`` into ``chunk id -> node ids``. + + Node order within a chunk follows the node order of ``chunk_links``, so the + passage edges are added in the same order as before. + """ + inverted: Dict[str, List[str]] = {} + for node, chunks in chunk_links.items(): + for chunk_id in chunks or (): + inverted.setdefault(chunk_id, []).append(node) + return inverted + + +def _damping(passage_nodes: bool) -> float: + """PageRank damping for a ranking mode: the value that mode was measured at.""" + return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY + + +def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]: + """Normalized restart distribution over ``nodes``. + + Weights are clamped at zero (a cosine distance above 1 yields a negative + seed weight) and normalized across the nodes actually present in the graph, + so restart mass can never leak to a node the subgraph does not contain. An + absent or all-zero personalization collapses to a uniform restart. + """ + if personalization: + weights = { + node: max(float(personalization.get(node, 0.0) or 0.0), 0.0) + for node in nodes + } + total = sum(weights.values()) + if total > 0: + return {node: weight / total for node, weight in weights.items()} + uniform = 1.0 / len(nodes) + return {node: uniform for node in nodes} + + +def _personalized_pagerank( + graph: nx.Graph, + personalization: Dict[Any, float] | None = None, + *, + weight: str = "weight", + alpha: float = 0.85, + max_iter: int = 100, + tol: float = 1.0e-6, +) -> Dict[Any, float]: + """Personalized PageRank by power iteration — no scipy. + + ``networkx.pagerank`` delegates to a scipy implementation, and scipy is not + a DocsGPT dependency: in a default install the import raises and every graph + retrieval silently degrades to the ClassicRAG fallback. This is the same + algorithm over the same row-normalized transition matrix, so the ranking is + unchanged where scipy happens to be installed. + + Args: + graph: Undirected graph whose edges may carry a ``weight`` attribute. + personalization: Node -> restart weight; ``None`` means uniform. + weight: Edge attribute holding the weight. + alpha: Damping factor. + max_iter: Iteration cap. The last iterate is returned if it is hit — + retrieval degrades to a slightly less converged ranking rather than + raising, which is what the library does. + tol: Convergence tolerance; iteration stops below ``len(graph) * tol``. + + Returns: + Node -> PageRank mass, summing to ~1.0. Empty dict for an empty graph. + """ + nodes = list(graph.nodes) + node_count = len(nodes) + if node_count == 0: + return {} + + restart = _restart_vector(nodes, personalization) + + # Row-normalized transitions. An undirected edge is traversable from both + # endpoints, so each node normalizes over its own incident weights. + transitions: Dict[Any, List[tuple[Any, float]]] = {} + for node in nodes: + neighbors = [] + total = 0.0 + for neighbor, data in graph[node].items(): + raw_weight = data.get(weight, 1.0) + # Default only a missing or null weight. ``or 1.0`` would also + # rewrite an explicit 0 — "these entities are not related" — into a + # full-strength transition, which changes the ranking. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + if edge_weight <= 0: + continue + neighbors.append((neighbor, edge_weight)) + total += edge_weight + transitions[node] = ( + [(n, w / total) for n, w in neighbors] if total > 0 else [] + ) + + # A node with no usable edge is dangling: its mass would vanish each pass, + # so it is redistributed along the restart vector instead. + dangling = [node for node in nodes if not transitions[node]] + + ranks = {node: 1.0 / node_count for node in nodes} + for _ in range(max_iter): + previous = ranks + ranks = dict.fromkeys(nodes, 0.0) + leaked = alpha * sum(previous[node] for node in dangling) + for node in nodes: + share = alpha * previous[node] + for neighbor, transition in transitions[node]: + ranks[neighbor] += share * transition + for node in nodes: + ranks[node] += (leaked + 1.0 - alpha) * restart[node] + if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: + break + else: + logging.debug( + "Personalized PageRank hit its %s-iteration cap on a %s-node " + "subgraph; ranking with the last iterate.", + max_iter, + node_count, + ) + return ranks + + class GraphRAGRetriever(BaseRetriever): """Per-source PPR retriever; falls back to ClassicRAG when a source has no graph.""" @@ -102,14 +240,9 @@ class GraphRAGRetriever(BaseRetriever): After PPR, each node's mass is scaled by ``1/log(2 + doc_freq)`` so a high-degree hub contributes less than a specific entity at equal mass. """ - graph = nx.Graph() - for node in subgraph.get("nodes", []): - graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) - for edge in subgraph.get("edges", []): - src, dst = edge["src_node_id"], edge["dst_node_id"] - if src in graph and dst in graph: - weight = float(edge.get("weight") or 1.0) - graph.add_edge(src, dst, weight=weight) + # Through the class, not ``self``: this method reads no instance state, + # and callers (and tests) rely on being able to invoke it unbound. + graph = GraphRAGRetriever._subgraph_graph(subgraph) if graph.number_of_nodes() == 0: return {} @@ -117,12 +250,34 @@ class GraphRAGRetriever(BaseRetriever): if not any(personalization.values()): personalization = None - ranks = nx.pagerank(graph, personalization=personalization, weight="weight") + ranks = _personalized_pagerank( + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=False), + ) return { node: rank * _idf(graph.nodes[node].get("doc_freq", 0)) for node, rank in ranks.items() } + @staticmethod + def _subgraph_graph(subgraph) -> "nx.Graph": + """The fetched subgraph as a weighted undirected graph.""" + graph = nx.Graph() + for node in subgraph.get("nodes", []): + graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) + for edge in subgraph.get("edges", []): + src, dst = edge["src_node_id"], edge["dst_node_id"] + if src in graph and dst in graph: + raw_weight = edge.get("weight") + # Default only a missing or null weight. Coercing an explicit 0 + # to 1.0 would make "these entities are not related" the + # strongest possible link. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + graph.add_edge(src, dst, weight=edge_weight) + return graph + def _rank_chunks(self, store, source_id, node_scores) -> List[str]: """Score chunks by summed (PPR mass x IDF) of their linked nodes; top candidates. @@ -140,6 +295,77 @@ class GraphRAGRetriever(BaseRetriever): candidates = max(self.chunks * 2, self.chunks + 5) return ranked[: max(1, candidates)] + def _rank_chunks_with_passages( + self, store, source_id, subgraph, seeds, query_embedding + ) -> List[str]: + """Rank chunks by walking a graph that contains the chunks themselves. + + :meth:`_rank_chunks` reads a chunk's score *off* its entities, summing + their PPR mass — so a chunk touching many mid-scoring generic entities + outranks one touching the few entities the question is about. Putting + the chunks in the walk instead, each joined to its own entities and + carrying a small share of the restart mass proportional to its own + vector similarity, makes a chunk reachable both ways: by being about the + question, and by being connected to what is. Graph retrieval then + contains vector retrieval rather than competing with it. + + Only an improvement when the seeds are good: measured across five + corpora it helped alongside richer seed embeddings and *hurt* with + bare-name seeds (0.73 -> 0.57 on one corpus). Graphs are now always + built with the richer seed text; one built before that change should be + rebuilt before this is relied on. + """ + node_ids = [node["id"] for node in subgraph.get("nodes", [])] + chunk_links = store.get_chunk_ids_for_nodes(source_id, node_ids) + candidate_ids = sorted({c for chunks in chunk_links.values() for c in chunks}) + if not candidate_ids: + return [] + + graph = self._subgraph_graph(subgraph) + similarities = store.chunk_similarities( + source_id, candidate_ids, query_embedding + ) + # Normalised so the passage share is a fixed fraction of the restart + # mass rather than whatever absolute cosine this embedding model emits. + scores = [similarities.get(c, 0.0) for c in candidate_ids] + low, high = (min(scores), max(scores)) if scores else (0.0, 0.0) + spread = high - low + + personalization = dict(seeds) + passage_of: Dict[str, str] = {} + # Inverted once: scanning every node's chunk list per candidate is + # quadratic, and at the candidate cap it cost more than the walk it + # feeds (76 ms against 0.5 ms measured). + nodes_by_chunk = _nodes_by_chunk(chunk_links) + for chunk_id in candidate_ids: + linked = [n for n in nodes_by_chunk.get(chunk_id, ()) if n in graph] + if not linked: + continue + passage_node = f"chunk::{chunk_id}" + passage_of[passage_node] = chunk_id + for node in linked: + graph.add_edge(passage_node, node, weight=1.0) + similarity = similarities.get(chunk_id, 0.0) + normalized = (similarity - low) / spread if spread > 0 else 0.0 + personalization[passage_node] = normalized * PASSAGE_NODE_WEIGHT + + if graph.number_of_nodes() == 0 or not any(personalization.values()): + return [] + + ranks = _personalized_pagerank( + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=True), + ) + chunk_scores = { + chunk_id: ranks.get(passage_node, 0.0) + for passage_node, chunk_id in passage_of.items() + } + ranked = sorted(chunk_scores, key=lambda c: chunk_scores[c], reverse=True) + candidates = max(self.chunks * 2, self.chunks + 5) + return ranked[: max(1, candidates)] + def _source_top_k(self, source_id) -> int: """How many chunks this source may contribute — its own top-k. @@ -159,6 +385,111 @@ class GraphRAGRetriever(BaseRetriever): base = self.base_chunks if self.base_chunks is not None else self.chunks return max(1, base // max(1, len(self.vectorstores))) + def _vector_ranking(self, source_id, query_embedding: List[float]) -> List[tuple]: + """The source's own vector ranking, as ``(text, metadata)`` in score order. + + Used only by the hybrid path. Vector hits carry no row id, so the fused + ranking is keyed on the chunk text itself — the one identifier both + rankings share — and the metadata travels with it so a hit the graph + never surfaced can still be emitted as a document. + """ + from docsgpt.vectorstore.vector_creator import VectorCreator + + store = None + try: + store = VectorCreator.create_vectorstore( + settings.VECTOR_STORE, source_id, settings.EMBEDDINGS_KEY + ) + hits = store.search( + self._classic._get_rephrased_question(), + k=max(self.chunks * 4, 20), + query_vector=query_embedding, + ) + except Exception as e: + logging.error( + "GraphRAG hybrid: vector ranking failed for %s: %s", source_id, e + ) + return [] + finally: + close = getattr(store, "close", None) + if close is not None: + try: + close() + except Exception as e: + logging.debug("Error closing hybrid vector store: %s", e) + ranked = [] + for hit in hits: + text = getattr(hit, "page_content", None) + metadata = getattr(hit, "metadata", None) + if text is None and isinstance(hit, dict): + text = hit.get("text") or hit.get("page_content") + metadata = hit.get("metadata") + if text: + ranked.append((text, metadata or {})) + return ranked + + @staticmethod + def _rrf_order(rankings: List[List[str]], k: int) -> Dict[str, float]: + """Reciprocal rank fusion over ranked lists of the same key type. + + Rank-based on purpose: PPR mass and cosine similarity are not on + comparable scales, and normalising either one invents a calibration + that does not exist. + """ + scores: Dict[str, float] = {} + for ranking in rankings: + for position, key in enumerate(ranking): + scores[key] = scores.get(key, 0.0) + 1.0 / (k + position + 1) + return scores + + def _graph_options(self, source_id) -> GraphRetrievalConfig: + """This source's graph retrieval options, or the recommended defaults. + + Options travel on the per-source retrieval config the Dispatcher hands + over. A request that carries no per-source detail gets the defaults, + which are the measured-best configuration rather than a neutral one. + """ + cfg = (getattr(self, "per_source_retrieval", None) or {}).get(source_id) + options = cfg.get("graph") if isinstance(cfg, dict) else getattr(cfg, "graph", None) + if isinstance(options, GraphRetrievalConfig): + return options + try: + return GraphRetrievalConfig.model_validate(options or {}) + except Exception: + return GraphRetrievalConfig() + + def _seed_rows( + self, store, source_id, query_embedding: List[float] + ) -> List[Dict[str, Any]]: + """The nodes the walk restarts from, per the source's ``seed_strategy``. + + Seeding decides more than ranking does: a walk that starts on the wrong + nodes cannot be rescued downstream. + + ``entities`` + Cosine NN over entity embeddings, built from each entity's name, + type and description so a whole question has something to match. + The default: best or tied-best on every corpus measured. + ``relationships`` + Cosine NN over relationship sentences ("A streams_to B: ..."), + seeding both endpoints of the best-matching facts. The only way to + start on an entity the question never names; strongest on + chain-structured content, weaker on ordinary prose. + + Relationship seeding falls back to entity matching for a source with no + fact embeddings (one built before they were recorded), so it still + retrieves rather than returning nothing. + """ + if self._graph_options(source_id).seed_strategy == "relationships": + by_fact = store.seed_nodes_from_facts( + source_id, query_embedding, fact_limit=FACT_SEED_FACTS, limit=SEED_NODES + ) + if by_fact: + return by_fact + return store.search_nodes_by_embedding( + source_id, query_embedding, k=SEED_NODES + ) + def _graph_docs_for_source( self, store, source_id, query_embedding: List[float] ) -> List[Dict[str, Any]]: @@ -170,9 +501,7 @@ class GraphRAGRetriever(BaseRetriever): query_embedding: Embedding of the rephrased question, computed once by the caller for the whole retrieval. """ - seed_rows = store.search_nodes_by_embedding( - source_id, query_embedding, k=SEED_NODES - ) + seed_rows = self._seed_rows(store, source_id, query_embedding) if not seed_rows: return [] @@ -186,26 +515,63 @@ class GraphRAGRetriever(BaseRetriever): for row in seed_rows } + options = self._graph_options(source_id) subgraph = store.get_subgraph(source_id, seed_ids, hops=SUBGRAPH_HOPS) - node_scores = self._ppr_scores(subgraph, seeds) - if not node_scores: + if options.passage_nodes: + chunk_ids = self._rank_chunks_with_passages( + store, source_id, subgraph, seeds, query_embedding + ) + else: + node_scores = self._ppr_scores(subgraph, seeds) + if not node_scores: + return [] + chunk_ids = self._rank_chunks(store, source_id, node_scores) + if not chunk_ids: return [] - chunk_ids = self._rank_chunks(store, source_id, node_scores) chunk_data = store.get_chunk_texts(source_id, chunk_ids) + # ``(text, metadata)`` in rank order. Chunk ids stop being the currency + # here: a hit contributed by the vector ranking has no graph chunk id, + # and keying on ids is what made an earlier version of this fusion able + # only to reorder the graph's own candidates. + candidates: List[tuple] = [] + for chunk_id in chunk_ids: + chunk = chunk_data.get(chunk_id) + text = chunk.get("text") if chunk else None + if text: + candidates.append((text, chunk.get("metadata"))) + + if options.blend_vector: + # The graph ranks by how much PPR mass landed on a chunk's + # entities, which says nothing about whether the chunk is about the + # question. Fusing with the source's own vector ranking keeps the + # graph's reach while letting plain relevance back in — including + # chunks the graph never surfaced, which is where most of the value + # is: no reordering can rescue a question whose answer the graph + # missed entirely. + vector_hits = self._vector_ranking(source_id, query_embedding) + if vector_hits: + metadata_by_text = {text: meta for text, meta in candidates} + for text, meta in vector_hits: + metadata_by_text.setdefault(text, meta) + fused = self._rrf_order( + [[t for t, _ in candidates], [t for t, _ in vector_hits]], + RRF_K, + ) + candidates = [ + (text, metadata_by_text.get(text)) + for text in sorted(fused, key=lambda t: fused[t], reverse=True) + ] + docs: List[Dict[str, Any]] = [] token_budget = max(int(self.doc_token_limit * 0.9), 100) cumulative_tokens = 0 source_top_k = self._source_top_k(source_id) - for chunk_id in chunk_ids: + for text, metadata in candidates: if len(docs) >= source_top_k: break - chunk = chunk_data.get(chunk_id) - text = chunk.get("text") if chunk else None - if not text: - continue - labels = labels_from_metadata(chunk.get("metadata"), text, source_id) + labels = labels_from_metadata(metadata, text, source_id) doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}") if cumulative_tokens + doc_tokens >= token_budget: break @@ -252,8 +618,9 @@ class GraphRAGRetriever(BaseRetriever): Graph sources keep their own slot in source order; every graphless source collapses into a single ClassicRAG run that occupies the slot of - the first graphless source. Sources whose PPR retrieval raises are - collected and retried as one more classic batch, appended at the end. + the first graphless source. Sources whose PPR retrieval raises, or + answers nothing, are collected and retried as one more classic batch, + appended at the end. """ try: counts = store.count_nodes_many(sources) @@ -279,7 +646,7 @@ class GraphRAGRetriever(BaseRetriever): segments.append([]) graphless.append(source_id) - failed: List[str] = [] + fallback: List[str] = [] query_embedding = None if graphed: # Embedded once for the whole retrieval, not once per graph source. @@ -292,26 +659,39 @@ class GraphRAGRetriever(BaseRetriever): f"GraphRAG query embedding failed, falling back: {e}", exc_info=True, ) - failed, graphed = list(graphed), [] + fallback, graphed = list(graphed), [] for source_id in graphed: try: - segments[graph_slots[source_id]] = self._graph_docs_for_source( - store, source_id, query_embedding - ) + docs = self._graph_docs_for_source(store, source_id, query_embedding) except Exception as e: logging.error( f"GraphRAG retrieval failed for {source_id}, falling back: {e}", exc_info=True, ) - failed.append(source_id) + fallback.append(source_id) + continue + if not docs: + # Empty is not an answer. Every graph read reports its own + # failure and returns nothing, so "no rows" covers a query that + # broke or a half-built graph as much as a walk that found + # nothing — and only a raise reaches the fallback, so the + # source would otherwise contribute nothing at all. Searching + # it classically is what a source with no graph already gets. + logging.info( + "GraphRAG retrieval returned nothing for %s, falling back", + source_id, + ) + fallback.append(source_id) + continue + segments[graph_slots[source_id]] = docs # Every remaining segment is a ClassicRAG fan-out, and each of its legs # checks out of the *same* per-DSN pool this store is holding. Hand the # graph connection back first, or concurrent GraphRAG retrievals occupy # every slot and then block on their own fallbacks until PoolTimeout. # ``close()`` nulls the connection, so ``_get_data``'s finally stays correct. - if graphless or failed: + if graphless or fallback: try: store.close() except Exception as e: @@ -319,8 +699,8 @@ class GraphRAGRetriever(BaseRetriever): if graphless: segments[classic_slot] = self._classic_for_sources(graphless) - if failed: - segments.append(self._classic_for_sources(failed)) + if fallback: + segments.append(self._classic_for_sources(fallback)) return [doc for segment in segments for doc in segment] diff --git a/docsgpt/storage/db/source_config.py b/docsgpt/storage/db/source_config.py index 20873b46..8f8d6209 100644 --- a/docsgpt/storage/db/source_config.py +++ b/docsgpt/storage/db/source_config.py @@ -12,7 +12,7 @@ reproduces today's chunking byte-for-byte. from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import BaseModel, ConfigDict, field_validator, model_validator @@ -90,6 +90,27 @@ class ChunkingConfig(BaseModel): duplicate_headers: bool = False +class GraphRetrievalConfig(BaseModel): + """How the graph retriever walks a graphrag source (live; no re-ingest). + + The defaults are the configuration that measured best across the corpora + tested rather than a neutral starting point: seed from entity matches, put + the passages in the walk, and blend with the source's own vector ranking. + """ + + model_config = ConfigDict(extra="forbid") + + # Where the walk starts: entities whose descriptions match the question, or + # relationships ("A streams_to B") that do. Relationships can start the + # walk on an entity the question never names. + seed_strategy: Literal["entities", "relationships"] = "entities" + # Chunks join the walk as nodes, so a passage is reachable both by being + # about the question and by being connected to what is. + passage_nodes: bool = True + # Fuse the graph ranking with plain vector search by reciprocal rank. + blend_vector: bool = True + + class RetrievalConfig(BaseModel): """Query-time retrieval knobs (live; no re-ingest needed).""" @@ -102,6 +123,7 @@ class RetrievalConfig(BaseModel): rephrase_query: bool = True # toggle ClassicRAG._rephrase_query side-call reranker: Optional[dict] = None # reserved: future cross-encoder/LLM reorder prescreen: Optional[dict] = None # None = off; else PreScreenConfig dict (D12) + graph: GraphRetrievalConfig = GraphRetrievalConfig() # graphrag retriever only @field_validator("chunks") @classmethod diff --git a/docsgpt/vectorstore/embeddings_delegated.py b/docsgpt/vectorstore/embeddings_delegated.py index 75f730b8..bc184f34 100644 --- a/docsgpt/vectorstore/embeddings_delegated.py +++ b/docsgpt/vectorstore/embeddings_delegated.py @@ -10,8 +10,9 @@ Celery and the vector comes back. The API pays a broker round trip per query and no resident model. Inside a worker there is nothing to delegate to -- dispatching would queue work -behind the task already running and wait on itself -- so a call made while a -task is executing runs locally, on a model this process loads once and caches. +behind the task already running and wait on itself -- so a call made anywhere in +a worker process, including from a thread a task started, runs locally, on a +model this process loads once and caches. ``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side. Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding @@ -79,11 +80,11 @@ def _forget(result) -> None: def _in_worker() -> bool: - """True when a Celery task is executing in this process.""" + """True anywhere in a Celery worker process -- on any thread, not only the task's.""" try: - from docsgpt.celery_init import celery + from docsgpt.celery_init import in_worker - return celery.current_worker_task is not None + return in_worker() except Exception: return False diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index e64ed52f..b914fbfd 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -268,6 +268,19 @@ }, "exposureHint": "Lade diese Quelle vorab in den Prompt oder lass den Agenten sie bei Bedarf als Werkzeug durchsuchen." }, + "graphRetrieval": { + "title": "Graph-Abruf", + "tag": "ohne Neuimport", + "seedStrategy": "Suche beginnt bei", + "seedStrategyHint": "Entitäten eignen sich für die meisten Dokumente. Beziehungen erreichen auch Entitäten, die in der Frage nicht vorkommen – ideal für Inhalte, die beschreiben, wie Dinge zusammenhängen.", + "seedEntities": "Entitäten (empfohlen)", + "seedRelationships": "Beziehungen", + "passageNodes": "Textabschnitte in die Suche einbeziehen", + "passageNodesHint": "Ein Abschnitt wird gefunden, wenn er zur Frage passt oder mit etwas Passendem verbunden ist. Am besten bei Graphen, die mit dieser Version erstellt wurden.", + "blendVector": "Mit Vektorsuche kombinieren", + "blendVectorHint": "Ergänzt Ergebnisse der Vektorsuche, damit kein Abschnitt verloren geht, den der Graph übersieht.", + "agentToolHint": "Agenten können diesen Beziehungen auch selbst folgen, wenn die Bereitstellung dieser Quelle „Suchwerkzeug auf Abruf“ ist oder ein agentischer Agent sie nutzt." + }, "prescreen": { "enable": "LLM-Vorfilterung aktivieren", "warning": "Ruft eine größere Kandidatenmenge ab und filtert sie mit einem LLM. Das erhöht Latenz und Kosten pro Anfrage.", diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 03f0f0dd..ba67a1cb 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -272,6 +272,19 @@ }, "exposureHint": "Pre-fetch this source into the prompt, or let the agent search it on demand as a tool." }, + "graphRetrieval": { + "title": "Graph retrieval", + "tag": "no re-ingest", + "seedStrategy": "Start the walk from", + "seedStrategyHint": "Entities suit most documents. Relationships can reach an entity the question never names, and suit content that describes how things connect.", + "seedEntities": "Entities (recommended)", + "seedRelationships": "Relationships", + "passageNodes": "Include passages in the walk", + "passageNodesHint": "Lets a passage be found both by matching the question and by being connected to what does. Works best on graphs built with this version.", + "blendVector": "Blend with vector search", + "blendVectorHint": "Adds plain vector search results, so a passage the graph misses is not lost.", + "agentToolHint": "Agents can also follow these relationships themselves when this source's exposure is “On-demand search tool”, or when an agentic agent uses it." + }, "prescreen": { "enable": "Enable LLM prescreen", "warning": "Fetches a larger candidate set and uses an LLM to filter it. This adds query-time latency and cost.", diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index c7b4f794..f30cadd7 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -268,6 +268,19 @@ }, "exposureHint": "Precarga esta fuente en el prompt, o deja que el agente la busque bajo demanda como herramienta." }, + "graphRetrieval": { + "title": "Recuperación por grafo", + "tag": "sin reingesta", + "seedStrategy": "Iniciar el recorrido desde", + "seedStrategyHint": "Las entidades funcionan para la mayoría de documentos. Las relaciones pueden llegar a una entidad que la pregunta no menciona; son ideales para contenido que describe cómo se conectan las cosas.", + "seedEntities": "Entidades (recomendado)", + "seedRelationships": "Relaciones", + "passageNodes": "Incluir fragmentos en el recorrido", + "passageNodesHint": "Un fragmento puede encontrarse por coincidir con la pregunta o por estar conectado con lo que coincide. Funciona mejor en grafos creados con esta versión.", + "blendVector": "Combinar con búsqueda vectorial", + "blendVectorHint": "Añade resultados de la búsqueda vectorial para no perder fragmentos que el grafo pase por alto.", + "agentToolHint": "Los agentes también pueden seguir estas relaciones por sí mismos cuando la exposición de esta fuente es «Herramienta de búsqueda bajo demanda» o cuando la usa un agente agéntico." + }, "prescreen": { "enable": "Habilitar preselección con LLM", "warning": "Obtiene un conjunto de candidatos más grande y usa un LLM para filtrarlo. Esto añade latencia y costo por consulta.", diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index dc022fe8..e475f954 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -268,6 +268,19 @@ }, "exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。" }, + "graphRetrieval": { + "title": "グラフ検索", + "tag": "再取り込み不要", + "seedStrategy": "探索の開始点", + "seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。", + "seedEntities": "エンティティ(推奨)", + "seedRelationships": "リレーション", + "passageNodes": "パッセージを探索に含める", + "passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。", + "blendVector": "ベクトル検索と組み合わせる", + "blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。", + "agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。" + }, "prescreen": { "enable": "LLMプリスクリーニングを有効にする", "warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。", diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index dd7aef1e..ba064d54 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -268,6 +268,19 @@ }, "exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту." }, + "graphRetrieval": { + "title": "Поиск по графу", + "tag": "без повторной загрузки", + "seedStrategy": "Начинать обход с", + "seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.", + "seedEntities": "Сущностей (рекомендуется)", + "seedRelationships": "Связей", + "passageNodes": "Включать фрагменты в обход", + "passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.", + "blendVector": "Сочетать с векторным поиском", + "blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.", + "agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент." + }, "prescreen": { "enable": "Включить предварительный отбор LLM", "warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.", diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 16236cec..c8c3e7e1 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -268,6 +268,19 @@ }, "exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。" }, + "graphRetrieval": { + "title": "圖譜檢索", + "tag": "無需重新匯入", + "seedStrategy": "走訪起點", + "seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。", + "seedEntities": "實體(建議)", + "seedRelationships": "關係", + "passageNodes": "將段落納入走訪", + "passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。", + "blendVector": "與向量檢索結合", + "blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。", + "agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。" + }, "prescreen": { "enable": "啟用 LLM 預篩選", "warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。", diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 5b4cf5a6..40e8fa0e 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -268,6 +268,19 @@ }, "exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。" }, + "graphRetrieval": { + "title": "图谱检索", + "tag": "无需重新导入", + "seedStrategy": "遍历起点", + "seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。", + "seedEntities": "实体(推荐)", + "seedRelationships": "关系", + "passageNodes": "将段落纳入遍历", + "passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。", + "blendVector": "与向量检索结合", + "blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。", + "agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。" + }, "prescreen": { "enable": "启用 LLM 预筛选", "warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。", diff --git a/frontend/src/models/misc.ts b/frontend/src/models/misc.ts index fc8b0782..c0c5efb2 100644 --- a/frontend/src/models/misc.ts +++ b/frontend/src/models/misc.ts @@ -32,6 +32,17 @@ export type SourcePrescreenConfig = { max_keep?: number; // default 8, <= candidate_k }; +// Where the graph walk starts: matching entities, or matching relationships +// ("A streams_to B"), which can reach an entity the question never names. +export type GraphSeedStrategy = 'entities' | 'relationships'; + +// Query-time graph retrieval knobs (graphrag only; live, no re-ingest). +export type SourceGraphRetrievalConfig = { + seed_strategy?: GraphSeedStrategy; // default 'entities' + passage_nodes?: boolean; // default true + blend_vector?: boolean; // default true +}; + // Query-time retrieval knobs (live; no re-ingest needed). export type SourceRetrievalConfig = { retriever?: string; // default 'classic' (only option for now) @@ -40,6 +51,7 @@ export type SourceRetrievalConfig = { score_threshold?: number | null; // default null rephrase_query?: boolean; // default true prescreen?: SourcePrescreenConfig | null; // null = off + graph?: SourceGraphRetrievalConfig; // graphrag retriever only }; // Ingest-time GraphRAG extraction knobs (only used when kind === 'graphrag'). diff --git a/frontend/src/settings/components/RetrievalOptions.test.tsx b/frontend/src/settings/components/RetrievalOptions.test.tsx index 92142b23..264c1794 100644 --- a/frontend/src/settings/components/RetrievalOptions.test.tsx +++ b/frontend/src/settings/components/RetrievalOptions.test.tsx @@ -146,6 +146,11 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { batch_size: 5, max_keep: 10, }, + graph: { + seed_strategy: 'relationships', + passage_nodes: false, + blend_vector: false, + }, }, graph: { extraction_model: null, @@ -157,6 +162,44 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { }); }); +describe('graph retrieval options', () => { + it('defaults to the measured-best configuration', () => { + expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('fills the defaults for a source saved before the options existed', () => { + const opts = configToOptions({ retrieval: { retriever: 'graphrag' } }); + expect(opts.retrieval.graph).toEqual( + DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph, + ); + }); + + it('honors stored options and fills only the missing ones', () => { + const opts = configToOptions({ + retrieval: { graph: { seed_strategy: 'relationships' } }, + }); + expect(opts.retrieval.graph).toEqual({ + seed_strategy: 'relationships', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('writes the options into the retrieval block', () => { + const v = clone(DEFAULT_RETRIEVAL_OPTIONS); + v.retrieval.graph.blend_vector = false; + expect(optionsToConfig(v).retrieval?.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: false, + }); + }); +}); + describe('isPrescreenConfigValid', () => { const withPrescreen = ( chunks: number, diff --git a/frontend/src/settings/components/RetrievalOptions.tsx b/frontend/src/settings/components/RetrievalOptions.tsx index e2723358..b9dec5d6 100644 --- a/frontend/src/settings/components/RetrievalOptions.tsx +++ b/frontend/src/settings/components/RetrievalOptions.tsx @@ -16,6 +16,7 @@ import { import { Switch } from '../../components/ui/switch'; import type { ChunkingStrategy, + GraphSeedStrategy, RetrievalExposure, SourceConfig, } from '../../models/misc'; @@ -59,6 +60,11 @@ export type RetrievalOptionsValue = { batch_size: number; max_keep: number; }; + graph: { + seed_strategy: GraphSeedStrategy; + passage_nodes: boolean; + blend_vector: boolean; + }; }; graph: { extraction_model: string | null; @@ -85,6 +91,12 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = { enabled: false, ...DEFAULT_PRESCREEN, }, + // The configuration that measured best across the corpora tested. + graph: { + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }, }, graph: { extraction_model: null, @@ -204,6 +216,7 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { const chunking = config?.chunking ?? {}; const retrieval = config?.retrieval ?? {}; const prescreen = retrieval.prescreen ?? null; + const retrievalGraph = retrieval.graph ?? {}; const graph = config?.graph ?? {}; const d = DEFAULT_RETRIEVAL_OPTIONS; return { @@ -228,6 +241,14 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { batch_size: prescreen?.batch_size ?? DEFAULT_PRESCREEN.batch_size, max_keep: prescreen?.max_keep ?? DEFAULT_PRESCREEN.max_keep, }, + graph: { + seed_strategy: + retrievalGraph.seed_strategy ?? d.retrieval.graph.seed_strategy, + passage_nodes: + retrievalGraph.passage_nodes ?? d.retrieval.graph.passage_nodes, + blend_vector: + retrievalGraph.blend_vector ?? d.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: graph.extraction_model ?? d.graph.extraction_model, @@ -271,6 +292,11 @@ export function optionsToConfig(value: RetrievalOptionsValue): SourceConfig { max_keep: ps.max_keep, } : null, + graph: { + seed_strategy: value.retrieval.graph.seed_strategy, + passage_nodes: value.retrieval.graph.passage_nodes, + blend_vector: value.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: value.graph.extraction_model?.trim() @@ -399,6 +425,12 @@ export default function RetrievalOptions({ }); }; + const setGraphRetrieval = ( + patch: Partial, + ) => { + setRetrieval({ graph: { ...value.retrieval.graph, ...patch } }); + }; + const modelOptions = useMemo(() => { const builtin: Model[] = []; const user: Model[] = []; @@ -611,6 +643,84 @@ export default function RetrievalOptions({ )} + {/* Graph retrieval group (graphrag only; live, so shown when testing too) */} + {isGraphRAG && ( +
+ +

+ {tr('graphRetrieval.agentToolHint')} +

+ +
+ + + + + + + setGraphRetrieval({ passage_nodes: checked }) + } + /> + + + + + setGraphRetrieval({ blend_vector: checked }) + } + /> + +
+
+ )} + {/* Graph extraction group (graphrag only; re-ingest required to apply) */} {isGraphRAG && !queryOnly && (
diff --git a/tests/agents/test_classic_agent.py b/tests/agents/test_classic_agent.py index b73e1145..6b10e79d 100644 --- a/tests/agents/test_classic_agent.py +++ b/tests/agents/test_classic_agent.py @@ -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 ): diff --git a/tests/agents/test_research_agent.py b/tests/agents/test_research_agent.py index 5b76fe84..cd9aabd9 100644 --- a/tests/agents/test_research_agent.py +++ b/tests/agents/test_research_agent.py @@ -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 ): diff --git a/tests/agents/tools/test_read_document_tool.py b/tests/agents/tools/test_read_document_tool.py index 8913495d..f0e7b177 100644 --- a/tests/agents/tools/test_read_document_tool.py +++ b/tests/agents/tools/test_read_document_tool.py @@ -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( diff --git a/tests/api/user/test_idempotency_decorator.py b/tests/api/user/test_idempotency_decorator.py index 35dfb9dd..cbed5f19 100644 --- a/tests/api/user/test_idempotency_decorator.py +++ b/tests/api/user/test_idempotency_decorator.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index 2d2e338b..876e6d89 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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() diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index ad2d670f..1125184f 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -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("\n").removesuffix("\n") + 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): diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py new file mode 100644 index 00000000..ca2f1330 --- /dev/null +++ b/tests/graphrag/test_graph_search_tool.py @@ -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 diff --git a/tests/graphrag/test_naming.py b/tests/graphrag/test_naming.py new file mode 100644 index 00000000..ff14f36a --- /dev/null +++ b/tests/graphrag/test_naming.py @@ -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") diff --git a/tests/graphrag/test_retriever_default_path.py b/tests/graphrag/test_retriever_default_path.py new file mode 100644 index 00000000..d50a2b16 --- /dev/null +++ b/tests/graphrag/test_retriever_default_path.py @@ -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 diff --git a/tests/graphrag/test_retriever_passages.py b/tests/graphrag/test_retriever_passages.py new file mode 100644 index 00000000..d2d9c29c --- /dev/null +++ b/tests/graphrag/test_retriever_passages.py @@ -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}) == {} diff --git a/tests/graphrag/test_retriever_seeding.py b/tests/graphrag/test_retriever_seeding.py new file mode 100644 index 00000000..0918aa4c --- /dev/null +++ b/tests/graphrag/test_retriever_seeding.py @@ -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() diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index d712e7fe..d9de877a 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -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) diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index e3fe2fae..5203eee1 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -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"] diff --git a/tests/test_celery.py b/tests/test_celery.py index c5b692df..e240871c 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -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 diff --git a/tests/test_dispatcher.py b/tests/test_dispatcher.py index 18816aa1..050664c1 100644 --- a/tests/test_dispatcher.py +++ b/tests/test_dispatcher.py @@ -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): diff --git a/tests/vectorstore/test_embeddings_delegated.py b/tests/vectorstore/test_embeddings_delegated.py index 060717ea..b9759510 100644 --- a/tests/vectorstore/test_embeddings_delegated.py +++ b/tests/vectorstore/test_embeddings_delegated.py @@ -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]]