fix(redis): do not crash on empty or None filters in search and list (#5446)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -142,8 +142,11 @@ class RedisDB(VectorStoreBase):
|
||||
self.index.load(data, id_field="memory_id")
|
||||
|
||||
def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None):
|
||||
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
|
||||
filter = reduce(lambda x, y: x & y, conditions)
|
||||
filter = None
|
||||
if filters:
|
||||
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
|
||||
if conditions:
|
||||
filter = reduce(lambda x, y: x & y, conditions)
|
||||
|
||||
v = VectorQuery(
|
||||
vector=np.array(vectors, dtype=np.float32).tobytes(),
|
||||
@@ -312,11 +315,14 @@ class RedisDB(VectorStoreBase):
|
||||
"""
|
||||
List all recent created memories from the vector store.
|
||||
"""
|
||||
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
|
||||
filter = reduce(lambda x, y: x & y, conditions)
|
||||
query = Query(str(filter)).sort_by("created_at", asc=False)
|
||||
filter = None
|
||||
if filters:
|
||||
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
|
||||
if conditions:
|
||||
filter = reduce(lambda x, y: x & y, conditions)
|
||||
query = Query(str(filter) if filter is not None else "*").sort_by("created_at", asc=False)
|
||||
if top_k is not None:
|
||||
query = Query(str(filter)).sort_by("created_at", asc=False).paging(0, top_k)
|
||||
query = query.paging(0, top_k)
|
||||
|
||||
results = self.index.search(query)
|
||||
return [
|
||||
|
||||
@@ -71,3 +71,45 @@ def test_update_with_vector_includes_embedding():
|
||||
)
|
||||
expected_bytes = np.array(vector, dtype=np.float32).tobytes()
|
||||
assert data_dict["embedding"] == expected_bytes
|
||||
|
||||
|
||||
def test_search_with_none_filters_does_not_crash():
|
||||
"""search() with the base-class default filters=None must run an unfiltered
|
||||
query, not raise. Regression: filters.items() on None raised AttributeError
|
||||
and reduce() over an empty list raised TypeError. keyword_search() in the
|
||||
same file already guards this; search() did not."""
|
||||
db, mock_index = _make_redis_db()
|
||||
mock_index.query.return_value = []
|
||||
|
||||
assert db.search("query", [0.1, 0.2, 0.3, 0.4], top_k=5, filters=None) == []
|
||||
mock_index.query.assert_called_once()
|
||||
|
||||
|
||||
def test_search_with_empty_or_all_none_filters_does_not_crash():
|
||||
db, mock_index = _make_redis_db()
|
||||
mock_index.query.return_value = []
|
||||
|
||||
assert db.search("query", [0.1, 0.2, 0.3, 0.4], filters={}) == []
|
||||
assert db.search("query", [0.1, 0.2, 0.3, 0.4], filters={"user_id": None}) == []
|
||||
|
||||
|
||||
def test_list_with_none_filters_matches_all():
|
||||
"""list() with no filters must issue a match-all ('*') query, not raise."""
|
||||
db, mock_index = _make_redis_db()
|
||||
mock_index.search.return_value = MagicMock(docs=[])
|
||||
|
||||
db.list(filters=None)
|
||||
|
||||
query = mock_index.search.call_args[0][0]
|
||||
assert query.query_string() == "*"
|
||||
|
||||
|
||||
def test_list_with_filter_builds_query():
|
||||
"""A real filter is still translated into a tag query (no regression)."""
|
||||
db, mock_index = _make_redis_db()
|
||||
mock_index.search.return_value = MagicMock(docs=[])
|
||||
|
||||
db.list(filters={"user_id": "alice"})
|
||||
|
||||
query = mock_index.search.call_args[0][0]
|
||||
assert query.query_string() == "@user_id:{alice}"
|
||||
|
||||
Reference in New Issue
Block a user