Compare commits

...

1 Commits

Author SHA1 Message Date
kartik-mem0 6d64bbc4cf refactor: add close() to Memory and thread-safe telemetry singleton 2026-03-23 11:21:27 +05:30
4 changed files with 380 additions and 23 deletions
+81
View File
@@ -1320,6 +1320,47 @@ class Memory(MemoryBase):
)
capture_event("mem0.reset", self, {"sync_type": "sync"})
def close(self):
"""Release all resources held by this Memory instance.
Shuts down telemetry background threads, closes the SQLite
history-store connection, and closes the vector store client
if it exposes a close() method.
Safe to call multiple times. Use as a context manager for
automatic cleanup::
with Memory.from_config(config) as m:
m.add(...)
"""
from mem0.memory.telemetry import shutdown_telemetry
try:
shutdown_telemetry()
except Exception:
logger.debug("Error shutting down telemetry", exc_info=True)
if hasattr(self, "db") and self.db is not None:
try:
self.db.close()
except Exception:
logger.debug("Error closing database", exc_info=True)
if hasattr(self, "vector_store") and self.vector_store is not None:
_close = getattr(self.vector_store, "close", None)
if callable(_close):
try:
_close()
except Exception:
logger.debug("Error closing vector store", exc_info=True)
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()
return False
def chat(self, query):
raise NotImplementedError("Chat function not implemented yet.")
@@ -2422,5 +2463,45 @@ class AsyncMemory(MemoryBase):
)
capture_event("mem0.reset", self, {"sync_type": "async"})
def close(self):
"""Release all resources held by this AsyncMemory instance.
Shuts down telemetry background threads, closes the SQLite
history-store connection, and closes the vector store client
if it exposes a close() method.
Safe to call multiple times. Use as an async context manager::
async with await AsyncMemory.from_config(config) as m:
await m.add(...)
"""
from mem0.memory.telemetry import shutdown_telemetry
try:
shutdown_telemetry()
except Exception:
logger.debug("Error shutting down telemetry", exc_info=True)
if hasattr(self, "db") and self.db is not None:
try:
self.db.close()
except Exception:
logger.debug("Error closing database", exc_info=True)
if hasattr(self, "vector_store") and self.vector_store is not None:
_close = getattr(self.vector_store, "close", None)
if callable(_close):
try:
_close()
except Exception:
logger.debug("Error closing vector store", exc_info=True)
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
self.close()
return False
async def chat(self, query):
raise NotImplementedError("Chat function not implemented yet.")
+54 -6
View File
@@ -1,7 +1,9 @@
import atexit
import logging
import os
import platform
import sys
import threading
from posthog import Posthog
@@ -57,18 +59,62 @@ class AnonymousTelemetry:
self.posthog.shutdown()
client_telemetry = AnonymousTelemetry()
# Thread-safe lazy singletons
_telemetry_lock = threading.Lock()
_oss_telemetry = None
_client_telemetry = None
def _get_oss_telemetry(vector_store=None):
"""Return the singleton OSS telemetry instance, creating it on first call.
The vector_store parameter is only used on the first call for
get_or_create_user_id(). Subsequent calls reuse the existing instance
since the user ID is stable once created.
"""
global _oss_telemetry
if _oss_telemetry is None:
with _telemetry_lock:
if _oss_telemetry is None:
_oss_telemetry = AnonymousTelemetry(vector_store=vector_store)
return _oss_telemetry
def _get_client_telemetry():
"""Return the singleton client telemetry instance, creating it on first call."""
global _client_telemetry
if _client_telemetry is None:
with _telemetry_lock:
if _client_telemetry is None:
_client_telemetry = AnonymousTelemetry()
return _client_telemetry
def shutdown_telemetry():
"""Shut down all telemetry singletons. Safe to call multiple times.
Singletons are lazily re-created on next use, so this is safe even
when multiple Memory instances exist.
"""
global _oss_telemetry, _client_telemetry
with _telemetry_lock:
if _oss_telemetry is not None:
_oss_telemetry.close()
_oss_telemetry = None
if _client_telemetry is not None:
_client_telemetry.close()
_client_telemetry = None
atexit.register(shutdown_telemetry)
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,
)
vector_store = getattr(memory_instance, "_telemetry_vector_store", None)
oss_telemetry = _get_oss_telemetry(vector_store=vector_store)
event_data = {
"collection": memory_instance.collection_name,
@@ -92,6 +138,8 @@ def capture_client_event(event_name, instance, additional_data=None):
if not MEM0_TELEMETRY:
return
client_telemetry = _get_client_telemetry()
event_data = {
"function": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
}
+164
View File
@@ -0,0 +1,164 @@
"""Tests for Memory resource cleanup: close() and context manager."""
from unittest.mock import MagicMock, patch
import pytest
@pytest.fixture
def memory_instance():
"""Create a Memory instance with all dependencies mocked."""
with (
patch("mem0.memory.main.EmbedderFactory") as mock_embedder,
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
patch("mem0.memory.main.LlmFactory") as mock_llm,
patch("mem0.memory.telemetry.capture_event"),
patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store,
):
mock_embedder.create.return_value = MagicMock()
mock_embedder.create.return_value.config.embedding_dims = 1536
mock_vector_store.create.return_value = MagicMock()
mock_llm.create.return_value = MagicMock()
mock_graph_store.create.return_value = MagicMock()
from mem0.memory.main import Memory
m = Memory()
yield m
class TestMemoryClose:
"""Verify Memory.close() releases resources."""
def test_close_shuts_down_telemetry(self, memory_instance):
with patch("mem0.memory.telemetry.shutdown_telemetry") as mock_shutdown:
memory_instance.close()
mock_shutdown.assert_called_once()
def test_close_closes_db(self, memory_instance):
memory_instance.db = MagicMock()
memory_instance.close()
memory_instance.db.close.assert_called_once()
def test_close_closes_vector_store_if_available(self, memory_instance):
mock_vs = MagicMock(spec=["close"])
memory_instance.vector_store = mock_vs
memory_instance.close()
mock_vs.close.assert_called_once()
def test_close_skips_vector_store_without_close(self, memory_instance):
mock_vs = MagicMock(spec=[])
memory_instance.vector_store = mock_vs
memory_instance.close() # should not raise
def test_close_is_idempotent(self, memory_instance):
memory_instance.close()
memory_instance.close() # should not raise
def test_close_continues_on_db_error(self, memory_instance):
memory_instance.db = MagicMock()
memory_instance.db.close.side_effect = Exception("db error")
memory_instance.close() # should not raise
class TestMemoryContextManager:
"""Verify Memory works as a context manager."""
def test_enter_returns_self(self, memory_instance):
assert memory_instance.__enter__() is memory_instance
def test_exit_calls_close(self, memory_instance):
memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
memory_instance.__exit__(None, None, None)
memory_instance.db.close.assert_called_once()
def test_does_not_suppress_exceptions(self, memory_instance):
assert memory_instance.__exit__(ValueError, ValueError("x"), None) is False
def test_with_statement(self, memory_instance):
memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
with memory_instance:
pass
memory_instance.db.close.assert_called_once()
def test_with_statement_on_exception(self, memory_instance):
memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
with pytest.raises(RuntimeError):
with memory_instance:
raise RuntimeError("boom")
memory_instance.db.close.assert_called_once()
@pytest.fixture
def async_memory_instance():
"""Create an AsyncMemory instance with all dependencies mocked."""
with (
patch("mem0.memory.main.EmbedderFactory") as mock_embedder,
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
patch("mem0.memory.main.LlmFactory") as mock_llm,
patch("mem0.memory.telemetry.capture_event"),
patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store,
):
mock_embedder.create.return_value = MagicMock()
mock_embedder.create.return_value.config.embedding_dims = 1536
mock_vector_store.create.return_value = MagicMock()
mock_llm.create.return_value = MagicMock()
mock_graph_store.create.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
m = AsyncMemory()
yield m
class TestAsyncMemoryClose:
"""Verify AsyncMemory.close() releases resources."""
def test_close_shuts_down_telemetry(self, async_memory_instance):
with patch("mem0.memory.telemetry.shutdown_telemetry") as mock_shutdown:
async_memory_instance.close()
mock_shutdown.assert_called_once()
def test_close_closes_db(self, async_memory_instance):
async_memory_instance.db = MagicMock()
async_memory_instance.close()
async_memory_instance.db.close.assert_called_once()
def test_close_is_idempotent(self, async_memory_instance):
async_memory_instance.close()
async_memory_instance.close() # should not raise
class TestAsyncMemoryContextManager:
"""Verify AsyncMemory works as an async context manager."""
@pytest.mark.asyncio
async def test_aenter_returns_self(self, async_memory_instance):
result = await async_memory_instance.__aenter__()
assert result is async_memory_instance
@pytest.mark.asyncio
async def test_aexit_calls_close(self, async_memory_instance):
async_memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
await async_memory_instance.__aexit__(None, None, None)
async_memory_instance.db.close.assert_called_once()
@pytest.mark.asyncio
async def test_async_with_statement(self, async_memory_instance):
async_memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
async with async_memory_instance:
pass
async_memory_instance.db.close.assert_called_once()
@pytest.mark.asyncio
async def test_async_with_statement_on_exception(self, async_memory_instance):
async_memory_instance.db = MagicMock()
with patch("mem0.memory.telemetry.shutdown_telemetry"):
with pytest.raises(RuntimeError):
async with async_memory_instance:
raise RuntimeError("boom")
async_memory_instance.db.close.assert_called_once()
+81 -17
View File
@@ -30,9 +30,10 @@ class TestTelemetryDisabled:
with patch.object(telemetry_module, "MEM0_TELEMETRY", False):
mock_instance = MagicMock()
mock_client_telemetry = MagicMock()
with patch.object(telemetry_module, "client_telemetry", mock_client_telemetry):
telemetry_module.capture_client_event("test.event", mock_instance)
mock_client_telemetry.capture_event.assert_not_called()
telemetry_module._client_telemetry = mock_client_telemetry
telemetry_module.capture_client_event("test.event", mock_instance)
mock_client_telemetry.capture_event.assert_not_called()
telemetry_module._client_telemetry = None
def test_instance_capture_event_noop_when_posthog_is_none(self):
"""AnonymousTelemetry.capture_event() should be a no-op when posthog is None."""
@@ -69,26 +70,89 @@ 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 singleton AnonymousTelemetry 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_memory = MagicMock()
mock_memory.config.graph_store.config = None
mock_memory.api_version = "v1"
telemetry_module.capture_event("test.event", mock_memory)
mock_at.capture_event.assert_called_once()
mock_at = MagicMock()
telemetry_module._oss_telemetry = mock_at
mock_memory = MagicMock()
mock_memory.config.graph_store.config = None
mock_memory.api_version = "v1"
telemetry_module.capture_event("test.event", mock_memory)
mock_at.capture_event.assert_called_once()
telemetry_module._oss_telemetry = None
def test_capture_client_event_sends_when_enabled(self):
"""capture_client_event() should call client_telemetry.capture_event when enabled."""
with patch.object(telemetry_module, "MEM0_TELEMETRY", True):
mock_client_telemetry = MagicMock()
with patch.object(telemetry_module, "client_telemetry", mock_client_telemetry):
mock_instance = MagicMock()
mock_instance.user_email = "test@example.com"
telemetry_module.capture_client_event("test.event", mock_instance)
mock_client_telemetry.capture_event.assert_called_once()
telemetry_module._client_telemetry = mock_client_telemetry
mock_instance = MagicMock()
mock_instance.user_email = "test@example.com"
telemetry_module.capture_client_event("test.event", mock_instance)
mock_client_telemetry.capture_event.assert_called_once()
telemetry_module._client_telemetry = None
def test_capture_event_reuses_singleton(self):
"""capture_event() must reuse a single AnonymousTelemetry, not create one per call."""
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
telemetry_module._oss_telemetry = None
mock_memory = MagicMock()
mock_memory.config.graph_store.config = None
mock_memory.api_version = "v1"
telemetry_module.capture_event("event1", mock_memory)
telemetry_module.capture_event("event2", mock_memory)
assert mock_cls.call_count == 1
assert mock_at.capture_event.call_count == 2
telemetry_module._oss_telemetry = None
class TestTelemetryShutdown:
"""Verify shutdown_telemetry() cleans up all singleton instances."""
def test_shutdown_calls_close_on_singletons(self):
mock_oss = MagicMock()
mock_client = MagicMock()
telemetry_module._oss_telemetry = mock_oss
telemetry_module._client_telemetry = mock_client
telemetry_module.shutdown_telemetry()
mock_oss.close.assert_called_once()
mock_client.close.assert_called_once()
assert telemetry_module._oss_telemetry is None
assert telemetry_module._client_telemetry is None
def test_shutdown_is_safe_when_no_singletons(self):
telemetry_module._oss_telemetry = None
telemetry_module._client_telemetry = None
telemetry_module.shutdown_telemetry() # should not raise
class TestNoThreadLeak:
"""Verify capture_event does not leak threads."""
def test_multiple_capture_events_single_posthog_instance(self):
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="test"):
telemetry_module._oss_telemetry = None
mock_memory = MagicMock()
mock_memory.config.graph_store.config = None
mock_memory.api_version = "v1"
for _ in range(10):
telemetry_module.capture_event("test.event", mock_memory)
assert mock_posthog.call_count == 1
telemetry_module._oss_telemetry = None
class TestTelemetryEnvVar: