From cf2d42a7fd22ca27a41e3f460beaf7f371f5d779 Mon Sep 17 00:00:00 2001 From: kartik-mem0 Date: Sat, 21 Mar 2026 20:13:20 +0530 Subject: [PATCH] =?UTF-8?q?refactor:=20add=20vector=20validation=20to=20Op?= =?UTF-8?q?enSearchDB=20to=20ensure=20non=E2=80=91null,=20non=E2=80=91empt?= =?UTF-8?q?y,=20and=20correct=E2=80=91dimension=20vectors?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- mem0/vector_stores/opensearch.py | 56 ++++++-- tests/vector_stores/test_opensearch.py | 178 +++++++++++++++++++------ 2 files changed, 182 insertions(+), 52 deletions(-) diff --git a/mem0/vector_stores/opensearch.py b/mem0/vector_stores/opensearch.py index deebae91e..aee256640 100644 --- a/mem0/vector_stores/opensearch.py +++ b/mem0/vector_stores/opensearch.py @@ -113,6 +113,25 @@ class OpenSearchDB(VectorStoreBase): if payloads is None: payloads = [{} for _ in range(len(vectors))] + for idx, vec in enumerate(vectors): + if vec is None: + raise ValueError( + f"Vector at index {idx} is null. " + f"This usually means the embedding model failed to generate an embedding. " + f"Check that your embedding model is configured correctly and returning valid vectors." + ) + if len(vec) == 0: + raise ValueError( + f"Vector at index {idx} is empty. " + f"Expected a vector of dimension {self.embedding_model_dims}, got an empty vector." + ) + if len(vec) != self.embedding_model_dims: + raise ValueError( + f"Vector at index {idx} has dimension {len(vec)}, " + f"but the index '{self.collection_name}' expects dimension {self.embedding_model_dims}. " + f"Ensure your embedding model's output dimensions match the vector store configuration." + ) + results = [] for i, (vec, id_) in enumerate(zip(vectors, ids)): body = { @@ -124,14 +143,16 @@ class OpenSearchDB(VectorStoreBase): self.client.index(index=self.collection_name, body=body) # Force refresh to make documents immediately searchable for tests self.client.indices.refresh(index=self.collection_name) - - results.append(OutputData( - id=id_, - score=1.0, # No score for inserts - payload=payloads[i] - )) + + results.append( + OutputData( + id=id_, + score=1.0, # No score for inserts + payload=payloads[i], + ) + ) except Exception as e: - logger.error(f"Error inserting vector {id_}: {e}") + logger.error(f"Error inserting vector {id_}: {e}", exc_info=True) raise return results @@ -179,7 +200,7 @@ class OpenSearchDB(VectorStoreBase): ] return results except Exception as e: - logger.error(f"Error during search: {e}") + logger.error(f"Error during search: {e}", exc_info=True) return [] def delete(self, vector_id: str) -> None: @@ -200,6 +221,15 @@ class OpenSearchDB(VectorStoreBase): def update(self, vector_id: str, vector: Optional[List[float]] = None, payload: Optional[Dict] = None) -> None: """Update a vector and its payload using the custom 'id' field.""" + if vector is not None: + if len(vector) == 0: + raise ValueError("Cannot update with an empty vector.") + if len(vector) != self.embedding_model_dims: + raise ValueError( + f"Update vector has dimension {len(vector)}, " + f"but the index '{self.collection_name}' expects dimension {self.embedding_model_dims}. " + f"Ensure your embedding model's output dimensions match the vector store configuration." + ) # First, find the document by custom ID search_query = {"query": {"term": {"id": vector_id}}} @@ -222,8 +252,9 @@ class OpenSearchDB(VectorStoreBase): if doc: try: response = self.client.update(index=self.collection_name, id=opensearch_id, body={"doc": doc}) - except Exception: - pass + except Exception as e: + logger.error(f"Error updating vector {vector_id}: {e}", exc_info=True) + raise def get(self, vector_id: str) -> Optional[OutputData]: """Retrieve a vector by ID.""" @@ -238,7 +269,7 @@ class OpenSearchDB(VectorStoreBase): return OutputData(id=hits[0]["_source"].get("id"), score=1.0, payload=hits[0]["_source"].get("payload", {})) except Exception as e: - logger.error(f"Error retrieving vector {vector_id}: {str(e)}") + logger.error(f"Error retrieving vector {vector_id}: {str(e)}", exc_info=True) return None def list_cols(self) -> List[str]: @@ -281,9 +312,8 @@ class OpenSearchDB(VectorStoreBase): ] return [results] # VectorStore expects tuple/list format except Exception as e: - logger.error(f"Error listing vectors: {e}") + logger.error(f"Error listing vectors: {e}", exc_info=True) return [] - def reset(self): """Reset the index by deleting and recreating it.""" diff --git a/tests/vector_stores/test_opensearch.py b/tests/vector_stores/test_opensearch.py index 82e2b41fe..626fdf048 100644 --- a/tests/vector_stores/test_opensearch.py +++ b/tests/vector_stores/test_opensearch.py @@ -18,25 +18,25 @@ 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), + "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") @@ -44,7 +44,7 @@ class MockOpenSearchConfig: 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") @@ -63,20 +63,18 @@ class MockOpenSearchConfig: 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") @@ -255,6 +253,104 @@ class TestOpenSearchDB(unittest.TestCase): self.os_db.delete_col() self.client_mock.indices.delete.assert_called_once_with(index="test_collection") + def test_insert_rejects_null_vectors(self): + """Vectors that are None should raise ValueError before hitting OpenSearch.""" + vectors = [None] + payloads = [{"key": "value"}] + ids = ["id1"] + + with self.assertRaises(ValueError) as ctx: + self.os_db.insert(vectors=vectors, payloads=payloads, ids=ids) + + self.assertIn("null", str(ctx.exception).lower()) + self.client_mock.index.assert_not_called() + + def test_insert_rejects_empty_vectors(self): + """Empty vector lists should raise ValueError.""" + vectors = [[]] + payloads = [{"key": "value"}] + ids = ["id1"] + + with self.assertRaises(ValueError) as ctx: + self.os_db.insert(vectors=vectors, payloads=payloads, ids=ids) + + self.assertIn("empty", str(ctx.exception).lower()) + self.client_mock.index.assert_not_called() + + def test_insert_rejects_dimension_mismatch(self): + """Vectors with wrong dimensions should raise ValueError with a clear message.""" + vectors = [[0.1] * 768] + payloads = [{"key": "value"}] + ids = ["id1"] + + with self.assertRaises(ValueError) as ctx: + self.os_db.insert(vectors=vectors, payloads=payloads, ids=ids) + + error_msg = str(ctx.exception) + self.assertIn("768", error_msg) + self.assertIn("1536", error_msg) + self.client_mock.index.assert_not_called() + + def test_update_rejects_empty_vector(self): + """Update with an empty vector should raise ValueError.""" + with self.assertRaises(ValueError) as ctx: + self.os_db.update("id1", vector=[], payload={"key": "value"}) + + self.assertIn("empty", str(ctx.exception).lower()) + self.client_mock.search.assert_not_called() + + def test_update_rejects_dimension_mismatch(self): + """Update with wrong vector dimensions should raise ValueError.""" + vector = [0.1] * 768 + payload = {"key": "value"} + + with self.assertRaises(ValueError) as ctx: + self.os_db.update("id1", vector=vector, payload=payload) + + error_msg = str(ctx.exception) + self.assertIn("768", error_msg) + self.assertIn("1536", error_msg) + self.client_mock.search.assert_not_called() + + @patch("mem0.vector_stores.opensearch.logger") + def test_update_error_logs_with_exc_info(self, mock_logger): + """Update errors should log with exc_info and re-raise.""" + mock_search_response = {"hits": {"hits": [{"_id": "doc1", "_source": {"id": "id1"}}]}} + self.client_mock.search.return_value = mock_search_response + self.client_mock.update.side_effect = Exception("Update failed") + + with self.assertRaises(Exception): + self.os_db.update("id1", vector=[0.1] * 1536, payload={"key": "value"}) + + mock_logger.error.assert_called_once() + call_kwargs = mock_logger.error.call_args + self.assertTrue(call_kwargs[1].get("exc_info"), "logger.error must be called with exc_info=True") + + @patch("mem0.vector_stores.opensearch.logger") + def test_insert_error_logs_with_exc_info(self, mock_logger): + """Error logging should include exc_info for full stack trace.""" + vectors = [[0.1] * 1536] + payloads = [{"key": "value"}] + ids = ["id1"] + self.client_mock.index.side_effect = Exception("Connection refused") + + with self.assertRaises(Exception): + self.os_db.insert(vectors=vectors, payloads=payloads, ids=ids) + + mock_logger.error.assert_called_once() + call_kwargs = mock_logger.error.call_args + self.assertTrue(call_kwargs[1].get("exc_info"), "logger.error must be called with exc_info=True") + + @patch("mem0.vector_stores.opensearch.logger") + def test_search_error_logs_with_exc_info(self, mock_logger): + """Search error logging should include exc_info for full stack trace.""" + self.client_mock.search.side_effect = Exception("Search failed") + results = self.os_db.search(query="", vectors=[[0.1] * 1536], limit=5) + self.assertEqual(results, []) + mock_logger.error.assert_called_once() + call_kwargs = mock_logger.error.call_args + self.assertTrue(call_kwargs[1].get("exc_info"), "logger.error must be called with exc_info=True") + def test_init_with_http_auth(self): mock_credentials = MagicMock() mock_signer = AWSV4SignerAuth(mock_credentials, "us-east-1", "es") @@ -282,11 +378,13 @@ class TestOpenSearchDB(unittest.TestCase): # 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): +@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() @@ -295,9 +393,9 @@ def test_safe_deepcopy_config_handles_opensearch_auth(mock_sqlite, mock_llm_fact 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) # Runtime auth objects must be preserved (Issue #3580) @@ -306,7 +404,7 @@ def test_safe_deepcopy_config_handles_opensearch_auth(mock_sqlite, mock_llm_fact assert safe_config.connection_class is not None # Credentials dict is a sensitive secret and should be redacted assert safe_config.credentials is None - + assert safe_config.collection_name == "opensearch_test" assert safe_config.host == "localhost" assert safe_config.port == 9200 @@ -315,10 +413,10 @@ def test_safe_deepcopy_config_handles_opensearch_auth(mock_sqlite, mock_llm_fact 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') +@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() @@ -328,12 +426,12 @@ def test_safe_deepcopy_config_normal_configs(mock_sqlite, mock_llm_factory, mock 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.collection_name == "normal_test" assert safe_config.host == "localhost" assert safe_config.port == 9200 assert safe_config.embedding_model_dims == 1536 @@ -341,13 +439,15 @@ def test_safe_deepcopy_config_normal_configs(mock_sqlite, mock_llm_factory, mock 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): +@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