From 1bc31b6ae389d8fd7cb2c6481cc5fbe00dbd6267 Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Wed, 13 Aug 2025 21:36:05 +0530 Subject: [PATCH] Added sanitation for better relationship mapping (#3300) --- mem0/graphs/tools.py | 12 ++++---- mem0/memory/graph_memory.py | 54 +++++++++++++++------------------- mem0/memory/main.py | 20 ++++--------- mem0/memory/memgraph_memory.py | 20 +++++++++---- mem0/memory/utils.py | 51 ++++++++++++++++++++++++++++++++ mem0/utils/factory.py | 2 +- 6 files changed, 100 insertions(+), 59 deletions(-) diff --git a/mem0/graphs/tools.py b/mem0/graphs/tools.py index 95bb32ad1..e27dc3f45 100644 --- a/mem0/graphs/tools.py +++ b/mem0/graphs/tools.py @@ -249,23 +249,23 @@ RELATIONS_STRUCT_TOOL = { "items": { "type": "object", "properties": { - "source_entity": { + "source": { "type": "string", "description": "The source entity of the relationship.", }, - "relatationship": { + "relationship": { "type": "string", "description": "The relationship between the source and destination entities.", }, - "destination_entity": { + "destination": { "type": "string", "description": "The destination entity of the relationship.", }, }, "required": [ - "source_entity", - "relatationship", - "destination_entity", + "source", + "relationship", + "destination", ], "additionalProperties": False, }, diff --git a/mem0/memory/graph_memory.py b/mem0/memory/graph_memory.py index a59c1c028..1ad28c060 100644 --- a/mem0/memory/graph_memory.py +++ b/mem0/memory/graph_memory.py @@ -1,8 +1,6 @@ import logging -import re -import unicodedata -from mem0.memory.utils import format_entities +from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher try: from langchain_neo4j import Neo4jGraph @@ -57,13 +55,20 @@ class MemoryGraph: except Exception: pass - self.llm_provider = "openai_structured" - if self.config.llm.provider: + # Default to openai if no specific provider is configured + self.llm_provider = "openai" + if self.config.llm and self.config.llm.provider: self.llm_provider = self.config.llm.provider - if self.config.graph_store.llm: + if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider: self.llm_provider = self.config.graph_store.llm.provider - self.llm = LlmFactory.create(self.llm_provider, self.config.graph_store.llm.config if self.config.graph_store.llm.config else self.config.llm.config) + # Get LLM config with proper null checks + llm_config = None + if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"): + llm_config = self.config.graph_store.llm.config + elif hasattr(self.config.llm, "config"): + llm_config = self.config.llm.config + self.llm = LlmFactory.create(self.llm_provider, llm_config) self.user_id = None self.threshold = 0.7 @@ -131,7 +136,7 @@ class MemoryGraph: if filters.get("run_id"): node_props.append("run_id: $run_id") node_props_str = ", ".join(node_props) - + cypher = f""" MATCH (n {self.node_label} {{{node_props_str}}}) DETACH DELETE n @@ -155,7 +160,7 @@ class MemoryGraph: - 'entities': A list of strings representing the nodes and relationships """ params = {"user_id": filters["user_id"], "limit": limit} - + # Build node properties based on filters node_props = ["user_id: $user_id"] if filters.get("agent_id"): @@ -265,7 +270,7 @@ class MemoryGraph: def _search_graph_db(self, node_list, filters, limit=100): """Search similar nodes among and their respective incoming and outgoing relations.""" result_relations = [] - + # Build node properties for filtering node_props = ["user_id: $user_id"] if filters.get("agent_id"): @@ -385,7 +390,7 @@ class MemoryGraph: dest_props.append("run_id: $run_id") source_props_str = ", ".join(source_props) dest_props_str = ", ".join(dest_props) - + # Delete the specific relationship between nodes cypher = f""" MATCH (n {self.node_label} {{{source_props_str}}}) @@ -600,26 +605,17 @@ class MemoryGraph: results.append(result) return results - def _sanitize_name(self, name): - normalized = unicodedata.normalize('NFKD', name) - ascii_text = normalized.encode('ascii', 'ignore').decode('ascii') - sanitized = re.sub(r'\W+', '_', ascii_text).lower() - return sanitized.strip('_') - - def _sanitize_entities(self, entity_list): + def _remove_spaces_from_entities(self, entity_list): for item in entity_list: - item["source"] = self._sanitize_name(item["source"]) - item["relationship"] = self._sanitize_name(item["relationship"]) - item["destination"] = self._sanitize_name(item["destination"]) + 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 def _search_source_node(self, source_embedding, filters, threshold=0.9): - # Build WHERE conditions - where_conditions = [ - "source_candidate.embedding IS NOT NULL", - "source_candidate.user_id = $user_id" - ] + where_conditions = ["source_candidate.embedding IS NOT NULL", "source_candidate.user_id = $user_id"] if filters.get("agent_id"): where_conditions.append("source_candidate.agent_id = $agent_id") if filters.get("run_id"): @@ -655,12 +651,8 @@ class MemoryGraph: return result def _search_destination_node(self, destination_embedding, filters, threshold=0.9): - # Build WHERE conditions - where_conditions = [ - "destination_candidate.embedding IS NOT NULL", - "destination_candidate.user_id = $user_id" - ] + where_conditions = ["destination_candidate.embedding IS NOT NULL", "destination_candidate.user_id = $user_id"] if filters.get("agent_id"): where_conditions.append("destination_candidate.agent_id = $agent_id") if filters.get("run_id"): diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 035c75a22..62df064ef 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -1384,31 +1384,21 @@ class AsyncMemory(MemoryBase): "mem0.get_all", self, {"limit": limit, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "async"} ) - vector_store_task = asyncio.create_task( - self._get_all_from_vector_store(effective_filters, limit) - ) + vector_store_task = asyncio.create_task(self._get_all_from_vector_store(effective_filters, limit)) graph_task = None if self.enable_graph: graph_get_all = getattr(self.graph, "get_all", None) if callable(graph_get_all): if asyncio.iscoroutinefunction(graph_get_all): - graph_task = asyncio.create_task( - graph_get_all(effective_filters, limit) - ) + graph_task = asyncio.create_task(graph_get_all(effective_filters, limit)) else: - graph_task = asyncio.create_task( - asyncio.to_thread(graph_get_all, effective_filters, limit) - ) + graph_task = asyncio.create_task(asyncio.to_thread(graph_get_all, effective_filters, limit)) results_dict = {} if graph_task: - vector_store_result, graph_entities_result = await asyncio.gather( - vector_store_task, graph_task - ) - results_dict.update( - {"results": vector_store_result, "relations": graph_entities_result} - ) + vector_store_result, graph_entities_result = await asyncio.gather(vector_store_task, graph_task) + results_dict.update({"results": vector_store_result, "relations": graph_entities_result}) else: results_dict.update({"results": await vector_store_task}) diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index 9aeaddc80..2414746c7 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 +from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher try: from langchain_memgraph.graphs.memgraph import Memgraph @@ -40,13 +40,20 @@ class MemoryGraph: {"enable_embeddings": True}, ) - self.llm_provider = "openai_structured" - if self.config.llm.provider: + # Default to openai if no specific provider is configured + self.llm_provider = "openai" + if self.config.llm and self.config.llm.provider: self.llm_provider = self.config.llm.provider - if self.config.graph_store.llm: + if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider: self.llm_provider = self.config.graph_store.llm.provider - self.llm = LlmFactory.create(self.llm_provider, self.config.llm.config) + # Get LLM config with proper null checks + llm_config = None + if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"): + llm_config = self.config.graph_store.llm.config + elif hasattr(self.config.llm, "config"): + llm_config = self.config.llm.config + self.llm = LlmFactory.create(self.llm_provider, llm_config) self.user_id = None self.threshold = 0.7 @@ -535,7 +542,8 @@ class MemoryGraph: 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(" ", "_") + # 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 diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index 00a0a36b1..466abdbe3 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -131,3 +131,54 @@ def process_telemetry_filters(filters): encoded_ids["run_id"] = hashlib.md5(filters["run_id"].encode()).hexdigest() return list(filters.keys()), encoded_ids + + +def sanitize_relationship_for_cypher(relationship) -> str: + """Sanitize relationship text for Cypher queries by replacing problematic characters.""" + char_map = { + "...": "_ellipsis_", + "…": "_ellipsis_", + "。": "_period_", + ",": "_comma_", + ";": "_semicolon_", + ":": "_colon_", + "!": "_exclamation_", + "?": "_question_", + "(": "_lparen_", + ")": "_rparen_", + "【": "_lbracket_", + "】": "_rbracket_", + "《": "_langle_", + "》": "_rangle_", + "'": "_apostrophe_", + '"': "_quote_", + "\\": "_backslash_", + "/": "_slash_", + "|": "_pipe_", + "&": "_ampersand_", + "=": "_equals_", + "+": "_plus_", + "*": "_asterisk_", + "^": "_caret_", + "%": "_percent_", + "$": "_dollar_", + "#": "_hash_", + "@": "_at_", + "!": "_bang_", + "?": "_question_", + "(": "_lparen_", + ")": "_rparen_", + "[": "_lbracket_", + "]": "_rbracket_", + "{": "_lbrace_", + "}": "_rbrace_", + "<": "_langle_", + ">": "_rangle_", + } + + # Apply replacements and clean up + sanitized = relationship + for old, new in char_map.items(): + sanitized = sanitized.replace(old, new) + + return re.sub(r"_+", "_", sanitized).strip("_") diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index dd798d44d..288cd8bbd 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -196,7 +196,7 @@ class GraphStoreFactory: Factory for creating MemoryGraph instances for different graph store providers. Usage: GraphStoreFactory.create(provider_name, config) """ - + provider_to_class = { "memgraph": "mem0.memory.memgraph_memory.MemoryGraph", "neptune": "mem0.graphs.neptune.main.MemoryGraph",