Compare commits

...

3 Commits

Author SHA1 Message Date
kartik-mem0 e92021f688 test: add regression tests for Redis vector store update() embedding corruption fix 2026-03-20 14:44:48 +05:30
DhilipBinny 269394d72d test: add Redis update tests for vector=None embedding guard 2026-03-18 10:40:36 +08:00
DhilipBinny f38b6b0ecc fix: prevent embedding corruption in Valkey and Redis when vector is None
When update() is called with vector=None (metadata-only update),
np.array(None) silently creates a 4-byte scalar instead of raising
an error, overwriting the real embedding and making memories
unsearchable. Skip the embedding field entirely when vector is None,
matching the pattern used by 9 other vector store implementations.

Fixes #4336
2026-03-17 01:17:14 +08:00
4 changed files with 127 additions and 2 deletions
+4 -1
View File
@@ -192,9 +192,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]
+4 -1
View File
@@ -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())
+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