Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e92021f688 | |||
| 269394d72d | |||
| f38b6b0ecc |
@@ -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]
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user