fixes memgraph async attirbute error (#3209)

This commit is contained in:
Parshva Daftari
2025-08-02 01:49:08 +05:30
committed by GitHub
parent 89b67e0834
commit 7d1d0ca806
2 changed files with 32 additions and 12 deletions
+10 -12
View File
@@ -31,7 +31,12 @@ from mem0.memory.utils import (
process_telemetry_filters,
remove_code_blocks,
)
from mem0.utils.factory import EmbedderFactory, LlmFactory, VectorStoreFactory
from mem0.utils.factory import (
EmbedderFactory,
GraphStoreFactory,
LlmFactory,
VectorStoreFactory,
)
def _build_filters_and_metadata(
@@ -136,14 +141,8 @@ class Memory(MemoryBase):
self.enable_graph = False
if self.config.graph_store.config:
if self.config.graph_store.provider == "memgraph":
from mem0.memory.memgraph_memory import MemoryGraph
elif self.config.graph_store.provider == "neptune":
from mem0.graphs.neptune.main import MemoryGraph
else:
from mem0.memory.graph_memory import MemoryGraph
self.graph = MemoryGraph(self.config)
provider = self.config.graph_store.provider
self.graph = GraphStoreFactory.create(provider, self.config)
self.enable_graph = True
else:
self.graph = None
@@ -989,9 +988,8 @@ class AsyncMemory(MemoryBase):
self.enable_graph = False
if self.config.graph_store.config:
from mem0.memory.graph_memory import MemoryGraph
self.graph = MemoryGraph(self.config)
provider = self.config.graph_store.provider
self.graph = GraphStoreFactory.create(provider, self.config)
self.enable_graph = True
else:
self.graph = None
+22
View File
@@ -106,3 +106,25 @@ class VectorStoreFactory:
def reset(cls, instance):
instance.reset()
return instance
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",
"default": "mem0.memory.graph_memory.MemoryGraph",
}
@classmethod
def create(cls, provider_name, config):
class_type = cls.provider_to_class.get(provider_name, cls.provider_to_class["default"])
try:
GraphClass = load_class(class_type)
except (ImportError, AttributeError) as e:
raise ImportError(f"Could not import MemoryGraph for provider '{provider_name}': {e}")
return GraphClass(config)