fixes memgraph async attirbute error (#3209)
This commit is contained in:
+10
-12
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user