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
332 lines
14 KiB
Python
332 lines
14 KiB
Python
import asyncio
|
|
import json
|
|
from typing import Optional, List, Type, Any, Union
|
|
|
|
from pydantic import BaseModel
|
|
|
|
from cognee.base_config import get_base_config
|
|
from cognee.modules.graph.cognee_graph.CogneeGraphElements import Edge
|
|
from cognee.modules.retrieval.utils.query_state import QueryState
|
|
from cognee.modules.retrieval.utils.validate_queries import validate_retriever_input
|
|
from cognee.shared.logging_utils import get_logger
|
|
from cognee.modules.retrieval.graph_completion_retriever import GraphCompletionRetriever
|
|
from cognee.modules.retrieval.utils.completion import (
|
|
batch_llm_completion,
|
|
generate_completion_batch,
|
|
)
|
|
from cognee.infrastructure.session.get_session_manager import get_session_manager
|
|
from cognee.context_global_variables import session_user
|
|
from cognee.infrastructure.llm.prompts import render_prompt, read_query_prompt
|
|
|
|
logger = get_logger()
|
|
|
|
|
|
def _as_answer_text(completion: Any) -> str:
|
|
"""Convert completion to human-readable text for validation and follow-up prompts."""
|
|
if isinstance(completion, str):
|
|
return completion
|
|
if isinstance(completion, BaseModel):
|
|
json_str = completion.model_dump_json(indent=2)
|
|
return f"[Structured Response]\n{json_str}"
|
|
try:
|
|
return json.dumps(completion, indent=2)
|
|
except TypeError:
|
|
return str(completion)
|
|
|
|
|
|
class GraphCompletionCotRetriever(GraphCompletionRetriever):
|
|
"""
|
|
Handles graph completion by generating responses based on a series of interactions with
|
|
a language model. This class extends from GraphCompletionRetriever and is designed to
|
|
manage the retrieval and validation process for user queries, integrating follow-up
|
|
questions based on reasoning. The public methods are:
|
|
|
|
- get_completion
|
|
|
|
Instance variables include:
|
|
- validation_system_prompt_path
|
|
- validation_user_prompt_path
|
|
- followup_system_prompt_path
|
|
- followup_user_prompt_path
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
user_prompt_path: str = "graph_context_for_question.txt",
|
|
system_prompt_path: str = "answer_simple_question.txt",
|
|
validation_user_prompt_path: str = "cot_validation_user_prompt.txt",
|
|
validation_system_prompt_path: str = "cot_validation_system_prompt.txt",
|
|
followup_system_prompt_path: str = "cot_followup_system_prompt.txt",
|
|
followup_user_prompt_path: str = "cot_followup_user_prompt.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,
|
|
max_iter: int = 4,
|
|
session_id: Optional[str] = None,
|
|
response_model: Type = str,
|
|
neighborhood_depth: Optional[int] = None,
|
|
neighborhood_seed_top_k: Optional[int] = 10,
|
|
include_references: bool = False,
|
|
):
|
|
super().__init__(
|
|
user_prompt_path=user_prompt_path,
|
|
system_prompt_path=system_prompt_path,
|
|
system_prompt=system_prompt,
|
|
top_k=top_k,
|
|
node_type=node_type,
|
|
node_name=node_name,
|
|
node_name_filter_operator=node_name_filter_operator,
|
|
wide_search_top_k=wide_search_top_k,
|
|
triplet_distance_penalty=triplet_distance_penalty,
|
|
feedback_influence=feedback_influence,
|
|
session_id=session_id,
|
|
response_model=response_model,
|
|
neighborhood_depth=neighborhood_depth,
|
|
neighborhood_seed_top_k=neighborhood_seed_top_k,
|
|
include_references=include_references,
|
|
)
|
|
self.validation_system_prompt_path = validation_system_prompt_path
|
|
self.validation_user_prompt_path = validation_user_prompt_path
|
|
self.followup_system_prompt_path = followup_system_prompt_path
|
|
self.followup_user_prompt_path = followup_user_prompt_path
|
|
self._cot_final_context = None
|
|
self.max_iter = max_iter
|
|
|
|
async def get_retrieved_objects(
|
|
self, query: Optional[str] = None, query_batch: Optional[List[str]] = None
|
|
) -> Union[List[Edge], List[List[Edge]]]:
|
|
"""
|
|
Run chain-of-thought completion with optional structured output.
|
|
|
|
Parameters:
|
|
-----------
|
|
- query: User query
|
|
|
|
Returns:
|
|
--------
|
|
- List of retrieved edges
|
|
"""
|
|
session_save = self._use_session_cache()
|
|
validate_retriever_input(query, query_batch, session_save)
|
|
|
|
conversation_history = ""
|
|
if session_save:
|
|
user = session_user.get()
|
|
user_id = getattr(user, "id", None)
|
|
if user_id:
|
|
sm = get_session_manager()
|
|
history = await sm.get_session(
|
|
user_id=str(user_id),
|
|
session_id=self.session_id,
|
|
formatted=True,
|
|
last_n=sm.session_history_last_n,
|
|
include_context=False,
|
|
)
|
|
conversation_history = history if isinstance(history, str) else ""
|
|
# Prepend the active session-context guidance block above the (history-only) CoT
|
|
# intermediate-round context. Gated + fail-open: never blocks CoT reasoning.
|
|
block = await self._maybe_active_context_block(sm, str(user_id), query)
|
|
if block:
|
|
conversation_history = (
|
|
block + "\n\n" + conversation_history if conversation_history else block
|
|
)
|
|
|
|
# Normalize single query to batch for uniform processing
|
|
effective_batch = [query] if query else query_batch
|
|
|
|
_completion, context_text, triplets = await self._run_cot_completion(
|
|
effective_batch, conversation_history, skip_final_completion=True
|
|
)
|
|
|
|
self._cot_final_context = context_text
|
|
|
|
if query:
|
|
return triplets[0]
|
|
return triplets
|
|
|
|
async def _maybe_active_context_block(self, sm, user_id: str, query: Optional[str]) -> str:
|
|
"""Render the active session-context block for the CoT intermediate rounds.
|
|
|
|
Gated on caching + auto_feedback. Fully fail-open: returns "" on any error or
|
|
when the layer is disabled, so it can never block chain-of-thought reasoning.
|
|
"""
|
|
try:
|
|
from cognee.infrastructure.databases.cache.config import CacheConfig
|
|
from cognee.infrastructure.session.session_context_builder import (
|
|
build_active_context_block,
|
|
)
|
|
|
|
cache_config = CacheConfig()
|
|
if not (cache_config.caching and cache_config.auto_feedback):
|
|
return ""
|
|
|
|
block, _served_ids = await build_active_context_block(
|
|
session_manager=sm,
|
|
user_id=user_id,
|
|
session_id=sm._resolve_session_id(self.session_id),
|
|
query=query or "",
|
|
)
|
|
return block or ""
|
|
except Exception as exc:
|
|
logger.warning("CoT active session-context block failed: %s", exc)
|
|
return ""
|
|
|
|
# -- CoT orchestrator --
|
|
|
|
async def _run_cot_completion(
|
|
self,
|
|
query_batch: List[str],
|
|
conversation_history: str = "",
|
|
skip_final_completion: bool = False,
|
|
) -> tuple[List[Any], List[str], List[List[Edge]]]:
|
|
"""
|
|
Run chain-of-thought completion with optional structured output.
|
|
|
|
Parameters:
|
|
-----------
|
|
- query_batch: Batch of user queries
|
|
- conversation_history: Optional conversation history string
|
|
- skip_final_completion: If True, do not run _generate_completions on the last
|
|
iteration; return [] for completions (final completion done in get_completion_from_context).
|
|
|
|
Returns:
|
|
--------
|
|
- completion_result: The generated completion (string or structured model), or [] if skip_final_completion.
|
|
- context_text: The resolved context text
|
|
- triplets: The list of triplets used
|
|
"""
|
|
states = {q: QueryState() for q in query_batch}
|
|
await self._fetch_initial_triplets_and_context(states)
|
|
await self._generate_completions(states, conversation_history)
|
|
|
|
for reasoning_iteration in range(self.max_iter):
|
|
followup_queries = await self._run_cot_round(states)
|
|
await self._merge_followup_triplets(states, followup_queries)
|
|
if not (skip_final_completion and reasoning_iteration == self.max_iter - 1):
|
|
await self._generate_completions(states, conversation_history)
|
|
|
|
return self._collect_results(states, query_batch, skip_final_completion)
|
|
|
|
# -- Helper methods called by the orchestrator --
|
|
|
|
async def _fetch_initial_triplets_and_context(self, states: dict):
|
|
"""Fetch triplets and resolve context text for all queries."""
|
|
queries = list(states.keys())
|
|
triplets_batch = await self.get_triplets_batch(queries)
|
|
context_batch = await asyncio.gather(
|
|
*[self.resolve_edges_to_text(t) for t in triplets_batch]
|
|
)
|
|
for q, triplets, context in zip(queries, triplets_batch, context_batch):
|
|
states[q].triplets = triplets
|
|
states[q].context_text = context
|
|
|
|
async def _generate_completions(self, states: dict, conversation_history: str):
|
|
"""Generate completions for all queries in parallel."""
|
|
queries = list(states.keys())
|
|
contexts = [states[q].context_text for q in queries]
|
|
completions = await generate_completion_batch(
|
|
query_batch=queries,
|
|
context=contexts,
|
|
user_prompt_path=self.user_prompt_path,
|
|
system_prompt_path=self.system_prompt_path,
|
|
system_prompt=self.system_prompt,
|
|
response_model=self.response_model,
|
|
conversation_history=conversation_history if conversation_history else None,
|
|
)
|
|
for q, comp in zip(queries, completions):
|
|
states[q].completion = comp
|
|
logger.info(f"Chain-of-thought: generated completions for {len(queries)} queries")
|
|
|
|
async def _run_cot_round(self, states: dict) -> List[str]:
|
|
"""Run one CoT round: validate answers, generate follow-up questions."""
|
|
validation_prompts, validation_system = self._build_validation_prompts(states)
|
|
reasoning_batch = await batch_llm_completion(validation_prompts, validation_system)
|
|
|
|
followup_prompts, followup_system = self._build_followup_prompts(states, reasoning_batch)
|
|
followup_questions = await batch_llm_completion(followup_prompts, followup_system)
|
|
|
|
logger.info(f"Chain-of-thought: follow-up questions: {followup_questions}")
|
|
return followup_questions
|
|
|
|
# -- Prompt builders --
|
|
|
|
def _build_cot_prompts(self, template_path, states, extras):
|
|
"""Build prompts with common query+answer fields and per-query extras."""
|
|
return [
|
|
render_prompt(
|
|
filename=template_path,
|
|
context={"query": q, "answer": _as_answer_text(states[q].completion), **extra},
|
|
)
|
|
for q, extra in zip(states.keys(), extras)
|
|
]
|
|
|
|
def _build_validation_prompts(self, states):
|
|
"""Build validation user prompts and load system prompt."""
|
|
system_prompt = read_query_prompt(prompt_file_name=self.validation_system_prompt_path)
|
|
user_prompts = self._build_cot_prompts(
|
|
self.validation_user_prompt_path,
|
|
states,
|
|
[{"context": s.context_text} for s in states.values()],
|
|
)
|
|
return user_prompts, system_prompt
|
|
|
|
def _build_followup_prompts(self, states, reasoning_batch):
|
|
"""Build followup user prompts and load system prompt."""
|
|
system_prompt = read_query_prompt(prompt_file_name=self.followup_system_prompt_path)
|
|
user_prompts = self._build_cot_prompts(
|
|
self.followup_user_prompt_path,
|
|
states,
|
|
[{"reasoning": r} for r in reasoning_batch],
|
|
)
|
|
return user_prompts, system_prompt
|
|
|
|
async def _merge_followup_triplets(self, states: dict, followup_questions: List[str]):
|
|
"""Fetch triplets for follow-up questions and merge with existing state."""
|
|
queries = list(states.keys())
|
|
new_triplets_batch = await self.get_triplets_batch(followup_questions)
|
|
|
|
for q, new_triplets in zip(queries, new_triplets_batch):
|
|
states[q].merge_triplets(new_triplets)
|
|
|
|
context_batch = await asyncio.gather(
|
|
*[self.resolve_edges_to_text(states[q].triplets) for q in queries]
|
|
)
|
|
for q, context in zip(queries, context_batch):
|
|
states[q].context_text = context
|
|
|
|
def _collect_results(
|
|
self,
|
|
states: dict,
|
|
query_batch: List[str],
|
|
skip_final_completion: bool = False,
|
|
) -> tuple[List[Any], List[str], List[List[Edge]]]:
|
|
"""Extract final completions, context texts, and triplets from states."""
|
|
completions = [] if skip_final_completion else [states[q].completion for q in query_batch]
|
|
contexts = [states[q].context_text for q in query_batch]
|
|
triplets = [states[q].triplets for q in query_batch]
|
|
return completions, contexts, triplets
|
|
|
|
async def get_context_from_objects(
|
|
self,
|
|
query: Optional[str] = None,
|
|
query_batch: Optional[List[str]] = None,
|
|
retrieved_objects=None,
|
|
) -> Union[str, List[str]]:
|
|
"""Return stored CoT final context when set; otherwise delegate to parent."""
|
|
cot_context = getattr(self, "_cot_final_context", None)
|
|
if cot_context is not None:
|
|
if query_batch:
|
|
return cot_context
|
|
return cot_context[0] if cot_context else ""
|
|
|
|
return await super().get_context_from_objects(
|
|
query=query,
|
|
query_batch=query_batch,
|
|
retrieved_objects=retrieved_objects,
|
|
)
|