Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 31 additions & 3 deletions cognee/modules/retrieval/summaries_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,12 +20,35 @@ class SummariesRetriever(BaseRetriever):

Instance variables:
- top_k: int - Number of top summaries to retrieve.
- node_name: Optional[List[str]] - Node set names the search is scoped to.
- node_name_filter_operator: str - How multiple node_name values combine.
"""

def __init__(self, top_k: int = 5, session_id: Optional[str] = None):
"""Initialize retriever with search parameters."""
def __init__(
self,
top_k: int = 5,
session_id: Optional[str] = None,
node_name: Optional[List[str]] = None,
node_name_filter_operator: str = "OR",
):
"""
Initialize retriever with search parameters.

Parameters:
-----------

- top_k (int): Maximum number of summaries to retrieve. Defaults to 5.
- session_id (Optional[str]): Session the search belongs to. Defaults to None.
- node_name (Optional[List[str]]): Node set names used to filter summaries by
their belongs_to_set relationship. Defaults to None, which applies no node
set filtering.
- node_name_filter_operator (str): Logical operator used when applying multiple
node_name filters, such as "OR" or "AND". Defaults to "OR".
"""
self.top_k = top_k
self.session_id = session_id
self.node_name = node_name
self.node_name_filter_operator = node_name_filter_operator

async def get_retrieved_objects(self, query: str) -> Any:
"""
Expand Down Expand Up @@ -53,7 +76,12 @@ async def get_retrieved_objects(self, query: str) -> Any:

try:
summaries_results = await vector_engine.search(
"TextSummary_text", query, limit=self.top_k, include_payload=True
"TextSummary_text",
query,
limit=self.top_k,
include_payload=True,
node_name=self.node_name,
node_name_filter_operator=self.node_name_filter_operator,
)
logger.info(f"Found {len(summaries_results)} summaries from vector search")

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,15 @@ async def get_search_type_retriever_instance(
CodeRetriever,
{"config": retriever_specific_config},
),
SearchType.SUMMARIES: (SummariesRetriever, {"top_k": top_k, "session_id": session_id}),
SearchType.SUMMARIES: (
SummariesRetriever,
{
"top_k": top_k,
"session_id": session_id,
"node_name": node_name,
"node_name_filter_operator": node_name_filter_operator,
},
),
SearchType.CHUNKS: (
ChunksRetriever,
{
Expand Down
78 changes: 76 additions & 2 deletions cognee/tests/unit/modules/retrieval/summaries_retriever_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,12 @@ async def test_get_context_success(mock_vector_engine):
assert completion[0]["text"] == "S.R."
assert completion[1]["text"] == "M.B."
mock_vector_engine.search.assert_awaited_once_with(
"TextSummary_text", "test query", limit=5, include_payload=True
"TextSummary_text",
"test query",
limit=5,
include_payload=True,
node_name=None,
node_name_filter_operator="OR",
)


Expand Down Expand Up @@ -101,7 +106,12 @@ async def test_get_objects_top_k_limit(mock_vector_engine):

assert len(objects) == 3
mock_vector_engine.search.assert_awaited_once_with(
"TextSummary_text", "test query", limit=3, include_payload=True
"TextSummary_text",
"test query",
limit=3,
include_payload=True,
node_name=None,
node_name_filter_operator="OR",
)


Expand Down Expand Up @@ -145,6 +155,8 @@ async def test_init_defaults():
retriever = SummariesRetriever()

assert retriever.top_k == 5
assert retriever.node_name is None
assert retriever.node_name_filter_operator == "OR"


@pytest.mark.asyncio
Expand Down Expand Up @@ -194,3 +206,65 @@ async def test_get_completion_with_session_id(mock_vector_engine):

assert len(completion) == 1
assert completion[0]["text"] == "S.R."


@pytest.mark.asyncio
async def test_get_objects_forwards_nodeset_filter_to_vector_search(mock_vector_engine):
"""Test that node_name filtering is passed through to the vector engine."""
mock_vector_engine.search.return_value = []

retriever = SummariesRetriever(
top_k=30,
node_name=["KEN", "src_type:figure"],
node_name_filter_operator="AND",
)

with patch(
"cognee.modules.retrieval.summaries_retriever.get_unified_engine",
return_value=_make_unified_mock(mock_vector_engine),
):
await retriever.get_retrieved_objects("land cover")

mock_vector_engine.search.assert_awaited_once_with(
"TextSummary_text",
"land cover",
limit=30,
include_payload=True,
node_name=["KEN", "src_type:figure"],
node_name_filter_operator="AND",
)


@pytest.mark.asyncio
async def test_scoped_search_that_matches_nothing_returns_no_summaries(mock_vector_engine):
"""A node set with no summaries is an empty result, not a missing collection.

NoDataError says the system holds no data, which stays reserved for the
CollectionNotFoundError path; a filter that matches nothing must not borrow it.
"""
mock_vector_engine.search.return_value = []

retriever = SummariesRetriever(node_name=["empty-node-set"])

with patch(
"cognee.modules.retrieval.summaries_retriever.get_unified_engine",
return_value=_make_unified_mock(mock_vector_engine),
):
objects = await retriever.get_retrieved_objects("test query")
context = await retriever.get_context_from_objects("test query", objects)
completion = await retriever.get_completion_from_context("test query", objects, context)

assert objects == []
assert context == ""
assert completion == []


@pytest.mark.asyncio
async def test_init_custom_nodeset_filter():
"""Test SummariesRetriever initialization with node set filtering."""
retriever = SummariesRetriever(
node_name=["tenant-a", "tenant-b"], node_name_filter_operator="AND"
)

assert retriever.node_name == ["tenant-a", "tenant-b"]
assert retriever.node_name_filter_operator == "AND"
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,25 @@ async def test_chunks_retriever_receives_nodeset_filter_arguments():
assert retriever_instance.node_name_filter_operator == "AND"


@pytest.mark.asyncio
async def test_summaries_retriever_receives_nodeset_filter_arguments():
import cognee.modules.search.methods.get_search_type_retriever_instance as mod
from cognee.modules.retrieval.summaries_retriever import SummariesRetriever

retriever_instance = await mod.get_search_type_retriever_instance(
SearchType.SUMMARIES,
query_text="land cover",
top_k=30,
node_name=["KEN", "src_type:figure"],
node_name_filter_operator="AND",
)

assert isinstance(retriever_instance, SummariesRetriever)
assert retriever_instance.top_k == 30
assert retriever_instance.node_name == ["KEN", "src_type:figure"]
assert retriever_instance.node_name_filter_operator == "AND"


@pytest.mark.asyncio
async def test_rag_completion_retriever_receives_nodeset_filter_arguments():
import cognee.modules.search.methods.get_search_type_retriever_instance as mod
Expand Down
Loading