diff --git a/mem0/graphs/neptune/base.py b/mem0/graphs/neptune/base.py index 2732ceab3..59aa9f8e6 100644 --- a/mem0/graphs/neptune/base.py +++ b/mem0/graphs/neptune/base.py @@ -1,7 +1,7 @@ import logging from abc import ABC, abstractmethod -from mem0.memory.utils import format_entities +from mem0.memory.utils import format_entities, remove_spaces_from_entities try: from rank_bm25 import BM25Okapi @@ -151,11 +151,7 @@ class NeptuneBase(ABC): return entities def _remove_spaces_from_entities(self, entity_list): - for item in entity_list: - item["source"] = item["source"].lower().replace(" ", "_") - item["relationship"] = item["relationship"].lower().replace(" ", "_") - item["destination"] = item["destination"].lower().replace(" ", "_") - return entity_list + return remove_spaces_from_entities(entity_list, sanitize_relationship=False) def _get_delete_entities_from_search_output(self, search_output, data, filters): """ diff --git a/mem0/memory/graph_memory.py b/mem0/memory/graph_memory.py index 86e49099e..80a3b905a 100644 --- a/mem0/memory/graph_memory.py +++ b/mem0/memory/graph_memory.py @@ -1,6 +1,6 @@ import logging -from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher +from mem0.memory.utils import format_entities, remove_spaces_from_entities try: from langchain_neo4j import Neo4jGraph @@ -657,12 +657,7 @@ class MemoryGraph: return results def _remove_spaces_from_entities(self, entity_list): - for item in entity_list: - item["source"] = item["source"].lower().replace(" ", "_") - # Use the sanitization function for relationships to handle special characters - item["relationship"] = sanitize_relationship_for_cypher(item["relationship"].lower().replace(" ", "_")) - item["destination"] = item["destination"].lower().replace(" ", "_") - return entity_list + return remove_spaces_from_entities(entity_list, sanitize_relationship=True) def _search_source_node(self, source_embedding, filters, threshold=0.9): # Build WHERE conditions diff --git a/mem0/memory/kuzu_memory.py b/mem0/memory/kuzu_memory.py index 1ed769fb4..0a9f1a4a1 100644 --- a/mem0/memory/kuzu_memory.py +++ b/mem0/memory/kuzu_memory.py @@ -1,6 +1,6 @@ import logging -from mem0.memory.utils import format_entities +from mem0.memory.utils import format_entities, remove_spaces_from_entities try: import kuzu @@ -654,11 +654,7 @@ class MemoryGraph: return results def _remove_spaces_from_entities(self, entity_list): - for item in entity_list: - item["source"] = item["source"].lower().replace(" ", "_") - item["relationship"] = item["relationship"].lower().replace(" ", "_") - item["destination"] = item["destination"].lower().replace(" ", "_") - return entity_list + return remove_spaces_from_entities(entity_list, sanitize_relationship=False) def _search_source_node(self, source_embedding, filters, threshold=0.9): params = { diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index c7f52ad79..3a29b9e71 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -1,6 +1,6 @@ import logging -from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher +from mem0.memory.utils import format_entities, remove_spaces_from_entities try: from langchain_memgraph.graphs.memgraph import Memgraph @@ -550,12 +550,7 @@ class MemoryGraph: return results def _remove_spaces_from_entities(self, entity_list): - for item in entity_list: - item["source"] = item["source"].lower().replace(" ", "_") - # Use the sanitization function for relationships to handle special characters - item["relationship"] = sanitize_relationship_for_cypher(item["relationship"].lower().replace(" ", "_")) - item["destination"] = item["destination"].lower().replace(" ", "_") - return entity_list + return remove_spaces_from_entities(entity_list, sanitize_relationship=True) def _search_source_node(self, source_embedding, filters, threshold=0.9): """Search for source nodes with similar embeddings.""" diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index 3a353144a..61e3863e3 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -1,6 +1,7 @@ import hashlib import logging import re +from typing import Any, Dict, List from mem0.configs.prompts import ( AGENT_MEMORY_EXTRACTION_PROMPT, @@ -265,3 +266,30 @@ def sanitize_relationship_for_cypher(relationship) -> str: return re.sub(r"_+", "_", sanitized).strip("_") + +def remove_spaces_from_entities( + entity_list: List[Any], + *, + sanitize_relationship: bool = True, +) -> List[Dict[str, Any]]: + """ + Normalize entity relation dicts from LLM/tool output: lowercase, spaces to underscores. + + Skips entries that are not non-empty dicts or that lack any of + ``source``, ``relationship``, or ``destination`` (avoids KeyError on ``[{}]`` + or partial dicts). + """ + required = ("source", "relationship", "destination") + cleaned: List[Dict[str, Any]] = [] + for item in entity_list: + if not isinstance(item, dict) or not item: + continue + if not all(key in item for key in required): + continue + item["source"] = item["source"].lower().replace(" ", "_") + rel = item["relationship"].lower().replace(" ", "_") + item["relationship"] = sanitize_relationship_for_cypher(rel) if sanitize_relationship else rel + item["destination"] = item["destination"].lower().replace(" ", "_") + cleaned.append(item) + return cleaned + diff --git a/tests/memory/test_memory_utils.py b/tests/memory/test_memory_utils.py new file mode 100644 index 000000000..41a02f76d --- /dev/null +++ b/tests/memory/test_memory_utils.py @@ -0,0 +1,57 @@ +import pytest +from mem0.memory.utils import remove_spaces_from_entities, sanitize_relationship_for_cypher + + +class TestRemoveSpacesFromEntities: + """ + Covers behavior used by Neo4j, Memgraph (sanitize_relationship=True), + Kuzu, and Neptune (sanitize_relationship=False). All backends delegate here. + """ + + @pytest.mark.parametrize( + "sanitize", + [True, False], + ids=["cypher_sanitized", "plain"], + ) + def test_filters_empty_and_incomplete_dicts(self, sanitize): + mixed = [ + {}, + {"source": "a"}, + {"source": "a", "relationship": "r"}, + {"source": "x", "relationship": "rel", "destination": "y"}, + ] + out = remove_spaces_from_entities(mixed, sanitize_relationship=sanitize) + assert len(out) == 1 + assert out[0]["source"] == "x" + assert out[0]["destination"] == "y" + + @pytest.mark.parametrize("sanitize", [True, False]) + def test_all_empty_returns_empty(self, sanitize): + assert remove_spaces_from_entities([{}, {}, {}], sanitize_relationship=sanitize) == [] + + def test_skips_non_dict_entries(self): + assert remove_spaces_from_entities([None, "not-a-dict", 123, {"source": "a", "relationship": "r", "destination": "b"}]) == [ + {"source": "a", "relationship": "r", "destination": "b"} + ] + + def test_sanitize_true_relationship_uses_sanitizer(self): + """Neo4j / Memgraph path: special characters mapped via sanitize_relationship_for_cypher.""" + entities = [{"source": "A", "relationship": "x/y", "destination": "B"}] + out = remove_spaces_from_entities(entities, sanitize_relationship=True) + assert out[0]["relationship"] == sanitize_relationship_for_cypher("x/y".lower().replace(" ", "_")) + + def test_sanitize_false_relationship_plain_only(self): + """Kuzu / Neptune path: only lowercase and spaces to underscores.""" + entities = [{"source": "A", "relationship": "Works At", "destination": "B Co"}] + out = remove_spaces_from_entities(entities, sanitize_relationship=False) + assert out[0]["relationship"] == "works_at" + assert out[0]["source"] == "a" + assert out[0]["destination"] == "b_co" + + def test_sanitize_true_vs_false_slash_in_relationship(self): + """Slash is rewritten when sanitizing (Cypher path); kept as-is for plain path.""" + base = {"source": "s", "relationship": "a/b", "destination": "d"} + t = remove_spaces_from_entities([dict(base)], sanitize_relationship=True)[0]["relationship"] + f = remove_spaces_from_entities([dict(base)], sanitize_relationship=False)[0]["relationship"] + assert t == sanitize_relationship_for_cypher("a/b") + assert f == "a/b"