Added sanitation for better relationship mapping (#3300)
This commit is contained in:
@@ -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
@@ -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
@@ -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})
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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("_")
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user