From 7d1d0ca8063d03e58265dacd7246229696b1260e Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Sat, 2 Aug 2025 01:49:08 +0530 Subject: [PATCH] fixes memgraph async attirbute error (#3209) --- mem0/memory/main.py | 22 ++++++++++------------ mem0/utils/factory.py | 22 ++++++++++++++++++++++ 2 files changed, 32 insertions(+), 12 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 562a13b2f..bef176699 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -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 diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 45491340b..5fe04fc6d 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -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)