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
This commit is contained in:
DhilipBinny
2026-03-17 01:17:14 +08:00
parent 69001d7b1f
commit f38b6b0ecc
3 changed files with 54 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())
+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