diff --git a/src/memos/graph_dbs/neo4j_community.py b/src/memos/graph_dbs/neo4j_community.py index b5c92f40a..08c30c5a8 100644 --- a/src/memos/graph_dbs/neo4j_community.py +++ b/src/memos/graph_dbs/neo4j_community.py @@ -1038,36 +1038,45 @@ def delete_node_by_prams( f"[delete_node_by_prams] Deleting nodes - memory_ids: {memory_ids}, file_ids: {file_ids}, filter: {filter}" ) - # First count matching nodes to get accurate count - count_query = f"MATCH (n:Memory) WHERE {ids_where} RETURN count(n) AS node_count" - logger.info(f"[delete_node_by_prams] count_query: {count_query}") - print(f"[delete_node_by_prams] count_query: {count_query}") + # Collect IDs before deletion so we can purge vectors from vec_db afterwards. + id_collect_query = f"MATCH (n:Memory) WHERE {ids_where} RETURN n.id AS id" + logger.info("[delete_node_by_prams] id_collect_query: %s", id_collect_query) - # Then delete nodes + # Delete nodes delete_query = f"MATCH (n:Memory) WHERE {ids_where} DETACH DELETE n" logger.info(f"[delete_node_by_prams] delete_query: {delete_query}") - print(f"[delete_node_by_prams] delete_query: {delete_query}") - print(f"[delete_node_by_prams] params: {params}") deleted_count = 0 + collected_ids: list[str] = [] try: with self.driver.session(database=self.db_name) as session: - # Count nodes before deletion - count_result = session.run(count_query, **params) - count_record = count_result.single() - expected_count = 0 - if count_record: - expected_count = count_record["node_count"] or 0 - - # Delete nodes + # Collect IDs of nodes that are about to be deleted + id_result = session.run(id_collect_query, **params) + collected_ids = [record["id"] for record in id_result if record["id"] is not None] + deleted_count = len(collected_ids) + + # Delete nodes from graph session.run(delete_query, **params) - # Use the count from before deletion as the actual deleted count - deleted_count = expected_count except Exception as e: logger.error(f"[delete_node_by_prams] Failed to delete nodes: {e}", exc_info=True) raise + # Purge corresponding embedding vectors so they are no longer searchable. + # This is the fix for #2331: the graph node was removed but stale vectors + # in vec_db caused deleted memories to resurface in search results. + if collected_ids: + try: + self.vec_db.delete(collected_ids) + logger.info( + "[delete_node_by_prams] Purged %d vectors from vec_db", len(collected_ids) + ) + except Exception as e: + logger.warning( + "[delete_node_by_prams] vec_db cleanup failed (graph deletion already succeeded): %s", + e, + ) + logger.info(f"[delete_node_by_prams] Successfully deleted {deleted_count} nodes") return deleted_count diff --git a/tests/graph_dbs/test_delete_vec_cleanup.py b/tests/graph_dbs/test_delete_vec_cleanup.py new file mode 100644 index 000000000..f5c90c1cb --- /dev/null +++ b/tests/graph_dbs/test_delete_vec_cleanup.py @@ -0,0 +1,141 @@ +""" +Regression tests for #2331: delete_node_by_prams must also purge vectors from vec_db. + +When a memory is deleted via Neo4jCommunityGraphDB.delete_node_by_prams the +corresponding embedding vector must be removed from vec_db so that subsequent +searches no longer surface the deleted node. +""" + +from unittest.mock import MagicMock + +import pytest + +from memos.configs.graph_db import Neo4jGraphDBConfig + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_db(config: Neo4jGraphDBConfig) -> "Neo4jCommunityGraphDB": # noqa: F821 + """Build a Neo4jCommunityGraphDB with all heavy dependencies mocked out. + + Uses __new__ to skip __init__ entirely so no real Neo4j driver or Qdrant + connection is attempted. + """ + from memos.graph_dbs.neo4j_community import Neo4jCommunityGraphDB + + db = Neo4jCommunityGraphDB.__new__(Neo4jCommunityGraphDB) + db.config = config + db.driver = MagicMock() + db.db_name = config.db_name + db.vec_db = MagicMock() + db._schema_ready = True + return db + + +@pytest.fixture +def community_config(): + return Neo4jGraphDBConfig( + uri="bolt://localhost:7687", + user="neo4j", + password="test", + db_name="test_db", + auto_create=False, + use_multi_db=False, + user_name="alice", + embedding_dimension=3, + ) + + +# --------------------------------------------------------------------------- +# Tests - vec_db cleanup is called for every delete mode +# --------------------------------------------------------------------------- + + +class TestDeleteNodeByPramsVecCleanup: + """delete_node_by_prams must remove vectors from vec_db after graph deletion.""" + + def _make_session_mock(self, db: "Neo4jCommunityGraphDB", ids_to_delete: list[str]): # noqa: F821 + """Wire a session that returns `ids_to_delete` from the pre-delete ID query. + + The fixed implementation makes exactly 2 session.run calls: + 1) id_collect_query — MATCH ... RETURN n.id AS id + 2) delete_query — MATCH ... DETACH DELETE n + """ + session_ctx = MagicMock() + session_ctx.__enter__ = MagicMock(return_value=session_ctx) + session_ctx.__exit__ = MagicMock(return_value=False) + + # First run: id_collect_query — yields records with record["id"] == the node id + id_records = [] + for nid in ids_to_delete: + record = MagicMock() + # capture nid via default argument to avoid late-binding closure + record.__getitem__ = MagicMock( + side_effect=lambda k, _id=nid: _id if k == "id" else None + ) + id_records.append(record) + id_result = MagicMock() + id_result.__iter__ = MagicMock(return_value=iter(id_records)) + + # Second run: delete_query — return value is not inspected + delete_result = MagicMock() + + session_ctx.run.side_effect = [id_result, delete_result] + db.driver.session.return_value = session_ctx + return session_ctx + + def test_delete_by_memory_ids_cleans_vec_db(self, community_config): + """Deleting by memory_ids must call vec_db.delete with those IDs.""" + db = _make_db(community_config) + ids = ["aaa-111", "bbb-222"] + self._make_session_mock(db, ids) + + db.delete_node_by_prams(memory_ids=ids) + + db.vec_db.delete.assert_called_once_with(ids) + + def test_delete_by_filter_cleans_vec_db(self, community_config): + """Deleting by filter must call vec_db.delete with all matched IDs.""" + db = _make_db(community_config) + matched_ids = ["ccc-333", "ddd-444"] + + # get_by_metadata is called internally for filter path + db.get_by_metadata = MagicMock(return_value=matched_ids) + self._make_session_mock(db, matched_ids) + + db.delete_node_by_prams(filter={"user_id": "alice"}) + + db.get_by_metadata.assert_called_once() + db.vec_db.delete.assert_called_once_with(matched_ids) + + def test_delete_by_memory_ids_empty_list_no_vec_call(self, community_config): + """Empty memory_ids must not call vec_db.delete (early-return path).""" + db = _make_db(community_config) + + result = db.delete_node_by_prams(memory_ids=[]) + + db.vec_db.delete.assert_not_called() + assert result == 0 + + def test_delete_no_args_no_vec_call(self, community_config): + """No delete args must return 0 and not touch vec_db.""" + db = _make_db(community_config) + + result = db.delete_node_by_prams() + + db.vec_db.delete.assert_not_called() + assert result == 0 + + def test_vec_db_delete_failure_does_not_raise(self, community_config): + """A vec_db failure during cleanup must log a warning, not crash the request.""" + db = _make_db(community_config) + ids = ["eee-555"] + self._make_session_mock(db, ids) + db.vec_db.delete.side_effect = RuntimeError("qdrant unavailable") + + # Should not raise — graph deletion already succeeded + result = db.delete_node_by_prams(memory_ids=ids) + assert result >= 0