fix(vector_stores): deep-copy Redis DEFAULT_FIELDS so instances keep distinct dims (#5633)
This commit is contained in:
@@ -7,7 +7,7 @@ when vector is None.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytz
|
||||
@@ -115,6 +115,58 @@ def test_list_with_filter_builds_query():
|
||||
assert query.query_string() == "@user_id:{alice}"
|
||||
|
||||
|
||||
def test_create_col_keeps_distinct_dims_across_instances():
|
||||
"""Building the schema for two collections with different embedding
|
||||
dimensions must keep them distinct. Regression: DEFAULT_FIELDS.copy() is a
|
||||
shallow copy, so the shared "embedding" field dict was mutated in place and
|
||||
the second create_col() clobbered the first one's dims (and the module
|
||||
global)."""
|
||||
import mem0.vector_stores.redis as redis_module
|
||||
|
||||
db, _ = _make_redis_db()
|
||||
db.client = MagicMock()
|
||||
|
||||
captured = []
|
||||
|
||||
def capture_schema(schema):
|
||||
captured.append(schema)
|
||||
return MagicMock()
|
||||
|
||||
with patch.object(redis_module.SearchIndex, "from_dict", side_effect=capture_schema):
|
||||
db.create_col(name="col_384", vector_size=384)
|
||||
db.create_col(name="col_1536", vector_size=1536)
|
||||
|
||||
assert captured[0]["fields"][-1]["attrs"]["dims"] == 384
|
||||
assert captured[1]["fields"][-1]["attrs"]["dims"] == 1536
|
||||
# the module-level default must never be mutated by building a schema
|
||||
assert "dims" not in redis_module.DEFAULT_FIELDS[-1]["attrs"]
|
||||
|
||||
|
||||
def test_init_keeps_distinct_dims_from_module_global():
|
||||
"""__init__ must stamp the requested dims into the index schema without
|
||||
mutating the shared module-level DEFAULT_FIELDS. The create_col test above
|
||||
covers the other deepcopy site; all other tests build RedisDB via __new__
|
||||
and so skip __init__, leaving this call site otherwise uncovered."""
|
||||
import mem0.vector_stores.redis as redis_module
|
||||
from mem0.vector_stores.redis import RedisDB
|
||||
|
||||
captured = []
|
||||
|
||||
def capture_schema(schema):
|
||||
captured.append(schema)
|
||||
return MagicMock()
|
||||
|
||||
with (
|
||||
patch("mem0.vector_stores.redis.redis.Redis.from_url", return_value=MagicMock()),
|
||||
patch.object(redis_module.SearchIndex, "from_dict", side_effect=capture_schema),
|
||||
):
|
||||
RedisDB("redis://localhost:6379", "col_384", 384)
|
||||
|
||||
assert captured[0]["fields"][-1]["attrs"]["dims"] == 384
|
||||
# the module-level default must never be mutated by constructing an instance
|
||||
assert "dims" not in redis_module.DEFAULT_FIELDS[-1]["attrs"]
|
||||
|
||||
|
||||
def test_get_returns_none_for_missing_id():
|
||||
"""get() must return None for a missing id.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user