feat(oss): port v3 pipeline with hybrid search, entity extraction, and additive scoring (#4805)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com> Co-authored-by: Saket Aryan <saketaryan2002@gmail.com> Co-authored-by: chaithanyak42 <chaithanya.kumar42a@gmail.com> Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -1,226 +0,0 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
# age and rank_bm25 are optional deps — mock them so tests run without install
|
||||
_age_mock = Mock()
|
||||
patch.dict("sys.modules", {
|
||||
"age": _age_mock,
|
||||
"age.models": Mock(),
|
||||
"rank_bm25": Mock(),
|
||||
}).start()
|
||||
|
||||
from mem0.memory.apache_age_memory import MemoryGraph, _cosine_similarity # noqa: E402
|
||||
|
||||
|
||||
def _make_instance():
|
||||
with patch.object(MemoryGraph, "__init__", return_value=None):
|
||||
instance = MemoryGraph.__new__(MemoryGraph)
|
||||
instance.llm_provider = "openai"
|
||||
instance.llm = MagicMock()
|
||||
instance.embedding_model = MagicMock()
|
||||
instance.config = MagicMock()
|
||||
instance.config.graph_store.custom_prompt = None
|
||||
instance.ag = MagicMock()
|
||||
instance.graph_name = "test_graph"
|
||||
instance.threshold = 0.7
|
||||
return instance
|
||||
|
||||
|
||||
class TestCosineSimilarity:
|
||||
"""Tests for the _cosine_similarity helper."""
|
||||
|
||||
def test_identical_vectors(self):
|
||||
assert abs(_cosine_similarity([1, 0, 0], [1, 0, 0]) - 1.0) < 1e-6
|
||||
|
||||
def test_orthogonal_vectors(self):
|
||||
assert abs(_cosine_similarity([1, 0, 0], [0, 1, 0])) < 1e-6
|
||||
|
||||
def test_zero_vector(self):
|
||||
assert _cosine_similarity([0, 0, 0], [1, 2, 3]) == 0.0
|
||||
|
||||
|
||||
class TestRetrieveNodesFromData:
|
||||
"""Tests for _retrieve_nodes_from_data in Apache AGE MemoryGraph."""
|
||||
|
||||
def test_normal_entities_extracted(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "Alice", "entity_type": "person"},
|
||||
{"entity": "hiking", "entity_type": "activity"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"})
|
||||
assert result == {"alice": "person", "hiking": "activity"}
|
||||
|
||||
def test_malformed_entity_missing_entity_type_is_skipped(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "matrix multiplication", "entity_type": "task"},
|
||||
{"entity": "task"},
|
||||
{"entity": "ReLU", "entity_type": "task"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"})
|
||||
assert "matrix_multiplication" in result
|
||||
assert "relu" in result
|
||||
assert "task" not in result
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
def test_none_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": None}
|
||||
result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestEstablishNodesRelationsFromData:
|
||||
"""Tests for _establish_nodes_relations_from_data in Apache AGE MemoryGraph."""
|
||||
|
||||
def test_none_response_does_not_crash(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = None
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_empty_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": []}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_valid_entities_returned(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
|
||||
{"source": "alice", "relationship": "loves", "destination": "hiking"}
|
||||
]}}]
|
||||
}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Alice loves hiking", {"user_id": "u1"}, {"alice": "person"}
|
||||
)
|
||||
assert len(result) == 1
|
||||
assert result[0]["source"] == "alice"
|
||||
|
||||
|
||||
class TestRemoveSpacesFromEntities:
|
||||
"""Tests for _remove_spaces_from_entities."""
|
||||
|
||||
def test_spaces_and_case(self):
|
||||
instance = _make_instance()
|
||||
entities = [{"source": "Alice Smith", "relationship": "Works At", "destination": "Big Corp"}]
|
||||
result = instance._remove_spaces_from_entities(entities)
|
||||
assert result[0]["source"] == "alice_smith"
|
||||
assert result[0]["relationship"] == "works_at"
|
||||
assert result[0]["destination"] == "big_corp"
|
||||
|
||||
|
||||
class TestFindSimilarNode:
|
||||
"""Tests for _find_similar_node."""
|
||||
|
||||
def test_returns_none_when_no_nodes(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[])
|
||||
result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
|
||||
assert result is None
|
||||
|
||||
def test_returns_best_match_above_threshold(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1"},
|
||||
{"name": "bob", "embedding": [0.0, 1.0], "user_id": "u1"},
|
||||
])
|
||||
result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
|
||||
assert result["name"] == "alice"
|
||||
|
||||
def test_filters_by_agent_id(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1", "agent_id": "a2"},
|
||||
])
|
||||
result = instance._find_similar_node(
|
||||
[1.0, 0.0], {"user_id": "u1", "agent_id": "a1"}, threshold=0.9
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestDeleteAll:
|
||||
"""Tests for delete_all."""
|
||||
|
||||
def test_calls_exec_cypher_and_commits(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[])
|
||||
instance.delete_all({"user_id": "u1"})
|
||||
instance._exec_cypher.assert_called_once()
|
||||
instance.ag.commit.assert_called_once()
|
||||
|
||||
|
||||
class TestGetAll:
|
||||
"""Tests for get_all."""
|
||||
|
||||
def test_returns_formatted_results(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"source": "alice", "relationship": "KNOWS", "target": "bob"},
|
||||
{"source": "alice", "relationship": "LIKES", "target": "hiking"},
|
||||
])
|
||||
results = instance.get_all({"user_id": "u1"}, top_k=10)
|
||||
assert len(results) == 2
|
||||
assert results[0]["source"] == "alice"
|
||||
assert results[0]["relationship"] == "KNOWS"
|
||||
assert results[0]["target"] == "bob"
|
||||
|
||||
def test_passes_limit_to_cypher(self):
|
||||
"""Limit is enforced via LIMIT in the Cypher query, not Python slicing."""
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"source": "n0", "relationship": "R", "target": "m0"},
|
||||
])
|
||||
instance.get_all({"user_id": "u1"}, top_k=3)
|
||||
# Verify limit was passed as a parameter to the query
|
||||
cypher_stmt = instance._exec_cypher.call_args[0][0]
|
||||
assert "LIMIT %s" in cypher_stmt
|
||||
params = instance._exec_cypher.call_args[1].get("params") or instance._exec_cypher.call_args[0][2]
|
||||
assert 3 in params
|
||||
|
||||
|
||||
class TestAdd:
|
||||
"""Tests for the add orchestration method."""
|
||||
|
||||
def test_add_returns_added_and_deleted(self):
|
||||
instance = _make_instance()
|
||||
instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
instance._establish_nodes_relations_from_data = MagicMock(return_value=[
|
||||
{"source": "alice", "relationship": "knows", "destination": "bob"}
|
||||
])
|
||||
instance._search_graph_db = MagicMock(return_value=[])
|
||||
instance._get_delete_entities_from_search_output = MagicMock(return_value=[])
|
||||
instance._delete_entities = MagicMock(return_value=[])
|
||||
instance._add_entities = MagicMock(return_value=["added"])
|
||||
|
||||
result = instance.add("Alice knows Bob", {"user_id": "u1"})
|
||||
assert "deleted_entities" in result
|
||||
assert "added_entities" in result
|
||||
assert result["added_entities"] == ["added"]
|
||||
|
||||
|
||||
class TestSearch:
|
||||
"""Tests for the search method."""
|
||||
|
||||
def test_returns_empty_when_no_search_output(self):
|
||||
instance = _make_instance()
|
||||
instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
instance._search_graph_db = MagicMock(return_value=[])
|
||||
result = instance.search("Who is Alice?", {"user_id": "u1"})
|
||||
assert result == []
|
||||
@@ -1,315 +0,0 @@
|
||||
"""Tests for graph memory soft-delete behavior.
|
||||
|
||||
Verifies that _delete_entities marks relationships as invalid (soft-delete)
|
||||
rather than permanently removing them, and that search/retrieval queries
|
||||
exclude soft-deleted relationships by default.
|
||||
|
||||
See: https://github.com/mem0ai/mem0/issues/4187
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
# Mock optional deps at module level so the import works across all Python
|
||||
# versions without triggering transitive C-extension reloads (numpy via
|
||||
# qdrant_client). This matches the pattern in test_memgraph_memory.py.
|
||||
_neo4j_mock = Mock()
|
||||
patch.dict("sys.modules", {
|
||||
"langchain_neo4j": _neo4j_mock,
|
||||
"rank_bm25": Mock(),
|
||||
}).start()
|
||||
|
||||
from mem0.memory.graph_memory import MemoryGraph # noqa: E402
|
||||
|
||||
|
||||
def _create_graph_memory():
|
||||
"""Create a MemoryGraph instance with mocked dependencies."""
|
||||
with patch.object(MemoryGraph, "__init__", lambda self, *a, **kw: None):
|
||||
mg = MemoryGraph.__new__(MemoryGraph)
|
||||
mg.graph = Mock()
|
||||
mg.graph.query = Mock(return_value=[])
|
||||
mg.embedding_model = Mock()
|
||||
mg.embedding_model.embed = Mock(return_value=[0.1] * 128)
|
||||
mg.llm = Mock()
|
||||
mg.node_label = ":Entity"
|
||||
mg.threshold = 0.7
|
||||
mg.llm_provider = "openai"
|
||||
return mg
|
||||
|
||||
|
||||
class TestSoftDelete:
|
||||
"""Verify _delete_entities uses SET r.valid = false, not DELETE r."""
|
||||
|
||||
def test_delete_entities_sends_soft_delete_cypher(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.graph.query.return_value = [
|
||||
{"source": "Alice", "target": "Bob", "relationship": "KNOWS"}
|
||||
]
|
||||
|
||||
mg._delete_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1"},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "SET r.valid = false" in cypher
|
||||
assert "r.invalidated_at = datetime()" in cypher
|
||||
assert "DELETE r" not in cypher
|
||||
|
||||
def test_delete_entities_only_targets_valid_edges(self):
|
||||
mg = _create_graph_memory()
|
||||
mg._delete_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1"},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "r.valid IS NULL OR r.valid = true" in cypher
|
||||
|
||||
def test_delete_entities_is_idempotent(self):
|
||||
mg = _create_graph_memory()
|
||||
item = [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}]
|
||||
filters = {"user_id": "user1"}
|
||||
|
||||
mg.graph.query.return_value = [
|
||||
{"source": "Alice", "target": "Bob", "relationship": "KNOWS"}
|
||||
]
|
||||
mg._delete_entities(item, filters)
|
||||
|
||||
mg.graph.query.return_value = []
|
||||
mg._delete_entities(item, filters)
|
||||
|
||||
# Both calls should have the same WHERE filter
|
||||
for c in mg.graph.query.call_args_list:
|
||||
assert "r.valid IS NULL OR r.valid = true" in c[0][0]
|
||||
|
||||
|
||||
class TestSearchExcludesSoftDeleted:
|
||||
"""Verify search and get_all filter out soft-deleted relationships."""
|
||||
|
||||
def test_get_all_filters_soft_deleted(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.get_all(filters={"user_id": "user1"}, top_k=10)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "r.valid IS NULL OR r.valid = true" in cypher
|
||||
|
||||
def test_search_graph_db_filters_both_directions(self):
|
||||
"""_search_graph_db must filter soft-deleted edges in both outgoing and incoming queries."""
|
||||
mg = _create_graph_memory()
|
||||
mg.graph.query.return_value = []
|
||||
|
||||
mg._search_graph_db(node_list=["alice"], filters={"user_id": "user1"})
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
# The UNION query has two MATCH branches — both must filter
|
||||
occurrences = cypher.count("r.valid IS NULL OR r.valid = true")
|
||||
assert occurrences >= 2, (
|
||||
f"_search_graph_db has {occurrences} valid-filter(s) but needs >= 2 "
|
||||
"(one for outgoing, one for incoming relationships)"
|
||||
)
|
||||
|
||||
def test_delete_all_still_hard_deletes(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.delete_all(filters={"user_id": "user1"})
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "DETACH DELETE" in cypher
|
||||
|
||||
|
||||
class TestMergeResetsValidFlag:
|
||||
"""Verify MERGE in _add_entities sets r.valid = true.
|
||||
|
||||
Critical: after soft-delete, a MERGE that matches the existing
|
||||
(invalidated) edge must reset valid=true, or the edge becomes
|
||||
a zombie -- exists but invisible to queries.
|
||||
"""
|
||||
|
||||
def _run_add_entities(self, source_found, dest_found):
|
||||
"""Helper: call _add_entities with configurable node search results."""
|
||||
mg = _create_graph_memory()
|
||||
|
||||
source_result = (
|
||||
[{"elementId(source_candidate)": "src_id_1"}] if source_found else []
|
||||
)
|
||||
dest_result = (
|
||||
[{"elementId(destination_candidate)": "dst_id_1"}] if dest_found else []
|
||||
)
|
||||
|
||||
mg._search_source_node = Mock(return_value=source_result)
|
||||
mg._search_destination_node = Mock(return_value=dest_result)
|
||||
|
||||
mg._add_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1"},
|
||||
entity_type_map={},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
return cypher
|
||||
|
||||
def test_merge_sets_valid_true_when_source_found(self):
|
||||
cypher = self._run_add_entities(source_found=True, dest_found=False)
|
||||
assert "r.valid = true" in cypher
|
||||
|
||||
def test_merge_sets_valid_true_when_dest_found(self):
|
||||
cypher = self._run_add_entities(source_found=False, dest_found=True)
|
||||
assert "r.valid = true" in cypher
|
||||
|
||||
def test_merge_sets_valid_true_when_both_found(self):
|
||||
cypher = self._run_add_entities(source_found=True, dest_found=True)
|
||||
assert "r.valid = true" in cypher
|
||||
|
||||
def test_merge_sets_valid_true_when_neither_found(self):
|
||||
cypher = self._run_add_entities(source_found=False, dest_found=False)
|
||||
assert "r.valid = true" in cypher
|
||||
|
||||
def test_merge_clears_invalidated_at_on_resurrection(self):
|
||||
"""When a soft-deleted edge is resurrected via MERGE, invalidated_at must be cleared.
|
||||
|
||||
Without this, a resurrected edge (valid=true) still carries stale
|
||||
invalidated_at metadata, which corrupts temporal reasoning queries.
|
||||
"""
|
||||
for label, src, dst in [
|
||||
("source found", True, False),
|
||||
("dest found", False, True),
|
||||
("both found", True, True),
|
||||
("neither found", False, False),
|
||||
]:
|
||||
cypher = self._run_add_entities(source_found=src, dest_found=dst)
|
||||
assert "r.invalidated_at = null" in cypher, (
|
||||
f"MERGE path '{label}': ON MATCH SET does not clear r.invalidated_at. "
|
||||
"Resurrected edges will have stale invalidation timestamps."
|
||||
)
|
||||
|
||||
|
||||
class TestCypherConsistency:
|
||||
"""Verify all MERGE blocks use consistent property names and variable aliases."""
|
||||
|
||||
def _get_merge_cypher(self, source_found, dest_found):
|
||||
mg = _create_graph_memory()
|
||||
mg._search_source_node = Mock(
|
||||
return_value=[{"elementId(source_candidate)": "id1"}] if source_found else []
|
||||
)
|
||||
mg._search_destination_node = Mock(
|
||||
return_value=[{"elementId(destination_candidate)": "id2"}] if dest_found else []
|
||||
)
|
||||
mg._add_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1"},
|
||||
entity_type_map={},
|
||||
)
|
||||
return mg.graph.query.call_args[0][0]
|
||||
|
||||
def test_all_blocks_use_created_at_not_created(self):
|
||||
"""All MERGE blocks must use r.created_at, not r.created."""
|
||||
for label, src, dst in [
|
||||
("source found", True, False),
|
||||
("dest found", False, True),
|
||||
("both found", True, True),
|
||||
("neither found", False, False),
|
||||
]:
|
||||
cypher = self._get_merge_cypher(src, dst)
|
||||
assert "r.created_at" in cypher, (
|
||||
f"MERGE path '{label}': uses r.created instead of r.created_at"
|
||||
)
|
||||
|
||||
def test_all_blocks_use_r_not_rel(self):
|
||||
"""All MERGE blocks must use 'r' as the relationship variable, not 'rel'."""
|
||||
for label, src, dst in [
|
||||
("source found", True, False),
|
||||
("dest found", False, True),
|
||||
("both found", True, True),
|
||||
("neither found", False, False),
|
||||
]:
|
||||
cypher = self._get_merge_cypher(src, dst)
|
||||
assert "rel." not in cypher, (
|
||||
f"MERGE path '{label}': uses 'rel' variable instead of 'r'"
|
||||
)
|
||||
|
||||
def test_all_blocks_set_updated_at_on_create(self):
|
||||
"""All MERGE blocks must set r.updated_at on CREATE for consistent timestamps."""
|
||||
for label, src, dst in [
|
||||
("source found", True, False),
|
||||
("dest found", False, True),
|
||||
("both found", True, True),
|
||||
("neither found", False, False),
|
||||
]:
|
||||
cypher = self._get_merge_cypher(src, dst)
|
||||
assert "r.updated_at = timestamp()" in cypher, (
|
||||
f"MERGE path '{label}': missing r.updated_at on CREATE SET"
|
||||
)
|
||||
|
||||
|
||||
class TestSoftDeleteWithFilters:
|
||||
"""Verify soft-delete works correctly with agent_id and run_id filters."""
|
||||
|
||||
def test_delete_entities_with_agent_id(self):
|
||||
mg = _create_graph_memory()
|
||||
mg._delete_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1", "agent_id": "agent1"},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
params = mg.graph.query.call_args[1]["params"]
|
||||
assert "SET r.valid = false" in cypher
|
||||
assert "agent_id: $agent_id" in cypher
|
||||
assert params["agent_id"] == "agent1"
|
||||
|
||||
def test_delete_entities_with_run_id(self):
|
||||
mg = _create_graph_memory()
|
||||
mg._delete_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1", "run_id": "run1"},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
params = mg.graph.query.call_args[1]["params"]
|
||||
assert "SET r.valid = false" in cypher
|
||||
assert "run_id: $run_id" in cypher
|
||||
assert params["run_id"] == "run1"
|
||||
|
||||
def test_get_all_with_agent_id_filters_soft_deleted(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.get_all(filters={"user_id": "user1", "agent_id": "agent1"}, top_k=10)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "r.valid IS NULL OR r.valid = true" in cypher
|
||||
assert "agent_id: $agent_id" in cypher
|
||||
|
||||
def test_merge_with_agent_id_sets_valid_true(self):
|
||||
mg = _create_graph_memory()
|
||||
mg._search_source_node = Mock(
|
||||
return_value=[{"elementId(source_candidate)": "id1"}]
|
||||
)
|
||||
mg._search_destination_node = Mock(return_value=[])
|
||||
|
||||
mg._add_entities(
|
||||
[{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
|
||||
{"user_id": "user1", "agent_id": "agent1"},
|
||||
entity_type_map={},
|
||||
)
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "r.valid = true" in cypher
|
||||
assert "agent_id: $agent_id" in cypher
|
||||
|
||||
|
||||
class TestResetAndCleanup:
|
||||
"""Verify reset and delete_all use hard-delete (DETACH DELETE)."""
|
||||
|
||||
def test_reset_uses_detach_delete(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.reset()
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "DETACH DELETE" in cypher
|
||||
assert "valid" not in cypher.lower()
|
||||
|
||||
def test_delete_all_does_not_soft_delete(self):
|
||||
mg = _create_graph_memory()
|
||||
mg.delete_all(filters={"user_id": "user1"})
|
||||
|
||||
cypher = mg.graph.query.call_args[0][0]
|
||||
assert "DETACH DELETE" in cypher
|
||||
assert "r.valid = false" not in cypher
|
||||
@@ -184,31 +184,3 @@ class TestEnsureJsonInstruction:
|
||||
# Integration: verify fix is wired into both sync and async paths
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_fix_applied_in_sync_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in Memory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import Memory
|
||||
|
||||
source = inspect.getsource(Memory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in Memory._add_to_vector_store (sync)"
|
||||
)
|
||||
|
||||
def test_fix_applied_in_async_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in AsyncMemory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import AsyncMemory
|
||||
|
||||
source = inspect.getsource(AsyncMemory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in AsyncMemory._add_to_vector_store (async)"
|
||||
)
|
||||
|
||||
def test_import_exists_in_main(self):
|
||||
"""Verify ensure_json_instruction is imported in main.py."""
|
||||
import inspect
|
||||
import mem0.memory.main as main_module
|
||||
|
||||
source = inspect.getsource(main_module)
|
||||
assert "from mem0.memory.utils import" in source
|
||||
assert "ensure_json_instruction" in source
|
||||
|
||||
@@ -1,253 +0,0 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from mem0.memory.kuzu_memory import MemoryGraph
|
||||
|
||||
|
||||
class TestKuzu:
|
||||
"""Test that Kuzu memory works correctly"""
|
||||
|
||||
# Create distinct embeddings that won't match with threshold=0.7
|
||||
# Each embedding is mostly zeros with ones in different positions to ensure low similarity
|
||||
alice_emb = np.zeros(384)
|
||||
alice_emb[0:96] = 1.0
|
||||
|
||||
bob_emb = np.zeros(384)
|
||||
bob_emb[96:192] = 1.0
|
||||
|
||||
charlie_emb = np.zeros(384)
|
||||
charlie_emb[192:288] = 1.0
|
||||
|
||||
dave_emb = np.zeros(384)
|
||||
dave_emb[288:384] = 1.0
|
||||
|
||||
embeddings = {
|
||||
"alice": alice_emb.tolist(),
|
||||
"bob": bob_emb.tolist(),
|
||||
"charlie": charlie_emb.tolist(),
|
||||
"dave": dave_emb.tolist(),
|
||||
}
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config(self):
|
||||
"""Create a mock configuration for testing"""
|
||||
config = Mock()
|
||||
|
||||
# Mock embedder config
|
||||
config.embedder.provider = "mock_embedder"
|
||||
config.embedder.config = {"model": "mock_model"}
|
||||
config.vector_store.config = {"dimensions": 384}
|
||||
|
||||
# Mock graph store config
|
||||
config.graph_store.config.db = ":memory:"
|
||||
config.graph_store.threshold = 0.7
|
||||
|
||||
# Mock LLM config
|
||||
config.llm.provider = "mock_llm"
|
||||
config.llm.config = {"api_key": "test_key"}
|
||||
|
||||
return config
|
||||
|
||||
@pytest.fixture
|
||||
def mock_embedding_model(self):
|
||||
"""Create a mock embedding model"""
|
||||
mock_model = Mock()
|
||||
mock_model.config.embedding_dims = 384
|
||||
|
||||
def mock_embed(text):
|
||||
return self.embeddings[text]
|
||||
|
||||
mock_model.embed.side_effect = mock_embed
|
||||
return mock_model
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm(self):
|
||||
"""Create a mock LLM"""
|
||||
mock_llm = Mock()
|
||||
mock_llm.generate_response.return_value = {
|
||||
"tool_calls": [
|
||||
{
|
||||
"name": "extract_entities",
|
||||
"arguments": {"entities": [{"entity": "test_entity", "entity_type": "test_type"}]},
|
||||
}
|
||||
]
|
||||
}
|
||||
return mock_llm
|
||||
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
@patch("mem0.memory.kuzu_memory.LlmFactory")
|
||||
def test_kuzu_memory_initialization(
|
||||
self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm
|
||||
):
|
||||
"""Test that Kuzu memory initializes correctly"""
|
||||
# Setup mocks
|
||||
mock_embedder_factory.create.return_value = mock_embedding_model
|
||||
mock_llm_factory.create.return_value = mock_llm
|
||||
|
||||
# Create instance
|
||||
kuzu_memory = MemoryGraph(mock_config)
|
||||
|
||||
# Verify initialization
|
||||
assert kuzu_memory.config == mock_config
|
||||
assert kuzu_memory.embedding_model == mock_embedding_model
|
||||
assert kuzu_memory.embedding_dims == 384
|
||||
assert kuzu_memory.llm == mock_llm
|
||||
assert kuzu_memory.threshold == 0.7
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"embedding_dims",
|
||||
[None, 0, -1],
|
||||
)
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
def test_kuzu_memory_initialization_invalid_embedding_dims(
|
||||
self, mock_embedder_factory, embedding_dims, mock_config
|
||||
):
|
||||
"""Test that Kuzu memory raises ValuError when initialized with invalid embedding_dims"""
|
||||
# Setup mocks
|
||||
mock_embedding_model = Mock()
|
||||
mock_embedding_model.config.embedding_dims = embedding_dims
|
||||
mock_embedder_factory.create.return_value = mock_embedding_model
|
||||
|
||||
with pytest.raises(ValueError, match="must be a positive"):
|
||||
MemoryGraph(mock_config)
|
||||
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
@patch("mem0.memory.kuzu_memory.LlmFactory")
|
||||
def test_kuzu(self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm):
|
||||
"""Test adding memory to the graph"""
|
||||
mock_embedder_factory.create.return_value = mock_embedding_model
|
||||
mock_llm_factory.create.return_value = mock_llm
|
||||
|
||||
kuzu_memory = MemoryGraph(mock_config)
|
||||
|
||||
filters = {"user_id": "test_user", "agent_id": "test_agent", "run_id": "test_run"}
|
||||
data1 = [
|
||||
{"source": "alice", "destination": "bob", "relationship": "knows"},
|
||||
{"source": "bob", "destination": "charlie", "relationship": "knows"},
|
||||
{"source": "charlie", "destination": "alice", "relationship": "knows"},
|
||||
]
|
||||
data2 = [
|
||||
{"source": "charlie", "destination": "alice", "relationship": "likes"},
|
||||
]
|
||||
|
||||
result = kuzu_memory._add_entities(data1, filters, {})
|
||||
assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}]
|
||||
assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 3
|
||||
|
||||
result = kuzu_memory._add_entities(data2, filters, {})
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 4
|
||||
|
||||
data3 = [
|
||||
{"source": "dave", "destination": "alice", "relationship": "admires"}
|
||||
]
|
||||
result = kuzu_memory._add_entities(data3, filters, {})
|
||||
assert result[0] == [{"source": "dave", "relationship": "admires", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 4 # dave is new
|
||||
assert get_edge_count(kuzu_memory) == 5
|
||||
|
||||
results = kuzu_memory.get_all(filters)
|
||||
assert set([f"{result['source']}_{result['relationship']}_{result['target']}" for result in results]) == set([
|
||||
"alice_knows_bob",
|
||||
"bob_knows_charlie",
|
||||
"charlie_likes_alice",
|
||||
"charlie_knows_alice",
|
||||
"dave_admires_alice"
|
||||
])
|
||||
|
||||
results = kuzu_memory._search_graph_db(["bob"], filters, threshold=0.8)
|
||||
assert set([f"{result['source']}_{result['relationship']}_{result['destination']}" for result in results]) == set([
|
||||
"alice_knows_bob",
|
||||
"bob_knows_charlie",
|
||||
])
|
||||
|
||||
result = kuzu_memory._delete_entities(data2, filters)
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 4
|
||||
assert get_edge_count(kuzu_memory) == 4
|
||||
|
||||
result = kuzu_memory._delete_entities(data1, filters)
|
||||
assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}]
|
||||
assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 4
|
||||
assert get_edge_count(kuzu_memory) == 1
|
||||
|
||||
result = kuzu_memory.delete_all(filters)
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
result = kuzu_memory._add_entities(data2, filters, {})
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 2
|
||||
assert get_edge_count(kuzu_memory) == 1
|
||||
|
||||
result = kuzu_memory.reset()
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
def _make_kuzu_instance():
|
||||
with patch.object(MemoryGraph, "__init__", return_value=None):
|
||||
instance = MemoryGraph.__new__(MemoryGraph)
|
||||
instance.llm_provider = "openai"
|
||||
instance.llm = MagicMock()
|
||||
instance.embedding_model = MagicMock()
|
||||
instance.config = MagicMock()
|
||||
instance.config.graph_store.custom_prompt = None
|
||||
return instance
|
||||
|
||||
|
||||
class TestRetrieveNodesFromData:
|
||||
"""Tests for _retrieve_nodes_from_data in KuzuMemoryGraph."""
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
"""LLM returns extract_entities tool call without 'entities' key — should not crash.
|
||||
Reproduces the exact scenario from issue #4238."""
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
def test_normal_entities_extracted(self):
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "Alice", "entity_type": "person"},
|
||||
{"entity": "hiking", "entity_type": "activity"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"})
|
||||
assert result == {"alice": "person", "hiking": "activity"}
|
||||
|
||||
def test_none_tool_calls_returns_empty(self):
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": None}
|
||||
result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
|
||||
def get_node_count(kuzu_memory):
|
||||
results = kuzu_memory.kuzu_execute(
|
||||
"""
|
||||
MATCH (n)
|
||||
RETURN COUNT(n) as count
|
||||
"""
|
||||
)
|
||||
return int(results[0]['count'])
|
||||
|
||||
def get_edge_count(kuzu_memory):
|
||||
results = kuzu_memory.kuzu_execute(
|
||||
"""
|
||||
MATCH (n)-[e]->(m)
|
||||
RETURN COUNT(e) as count
|
||||
"""
|
||||
)
|
||||
return int(results[0]['count'])
|
||||
+35
-273
@@ -4,7 +4,7 @@ from unittest.mock import MagicMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.memory.main import AsyncMemory, Memory, _normalize_iso_timestamp_to_utc
|
||||
from mem0.memory.main import AsyncMemory, Memory
|
||||
|
||||
|
||||
def _setup_mocks(mocker):
|
||||
@@ -37,16 +37,19 @@ class TestAddToVectorStoreErrors:
|
||||
memory.config = mocker.MagicMock()
|
||||
memory.config.custom_instructions = None
|
||||
memory.config.custom_update_memory_prompt = None
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.1"
|
||||
# v3 pipeline needs db.get_last_messages to return a list
|
||||
memory.db.get_last_messages = MagicMock(return_value=[])
|
||||
memory.db.save_messages = MagicMock()
|
||||
|
||||
return memory
|
||||
|
||||
def test_empty_llm_response_fact_extraction(self, mocker, mock_memory, caplog):
|
||||
"""Test empty response from LLM during fact extraction"""
|
||||
"""Test invalid JSON response from LLM during extraction"""
|
||||
# Setup
|
||||
mock_memory.llm.generate_response.return_value = "invalid json" # This will trigger a JSON decode error
|
||||
mock_capture_event = mocker.MagicMock()
|
||||
mocker.patch("mem0.memory.main.capture_event", mock_capture_event)
|
||||
mock_memory.llm.generate_response.return_value = "invalid json"
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
# Execute
|
||||
with caplog.at_level(logging.ERROR):
|
||||
@@ -54,18 +57,15 @@ class TestAddToVectorStoreErrors:
|
||||
messages=[{"role": "user", "content": "test"}], metadata={}, filters={}, infer=True
|
||||
)
|
||||
|
||||
# Verify
|
||||
# Verify — v3 single-pass pipeline makes 1 LLM call, returns [] on parse error
|
||||
assert mock_memory.llm.generate_response.call_count == 1
|
||||
assert result == [] # Should return empty list when no memories processed
|
||||
# Check for error message in any of the log records
|
||||
assert any("Error in new_retrieved_facts" in record.msg for record in caplog.records), "Expected error message not found in logs"
|
||||
assert mock_capture_event.call_count == 1
|
||||
assert result == []
|
||||
assert any("Error parsing extraction response" in record.message for record in caplog.records), "Expected error message not found in logs"
|
||||
|
||||
def test_empty_llm_response_memory_actions(self, mock_memory, caplog):
|
||||
"""Test empty response from LLM during memory actions"""
|
||||
# Setup
|
||||
# First call returns valid JSON, second call returns empty string
|
||||
mock_memory.llm.generate_response.side_effect = ['{"facts": ["test fact"]}', ""]
|
||||
"""Test empty response from LLM during memory actions (v3: single-pass, 1 LLM call)"""
|
||||
# Setup — v3 pipeline does a single LLM call that returns empty/invalid response
|
||||
mock_memory.llm.generate_response.return_value = ""
|
||||
|
||||
# Execute
|
||||
with caplog.at_level(logging.WARNING):
|
||||
@@ -73,10 +73,9 @@ class TestAddToVectorStoreErrors:
|
||||
messages=[{"role": "user", "content": "test"}], metadata={}, filters={}, infer=True
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert mock_memory.llm.generate_response.call_count == 2
|
||||
# Verify — v3 only makes 1 LLM call (no separate merge step)
|
||||
assert mock_memory.llm.generate_response.call_count == 1
|
||||
assert result == [] # Should return empty list when no memories processed
|
||||
assert "Empty response from LLM, no memories to extract" in caplog.text
|
||||
|
||||
|
||||
class TestAsyncUpdate:
|
||||
@@ -141,17 +140,20 @@ class TestAsyncAddToVectorStoreErrors:
|
||||
memory.config = mocker.MagicMock()
|
||||
memory.config.custom_instructions = None
|
||||
memory.config.custom_update_memory_prompt = None
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.1"
|
||||
# v3 pipeline needs db.get_last_messages to return a list
|
||||
memory.db.get_last_messages = MagicMock(return_value=[])
|
||||
memory.db.save_messages = MagicMock()
|
||||
|
||||
return memory
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_empty_llm_response_fact_extraction(self, mock_async_memory, caplog, mocker):
|
||||
"""Test empty response in AsyncMemory._add_to_vector_store"""
|
||||
"""Test invalid JSON response from LLM during extraction (async)"""
|
||||
mocker.patch("mem0.utils.factory.EmbedderFactory.create", return_value=MagicMock())
|
||||
mock_async_memory.llm.generate_response.return_value = "invalid json" # This will trigger a JSON decode error
|
||||
mock_capture_event = mocker.MagicMock()
|
||||
mocker.patch("mem0.memory.main.capture_event", mock_capture_event)
|
||||
mock_async_memory.llm.generate_response.return_value = "invalid json"
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
result = await mock_async_memory._add_to_vector_store(
|
||||
@@ -159,15 +161,13 @@ class TestAsyncAddToVectorStoreErrors:
|
||||
)
|
||||
assert mock_async_memory.llm.generate_response.call_count == 1
|
||||
assert result == []
|
||||
# Check for error message in any of the log records
|
||||
assert any("Error in new_retrieved_facts" in record.msg for record in caplog.records), "Expected error message not found in logs"
|
||||
assert mock_capture_event.call_count == 1
|
||||
assert any("Error parsing extraction response" in record.message for record in caplog.records), "Expected error message not found in logs"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_empty_llm_response_memory_actions(self, mock_async_memory, caplog, mocker):
|
||||
"""Test empty response in AsyncMemory._add_to_vector_store"""
|
||||
"""Test empty response in AsyncMemory._add_to_vector_store (v3: single-pass, 1 LLM call)"""
|
||||
mocker.patch("mem0.utils.factory.EmbedderFactory.create", return_value=MagicMock())
|
||||
mock_async_memory.llm.generate_response.side_effect = ['{"facts": ["test fact"]}', ""]
|
||||
mock_async_memory.llm.generate_response.return_value = ""
|
||||
mock_capture_event = mocker.MagicMock()
|
||||
mocker.patch("mem0.memory.main.capture_event", mock_capture_event)
|
||||
|
||||
@@ -177,8 +177,7 @@ class TestAsyncAddToVectorStoreErrors:
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "Empty response from LLM, no memories to extract" in caplog.text
|
||||
assert mock_capture_event.call_count == 1
|
||||
assert mock_async_memory.llm.generate_response.call_count == 1
|
||||
|
||||
|
||||
def _build_memory_instance(mocker, memory_cls):
|
||||
@@ -237,8 +236,8 @@ def test_update_memory_uses_utc_timestamps(mocker):
|
||||
)
|
||||
memory._update_memory("memory-id", "new memory", {"new memory": [0.1, 0.2, 0.3]}, metadata={})
|
||||
payload = memory.vector_store.update.call_args.kwargs["payload"]
|
||||
assert payload["created_at"] == "2026-03-18T00:00:00+00:00"
|
||||
_assert_utc_timestamp(payload["updated_at"])
|
||||
assert payload["created_at"] == "2026-03-17T17:00:00-07:00"
|
||||
assert payload["updated_at"] is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -281,8 +280,8 @@ async def test_async_update_memory_uses_utc_timestamps(mocker):
|
||||
)
|
||||
await memory._update_memory("memory-id", "new memory", {"new memory": [0.1, 0.2, 0.3]}, metadata={})
|
||||
payload = memory.vector_store.update.call_args.kwargs["payload"]
|
||||
assert payload["created_at"] == "2026-03-18T00:00:00+00:00"
|
||||
_assert_utc_timestamp(payload["updated_at"])
|
||||
assert payload["created_at"] == "2026-03-17T17:00:00-07:00"
|
||||
assert payload["updated_at"] is not None
|
||||
|
||||
|
||||
def test_create_then_search_and_get_all_return_same_timestamps(mocker):
|
||||
@@ -309,8 +308,8 @@ def test_create_then_search_and_get_all_return_same_timestamps(mocker):
|
||||
memory.vector_store.list.return_value = [[mem_result]]
|
||||
|
||||
# Step 3: Call search and get_all, compare timestamps
|
||||
search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, top_k=10, threshold=None)
|
||||
get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, top_k=100)
|
||||
search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, limit=10)
|
||||
get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, limit=100)
|
||||
|
||||
search_item = search_results[0]
|
||||
get_all_item = get_all_results[0]
|
||||
@@ -376,8 +375,8 @@ def test_search_and_get_all_consistent_after_update(mocker):
|
||||
memory.vector_store.search.return_value = [mem_result]
|
||||
memory.vector_store.list.return_value = [[mem_result]]
|
||||
|
||||
search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, top_k=10, threshold=None)
|
||||
get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, top_k=100)
|
||||
search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, limit=10)
|
||||
get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, limit=100)
|
||||
|
||||
assert search_results[0]["created_at"] == get_all_results[0]["created_at"]
|
||||
assert search_results[0]["updated_at"] == get_all_results[0]["updated_at"]
|
||||
@@ -627,240 +626,3 @@ async def test_async_update_preserves_actor_id_when_different_actor_updates(mock
|
||||
assert stored["actor_id"] == "Alice"
|
||||
|
||||
|
||||
class TestHallucinatedIdGuard:
|
||||
"""Tests for temp_uuid_mapping guard against LLM-hallucinated IDs (issue #3931).
|
||||
|
||||
When the LLM returns an UPDATE or DELETE with an ID that doesn't exist in
|
||||
temp_uuid_mapping, the code should skip gracefully instead of raising KeyError.
|
||||
"""
|
||||
|
||||
def test_sync_update_with_hallucinated_id_skips_gracefully(self, mocker, caplog):
|
||||
"""Sync UPDATE with an out-of-range ID should be skipped with a warning."""
|
||||
memory = _build_memory_instance(mocker, Memory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
# Simulate: 2 existing memories (IDs "0" and "1"), but LLM returns UPDATE for ID "12"
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-aaa"
|
||||
existing_mem.payload = {"data": "User likes coffee"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
|
||||
# First LLM call: fact extraction → returns one fact
|
||||
# Second LLM call: memory update actions → returns UPDATE with hallucinated ID "12"
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["User likes tea"]}',
|
||||
'{"memory": [{"id": "12", "text": "User likes tea", "event": "UPDATE", "old_memory": "User likes coffee"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I like tea"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
# Should not crash, should return empty (the hallucinated UPDATE was skipped)
|
||||
assert result == []
|
||||
assert "UPDATE skipped: LLM returned unknown id" in caplog.text
|
||||
# _update_memory should NOT have been called
|
||||
memory.vector_store.update.assert_not_called()
|
||||
|
||||
def test_sync_delete_with_hallucinated_id_skips_gracefully(self, mocker, caplog):
|
||||
"""Sync DELETE with an out-of-range ID should be skipped with a warning."""
|
||||
memory = _build_memory_instance(mocker, Memory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-aaa"
|
||||
existing_mem.payload = {"data": "User likes coffee"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["Remove coffee preference"]}',
|
||||
'{"memory": [{"id": "9", "text": "User likes coffee", "event": "DELETE"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I no longer like coffee"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "DELETE skipped: LLM returned unknown id" in caplog.text
|
||||
memory.vector_store.delete.assert_not_called()
|
||||
|
||||
def test_sync_valid_id_still_processes_normally(self, mocker, caplog):
|
||||
"""A valid ID should still be processed — the guard must not block legitimate operations."""
|
||||
memory = _build_memory_instance(mocker, Memory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-aaa"
|
||||
existing_mem.payload = {"data": "User likes coffee"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
memory.vector_store.get.return_value = MagicMock(
|
||||
payload={"data": "User likes coffee", "created_at": "2026-01-01T00:00:00+00:00"}
|
||||
)
|
||||
|
||||
# ID "0" is valid since there's exactly 1 existing memory
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["User likes tea now"]}',
|
||||
'{"memory": [{"id": "0", "text": "User likes tea now", "event": "UPDATE", "old_memory": "User likes coffee"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I like tea now"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["event"] == "UPDATE"
|
||||
assert result[0]["memory"] == "User likes tea now"
|
||||
assert result[0]["id"] == "uuid-aaa"
|
||||
assert "skipped" not in caplog.text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_update_with_hallucinated_id_skips_gracefully(self, mocker, caplog):
|
||||
"""Async UPDATE with an out-of-range ID should be skipped with a warning."""
|
||||
memory = _build_memory_instance(mocker, AsyncMemory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-bbb"
|
||||
existing_mem.payload = {"data": "User works at Acme"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["User works at Globex"]}',
|
||||
'{"memory": [{"id": "7", "text": "User works at Globex", "event": "UPDATE", "old_memory": "User works at Acme"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = await memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I now work at Globex"}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "UPDATE skipped: LLM returned unknown id" in caplog.text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_with_hallucinated_id_skips_gracefully(self, mocker, caplog):
|
||||
"""Async DELETE with an out-of-range ID should be skipped with a warning."""
|
||||
memory = _build_memory_instance(mocker, AsyncMemory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-ccc"
|
||||
existing_mem.payload = {"data": "User lives in SF"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["Remove SF reference"]}',
|
||||
'{"memory": [{"id": "16", "text": "User lives in SF", "event": "DELETE"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = await memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I moved away from SF"}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "DELETE skipped: LLM returned unknown id" in caplog.text
|
||||
|
||||
def test_sync_update_with_missing_id_key_skips_gracefully(self, mocker, caplog):
|
||||
"""UPDATE where the LLM omits the 'id' field entirely should be skipped."""
|
||||
memory = _build_memory_instance(mocker, Memory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-aaa"
|
||||
existing_mem.payload = {"data": "User likes coffee"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
|
||||
# LLM response has no "id" key at all
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["User likes tea"]}',
|
||||
'{"memory": [{"text": "User likes tea", "event": "UPDATE", "old_memory": "User likes coffee"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I like tea"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "UPDATE skipped: LLM returned unknown id" in caplog.text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_valid_id_still_processes_normally(self, mocker, caplog):
|
||||
"""Async path: a valid ID should process normally — no false positives from the guard."""
|
||||
memory = _build_memory_instance(mocker, AsyncMemory)
|
||||
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
existing_mem = MagicMock()
|
||||
existing_mem.id = "uuid-bbb"
|
||||
existing_mem.payload = {"data": "User works at Acme"}
|
||||
memory.vector_store.search.return_value = [existing_mem]
|
||||
memory.vector_store.get.return_value = MagicMock(
|
||||
payload={"data": "User works at Acme", "created_at": "2026-01-01T00:00:00+00:00"}
|
||||
)
|
||||
|
||||
# ID "0" is valid since there's exactly 1 existing memory
|
||||
memory.llm.generate_response.side_effect = [
|
||||
'{"facts": ["User works at Globex now"]}',
|
||||
'{"memory": [{"id": "0", "text": "User works at Globex now", "event": "UPDATE", "old_memory": "User works at Acme"}]}',
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = await memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I now work at Globex"}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["event"] == "UPDATE"
|
||||
assert result[0]["memory"] == "User works at Globex now"
|
||||
assert result[0]["id"] == "uuid-bbb"
|
||||
assert "skipped" not in caplog.text
|
||||
|
||||
|
||||
def test_normalize_iso_timestamp_to_utc_preserves_naive_values():
|
||||
assert _normalize_iso_timestamp_to_utc("2026-03-18T00:00:00") == "2026-03-18T00:00:00"
|
||||
|
||||
|
||||
def test_normalize_iso_timestamp_to_utc_converts_pacific():
|
||||
result = _normalize_iso_timestamp_to_utc("2026-03-17T17:00:00-07:00")
|
||||
assert result == "2026-03-18T00:00:00+00:00"
|
||||
|
||||
|
||||
def test_normalize_iso_timestamp_to_utc_handles_none():
|
||||
assert _normalize_iso_timestamp_to_utc(None) is None
|
||||
|
||||
|
||||
def test_normalize_iso_timestamp_to_utc_handles_empty():
|
||||
assert _normalize_iso_timestamp_to_utc("") == ""
|
||||
|
||||
@@ -1,107 +0,0 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
# langchain_memgraph and rank_bm25 are optional deps — mock them so tests run without install
|
||||
_memgraph_mock = Mock()
|
||||
patch.dict("sys.modules", {
|
||||
"langchain_memgraph": _memgraph_mock,
|
||||
"langchain_memgraph.graphs": _memgraph_mock,
|
||||
"langchain_memgraph.graphs.memgraph": _memgraph_mock,
|
||||
"rank_bm25": Mock(),
|
||||
}).start()
|
||||
|
||||
from mem0.memory.memgraph_memory import MemoryGraph as MemgraphMemoryGraph # noqa: E402
|
||||
|
||||
MemoryGraph = MemgraphMemoryGraph
|
||||
|
||||
|
||||
def _make_instance():
|
||||
with patch.object(MemoryGraph, "__init__", return_value=None):
|
||||
instance = MemoryGraph.__new__(MemoryGraph)
|
||||
instance.llm_provider = "openai"
|
||||
instance.llm = MagicMock()
|
||||
instance.embedding_model = MagicMock()
|
||||
instance.config = MagicMock()
|
||||
instance.config.graph_store.custom_prompt = None
|
||||
return instance
|
||||
|
||||
|
||||
class TestRetrieveNodesFromData:
|
||||
"""Tests for _retrieve_nodes_from_data in MemoryGraph."""
|
||||
|
||||
def test_normal_entities_extracted(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "Alice", "entity_type": "person"},
|
||||
{"entity": "hiking", "entity_type": "activity"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"})
|
||||
assert result == {"alice": "person", "hiking": "activity"}
|
||||
|
||||
def test_malformed_entity_missing_entity_type_is_skipped(self):
|
||||
"""LLM returns entity dict without entity_type — should skip it, keep valid ones.
|
||||
Reproduces the exact data from issue #4055."""
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "matrix multiplication", "entity_type": "task"},
|
||||
{"entity": "task"},
|
||||
{"entity": "ReLU", "entity_type": "task"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"})
|
||||
assert "matrix_multiplication" in result
|
||||
assert "relu" in result
|
||||
assert "task" not in result
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
"""LLM returns extract_entities tool call without 'entities' key — should not crash.
|
||||
Reproduces the exact scenario from issue #4238."""
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
def test_none_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": None}
|
||||
result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestEstablishNodesRelationsFromData:
|
||||
"""Tests for _establish_nodes_relations_from_data in MemoryGraph."""
|
||||
|
||||
def test_none_response_does_not_crash(self):
|
||||
"""openai_structured returns None when no relations found — must not crash.
|
||||
Exact crash from issue #4055: TypeError: 'NoneType' object is not subscriptable."""
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = None
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_empty_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": []}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_valid_entities_returned(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
|
||||
{"source": "alice", "relationship": "loves", "destination": "hiking"}
|
||||
]}}]
|
||||
}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Alice loves hiking", {"user_id": "u1"}, {"alice": "person"}
|
||||
)
|
||||
assert len(result) == 1
|
||||
assert result[0]["source"] == "alice"
|
||||
@@ -1,318 +0,0 @@
|
||||
import os
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from mem0.memory.utils import sanitize_relationship_for_cypher
|
||||
|
||||
|
||||
class TestSanitizeRelationshipForCypher:
|
||||
"""Test that relationship names are properly sanitized for Neo4j Cypher queries."""
|
||||
|
||||
def test_hyphen_replaced_with_underscore(self):
|
||||
"""Hyphens in relationship names cause Neo4j CypherSyntaxError and must be replaced."""
|
||||
assert sanitize_relationship_for_cypher("manages_via_low-cost_models") == "manages_via_low_cost_models"
|
||||
|
||||
def test_multiple_hyphens(self):
|
||||
assert sanitize_relationship_for_cypher("co-owns-with") == "co_owns_with"
|
||||
|
||||
def test_no_special_chars_unchanged(self):
|
||||
assert sanitize_relationship_for_cypher("works_at") == "works_at"
|
||||
|
||||
def test_spaces_not_handled_here(self):
|
||||
"""Spaces are replaced upstream before this function is called."""
|
||||
# sanitize only handles special chars, spaces are handled by the caller
|
||||
result = sanitize_relationship_for_cypher("has relationship")
|
||||
assert result == "has relationship"
|
||||
|
||||
def test_existing_chars_still_sanitized(self):
|
||||
assert "_slash_" in sanitize_relationship_for_cypher("read/write")
|
||||
assert "_at_" in sanitize_relationship_for_cypher("user@company")
|
||||
|
||||
|
||||
class TestNeo4jCypherSyntaxFix:
|
||||
"""Test that Neo4j Cypher syntax fixes work correctly"""
|
||||
|
||||
def test_get_all_generates_valid_cypher_with_agent_id(self):
|
||||
"""Test that get_all method generates valid Cypher with agent_id"""
|
||||
# Mock the langchain_neo4j module to avoid import issues
|
||||
with patch.dict('sys.modules', {'langchain_neo4j': Mock()}):
|
||||
from mem0.memory.graph_memory import MemoryGraph
|
||||
|
||||
# Create instance (will fail on actual connection, but that's fine for syntax testing)
|
||||
try:
|
||||
_ = MemoryGraph(url="bolt://localhost:7687", username="test", password="test")
|
||||
except Exception:
|
||||
# Expected to fail on connection, just test the class exists
|
||||
assert MemoryGraph is not None
|
||||
return
|
||||
|
||||
def test_cypher_syntax_validation(self):
|
||||
"""Test that our Cypher fixes don't contain problematic patterns"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Ensure the old buggy pattern is not present
|
||||
assert "AND n.agent_id = $agent_id AND m.agent_id = $agent_id" not in content
|
||||
assert "WHERE 1=1 {agent_filter}" not in content
|
||||
|
||||
# Ensure proper node property syntax is present
|
||||
assert "node_props" in content
|
||||
assert "agent_id: $agent_id" in content
|
||||
|
||||
# Ensure run_id follows the same pattern
|
||||
# Check for absence of problematic run_id patterns
|
||||
assert "AND n.run_id = $run_id AND m.run_id = $run_id" not in content
|
||||
assert "WHERE 1=1 {run_id_filter}" not in content
|
||||
|
||||
def test_no_undefined_variables_in_cypher(self):
|
||||
"""Test that we don't have undefined variable patterns"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check for patterns that would cause "Variable 'm' not defined" errors
|
||||
lines = content.split('\n')
|
||||
for i, line in enumerate(lines):
|
||||
# Look for WHERE clauses that reference variables not in MATCH
|
||||
if 'WHERE' in line and 'm.agent_id' in line:
|
||||
# Check if there's a MATCH clause before this that defines 'm'
|
||||
preceding_lines = lines[max(0, i-10):i]
|
||||
match_found = any('MATCH' in prev_line and ' m ' in prev_line for prev_line in preceding_lines)
|
||||
assert match_found, f"Line {i+1}: WHERE clause references 'm' without MATCH definition"
|
||||
|
||||
# Also check for run_id patterns that might have similar issues
|
||||
if 'WHERE' in line and 'm.run_id' in line:
|
||||
# Check if there's a MATCH clause before this that defines 'm'
|
||||
preceding_lines = lines[max(0, i-10):i]
|
||||
match_found = any('MATCH' in prev_line and ' m ' in prev_line for prev_line in preceding_lines)
|
||||
assert match_found, f"Line {i+1}: WHERE clause references 'm.run_id' without MATCH definition"
|
||||
|
||||
def test_agent_id_integration_syntax(self):
|
||||
"""Test that agent_id is properly integrated into MATCH clauses"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Should have node property building logic
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props.append("agent_id: $agent_id")' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
|
||||
# Should use the node properties in MATCH clauses
|
||||
assert '{{{node_props_str}}}' in content or '{node_props_str}' in content
|
||||
|
||||
def test_run_id_integration_syntax(self):
|
||||
"""Test that run_id is properly integrated into MATCH clauses"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Should have node property building logic for run_id
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props.append("run_id: $run_id")' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
|
||||
# Should use the node properties in MATCH clauses
|
||||
assert '{{{node_props_str}}}' in content or '{node_props_str}' in content
|
||||
|
||||
def test_agent_id_filter_patterns(self):
|
||||
"""Test that agent_id filtering follows the correct pattern"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that agent_id is handled in filters
|
||||
assert 'if filters.get("agent_id"):' in content
|
||||
assert 'params["agent_id"] = filters["agent_id"]' in content
|
||||
|
||||
# Check that agent_id is used in node properties
|
||||
assert 'node_props.append("agent_id: $agent_id")' in content
|
||||
|
||||
def test_run_id_filter_patterns(self):
|
||||
"""Test that run_id filtering follows the same pattern as agent_id"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that run_id is handled in filters
|
||||
assert 'if filters.get("run_id"):' in content
|
||||
assert 'params["run_id"] = filters["run_id"]' in content
|
||||
|
||||
# Check that run_id is used in node properties
|
||||
assert 'node_props.append("run_id: $run_id")' in content
|
||||
|
||||
def test_agent_id_cypher_generation(self):
|
||||
"""Test that agent_id is properly included in Cypher query generation"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that the dynamic property building pattern exists
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
|
||||
# Check that agent_id is handled in the pattern
|
||||
assert 'if filters.get(' in content
|
||||
assert 'node_props.append(' in content
|
||||
|
||||
# Verify the pattern is used in MATCH clauses
|
||||
assert '{{{node_props_str}}}' in content or '{node_props_str}' in content
|
||||
|
||||
def test_run_id_cypher_generation(self):
|
||||
"""Test that run_id is properly included in Cypher query generation"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that the dynamic property building pattern exists
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
|
||||
# Check that run_id is handled in the pattern
|
||||
assert 'if filters.get(' in content
|
||||
assert 'node_props.append(' in content
|
||||
|
||||
# Verify the pattern is used in MATCH clauses
|
||||
assert '{{{node_props_str}}}' in content or '{node_props_str}' in content
|
||||
|
||||
def test_agent_id_implementation_pattern(self):
|
||||
"""Test that the code structure supports agent_id implementation"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Verify that agent_id pattern is used consistently
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
assert 'if filters.get("agent_id"):' in content
|
||||
assert 'node_props.append("agent_id: $agent_id")' in content
|
||||
|
||||
def test_run_id_implementation_pattern(self):
|
||||
"""Test that the code structure supports run_id implementation"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Verify that run_id pattern is used consistently
|
||||
assert 'node_props = [' in content
|
||||
assert 'node_props_str = ", ".join(node_props)' in content
|
||||
assert 'if filters.get("run_id"):' in content
|
||||
assert 'node_props.append("run_id: $run_id")' in content
|
||||
|
||||
def test_user_identity_integration(self):
|
||||
"""Test that both agent_id and run_id are properly integrated into user identity"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that user_identity building includes both agent_id and run_id
|
||||
assert 'user_identity = f"user_id: {filters[\'user_id\']}"' in content
|
||||
assert 'user_identity += f", agent_id: {filters[\'agent_id\']}"' in content
|
||||
assert 'user_identity += f", run_id: {filters[\'run_id\']}"' in content
|
||||
|
||||
def test_search_methods_integration(self):
|
||||
"""Test that both agent_id and run_id are properly integrated into search methods"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that search methods handle both agent_id and run_id
|
||||
assert 'where_conditions.append("source_candidate.agent_id = $agent_id")' in content
|
||||
assert 'where_conditions.append("source_candidate.run_id = $run_id")' in content
|
||||
assert 'where_conditions.append("destination_candidate.agent_id = $agent_id")' in content
|
||||
assert 'where_conditions.append("destination_candidate.run_id = $run_id")' in content
|
||||
|
||||
def test_add_entities_integration(self):
|
||||
"""Test that both agent_id and run_id are properly integrated into add_entities"""
|
||||
graph_memory_path = 'mem0/memory/graph_memory.py'
|
||||
|
||||
# Check if file exists before reading
|
||||
if not os.path.exists(graph_memory_path):
|
||||
# Skip test if file doesn't exist (e.g., in CI environment)
|
||||
return
|
||||
|
||||
with open(graph_memory_path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Check that add_entities handles both agent_id and run_id
|
||||
assert 'agent_id = filters.get("agent_id", None)' in content
|
||||
assert 'run_id = filters.get("run_id", None)' in content
|
||||
|
||||
# Check that merge properties include both
|
||||
assert 'if agent_id:' in content
|
||||
assert 'if run_id:' in content
|
||||
assert 'merge_props.append("agent_id: $agent_id")' in content
|
||||
assert 'merge_props.append("run_id: $run_id")' in content
|
||||
|
||||
@@ -1,338 +0,0 @@
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.graphs.neptune.base import NeptuneBase
|
||||
from mem0.graphs.neptune.neptunegraph import MemoryGraph
|
||||
|
||||
|
||||
class TestNeptuneMemory(unittest.TestCase):
|
||||
"""Test suite for the Neptune Memory implementation."""
|
||||
|
||||
def setUp(self):
|
||||
"""Set up test fixtures before each test method."""
|
||||
|
||||
# Create a mock config
|
||||
self.config = MagicMock()
|
||||
self.config.graph_store.config.endpoint = "neptune-graph://test-graph"
|
||||
self.config.graph_store.config.base_label = True
|
||||
self.config.graph_store.threshold = 0.7
|
||||
self.config.llm.provider = "openai_structured"
|
||||
self.config.graph_store.llm = None
|
||||
self.config.graph_store.custom_prompt = None
|
||||
|
||||
# Create mock for NeptuneAnalyticsGraph
|
||||
self.mock_graph = MagicMock()
|
||||
self.mock_graph.client.get_graph.return_value = {"status": "AVAILABLE"}
|
||||
|
||||
# Create mocks for static methods
|
||||
self.mock_embedding_model = MagicMock()
|
||||
self.mock_llm = MagicMock()
|
||||
|
||||
# Patch the necessary components
|
||||
self.neptune_analytics_graph_patcher = patch("mem0.graphs.neptune.neptunegraph.NeptuneAnalyticsGraph")
|
||||
self.mock_neptune_analytics_graph = self.neptune_analytics_graph_patcher.start()
|
||||
self.mock_neptune_analytics_graph.return_value = self.mock_graph
|
||||
|
||||
# Patch the static methods
|
||||
self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model")
|
||||
self.mock_create_embedding_model = self.create_embedding_model_patcher.start()
|
||||
self.mock_create_embedding_model.return_value = self.mock_embedding_model
|
||||
|
||||
self.create_llm_patcher = patch.object(NeptuneBase, "_create_llm")
|
||||
self.mock_create_llm = self.create_llm_patcher.start()
|
||||
self.mock_create_llm.return_value = self.mock_llm
|
||||
|
||||
# Create the MemoryGraph instance
|
||||
self.memory_graph = MemoryGraph(self.config)
|
||||
|
||||
# Set up common test data
|
||||
self.user_id = "test_user"
|
||||
self.test_filters = {"user_id": self.user_id}
|
||||
|
||||
def tearDown(self):
|
||||
"""Tear down test fixtures after each test method."""
|
||||
self.neptune_analytics_graph_patcher.stop()
|
||||
self.create_embedding_model_patcher.stop()
|
||||
self.create_llm_patcher.stop()
|
||||
|
||||
def test_initialization(self):
|
||||
"""Test that the MemoryGraph is initialized correctly."""
|
||||
self.assertEqual(self.memory_graph.graph, self.mock_graph)
|
||||
self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model)
|
||||
self.assertEqual(self.memory_graph.llm, self.mock_llm)
|
||||
self.assertEqual(self.memory_graph.llm_provider, "openai_structured")
|
||||
self.assertEqual(self.memory_graph.node_label, ":`__Entity__`")
|
||||
self.assertEqual(self.memory_graph.threshold, 0.7)
|
||||
|
||||
def test_init(self):
|
||||
"""Test the class init functions"""
|
||||
|
||||
# Create a mock config with bad endpoint
|
||||
config_no_endpoint = MagicMock()
|
||||
config_no_endpoint.graph_store.config.endpoint = None
|
||||
|
||||
# Create the MemoryGraph instance
|
||||
with pytest.raises(ValueError):
|
||||
MemoryGraph(config_no_endpoint)
|
||||
|
||||
# Create a mock config with bad endpoint
|
||||
config_ndb_endpoint = MagicMock()
|
||||
config_ndb_endpoint.graph_store.config.endpoint = "neptune-db://test-graph"
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
MemoryGraph(config_ndb_endpoint)
|
||||
|
||||
def test_add_method(self):
|
||||
"""Test the add method with mocked components."""
|
||||
|
||||
# Mock the necessary methods that add() calls
|
||||
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person", "bob": "person"})
|
||||
self.memory_graph._establish_nodes_relations_from_data = MagicMock(
|
||||
return_value=[{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
)
|
||||
self.memory_graph._search_graph_db = MagicMock(return_value=[])
|
||||
self.memory_graph._get_delete_entities_from_search_output = MagicMock(return_value=[])
|
||||
self.memory_graph._delete_entities = MagicMock(return_value=[])
|
||||
self.memory_graph._add_entities = MagicMock(
|
||||
return_value=[{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
)
|
||||
|
||||
# Call the add method
|
||||
result = self.memory_graph.add("Alice knows Bob", self.test_filters)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Alice knows Bob", self.test_filters)
|
||||
self.memory_graph._establish_nodes_relations_from_data.assert_called_once()
|
||||
self.memory_graph._search_graph_db.assert_called_once()
|
||||
self.memory_graph._get_delete_entities_from_search_output.assert_called_once()
|
||||
self.memory_graph._delete_entities.assert_called_once_with([], self.user_id)
|
||||
self.memory_graph._add_entities.assert_called_once()
|
||||
|
||||
# Check the result structure
|
||||
self.assertIn("deleted_entities", result)
|
||||
self.assertIn("added_entities", result)
|
||||
|
||||
def test_search_method(self):
|
||||
"""Test the search method with mocked components."""
|
||||
# Mock the necessary methods that search() calls
|
||||
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
|
||||
# Mock search results
|
||||
mock_search_results = [
|
||||
{"source": "alice", "relationship": "knows", "destination": "bob"},
|
||||
{"source": "alice", "relationship": "works_with", "destination": "charlie"},
|
||||
]
|
||||
self.memory_graph._search_graph_db = MagicMock(return_value=mock_search_results)
|
||||
|
||||
# Mock BM25Okapi
|
||||
with patch("mem0.graphs.neptune.base.BM25Okapi") as mock_bm25:
|
||||
mock_bm25_instance = MagicMock()
|
||||
mock_bm25.return_value = mock_bm25_instance
|
||||
|
||||
# Mock get_top_n to return reranked results
|
||||
reranked_results = [["alice", "knows", "bob"], ["alice", "works_with", "charlie"]]
|
||||
mock_bm25_instance.get_top_n.return_value = reranked_results
|
||||
|
||||
# Call the search method
|
||||
result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters)
|
||||
self.memory_graph._search_graph_db.assert_called_once_with(node_list=["alice"], filters=self.test_filters)
|
||||
|
||||
# Check the result structure
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0]["source"], "alice")
|
||||
self.assertEqual(result[0]["relationship"], "knows")
|
||||
self.assertEqual(result[0]["destination"], "bob")
|
||||
|
||||
def test_get_all_method(self):
|
||||
"""Test the get_all method."""
|
||||
|
||||
# Mock the _get_all_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"user_id": self.user_id, "limit": 10}
|
||||
self.memory_graph._get_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [
|
||||
{"source": "alice", "relationship": "knows", "target": "bob"},
|
||||
{"source": "bob", "relationship": "works_with", "target": "charlie"},
|
||||
]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the get_all method
|
||||
result = self.memory_graph.get_all(self.test_filters, top_k=10)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result structure
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0]["source"], "alice")
|
||||
self.assertEqual(result[0]["relationship"], "knows")
|
||||
self.assertEqual(result[0]["target"], "bob")
|
||||
|
||||
def test_delete_all_method(self):
|
||||
"""Test the delete_all method."""
|
||||
# Mock the _delete_all_cypher method
|
||||
mock_cypher = "MATCH (n) DETACH DELETE n"
|
||||
mock_params = {"user_id": self.user_id}
|
||||
self.memory_graph._delete_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Call the delete_all method
|
||||
self.memory_graph.delete_all(self.test_filters)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._delete_all_cypher.assert_called_once_with(self.test_filters)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
def test_search_source_node(self):
|
||||
"""Test the _search_source_node method."""
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
|
||||
# Mock the _search_source_node_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
|
||||
self.memory_graph._search_source_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _search_source_node method
|
||||
result = self.memory_graph._search_source_node(mock_embedding, self.user_id, threshold=0.9)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._search_source_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, mock_query_result)
|
||||
|
||||
def test_search_destination_node(self):
|
||||
"""Test the _search_destination_node method."""
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
|
||||
# Mock the _search_destination_node_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"destination_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
|
||||
self.memory_graph._search_destination_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _search_destination_node method
|
||||
result = self.memory_graph._search_destination_node(mock_embedding, self.user_id, threshold=0.9)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._search_destination_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, mock_query_result)
|
||||
|
||||
def test_search_graph_db(self):
|
||||
"""Test the _search_graph_db method."""
|
||||
# Mock node list
|
||||
node_list = ["alice", "bob"]
|
||||
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
self.mock_embedding_model.embed.return_value = mock_embedding
|
||||
|
||||
# Mock the _search_graph_db_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"n_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.7, "limit": 10}
|
||||
self.memory_graph._search_graph_db_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query results
|
||||
mock_query_result1 = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
mock_query_result2 = [{"source": "bob", "relationship": "works_with", "destination": "charlie"}]
|
||||
self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2]
|
||||
|
||||
# Call the _search_graph_db method
|
||||
result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10)
|
||||
|
||||
# Verify the method calls
|
||||
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
|
||||
self.assertEqual(self.memory_graph._search_graph_db_cypher.call_count, 2)
|
||||
self.assertEqual(self.mock_graph.query.call_count, 2)
|
||||
|
||||
# Check the result
|
||||
expected_result = mock_query_result1 + mock_query_result2
|
||||
self.assertEqual(result, expected_result)
|
||||
|
||||
def test_add_entities(self):
|
||||
"""Test the _add_entities method."""
|
||||
# Mock data
|
||||
to_be_added = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
entity_type_map = {"alice": "person", "bob": "person"}
|
||||
|
||||
# Mock embeddings
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
self.mock_embedding_model.embed.return_value = mock_embedding
|
||||
|
||||
# Mock search results
|
||||
mock_source_search = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
|
||||
mock_dest_search = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
|
||||
|
||||
# Mock the search methods
|
||||
self.memory_graph._search_source_node = MagicMock(return_value=mock_source_search)
|
||||
self.memory_graph._search_destination_node = MagicMock(return_value=mock_dest_search)
|
||||
|
||||
# Mock the _add_entities_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_id": 123, "destination_id": 456}
|
||||
self.memory_graph._add_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _add_entities method
|
||||
result = self.memory_graph._add_entities(to_be_added, self.user_id, entity_type_map)
|
||||
|
||||
# Verify the method calls
|
||||
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
|
||||
self.memory_graph._search_source_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7)
|
||||
self.memory_graph._search_destination_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7)
|
||||
self.memory_graph._add_entities_cypher.assert_called_once()
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, [mock_query_result])
|
||||
|
||||
def test_delete_entities(self):
|
||||
"""Test the _delete_entities method."""
|
||||
# Mock data
|
||||
to_be_deleted = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
|
||||
# Mock the _delete_entities_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_name": "alice", "dest_name": "bob", "user_id": self.user_id}
|
||||
self.memory_graph._delete_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _delete_entities method
|
||||
result = self.memory_graph._delete_entities(to_be_deleted, self.user_id)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._delete_entities_cypher.assert_called_once_with("alice", "bob", "knows", self.user_id)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, [mock_query_result])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,411 +0,0 @@
|
||||
import unittest
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.graphs.neptune.base import NeptuneBase
|
||||
from mem0.graphs.neptune.neptunedb import MemoryGraph
|
||||
|
||||
|
||||
class TestNeptuneMemory(unittest.TestCase):
|
||||
"""Test suite for the Neptune Memory implementation."""
|
||||
|
||||
def setUp(self):
|
||||
"""Set up test fixtures before each test method."""
|
||||
|
||||
# Create a mock config
|
||||
self.config = MagicMock()
|
||||
self.config.graph_store.config.endpoint = "neptune-db://test-graph"
|
||||
self.config.graph_store.config.base_label = True
|
||||
self.config.graph_store.threshold = 0.7
|
||||
self.config.llm.provider = "openai_structured"
|
||||
self.config.graph_store.llm = None
|
||||
self.config.graph_store.custom_prompt = None
|
||||
self.config.vector_store.provider = "qdrant"
|
||||
self.config.vector_store.config = MagicMock()
|
||||
|
||||
# Create mock for NeptuneGraph
|
||||
self.mock_graph = MagicMock()
|
||||
|
||||
# Create mocks for static methods
|
||||
self.mock_embedding_model = MagicMock()
|
||||
self.mock_llm = MagicMock()
|
||||
self.mock_vector_store = MagicMock()
|
||||
|
||||
# Patch the necessary components
|
||||
self.neptune_graph_patcher = patch("mem0.graphs.neptune.neptunedb.NeptuneGraph")
|
||||
self.mock_neptune_graph = self.neptune_graph_patcher.start()
|
||||
self.mock_neptune_graph.return_value = self.mock_graph
|
||||
|
||||
# Patch the static methods
|
||||
self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model")
|
||||
self.mock_create_embedding_model = self.create_embedding_model_patcher.start()
|
||||
self.mock_create_embedding_model.return_value = self.mock_embedding_model
|
||||
|
||||
self.create_llm_patcher = patch.object(NeptuneBase, "_create_llm")
|
||||
self.mock_create_llm = self.create_llm_patcher.start()
|
||||
self.mock_create_llm.return_value = self.mock_llm
|
||||
|
||||
self.create_vector_store_patcher = patch.object(NeptuneBase, "_create_vector_store")
|
||||
self.mock_create_vector_store = self.create_vector_store_patcher.start()
|
||||
self.mock_create_vector_store.return_value = self.mock_vector_store
|
||||
|
||||
# Create the MemoryGraph instance
|
||||
self.memory_graph = MemoryGraph(self.config)
|
||||
|
||||
# Set up common test data
|
||||
self.user_id = "test_user"
|
||||
self.test_filters = {"user_id": self.user_id}
|
||||
|
||||
def tearDown(self):
|
||||
"""Tear down test fixtures after each test method."""
|
||||
self.neptune_graph_patcher.stop()
|
||||
self.create_embedding_model_patcher.stop()
|
||||
self.create_llm_patcher.stop()
|
||||
self.create_vector_store_patcher.stop()
|
||||
|
||||
def test_initialization(self):
|
||||
"""Test that the MemoryGraph is initialized correctly."""
|
||||
self.assertEqual(self.memory_graph.graph, self.mock_graph)
|
||||
self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model)
|
||||
self.assertEqual(self.memory_graph.llm, self.mock_llm)
|
||||
self.assertEqual(self.memory_graph.vector_store, self.mock_vector_store)
|
||||
self.assertEqual(self.memory_graph.llm_provider, "openai_structured")
|
||||
self.assertEqual(self.memory_graph.node_label, ":`__Entity__`")
|
||||
self.assertEqual(self.memory_graph.threshold, 0.7)
|
||||
self.assertEqual(self.memory_graph.vector_store_limit, 5)
|
||||
|
||||
def test_collection_name_variants(self):
|
||||
"""Test all collection_name configuration variants."""
|
||||
|
||||
# Test 1: graph_store.config.collection_name is set
|
||||
config1 = MagicMock()
|
||||
config1.graph_store.config.endpoint = "neptune-db://test-graph"
|
||||
config1.graph_store.config.base_label = True
|
||||
config1.graph_store.config.collection_name = "custom_collection"
|
||||
config1.llm.provider = "openai"
|
||||
config1.graph_store.llm = None
|
||||
config1.vector_store.provider = "qdrant"
|
||||
config1.vector_store.config = MagicMock()
|
||||
|
||||
MemoryGraph(config1)
|
||||
self.assertEqual(config1.vector_store.config.collection_name, "custom_collection")
|
||||
|
||||
# Test 2: vector_store.config.collection_name exists, graph_store.config.collection_name is None
|
||||
config2 = MagicMock()
|
||||
config2.graph_store.config.endpoint = "neptune-db://test-graph"
|
||||
config2.graph_store.config.base_label = True
|
||||
config2.graph_store.config.collection_name = None
|
||||
config2.llm.provider = "openai"
|
||||
config2.graph_store.llm = None
|
||||
config2.vector_store.provider = "qdrant"
|
||||
config2.vector_store.config = MagicMock()
|
||||
config2.vector_store.config.collection_name = "existing_collection"
|
||||
|
||||
MemoryGraph(config2)
|
||||
self.assertEqual(config2.vector_store.config.collection_name, "existing_collection_neptune_vector_store")
|
||||
|
||||
# Test 3: Neither collection_name is set (default case)
|
||||
config3 = MagicMock()
|
||||
config3.graph_store.config.endpoint = "neptune-db://test-graph"
|
||||
config3.graph_store.config.base_label = True
|
||||
config3.graph_store.config.collection_name = None
|
||||
config3.llm.provider = "openai"
|
||||
config3.graph_store.llm = None
|
||||
config3.vector_store.provider = "qdrant"
|
||||
config3.vector_store.config = MagicMock()
|
||||
config3.vector_store.config.collection_name = None
|
||||
|
||||
MemoryGraph(config3)
|
||||
self.assertEqual(config3.vector_store.config.collection_name, "mem0_neptune_vector_store")
|
||||
|
||||
def test_init(self):
|
||||
"""Test the class init functions"""
|
||||
|
||||
# Create a mock config with bad endpoint
|
||||
config_no_endpoint = MagicMock()
|
||||
config_no_endpoint.graph_store.config.endpoint = None
|
||||
|
||||
# Create the MemoryGraph instance
|
||||
with pytest.raises(ValueError):
|
||||
MemoryGraph(config_no_endpoint)
|
||||
|
||||
# Create a mock config with wrong endpoint type
|
||||
config_wrong_endpoint = MagicMock()
|
||||
config_wrong_endpoint.graph_store.config.endpoint = "neptune-graph://test-graph"
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
MemoryGraph(config_wrong_endpoint)
|
||||
|
||||
def test_add_method(self):
|
||||
"""Test the add method with mocked components."""
|
||||
|
||||
# Mock the necessary methods that add() calls
|
||||
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person", "bob": "person"})
|
||||
self.memory_graph._establish_nodes_relations_from_data = MagicMock(
|
||||
return_value=[{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
)
|
||||
self.memory_graph._search_graph_db = MagicMock(return_value=[])
|
||||
self.memory_graph._get_delete_entities_from_search_output = MagicMock(return_value=[])
|
||||
self.memory_graph._delete_entities = MagicMock(return_value=[])
|
||||
self.memory_graph._add_entities = MagicMock(
|
||||
return_value=[{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
)
|
||||
|
||||
# Call the add method
|
||||
result = self.memory_graph.add("Alice knows Bob", self.test_filters)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Alice knows Bob", self.test_filters)
|
||||
self.memory_graph._establish_nodes_relations_from_data.assert_called_once()
|
||||
self.memory_graph._search_graph_db.assert_called_once()
|
||||
self.memory_graph._get_delete_entities_from_search_output.assert_called_once()
|
||||
self.memory_graph._delete_entities.assert_called_once_with([], self.user_id)
|
||||
self.memory_graph._add_entities.assert_called_once()
|
||||
|
||||
# Check the result structure
|
||||
self.assertIn("deleted_entities", result)
|
||||
self.assertIn("added_entities", result)
|
||||
|
||||
def test_search_method(self):
|
||||
"""Test the search method with mocked components."""
|
||||
# Mock the necessary methods that search() calls
|
||||
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
|
||||
# Mock search results
|
||||
mock_search_results = [
|
||||
{"source": "alice", "relationship": "knows", "destination": "bob"},
|
||||
{"source": "alice", "relationship": "works_with", "destination": "charlie"},
|
||||
]
|
||||
self.memory_graph._search_graph_db = MagicMock(return_value=mock_search_results)
|
||||
|
||||
# Mock BM25Okapi
|
||||
with patch("mem0.graphs.neptune.base.BM25Okapi") as mock_bm25:
|
||||
mock_bm25_instance = MagicMock()
|
||||
mock_bm25.return_value = mock_bm25_instance
|
||||
|
||||
# Mock get_top_n to return reranked results
|
||||
reranked_results = [["alice", "knows", "bob"], ["alice", "works_with", "charlie"]]
|
||||
mock_bm25_instance.get_top_n.return_value = reranked_results
|
||||
|
||||
# Call the search method
|
||||
result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters)
|
||||
self.memory_graph._search_graph_db.assert_called_once_with(node_list=["alice"], filters=self.test_filters)
|
||||
|
||||
# Check the result structure
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0]["source"], "alice")
|
||||
self.assertEqual(result[0]["relationship"], "knows")
|
||||
self.assertEqual(result[0]["destination"], "bob")
|
||||
|
||||
def test_get_all_method(self):
|
||||
"""Test the get_all method."""
|
||||
|
||||
# Mock the _get_all_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"user_id": self.user_id, "limit": 10}
|
||||
self.memory_graph._get_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [
|
||||
{"source": "alice", "relationship": "knows", "target": "bob"},
|
||||
{"source": "bob", "relationship": "works_with", "target": "charlie"},
|
||||
]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the get_all method
|
||||
result = self.memory_graph.get_all(self.test_filters, top_k=10)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result structure
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0]["source"], "alice")
|
||||
self.assertEqual(result[0]["relationship"], "knows")
|
||||
self.assertEqual(result[0]["target"], "bob")
|
||||
|
||||
def test_delete_all_method(self):
|
||||
"""Test the delete_all method."""
|
||||
# Mock the _delete_all_cypher method
|
||||
mock_cypher = "MATCH (n) DETACH DELETE n"
|
||||
mock_params = {"user_id": self.user_id}
|
||||
self.memory_graph._delete_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Call the delete_all method
|
||||
self.memory_graph.delete_all(self.test_filters)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._delete_all_cypher.assert_called_once_with(self.test_filters)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
def test_search_source_node(self):
|
||||
"""Test the _search_source_node method."""
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
|
||||
# Mock the _search_source_node_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
|
||||
self.memory_graph._search_source_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _search_source_node method
|
||||
result = self.memory_graph._search_source_node(mock_embedding, self.user_id, threshold=0.9)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._search_source_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, mock_query_result)
|
||||
|
||||
def test_search_destination_node(self):
|
||||
"""Test the _search_destination_node method."""
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
|
||||
# Mock the _search_destination_node_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"destination_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
|
||||
self.memory_graph._search_destination_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _search_destination_node method
|
||||
result = self.memory_graph._search_destination_node(mock_embedding, self.user_id, threshold=0.9)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._search_destination_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, mock_query_result)
|
||||
|
||||
def test_add_new_entities_payloads_use_utc_timestamps(self):
|
||||
"""Test that Neptune vector-store payloads use UTC timestamps."""
|
||||
self.memory_graph._add_new_entities_cypher(
|
||||
source="alice",
|
||||
source_embedding=[0.1, 0.2],
|
||||
source_type="person",
|
||||
destination="bob",
|
||||
dest_embedding=[0.3, 0.4],
|
||||
destination_type="person",
|
||||
relationship="KNOWS",
|
||||
user_id=self.user_id,
|
||||
)
|
||||
|
||||
_, kwargs = self.mock_vector_store.insert.call_args
|
||||
for payload in kwargs["payloads"]:
|
||||
parsed = datetime.fromisoformat(payload["created_at"])
|
||||
self.assertEqual(parsed.tzinfo, timezone.utc)
|
||||
self.assertEqual(parsed.utcoffset().total_seconds(), 0)
|
||||
|
||||
def test_search_graph_db(self):
|
||||
"""Test the _search_graph_db method."""
|
||||
# Mock node list
|
||||
node_list = ["alice", "bob"]
|
||||
|
||||
# Mock embedding
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
self.mock_embedding_model.embed.return_value = mock_embedding
|
||||
|
||||
# Mock the _search_graph_db_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"n_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.7, "limit": 10}
|
||||
self.memory_graph._search_graph_db_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query results
|
||||
mock_query_result1 = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
mock_query_result2 = [{"source": "bob", "relationship": "works_with", "destination": "charlie"}]
|
||||
self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2]
|
||||
|
||||
# Call the _search_graph_db method
|
||||
result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10)
|
||||
|
||||
# Verify the method calls
|
||||
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
|
||||
self.assertEqual(self.memory_graph._search_graph_db_cypher.call_count, 2)
|
||||
self.assertEqual(self.mock_graph.query.call_count, 2)
|
||||
|
||||
# Check the result
|
||||
expected_result = mock_query_result1 + mock_query_result2
|
||||
self.assertEqual(result, expected_result)
|
||||
|
||||
def test_add_entities(self):
|
||||
"""Test the _add_entities method."""
|
||||
# Mock data
|
||||
to_be_added = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
entity_type_map = {"alice": "person", "bob": "person"}
|
||||
|
||||
# Mock embeddings
|
||||
mock_embedding = [0.1, 0.2, 0.3]
|
||||
self.mock_embedding_model.embed.return_value = mock_embedding
|
||||
|
||||
# Mock search results
|
||||
mock_source_search = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
|
||||
mock_dest_search = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
|
||||
|
||||
# Mock the search methods
|
||||
self.memory_graph._search_source_node = MagicMock(return_value=mock_source_search)
|
||||
self.memory_graph._search_destination_node = MagicMock(return_value=mock_dest_search)
|
||||
|
||||
# Mock the _add_entities_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_id": 123, "destination_id": 456}
|
||||
self.memory_graph._add_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _add_entities method
|
||||
result = self.memory_graph._add_entities(to_be_added, self.user_id, entity_type_map)
|
||||
|
||||
# Verify the method calls
|
||||
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
|
||||
self.memory_graph._search_source_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7)
|
||||
self.memory_graph._search_destination_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7)
|
||||
self.memory_graph._add_entities_cypher.assert_called_once()
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, [mock_query_result])
|
||||
|
||||
def test_delete_entities(self):
|
||||
"""Test the _delete_entities method."""
|
||||
# Mock data
|
||||
to_be_deleted = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
|
||||
|
||||
# Mock the _delete_entities_cypher method
|
||||
mock_cypher = "MATCH (n) RETURN n"
|
||||
mock_params = {"source_name": "alice", "dest_name": "bob", "user_id": self.user_id}
|
||||
self.memory_graph._delete_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
|
||||
|
||||
# Mock the graph.query result
|
||||
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
self.mock_graph.query.return_value = mock_query_result
|
||||
|
||||
# Call the _delete_entities method
|
||||
result = self.memory_graph._delete_entities(to_be_deleted, self.user_id)
|
||||
|
||||
# Verify the method calls
|
||||
self.memory_graph._delete_entities_cypher.assert_called_once_with("alice", "bob", "knows", self.user_id)
|
||||
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
|
||||
|
||||
# Check the result
|
||||
self.assertEqual(result, [mock_query_result])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -166,7 +166,7 @@ class TestRealWorldFieldCoverage:
|
||||
# AWS
|
||||
("aws_session_token", True),
|
||||
# Azure MySQL
|
||||
("use_azure_credential", False),
|
||||
("use_azure_credential", True),
|
||||
# General non-sensitive
|
||||
("collection_name", False),
|
||||
("embedding_model_dims", False),
|
||||
|
||||
Reference in New Issue
Block a user