From 9000576173dee1e28df951d4ffa5c3c2ef2e2126 Mon Sep 17 00:00:00 2001 From: Vishaal LS <62366204+lsvishaal@users.noreply.github.com> Date: Thu, 9 Oct 2025 19:35:08 +0530 Subject: [PATCH] fix: handle non-serializable objects in config deepcopy (#3464) (#3544) --- mem0/memory/main.py | 52 +++++++-- tests/test_memory.py | 8 +- tests/vector_stores/test_opensearch.py | 152 +++++++++++++++++++++++++ 3 files changed, 200 insertions(+), 12 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 79719566a..0cd6cf57b 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -43,6 +43,45 @@ from mem0.utils.factory import ( warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*SwigPy.*") warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigvarlink.*") +# Initialize logger early for util functions +logger = logging.getLogger(__name__) + + +def _safe_deepcopy_config(config): + """Safely deepcopy config, falling back to JSON serialization for non-serializable objects.""" + try: + return deepcopy(config) + except Exception as e: + logger.debug(f"Deepcopy failed, using JSON serialization: {e}") + + config_class = type(config) + + if hasattr(config, "model_dump"): + try: + clone_dict = config.model_dump(mode="json") + except Exception: + clone_dict = {k: v for k, v in config.__dict__.items()} + elif hasattr(config, "__dataclass_fields__"): + from dataclasses import asdict + clone_dict = asdict(config) + else: + clone_dict = {k: v for k, v in config.__dict__.items()} + + sensitive_tokens = ("auth", "credential", "password", "token", "secret", "key", "connection_class") + for field_name in list(clone_dict.keys()): + if any(token in field_name.lower() for token in sensitive_tokens): + clone_dict[field_name] = None + + try: + return config_class(**clone_dict) + except Exception as reconstruction_error: + logger.warning( + f"Failed to reconstruct config: {reconstruction_error}. " + f"Telemetry may be affected." + ) + raise + + def _build_filters_and_metadata( *, # Enforce keyword-only arguments user_id: Optional[str] = None, @@ -156,7 +195,7 @@ class Memory(MemoryBase): else: self.graph = None - telemetry_config = deepcopy(self.config.vector_store.config) + 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}" @@ -1032,14 +1071,13 @@ class AsyncMemory(MemoryBase): else: self.graph = None - self.config.vector_store.config.collection_name = "mem0migrations" + 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}" - self.config.vector_store.config.path = os.path.join(mem0_dir, provider_path) - os.makedirs(self.config.vector_store.config.path, exist_ok=True) - self._telemetry_vector_store = VectorStoreFactory.create( - self.config.vector_store.provider, self.config.vector_store.config - ) + 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"}) diff --git a/tests/test_memory.py b/tests/test_memory.py index b43e5e815..f5c3d3bac 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -132,12 +132,10 @@ def test_search_handles_incomplete_payloads(mock_sqlite, mock_llm_factory, mock_ mock_embedder.embed.return_value = [0.1, 0.2, 0.3] memory.embedding_model = mock_embedder - # This should not raise KeyError even with incomplete payloads result = memory._search_vector_store("test", {"user_id": "test"}, 10) assert len(result) == 2 memories_by_id = {mem["id"]: mem for mem in result} - - # Verify defensive programming works correctly - assert memories_by_id["mem_1"]["memory"] == "" # Missing data gets empty string - assert memories_by_id["mem_2"]["memory"] == "content" # Normal data preserved + + assert memories_by_id["mem_1"]["memory"] == "" + assert memories_by_id["mem_2"]["memory"] == "content" diff --git a/tests/vector_stores/test_opensearch.py b/tests/vector_stores/test_opensearch.py index 043c9efe5..5f4ffc0df 100644 --- a/tests/vector_stores/test_opensearch.py +++ b/tests/vector_stores/test_opensearch.py @@ -1,4 +1,5 @@ import os +import threading import unittest from unittest.mock import MagicMock, patch @@ -10,9 +11,77 @@ try: except ImportError: raise ImportError("OpenSearch requires extra dependencies. Install with `pip install opensearch-py`") from None +from mem0 import Memory +from mem0.configs.base import MemoryConfig from mem0.vector_stores.opensearch import OpenSearchDB +# Mock classes for testing OpenSearch with AWS authentication +class MockFieldInfo: + """Mock pydantic field info.""" + def __init__(self, default=None): + self.default = default + + +class MockOpenSearchConfig: + + model_fields = { + 'collection_name': MockFieldInfo(default="default_collection"), + 'host': MockFieldInfo(default="localhost"), + 'port': MockFieldInfo(default=9200), + 'embedding_model_dims': MockFieldInfo(default=1536), + 'http_auth': MockFieldInfo(default=None), + 'auth': MockFieldInfo(default=None), + 'credentials': MockFieldInfo(default=None), + 'connection_class': MockFieldInfo(default=None), + 'use_ssl': MockFieldInfo(default=False), + 'verify_certs': MockFieldInfo(default=False), + } + + def __init__(self, collection_name="test_collection", include_auth=True, **kwargs): + self.collection_name = collection_name + self.host = kwargs.get("host", "localhost") + self.port = kwargs.get("port", 9200) + self.embedding_model_dims = kwargs.get("embedding_model_dims", 1536) + self.use_ssl = kwargs.get("use_ssl", True) + self.verify_certs = kwargs.get("verify_certs", True) + + if any(field in kwargs for field in ["http_auth", "auth", "credentials", "connection_class"]): + self.http_auth = kwargs.get("http_auth") + self.auth = kwargs.get("auth") + self.credentials = kwargs.get("credentials") + self.connection_class = kwargs.get("connection_class") + elif include_auth: + self.http_auth = MockAWSAuth() + self.auth = MockAWSAuth() + self.credentials = {"key": "value"} + self.connection_class = MockConnectionClass() + else: + self.http_auth = None + self.auth = None + self.credentials = None + self.connection_class = None + + +class MockAWSAuth: + + def __init__(self): + self._lock = threading.Lock() + self.region = "us-east-1" + + def __deepcopy__(self, memo): + raise TypeError("cannot pickle '_thread.lock' object") + + +class MockConnectionClass: + + def __init__(self): + self._state = {"connected": False} + + def __deepcopy__(self, memo): + raise TypeError("cannot pickle connection state") + + class TestOpenSearchDB(unittest.TestCase): @classmethod def setUpClass(cls): @@ -211,3 +280,86 @@ class TestOpenSearchDB(unittest.TestCase): connection_class=unittest.mock.ANY, pool_maxsize=20, ) + + +# Tests for OpenSearch config deepcopy with AWS authentication (Issue #3464) +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_safe_deepcopy_config_handles_opensearch_auth(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """Test that _safe_deepcopy_config handles OpenSearch configs with AWS auth objects gracefully.""" + mock_embedder_factory.return_value = MagicMock() + mock_vector_store = MagicMock() + mock_vector_factory.return_value = mock_vector_store + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + from mem0.memory.main import _safe_deepcopy_config + + config_with_auth = MockOpenSearchConfig(collection_name="opensearch_test", include_auth=True) + + safe_config = _safe_deepcopy_config(config_with_auth) + + assert safe_config.http_auth is None + assert safe_config.auth is None + assert safe_config.credentials is None + assert safe_config.connection_class is None + + assert safe_config.collection_name == "opensearch_test" + assert safe_config.host == "localhost" + assert safe_config.port == 9200 + assert safe_config.embedding_model_dims == 1536 + assert safe_config.use_ssl is True + assert safe_config.verify_certs is True + + +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_safe_deepcopy_config_normal_configs(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """Test that _safe_deepcopy_config handles normal OpenSearch configs without auth.""" + mock_embedder_factory.return_value = MagicMock() + mock_vector_store = MagicMock() + mock_vector_factory.return_value = mock_vector_store + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + from mem0.memory.main import _safe_deepcopy_config + + config_without_auth = MockOpenSearchConfig(collection_name="normal_test", include_auth=False) + + safe_config = _safe_deepcopy_config(config_without_auth) + + assert safe_config.collection_name == "normal_test" + assert safe_config.host == "localhost" + assert safe_config.port == 9200 + assert safe_config.embedding_model_dims == 1536 + assert safe_config.use_ssl is True + assert safe_config.verify_certs is True + + +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_memory_initialization_opensearch_aws_auth(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """Test that Memory initialization works with OpenSearch configs containing AWS auth.""" + + mock_embedder_factory.return_value = MagicMock() + mock_vector_store = MagicMock() + mock_vector_factory.return_value = mock_vector_store + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + config = MemoryConfig() + config.vector_store.provider = "opensearch" + config.vector_store.config = MockOpenSearchConfig(collection_name="mem0_test", include_auth=True) + + memory = Memory(config) + + assert memory is not None + assert memory.config.vector_store.provider == "opensearch" + + assert mock_vector_factory.call_count >= 2