fix: prevent embedding corruption in Valkey and Redis when vector is None (#4336) (#4362)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
dhilip_binny
2026-03-20 17:37:17 +08:00
committed by GitHub
parent 73038900f5
commit 401754ca65
4 changed files with 127 additions and 2 deletions
+73
View File
@@ -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
+46
View File
@@ -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