diff --git a/src/memos/mem_scheduler/memory_manage_modules/memory_filter.py b/src/memos/mem_scheduler/memory_manage_modules/memory_filter.py index 25b9a98f3..a11ee88ea 100644 --- a/src/memos/mem_scheduler/memory_manage_modules/memory_filter.py +++ b/src/memos/mem_scheduler/memory_manage_modules/memory_filter.py @@ -234,10 +234,6 @@ def filter_unrelated_and_redundant_memories( logger.info("No query history provided - keeping all memories") return memories, True - if len(memories) <= 1: - logger.info("Only one memory - no filtering needed") - return memories, True - logger.info( f"Starting combined unrelated and redundant filtering for {len(memories)} memories against {len(query_history)} queries" ) diff --git a/tests/mem_scheduler/test_retriever.py b/tests/mem_scheduler/test_retriever.py index 35c8b7f3a..a40203316 100644 --- a/tests/mem_scheduler/test_retriever.py +++ b/tests/mem_scheduler/test_retriever.py @@ -360,3 +360,54 @@ def test_filter_unrelated_memories_conservative_filtering(self): # Should return all memories self.assertEqual(result, memories) self.assertTrue(success_flag) + self.assertTrue(self.llm.generate.called) + + def test_combined_filtering_still_filters_a_single_unrelated_memory(self): + """A lone memory must still go through unrelated filtering. + + The combined filter used to return early on `len(memories) <= 1`, a + guard that only makes sense for redundancy (one memory cannot be + redundant with itself). Unrelated filtering is per-memory, so skipping + it left an off-topic memory sitting in working memory. + """ + query_history = ["What is my deployment pipeline?"] + memories = [TextualMemoryItem(memory="The user's cat is named Whiskers")] + + self.llm.generate.return_value = json.dumps( + { + "kept_memories": [], + "unrelated_removed_count": 1, + "redundant_removed_count": 0, + "reasoning": "No semantic connection to the deployment query", + } + ) + + result, success_flag = self.retriever.filter_unrelated_and_redundant_memories( + query_history=query_history, memories=memories + ) + + self.assertEqual(result, []) + self.assertTrue(success_flag) + # The LLM must actually be consulted, not short-circuited. + self.assertTrue(self.llm.generate.called) + + def test_combined_filtering_keeps_a_single_relevant_memory(self): + """The counterpart: one on-topic memory must survive the filter.""" + query_history = ["What is my deployment pipeline?"] + memories = [TextualMemoryItem(memory="Deployment uses a blue-green pipeline")] + + self.llm.generate.return_value = json.dumps( + { + "kept_memories": [0], + "unrelated_removed_count": 0, + "redundant_removed_count": 0, + "reasoning": "Directly answers the deployment query", + } + ) + + result, success_flag = self.retriever.filter_unrelated_and_redundant_memories( + query_history=query_history, memories=memories + ) + + self.assertEqual(result, memories) + self.assertTrue(success_flag)