Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 6d64bbc4cf |
@@ -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.")
|
||||
|
||||
@@ -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__}",
|
||||
}
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user