Skip to content
Closed
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
43 changes: 26 additions & 17 deletions src/memos/graph_dbs/neo4j_community.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
141 changes: 141 additions & 0 deletions tests/graph_dbs/test_delete_vec_cleanup.py
Original file line number Diff line number Diff line change
@@ -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
Loading