diff --git a/mem0/vector_stores/redis.py b/mem0/vector_stores/redis.py index fc2048e35..6e2544a0a 100644 --- a/mem0/vector_stores/redis.py +++ b/mem0/vector_stores/redis.py @@ -191,9 +191,12 @@ class RedisDB(VectorStoreBase): "memory": payload["data"], "created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()), "updated_at": int(datetime.fromisoformat(payload["updated_at"]).timestamp()), - "embedding": np.array(vector, dtype=np.float32).tobytes(), } + # Only update embedding if vector is provided + if vector is not None: + data["embedding"] = np.array(vector, dtype=np.float32).tobytes() + for field in ["agent_id", "run_id", "user_id"]: if field in payload: data[field] = payload[field] diff --git a/mem0/vector_stores/valkey.py b/mem0/vector_stores/valkey.py index c4539dcd2..273aaba11 100644 --- a/mem0/vector_stores/valkey.py +++ b/mem0/vector_stores/valkey.py @@ -491,9 +491,12 @@ class ValkeyDB(VectorStoreBase): "hash": payload.get("hash", f"hash_{vector_id}"), # Use a default hash if not provided "memory": payload.get("data", f"data_{vector_id}"), # Use a default data if not provided "created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()), - "embedding": np.array(vector, dtype=np.float32).tobytes(), } + # Only update embedding if vector is provided + if vector is not None: + hash_data["embedding"] = np.array(vector, dtype=np.float32).tobytes() + # Add updated_at if available if "updated_at" in payload: hash_data["updated_at"] = int(datetime.fromisoformat(payload["updated_at"]).timestamp()) diff --git a/tests/vector_stores/test_redis.py b/tests/vector_stores/test_redis.py new file mode 100644 index 000000000..354d1c4b1 --- /dev/null +++ b/tests/vector_stores/test_redis.py @@ -0,0 +1,73 @@ +"""Tests for Redis vector store update() — embedding corruption fix. + +Regression tests for #4336: when update() is called with vector=None +(metadata-only update), np.array(None) silently creates a 4-byte scalar, +overwriting the real embedding. The fix skips the embedding field entirely +when vector is None. +""" + +from datetime import datetime +from unittest.mock import MagicMock + +import numpy as np +import pytz + + +def _make_redis_db(): + """Create a RedisDB instance with mocked internals, bypassing __init__ + to avoid the redis module name collision with mem0.vector_stores.redis.""" + from mem0.vector_stores.redis import RedisDB + + db = RedisDB.__new__(RedisDB) + mock_index = MagicMock() + db.index = mock_index + db.schema = {"index": {"prefix": "mem0:test"}} + return db, mock_index + + +def test_update_with_none_vector_preserves_embedding(): + """update() with vector=None should not include embedding in the data.""" + db, mock_index = _make_redis_db() + + payload = { + "hash": "test_hash", + "data": "updated_data", + "created_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "updated_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "user_id": "test_user", + } + + db.update(vector_id="test_id", vector=None, payload=payload) + + mock_index.load.assert_called_once() + call_kwargs = mock_index.load.call_args + data_dict = call_kwargs[1]["data"][0] if "data" in call_kwargs[1] else call_kwargs[0][0][0] + assert "embedding" not in data_dict, ( + "embedding should not be in data when vector is None" + ) + assert data_dict["memory_id"] == "test_id" + + +def test_update_with_vector_includes_embedding(): + """update() with a real vector should include embedding in the data.""" + db, mock_index = _make_redis_db() + + vector = np.random.rand(1536).tolist() + payload = { + "hash": "test_hash", + "data": "updated_data", + "created_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "updated_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "user_id": "test_user", + } + + db.update(vector_id="test_id", vector=vector, payload=payload) + + mock_index.load.assert_called_once() + call_kwargs = mock_index.load.call_args + data_dict = call_kwargs[1]["data"][0] if "data" in call_kwargs[1] else call_kwargs[0][0][0] + assert "embedding" in data_dict, ( + "embedding should be in data when vector is provided" + ) + expected_bytes = np.array(vector, dtype=np.float32).tobytes() + assert data_dict["embedding"] == expected_bytes diff --git a/tests/vector_stores/test_valkey.py b/tests/vector_stores/test_valkey.py index 482f9b83d..65fa89fbb 100644 --- a/tests/vector_stores/test_valkey.py +++ b/tests/vector_stores/test_valkey.py @@ -204,6 +204,52 @@ def test_update_handles_missing_created_at(valkey_db, mock_valkey_client): assert "created_at" in kwargs["mapping"] # Should be added automatically +def test_update_with_none_vector_preserves_embedding(valkey_db, mock_valkey_client): + """Test that update with vector=None does not corrupt the stored embedding. + + Regression test for #4336: when vector=None is passed (metadata-only update), + np.array(None) silently creates a 4-byte scalar, overwriting the real embedding. + The fix skips the embedding field entirely so the existing value is preserved. + """ + payload = { + "hash": "test_hash", + "data": "updated_data", + "created_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "user_id": "test_user", + } + + valkey_db.update(vector_id="test_id", vector=None, payload=payload) + + mock_valkey_client.hset.assert_called_once() + args, kwargs = mock_valkey_client.hset.call_args + assert "embedding" not in kwargs["mapping"], ( + "embedding should not be in hash_data when vector is None" + ) + assert kwargs["mapping"]["memory_id"] == "test_id" + assert kwargs["mapping"]["memory"] == "updated_data" + + +def test_update_with_vector_includes_embedding(valkey_db, mock_valkey_client): + """Test that update with a real vector includes the embedding in hash_data.""" + vector = np.random.rand(1536).tolist() + payload = { + "hash": "test_hash", + "data": "updated_data", + "created_at": datetime.now(pytz.timezone("UTC")).isoformat(), + "user_id": "test_user", + } + + valkey_db.update(vector_id="test_id", vector=vector, payload=payload) + + mock_valkey_client.hset.assert_called_once() + args, kwargs = mock_valkey_client.hset.call_args + assert "embedding" in kwargs["mapping"], ( + "embedding should be in hash_data when vector is provided" + ) + expected_bytes = np.array(vector, dtype=np.float32).tobytes() + assert kwargs["mapping"]["embedding"] == expected_bytes + + def test_get(valkey_db, mock_valkey_client): """Test getting a vector.""" # Mock hgetall to return a vector