diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 35850078a..1804e7831 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -304,6 +304,23 @@ class Memory(MemoryBase): ) capture_event("mem0.init", self, {"sync_type": "sync"}) + def close(self): + """Release resources held by this Memory instance (SQLite connections, etc.). + + The global telemetry singleton is intentionally *not* shut down here + because it is shared across all Memory instances in the process. It is + cleaned up automatically at process exit via an ``atexit`` handler. + """ + if hasattr(self, "db") and self.db is not None: + self.db.close() + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + self.close() + return False + @classmethod def from_config(cls, config_dict: Dict[str, Any]): try: @@ -1425,6 +1442,18 @@ class AsyncMemory(MemoryBase): self._telemetry_vector_store = VectorStoreFactory.create(self.config.vector_store.provider, telemetry_config) capture_event("mem0.init", self, {"sync_type": "async"}) + def close(self): + """Release resources held by this AsyncMemory instance.""" + if hasattr(self, "db") and self.db is not None: + self.db.close() + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + self.close() + return False + @classmethod def from_config(cls, config_dict: Dict[str, Any]): try: diff --git a/mem0/memory/telemetry.py b/mem0/memory/telemetry.py index 45b2f3923..e35518b92 100644 --- a/mem0/memory/telemetry.py +++ b/mem0/memory/telemetry.py @@ -1,7 +1,9 @@ +import atexit import logging import os import platform import sys +import threading from posthog import Posthog @@ -55,20 +57,62 @@ class AnonymousTelemetry: def close(self): if self.posthog is not None: self.posthog.shutdown() + self.posthog = None +# Thread-safe lazy singleton for OSS telemetry. +# A single AnonymousTelemetry instance (and its underlying PostHog client / +# background thread) is reused for the lifetime of the process instead of +# creating a new one on every capture_event() call. The singleton is shut down +# once at process exit via an atexit handler. +_oss_telemetry_instance = None +_oss_telemetry_lock = threading.Lock() +_oss_telemetry_shutting_down = False + + +def _get_oss_telemetry(): + """Return the process-wide AnonymousTelemetry singleton, creating it on first call. + + Returns None after _shutdown_oss_telemetry() has run (interpreter exit). + """ + global _oss_telemetry_instance + if _oss_telemetry_shutting_down: + return None + if _oss_telemetry_instance is not None: + return _oss_telemetry_instance + + with _oss_telemetry_lock: + if _oss_telemetry_shutting_down: + return None + # Double-checked locking + if _oss_telemetry_instance is not None: + return _oss_telemetry_instance + _oss_telemetry_instance = AnonymousTelemetry() + atexit.register(_shutdown_oss_telemetry) + return _oss_telemetry_instance + + +def _shutdown_oss_telemetry(): + global _oss_telemetry_instance, _oss_telemetry_shutting_down + with _oss_telemetry_lock: + _oss_telemetry_shutting_down = True + if _oss_telemetry_instance is not None: + _oss_telemetry_instance.close() + _oss_telemetry_instance = None + + +# Module-level client telemetry singleton (used by capture_client_event). client_telemetry = AnonymousTelemetry() +atexit.register(client_telemetry.close) def capture_event(event_name, memory_instance, additional_data=None): if not MEM0_TELEMETRY: return - oss_telemetry = AnonymousTelemetry( - vector_store=memory_instance._telemetry_vector_store - if hasattr(memory_instance, "_telemetry_vector_store") - else None, - ) + oss_telemetry = _get_oss_telemetry() + if oss_telemetry is None: + return event_data = { "collection": memory_instance.collection_name, diff --git a/tests/test_telemetry.py b/tests/test_telemetry.py index b3b537104..44e8191ad 100644 --- a/tests/test_telemetry.py +++ b/tests/test_telemetry.py @@ -1,3 +1,4 @@ +import asyncio import threading from unittest.mock import MagicMock, patch @@ -19,11 +20,16 @@ class TestTelemetryDisabled: assert at.user_id is None def test_capture_event_noop_when_disabled(self): - """capture_event() should return immediately without creating AnonymousTelemetry.""" - with patch.object(telemetry_module, "MEM0_TELEMETRY", False): - with patch("mem0.memory.telemetry.AnonymousTelemetry") as mock_cls: + """capture_event() should return immediately without touching the singleton.""" + saved = telemetry_module._oss_telemetry_instance + try: + telemetry_module._oss_telemetry_instance = None + with patch.object(telemetry_module, "MEM0_TELEMETRY", False): telemetry_module.capture_event("test.event", MagicMock()) - mock_cls.assert_not_called() + # Singleton must not have been initialised + assert telemetry_module._oss_telemetry_instance is None + finally: + telemetry_module._oss_telemetry_instance = saved def test_capture_client_event_noop_when_disabled(self): """capture_client_event() should return immediately without calling posthog.""" @@ -69,11 +75,10 @@ class TestTelemetryEnabled: assert at.user_id == "test-user" def test_capture_event_sends_when_enabled(self): - """capture_event() should create AnonymousTelemetry and call capture when enabled.""" + """capture_event() should use the singleton and call capture when enabled.""" with patch.object(telemetry_module, "MEM0_TELEMETRY", True): - with patch("mem0.memory.telemetry.AnonymousTelemetry") as mock_cls: - mock_at = MagicMock() - mock_cls.return_value = mock_at + mock_at = MagicMock() + with patch.object(telemetry_module, "_oss_telemetry_instance", mock_at): mock_memory = MagicMock() mock_memory.config.graph_store.config = None mock_memory.api_version = "v1" @@ -91,6 +96,297 @@ class TestTelemetryEnabled: mock_client_telemetry.capture_event.assert_called_once() +class TestAnonymousTelemetryClose: + """Verify AnonymousTelemetry.close() edge cases.""" + + def test_close_calls_posthog_shutdown(self): + """close() should call posthog.shutdown().""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls: + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + at = telemetry_module.AnonymousTelemetry() + mock_ph = mock_posthog_cls.return_value + at.close() + mock_ph.shutdown.assert_called_once() + + def test_close_sets_posthog_to_none(self): + """close() should set posthog to None to prevent double-shutdown.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog"): + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + at = telemetry_module.AnonymousTelemetry() + at.close() + assert at.posthog is None + + def test_double_close_is_safe(self): + """Calling close() twice should not raise.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls: + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + at = telemetry_module.AnonymousTelemetry() + mock_ph = mock_posthog_cls.return_value + at.close() + at.close() # should not raise + # shutdown only called once because posthog was set to None after first close + mock_ph.shutdown.assert_called_once() + + def test_capture_after_close_is_noop(self): + """capture_event() should be a no-op after close() (posthog is None).""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls: + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + at = telemetry_module.AnonymousTelemetry() + mock_ph = mock_posthog_cls.return_value + at.close() + mock_ph.reset_mock() + at.capture_event("test.event", {"key": "value"}) + mock_ph.capture.assert_not_called() + + +class TestTelemetrySingleton: + """Verify the OSS telemetry singleton behaviour.""" + + def setup_method(self): + # Reset singleton state before each test + telemetry_module._oss_telemetry_instance = None + telemetry_module._oss_telemetry_shutting_down = False + + def teardown_method(self): + telemetry_module._oss_telemetry_instance = None + telemetry_module._oss_telemetry_shutting_down = False + + def test_singleton_reuses_instance(self): + """_get_oss_telemetry() should return the same instance on repeated calls.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog"): + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register"): + first = telemetry_module._get_oss_telemetry() + second = telemetry_module._get_oss_telemetry() + assert first is second + + def test_singleton_created_only_once_across_threads(self): + """Only one AnonymousTelemetry should be created even under concurrent access.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog") as mock_posthog: + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register"): + instances = [] + + def grab(): + instances.append(telemetry_module._get_oss_telemetry()) + + threads = [threading.Thread(target=grab) for _ in range(20)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert len(set(id(i) for i in instances)) == 1 + assert mock_posthog.call_count == 1 + + def test_atexit_registered_once(self): + """atexit.register should be called exactly once for the singleton.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog"): + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register") as mock_atexit: + telemetry_module._get_oss_telemetry() + telemetry_module._get_oss_telemetry() + telemetry_module._get_oss_telemetry() + # Only one atexit registration despite multiple calls + mock_atexit.assert_called_once_with(telemetry_module._shutdown_oss_telemetry) + + def test_capture_event_does_not_create_new_instance_each_call(self): + """capture_event() should not create a new AnonymousTelemetry per call.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog"): + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register"): + mock_memory = MagicMock() + mock_memory.config.graph_store.config = None + mock_memory.api_version = "v1" + + telemetry_module.capture_event("e1", mock_memory) + first = telemetry_module._oss_telemetry_instance + + telemetry_module.capture_event("e2", mock_memory) + second = telemetry_module._oss_telemetry_instance + + assert first is second + + def test_posthog_constructed_once_across_many_capture_event_calls(self): + """The core leak fix: Posthog() should only be called once no matter how + many times capture_event() is invoked.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls: + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register"): + mock_memory = MagicMock() + mock_memory.config.graph_store.config = None + mock_memory.api_version = "v1" + + for i in range(50): + telemetry_module.capture_event(f"event_{i}", mock_memory) + + # Only ONE Posthog client created, not 50 + assert mock_posthog_cls.call_count == 1 + + def test_shutdown_clears_singleton(self): + """_shutdown_oss_telemetry() should close and clear the singleton.""" + mock_at = MagicMock() + telemetry_module._oss_telemetry_instance = mock_at + + telemetry_module._shutdown_oss_telemetry() + + mock_at.close.assert_called_once() + assert telemetry_module._oss_telemetry_instance is None + + def test_shutdown_idempotent(self): + """Calling _shutdown_oss_telemetry() twice should not raise.""" + mock_at = MagicMock() + telemetry_module._oss_telemetry_instance = mock_at + + telemetry_module._shutdown_oss_telemetry() + telemetry_module._shutdown_oss_telemetry() # should not raise + + mock_at.close.assert_called_once() + + def test_shutdown_noop_when_no_instance(self): + """_shutdown_oss_telemetry() should be a no-op when singleton was never created.""" + assert telemetry_module._oss_telemetry_instance is None + telemetry_module._shutdown_oss_telemetry() # should not raise + + def test_capture_event_noop_when_disabled_with_singleton(self): + """capture_event() should not initialise the singleton when telemetry is off.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", False): + telemetry_module.capture_event("test.event", MagicMock()) + assert telemetry_module._oss_telemetry_instance is None + + def test_no_new_instance_after_shutdown(self): + """After _shutdown_oss_telemetry(), _get_oss_telemetry() should not create a new instance.""" + with patch.object(telemetry_module, "MEM0_TELEMETRY", True): + with patch("mem0.memory.telemetry.Posthog"): + with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="u"): + with patch("atexit.register"): + # Create and then shut down the singleton + telemetry_module._get_oss_telemetry() + telemetry_module._shutdown_oss_telemetry() + + # After shutdown, getting telemetry should return None, not a new instance + result = telemetry_module._get_oss_telemetry() + assert result is None + assert telemetry_module._oss_telemetry_instance is None + + +class TestMemoryLifecycle: + """Verify Memory.close() and context manager support.""" + + def _make_mock_memory(self): + """Create a Memory-like object with a mock db for testing close().""" + from mem0.memory.main import Memory + + with patch.object(Memory, "__init__", lambda self: None): + m = Memory.__new__(Memory) + m.db = MagicMock() + return m + + def test_close_calls_db_close(self): + m = self._make_mock_memory() + m.close() + m.db.close.assert_called_once() + + def test_double_close_is_safe(self): + """close() should be safe to call twice (SQLiteManager.close sets connection=None).""" + m = self._make_mock_memory() + m.close() + # After first close, db is still the mock but SQLiteManager.close() handles + # the None-connection case internally. Simulate that by making db.close a no-op. + m.close() # should not raise + + def test_close_when_db_is_none(self): + """close() should not raise if db was already None.""" + m = self._make_mock_memory() + m.db = None + m.close() # should not raise + + def test_close_when_db_not_set(self): + """close() should not raise if __init__ failed before setting db.""" + from mem0.memory.main import Memory + + with patch.object(Memory, "__init__", lambda self: None): + m = Memory.__new__(Memory) + # db attribute not set at all + m.close() # should not raise due to hasattr guard + + def test_context_manager(self): + """Memory should support `with` statement and close on exit.""" + m = self._make_mock_memory() + with m: + pass + m.db.close.assert_called_once() + + def test_context_manager_closes_on_exception(self): + """Memory should close even if the with-block raises.""" + m = self._make_mock_memory() + with pytest.raises(ValueError): + with m: + raise ValueError("boom") + m.db.close.assert_called_once() + + +class TestAsyncMemoryLifecycle: + """Verify AsyncMemory.close() and async context manager support.""" + + def _make_mock_async_memory(self): + from mem0.memory.main import AsyncMemory + + with patch.object(AsyncMemory, "__init__", lambda self: None): + m = AsyncMemory.__new__(AsyncMemory) + m.db = MagicMock() + return m + + def test_close_calls_db_close(self): + m = self._make_mock_async_memory() + m.close() + m.db.close.assert_called_once() + + def test_close_when_db_is_none(self): + m = self._make_mock_async_memory() + m.db = None + m.close() # should not raise + + def test_close_when_db_not_set(self): + from mem0.memory.main import AsyncMemory + + with patch.object(AsyncMemory, "__init__", lambda self: None): + m = AsyncMemory.__new__(AsyncMemory) + m.close() # should not raise + + def test_async_context_manager(self): + """AsyncMemory should support `async with` and close on exit.""" + m = self._make_mock_async_memory() + + async def run(): + async with m: + pass + + asyncio.run(run()) + m.db.close.assert_called_once() + + def test_async_context_manager_closes_on_exception(self): + """AsyncMemory should close even if the async with-block raises.""" + m = self._make_mock_async_memory() + + async def run(): + async with m: + raise ValueError("boom") + + with pytest.raises(ValueError): + asyncio.run(run()) + m.db.close.assert_called_once() + + class TestTelemetryEnvVar: """Verify the MEM0_TELEMETRY env var parsing logic."""