topoteretes--cognee
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
167 行
7.0 KiB
Python
167 行
7.0 KiB
Python
from typing import Any, Dict, List, Optional, Type
|
|
|
|
from cognee.shared.logging_utils import get_logger
|
|
from cognee.infrastructure.databases.vector import get_vector_engine_async
|
|
from cognee.modules.retrieval.utils.completion import generate_completion
|
|
from cognee.infrastructure.session.get_session_manager import get_session_manager
|
|
from cognee.modules.retrieval.base_retriever import BaseRetriever
|
|
from cognee.modules.retrieval.utils.used_graph_elements import extract_from_scored_results
|
|
from cognee.modules.retrieval.exceptions.exceptions import NoDataError
|
|
from cognee.infrastructure.databases.vector.exceptions import CollectionNotFoundError
|
|
from cognee.context_global_variables import session_user
|
|
from cognee.infrastructure.databases.cache.config import CacheConfig
|
|
from cognee.modules.retrieval.utils.references import append_chunk_evidence
|
|
|
|
logger = get_logger("CompletionRetriever")
|
|
|
|
|
|
class CompletionRetriever(BaseRetriever):
|
|
"""
|
|
Retriever for handling LLM-based completion searches.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
user_prompt_path: str = "context_for_question.txt",
|
|
system_prompt_path: str = "answer_simple_question.txt",
|
|
system_prompt: Optional[str] = None,
|
|
top_k: Optional[int] = 1,
|
|
session_id: Optional[str] = None,
|
|
response_model: Type = str,
|
|
include_references: bool = False,
|
|
):
|
|
"""Initialize retriever with optional custom prompt paths."""
|
|
self.user_prompt_path = user_prompt_path
|
|
self.system_prompt_path = system_prompt_path
|
|
self.top_k = top_k if top_k is not None else 1
|
|
self.system_prompt = system_prompt
|
|
self.session_id = session_id
|
|
self.response_model = response_model
|
|
self.include_references = include_references
|
|
|
|
async def get_retrieved_objects(self, query: str) -> Any:
|
|
vector_engine = await get_vector_engine_async()
|
|
|
|
try:
|
|
found_chunks = await vector_engine.search(
|
|
"DocumentChunk_text", query, limit=self.top_k, include_payload=True
|
|
)
|
|
|
|
return found_chunks
|
|
except CollectionNotFoundError as error:
|
|
logger.error("DocumentChunk_text collection not found")
|
|
raise NoDataError("No data found in the system, please add data first.") from error
|
|
|
|
def _extract_context_object_ids(self, retrieved_objects: Any) -> Optional[Dict[str, List[str]]]:
|
|
"""Extract node_ids from ScoredResult-like list for session QA."""
|
|
if isinstance(retrieved_objects, list) and retrieved_objects:
|
|
return extract_from_scored_results(retrieved_objects)
|
|
return None
|
|
|
|
async def get_context_from_objects(self, query: str, retrieved_objects: Any) -> str:
|
|
"""
|
|
Retrieves relevant document chunks as context.
|
|
|
|
Fetches document chunks based on a query from a vector engine and combines their text.
|
|
Returns empty string if no chunks are found. Raises NoDataError if the collection is not
|
|
found.
|
|
|
|
Parameters:
|
|
-----------
|
|
|
|
- query (str): The query string used to search for relevant document chunks.
|
|
|
|
Returns:
|
|
--------
|
|
|
|
- str: A string containing the combined text of the retrieved document chunks, or an
|
|
empty string if none are found.
|
|
"""
|
|
if retrieved_objects:
|
|
# Combine all chunks text returned from vector search (number of chunks is determined by top_k)
|
|
chunks_payload = [found_chunk.payload["text"] for found_chunk in retrieved_objects]
|
|
combined_context = "\n".join(chunks_payload)
|
|
return combined_context
|
|
return ""
|
|
|
|
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: str, context: str) -> List[Any]:
|
|
"""Generate completion without session; returns list of one completion."""
|
|
kwargs = self._completion_kwargs(context)
|
|
completion = await generate_completion(query=query, **kwargs)
|
|
return [completion]
|
|
|
|
async def get_completion_from_context(
|
|
self,
|
|
query: str,
|
|
retrieved_objects: Any,
|
|
context: Optional[Any] = None,
|
|
effective_query: Optional[str] = None,
|
|
turn_preparation=None,
|
|
) -> List[Any]:
|
|
"""
|
|
Generates an LLM completion using the context.
|
|
|
|
Retrieves context if not provided and generates a completion based on the query and
|
|
context using an external completion generator.
|
|
|
|
Parameters:
|
|
-----------
|
|
|
|
- query (str): The query string to be used for generating a completion.
|
|
- context (Optional[Any]): Optional pre-fetched context to use for generating the
|
|
completion; if None, it retrieves the context for the query. (default None)
|
|
- session_id (Optional[str]): Optional session identifier for caching. If None,
|
|
defaults to 'default_session'. (default None)
|
|
- response_model (Type): The Pydantic model type for structured output. (default str)
|
|
|
|
Returns:
|
|
--------
|
|
|
|
- Any: The generated completion based on the provided query and context.
|
|
"""
|
|
cache_config = CacheConfig()
|
|
user = session_user.get()
|
|
user_id = getattr(user, "id", None)
|
|
use_session = user_id and cache_config.caching
|
|
|
|
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, context)
|
|
|
|
# Both the session/cache branch and the non-session branch rejoin here so
|
|
# logged-in/cached calls also receive references. Evidence is grounded in
|
|
# each completion's own text, so a cache-hit answer never cites chunks
|
|
# that share nothing with it.
|
|
return append_chunk_evidence(
|
|
completions,
|
|
retrieved_objects,
|
|
enabled=self.include_references and self.response_model is str,
|
|
)
|