refactor: add vector validation to OpenSearchDB to ensure non‑null, non‑empty, and correct‑dimension vectors

This commit is contained in:
kartik-mem0
2026-03-21 20:13:20 +05:30
parent 06c25eb00b
commit cf2d42a7fd
2 changed files with 182 additions and 52 deletions
+43 -13
View File
@@ -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."""
+139 -39
View File
@@ -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