c889a57b6b
Test Suites / Build CI Environment (push) Has been cancelled
Test Suites / Basic Tests (push) Has been cancelled
Test Suites / End-to-End Tests (push) Has been cancelled
Test Suites / CLI Tests (push) Has been cancelled
Test Suites / Slow End-to-End Tests (push) Has been cancelled
Test Suites / Graph Database Tests (push) Has been cancelled
Test Suites / Vector DB Tests (push) Has been cancelled
Test Suites / Temporal Graph Test (push) Has been cancelled
Test Suites / Search Test on Different DBs (push) Has been cancelled
Test Suites / Example Tests (push) Has been cancelled
Test Suites / Notebook Tests (push) Has been cancelled
Test Suites / OS and Python Tests Ubuntu (push) Has been cancelled
Test Suites / OS and Python Tests Extended (push) Has been cancelled
Test Suites / LLM Test Suite (push) Has been cancelled
Test Suites / S3 File Storage Test (push) Has been cancelled
Test Suites / Run Integration Tests (push) Has been cancelled
Test Suites / MCP Tests (push) Has been cancelled
Test Suites / Docker Compose Test (push) Has been cancelled
Test Suites / Docker CI test (push) Has been cancelled
Test Suites / Relational DB Migration Tests (push) Has been cancelled
Test Suites / Distributed Cognee Test (push) Has been cancelled
Test Suites / DB Examples Tests (push) Has been cancelled
Test Suites / Test Completion Status (push) Has been cancelled
Test Suites / Claude Code Review (push) Has been cancelled
Test Suites / basic checks (push) Has been cancelled
build | Build and Push Cognee MCP Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
build | Build and Push Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.11) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.12) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (kuzu, kuzu) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (neo4j, neo4j) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Examples (push) Has been cancelled
Weighted Edges Tests / Code Quality for Weighted Edges (push) Has been cancelled
423 lines
17 KiB
Python
423 lines
17 KiB
Python
import asyncio
|
|
from typing import Any, Dict, List, Optional, Type, Union
|
|
|
|
from cognee.base_config import get_base_config
|
|
from cognee.infrastructure.engine import DataPoint
|
|
from cognee.modules.graph.cognee_graph.CogneeGraphElements import Edge
|
|
from cognee.modules.retrieval.utils.validate_queries import validate_retriever_input
|
|
from cognee.modules.graph.utils import resolve_edges_to_text
|
|
from cognee.modules.graph.utils.convert_node_to_data_point import get_all_subclasses
|
|
from cognee.modules.retrieval.base_retriever import BaseRetriever
|
|
from cognee.modules.retrieval.utils.brute_force_triplet_search import brute_force_triplet_search
|
|
from cognee.modules.retrieval.utils.global_context import (
|
|
format_global_context_prelude,
|
|
load_root_text,
|
|
search_top_global_context_summaries,
|
|
)
|
|
from cognee.modules.retrieval.utils.used_graph_elements import (
|
|
is_edge_list,
|
|
extract_from_edges,
|
|
)
|
|
from cognee.modules.retrieval.utils.references import append_answer_grounded_evidence
|
|
from cognee.modules.retrieval.utils.completion import (
|
|
generate_completion,
|
|
generate_completion_batch,
|
|
)
|
|
from cognee.infrastructure.session.get_session_manager import get_session_manager
|
|
from cognee.shared.logging_utils import get_logger
|
|
from cognee.infrastructure.databases.unified import get_unified_engine
|
|
from cognee.context_global_variables import session_user
|
|
from cognee.infrastructure.databases.cache.config import CacheConfig
|
|
|
|
logger = get_logger("GraphCompletionRetriever")
|
|
|
|
|
|
class GraphCompletionRetriever(BaseRetriever):
|
|
"""
|
|
Retriever for handling graph-based completion searches.
|
|
|
|
This class implements the retrieval pipeline by searching for graph triplets (get_retrieved_objects function),
|
|
resolving those triplets into human-readable text context (get_context_from_objects function), and generating
|
|
LLM completions using the retrieved graph data (get_completion_from_context function).
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
user_prompt_path: str = "graph_context_for_question.txt",
|
|
system_prompt_path: str = "answer_simple_question.txt",
|
|
system_prompt: Optional[str] = None,
|
|
top_k: Optional[int] = 5,
|
|
node_type: Optional[Type] = None,
|
|
node_name: Optional[List[str]] = None,
|
|
node_name_filter_operator: str = "OR",
|
|
wide_search_top_k: Optional[int] = 100,
|
|
triplet_distance_penalty: Optional[float] = 6.5,
|
|
feedback_influence: float = get_base_config().default_feedback_influence,
|
|
session_id: Optional[str] = None,
|
|
response_model: Type = str,
|
|
neighborhood_depth: Optional[int] = None,
|
|
neighborhood_seed_top_k: Optional[int] = 10,
|
|
include_global_context_index: bool = False,
|
|
global_context_index_top_k: int = 3,
|
|
include_references: bool = False,
|
|
):
|
|
"""Initialize retriever with prompt paths and search parameters."""
|
|
self.user_prompt_path = user_prompt_path
|
|
self.system_prompt_path = system_prompt_path
|
|
self.system_prompt = system_prompt
|
|
self.top_k = top_k if top_k is not None else 5
|
|
self.wide_search_top_k = wide_search_top_k
|
|
self.node_type = node_type
|
|
self.node_name = node_name
|
|
self.node_name_filter_operator = node_name_filter_operator
|
|
self.triplet_distance_penalty = triplet_distance_penalty
|
|
self.feedback_influence = feedback_influence
|
|
# session_id (Optional[str]): Identifier for managing conversation history.
|
|
self.session_id = session_id
|
|
# response_model (Type): The Pydantic model or type for the expected response.
|
|
self.response_model = response_model
|
|
self.neighborhood_depth = neighborhood_depth
|
|
self.neighborhood_seed_top_k = neighborhood_seed_top_k
|
|
self.include_global_context_index = include_global_context_index
|
|
self.global_context_index_top_k = global_context_index_top_k
|
|
self.include_references = include_references
|
|
|
|
def _use_session_cache(self) -> bool:
|
|
"""Check if session caching is enabled for the current user."""
|
|
user = session_user.get()
|
|
user_id = getattr(user, "id", None)
|
|
return bool(user_id and CacheConfig().caching)
|
|
|
|
@staticmethod
|
|
def _get_vector_index_collections() -> List[str]:
|
|
"""Collect vector index collection names from all DataPoint subclasses."""
|
|
collections = []
|
|
for subclass in get_all_subclasses(DataPoint):
|
|
metadata = subclass.model_fields.get("metadata")
|
|
if metadata is None:
|
|
continue
|
|
default = getattr(metadata, "default", None)
|
|
if isinstance(default, dict):
|
|
for field_name in default.get("index_fields", []):
|
|
collections.append(f"{subclass.__name__}_{field_name}")
|
|
return collections
|
|
|
|
async def get_retrieved_objects(
|
|
self, query: Optional[str] = None, query_batch: Optional[List[str]] = None
|
|
) -> Union[List[Edge], List[List[Edge]]]:
|
|
"""
|
|
Performs a brute-force triplet search on the graph and updates access timestamps.
|
|
|
|
Args:
|
|
query (str): The search query to find relevant graph triplets.
|
|
query_batch (str): The batch of search queries to find relevant graph triplets.
|
|
|
|
Returns:
|
|
List[Edge]: A list of retrieved Edge objects (triplets).
|
|
Returns an empty list if the graph is empty or no results are found.
|
|
"""
|
|
|
|
validate_retriever_input(query, query_batch, self._use_session_cache())
|
|
|
|
self._unified_engine = await get_unified_engine()
|
|
is_empty = await self._unified_engine.graph.is_empty()
|
|
|
|
if is_empty:
|
|
logger.warning("Search attempt on an empty knowledge graph")
|
|
return []
|
|
|
|
triplets = await self.get_triplets(query, query_batch)
|
|
|
|
# Check if all triplets are empty, in case of batch queries
|
|
if query_batch and all(len(batched_triplets) == 0 for batched_triplets in triplets):
|
|
logger.warning("Empty context was provided to the completion")
|
|
return []
|
|
|
|
if len(triplets) == 0:
|
|
logger.warning("Empty context was provided to the completion")
|
|
return []
|
|
|
|
return triplets
|
|
|
|
async def resolve_edges_to_text(self, retrieved_edges: list) -> str:
|
|
"""
|
|
Converts retrieved graph edges into a human-readable string format.
|
|
|
|
Parameters:
|
|
-----------
|
|
|
|
- retrieved_edges (list): A list of edges retrieved from the graph.
|
|
|
|
Returns:
|
|
--------
|
|
|
|
- str: A formatted string representation of the nodes and their connections.
|
|
"""
|
|
return await resolve_edges_to_text(retrieved_edges)
|
|
|
|
async def get_triplets(
|
|
self,
|
|
query: Optional[str] = None,
|
|
query_batch: Optional[List[str]] = None,
|
|
) -> Union[List[Edge], List[List[Edge]]]:
|
|
"""
|
|
Retrieves relevant graph triplets based on a query string.
|
|
|
|
Parameters:
|
|
-----------
|
|
|
|
- query (str): The query string used to search for relevant triplets in the graph.
|
|
|
|
Returns:
|
|
--------
|
|
|
|
- list: A list of found triplets that match the query.
|
|
"""
|
|
collections = self._get_vector_index_collections()
|
|
unified_engine = getattr(self, "_unified_engine", None)
|
|
return await brute_force_triplet_search(
|
|
query,
|
|
query_batch,
|
|
top_k=self.top_k,
|
|
collections=collections or None,
|
|
node_type=self.node_type,
|
|
node_name=self.node_name,
|
|
node_name_filter_operator=self.node_name_filter_operator,
|
|
wide_search_top_k=self.wide_search_top_k,
|
|
triplet_distance_penalty=self.triplet_distance_penalty,
|
|
feedback_influence=self.feedback_influence,
|
|
unified_engine=unified_engine,
|
|
neighborhood_depth=self.neighborhood_depth,
|
|
neighborhood_seed_top_k=self.neighborhood_seed_top_k,
|
|
)
|
|
|
|
async def get_triplets_batch(
|
|
self,
|
|
queries: List[str],
|
|
) -> List[List[Edge]]:
|
|
"""
|
|
Retrieves triplets for a list of queries, using single-query mode when
|
|
possible to enable ID-filtered graph projection.
|
|
|
|
When there is only one query, delegates to single-query mode (query=)
|
|
which computes relevant node IDs and filters the graph projection.
|
|
For multiple queries, uses batch mode (query_batch=).
|
|
|
|
Returns:
|
|
List[List[Edge]]: One list of edges per query.
|
|
"""
|
|
if len(queries) == 1:
|
|
triplets = await self.get_triplets(query=queries[0])
|
|
return [triplets]
|
|
return await self.get_triplets(query_batch=queries)
|
|
|
|
async def get_context_from_objects(
|
|
self,
|
|
query: Optional[str] = None,
|
|
query_batch: Optional[List[str]] = None,
|
|
retrieved_objects=None,
|
|
) -> Union[str, List[str]]:
|
|
"""
|
|
Transforms raw retrieved graph triplets into a textual context string.
|
|
|
|
Args:
|
|
query (str): The original search query.
|
|
query_batch (List[str]): The batch of original search queries.
|
|
retrieved_objects (List[Edge]): The raw triplets returned from the search.
|
|
Output of the get_retrieved_objects method.
|
|
|
|
Returns:
|
|
str: A string representing the resolved graph context.
|
|
Returns an empty list (as string) if no triplets are provided.
|
|
|
|
Note: To avoid duplicate retrievals, ensure that retrieved_objects
|
|
are provided from get_retrieved_objects method call.
|
|
"""
|
|
|
|
triplets = retrieved_objects
|
|
|
|
if query_batch:
|
|
# Check if all triplets are empty, in case of batch queries
|
|
if not triplets or all(len(batched_triplets) == 0 for batched_triplets in triplets):
|
|
logger.warning("Empty context was provided to the completion")
|
|
return ["" for _ in query_batch]
|
|
|
|
return await asyncio.gather(
|
|
*[self.resolve_edges_to_text(batched_triplets) for batched_triplets in triplets]
|
|
)
|
|
|
|
graph_context = await self.resolve_edges_to_text(triplets) if triplets else ""
|
|
|
|
if not self.include_global_context_index:
|
|
if not triplets:
|
|
logger.warning("Empty context was provided to the completion")
|
|
return ""
|
|
return graph_context
|
|
|
|
prelude = await self._build_global_context_prelude(query)
|
|
if not prelude and not graph_context:
|
|
logger.warning("Empty context was provided to the completion")
|
|
return ""
|
|
if not prelude:
|
|
return graph_context
|
|
if not graph_context:
|
|
return prelude
|
|
return f"{prelude}\n\n{graph_context}"
|
|
|
|
async def _build_global_context_prelude(self, query: Optional[str]) -> str:
|
|
if not query:
|
|
return ""
|
|
if getattr(self, "_unified_engine", None) is None:
|
|
self._unified_engine = await get_unified_engine()
|
|
root_text = await load_root_text()
|
|
top_summaries = await search_top_global_context_summaries(
|
|
query, self.global_context_index_top_k, self._unified_engine.vector
|
|
)
|
|
return format_global_context_prelude(root_text, top_summaries)
|
|
|
|
def _extract_context_object_ids(self, retrieved_objects: Any) -> Optional[Dict[str, List[str]]]:
|
|
"""Extract node_ids and edge_ids from list of Edge. Only used for single-query session path."""
|
|
if not isinstance(retrieved_objects, list) or not retrieved_objects:
|
|
return None
|
|
if not is_edge_list(retrieved_objects):
|
|
return None
|
|
return extract_from_edges(retrieved_objects)
|
|
|
|
def _completion_kwargs(self, context: str) -> dict:
|
|
"""Common kwargs for completion calls (no session)."""
|
|
return {
|
|
"context": context,
|
|
"user_prompt_path": self.user_prompt_path,
|
|
"system_prompt_path": self.system_prompt_path,
|
|
"system_prompt": self.system_prompt,
|
|
"response_model": self.response_model,
|
|
}
|
|
|
|
async def _generate_completion_without_session(
|
|
self,
|
|
query: Optional[str],
|
|
query_batch: Optional[List[str]],
|
|
context: str,
|
|
) -> List[Any]:
|
|
"""Generate completion(s) without session; returns list of completions."""
|
|
kwargs = self._completion_kwargs(context)
|
|
if query_batch:
|
|
return await generate_completion_batch(query_batch=query_batch, **kwargs)
|
|
completion = await generate_completion(query=query, **kwargs)
|
|
return [completion]
|
|
|
|
async def _append_graph_evidence(self, completions: List[Any]) -> List[Any]:
|
|
"""Append an answer-grounded chunk Evidence block to string completions.
|
|
|
|
Each answer is run as a vector query against the chunk index, so the
|
|
Evidence bullets reflect where the answer text is grounded in the corpus
|
|
rather than which graph elements happened to be retrieved. Evidence is
|
|
appended only when references are enabled and the completion is a plain
|
|
string (never corrupt a structured response_model); search failures
|
|
degrade to no Evidence.
|
|
"""
|
|
return await append_answer_grounded_evidence(
|
|
completions,
|
|
enabled=self.include_references and self.response_model is str,
|
|
)
|
|
|
|
async def get_completion_from_context(
|
|
self,
|
|
query: Optional[str] = None,
|
|
query_batch: Optional[List[str]] = None,
|
|
retrieved_objects: Optional[List[Edge]] = None,
|
|
context: str = None,
|
|
effective_query: Optional[str] = None,
|
|
turn_preparation=None,
|
|
) -> List[Any]:
|
|
"""
|
|
Generates an LLM response based on the query, context, and conversation history.
|
|
Optionally saves the interaction and updates the session cache.
|
|
|
|
Args:
|
|
query (str): The user's question or prompt.
|
|
query_batch (List[str]): The batch of user queries.
|
|
retrieved_objects (Optional[List[Edge]]): Raw triplets used for interaction mapping.
|
|
Output of get_retrieved_objects method.
|
|
context (str): The text-resolved graph context.
|
|
Output of the get_context_from_objects method.
|
|
|
|
Returns:
|
|
List[Any]: A list containing the generated response (completion).
|
|
|
|
Note: To avoid duplicate retrievals, ensure that retrieved_objects and context
|
|
are provided from previous method calls.
|
|
"""
|
|
use_session = self._use_session_cache() and not query_batch
|
|
if use_session:
|
|
sm = get_session_manager()
|
|
used_graph_element_ids = self._extract_context_object_ids(retrieved_objects)
|
|
completion = await sm.generate_completion_with_session(
|
|
session_id=self.session_id,
|
|
query=query,
|
|
context=context,
|
|
user_prompt_path=self.user_prompt_path,
|
|
system_prompt_path=self.system_prompt_path,
|
|
system_prompt=self.system_prompt,
|
|
response_model=self.response_model,
|
|
summarize_context=False,
|
|
used_graph_element_ids=used_graph_element_ids,
|
|
max_context_chars=getattr(self, "max_context_chars", None),
|
|
effective_query=effective_query,
|
|
turn_preparation=turn_preparation,
|
|
)
|
|
completions = [completion]
|
|
else:
|
|
completions = await self._generate_completion_without_session(
|
|
query, query_batch, context
|
|
)
|
|
|
|
# Session and non-session branches rejoin here so every variant that calls
|
|
# this method (including via super()) appends references once. Evidence is
|
|
# grounded in each completion's own text, so a cache-hit answer never
|
|
# cites chunks that share nothing with it.
|
|
return await self._append_graph_evidence(completions)
|
|
|
|
async def get_completion(
|
|
self, query: Optional[str] = None, query_batch: Optional[List[str]] = None
|
|
) -> List[Any]:
|
|
"""
|
|
Generates a final output or answer based on the query and retrieved context.
|
|
|
|
Args:
|
|
query (str): The original user query.
|
|
query_batch (List[str]): The batch of user queries.
|
|
|
|
Returns:
|
|
List[Any]: A list containing the generated completions or response objects.
|
|
"""
|
|
validate_retriever_input(query, query_batch)
|
|
|
|
effective_query = query
|
|
turn_preparation = None
|
|
if query is not None and not query_batch:
|
|
turn_preparation = await self.prepare_session_turn_for_retrieval(query)
|
|
if not turn_preparation.should_answer:
|
|
return [turn_preparation.response_to_user or "Got it."]
|
|
effective_query = turn_preparation.effective_query or query
|
|
|
|
retrieved_objects = await self.get_retrieved_objects(
|
|
query=effective_query,
|
|
query_batch=query_batch,
|
|
)
|
|
context = await self.get_context_from_objects(
|
|
query=effective_query,
|
|
query_batch=query_batch,
|
|
retrieved_objects=retrieved_objects,
|
|
)
|
|
completion = await self.get_completion_from_context(
|
|
query=query,
|
|
query_batch=query_batch,
|
|
retrieved_objects=retrieved_objects,
|
|
context=context,
|
|
effective_query=effective_query,
|
|
turn_preparation=turn_preparation,
|
|
)
|
|
|
|
return completion
|