项目文件夹

文件
wehub-resource-sync 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
chore: import upstream snapshot with attribution
2026-07-13 13:02:24 +08:00

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,
)