fix: handle non-serializable objects in config deepcopy (#3464) (#3544)

This commit is contained in:
Vishaal LS
2025-10-09 19:35:08 +05:30
committed by GitHub
parent 922471f43b
commit 9000576173
3 changed files with 200 additions and 12 deletions
+45 -7
View File
@@ -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"})
+3 -5
View File
@@ -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"
+152
View File
@@ -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