fix(memory): centralize entity cleanup and skip malformed LLM relation dicts (#4515)
This commit is contained in:
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user