diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 99bc0b1e5..06e6b5976 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -24,7 +24,7 @@ from mem0.exceptions import ValidationError as Mem0ValidationError from mem0.memory.base import MemoryBase from mem0.memory.setup import mem0_dir, setup_config from mem0.memory.storage import SQLiteManager -from mem0.memory.telemetry import capture_event +from mem0.memory.telemetry import MEM0_TELEMETRY, capture_event from mem0.memory.utils import ( extract_json, get_fact_retrieval_messages, @@ -37,8 +37,8 @@ from mem0.utils.factory import ( EmbedderFactory, GraphStoreFactory, LlmFactory, - VectorStoreFactory, RerankerFactory, + VectorStoreFactory, ) # Suppress SWIG deprecation warnings globally @@ -204,32 +204,33 @@ class Memory(MemoryBase): self.enable_graph = True else: self.graph = None - # Create telemetry config manually to avoid deepcopy issues with thread locks - telemetry_config_dict = {} - if hasattr(self.config.vector_store.config, 'model_dump'): - # For pydantic models - telemetry_config_dict = self.config.vector_store.config.model_dump() - else: - # For other objects, manually copy common attributes - for attr in ['host', 'port', 'path', 'api_key', 'index_name', 'dimension', 'metric']: - if hasattr(self.config.vector_store.config, attr): - telemetry_config_dict[attr] = getattr(self.config.vector_store.config, attr) + if MEM0_TELEMETRY: + # Create telemetry config manually to avoid deepcopy issues with thread locks + telemetry_config_dict = {} + if hasattr(self.config.vector_store.config, 'model_dump'): + # For pydantic models + telemetry_config_dict = self.config.vector_store.config.model_dump() + else: + # For other objects, manually copy common attributes + for attr in ['host', 'port', 'path', 'api_key', 'index_name', 'dimension', 'metric']: + if hasattr(self.config.vector_store.config, attr): + telemetry_config_dict[attr] = getattr(self.config.vector_store.config, attr) - # Override collection name for telemetry - telemetry_config_dict['collection_name'] = "mem0migrations" + # Override collection name for telemetry + telemetry_config_dict['collection_name'] = "mem0migrations" - # Set path for file-based vector stores - telemetry_config = _safe_deepcopy_config(self.config.vector_store.config) - if self.config.vector_store.provider in ["faiss", "qdrant"]: - provider_path = f"migrations_{self.config.vector_store.provider}" - telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path) - os.makedirs(telemetry_config_dict['path'], exist_ok=True) + # Set path for file-based vector stores + telemetry_config = _safe_deepcopy_config(self.config.vector_store.config) + if self.config.vector_store.provider in ["faiss", "qdrant"]: + provider_path = f"migrations_{self.config.vector_store.provider}" + telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path) + os.makedirs(telemetry_config_dict['path'], exist_ok=True) - # Create the config object using the same class as the original - telemetry_config = self.config.vector_store.config.__class__(**telemetry_config_dict) - self._telemetry_vector_store = VectorStoreFactory.create( - self.config.vector_store.provider, telemetry_config - ) + # Create the config object using the same class as the original + telemetry_config = self.config.vector_store.config.__class__(**telemetry_config_dict) + self._telemetry_vector_store = VectorStoreFactory.create( + self.config.vector_store.provider, telemetry_config + ) capture_event("mem0.init", self, {"sync_type": "sync"}) @classmethod @@ -1272,14 +1273,14 @@ class AsyncMemory(MemoryBase): else: self.graph = None - telemetry_config = _safe_deepcopy_config(self.config.vector_store.config) - telemetry_config.collection_name = "mem0migrations" - if self.config.vector_store.provider in ["faiss", "qdrant"]: - provider_path = f"migrations_{self.config.vector_store.provider}" - telemetry_config.path = os.path.join(mem0_dir, provider_path) - os.makedirs(telemetry_config.path, exist_ok=True) - self._telemetry_vector_store = VectorStoreFactory.create(self.config.vector_store.provider, telemetry_config) - + if MEM0_TELEMETRY: + telemetry_config = _safe_deepcopy_config(self.config.vector_store.config) + telemetry_config.collection_name = "mem0migrations" + if self.config.vector_store.provider in ["faiss", "qdrant"]: + provider_path = f"migrations_{self.config.vector_store.provider}" + telemetry_config.path = os.path.join(mem0_dir, provider_path) + os.makedirs(telemetry_config.path, exist_ok=True) + self._telemetry_vector_store = VectorStoreFactory.create(self.config.vector_store.provider, telemetry_config) capture_event("mem0.init", self, {"sync_type": "async"}) @classmethod diff --git a/tests/test_main.py b/tests/test_main.py index 2f548e315..4a38f6b3f 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -296,3 +296,43 @@ def test_custom_prompts(memory_custom_instance): messages=[{"role": "user", "content": mock_get_update_memory_messages.return_value}], response_format={"type": "json_object"}, ) + + +def test_no_telemetry_vector_store_when_disabled(): + """VectorStoreFactory should only be called once (for user data) when telemetry is disabled.""" + with ( + patch("mem0.memory.main.MEM0_TELEMETRY", False), + patch("mem0.utils.factory.EmbedderFactory") as mock_embedder, + patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store, + patch("mem0.utils.factory.LlmFactory") as mock_llm, + patch("mem0.memory.telemetry.capture_event"), + ): + mock_embedder.create.return_value = Mock() + mock_vector_store.create.return_value = Mock() + mock_llm.create.return_value = Mock() + + config = MemoryConfig(version="v1.1") + Memory(config) + + # VectorStoreFactory.create should be called exactly once — for user data only, not telemetry + assert mock_vector_store.create.call_count == 1 + + +def test_telemetry_vector_store_created_when_enabled(): + """VectorStoreFactory should be called twice (user data + telemetry) when telemetry is enabled.""" + with ( + patch("mem0.memory.main.MEM0_TELEMETRY", True), + patch("mem0.utils.factory.EmbedderFactory") as mock_embedder, + patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store, + patch("mem0.utils.factory.LlmFactory") as mock_llm, + patch("mem0.memory.telemetry.capture_event"), + ): + mock_embedder.create.return_value = Mock() + mock_vector_store.create.return_value = Mock() + mock_llm.create.return_value = Mock() + + config = MemoryConfig(version="v1.1") + Memory(config) + + # VectorStoreFactory.create should be called twice — user data + telemetry + assert mock_vector_store.create.call_count == 2