Added sanitation for better relationship mapping (#3300)

This commit is contained in:
Parshva Daftari
2025-08-13 21:36:05 +05:30
committed by GitHub
parent 7159221d2e
commit 1bc31b6ae3
6 changed files with 100 additions and 59 deletions
+6 -6
View File
@@ -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,
},
+23 -31
View File
@@ -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"):
+5 -15
View File
@@ -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})
+14 -6
View File
@@ -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
+51
View File
@@ -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("_")
+1 -1
View File
@@ -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",