fix(vector_stores/opensearch): honor all filter keys instead of a hardcoded identity-key list (#6454)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -16,6 +16,7 @@ from mem0.vector_stores.base import VectorStoreBase
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SAFE_FILTER_KEY = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_.]*$")
|
||||
_IDENTITY_FILTER_KEYS = ("user_id", "agent_id", "run_id")
|
||||
|
||||
|
||||
def _validate_filter(key: str, value) -> None:
|
||||
@@ -28,6 +29,21 @@ def _validate_filter(key: str, value) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _build_filter_clauses(filters):
|
||||
"""Build term clauses from every filter key, not just the identity keys."""
|
||||
filter_clauses = []
|
||||
for key, value in (filters or {}).items():
|
||||
if value is None:
|
||||
continue
|
||||
if key not in _IDENTITY_FILTER_KEYS and (not isinstance(value, (str, int, float, bool)) or value == "*"):
|
||||
logger.debug(f"Ignoring non-scalar or wildcard filter value for key {key!r}")
|
||||
continue
|
||||
_validate_filter(key, value)
|
||||
field = f"payload.{key}.keyword" if isinstance(value, str) else f"payload.{key}"
|
||||
filter_clauses.append({"term": {field: value}})
|
||||
return filter_clauses
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: str
|
||||
score: float
|
||||
@@ -204,13 +220,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
query_body = {"size": top_k * 2, "query": None}
|
||||
|
||||
# Prepare filter conditions if applicable
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
# Combine knn with filters if needed
|
||||
if filter_clauses:
|
||||
@@ -255,13 +265,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
}
|
||||
|
||||
# Apply filters consistently with the existing search() method
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
if filter_clauses:
|
||||
bool_query["filter"] = filter_clauses
|
||||
@@ -370,13 +374,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
"""List all memories with optional filters."""
|
||||
query: Dict = {"query": {"match_all": {}}}
|
||||
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
if filter_clauses:
|
||||
query["query"] = {"bool": {"filter": filter_clauses}}
|
||||
|
||||
@@ -594,3 +594,62 @@ class TestOpenSearchFilterValidation(unittest.TestCase):
|
||||
results = self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice"})
|
||||
self.assertEqual(results, [])
|
||||
self.client_mock.search.assert_called_once()
|
||||
|
||||
|
||||
class TestOpenSearchCustomFilters(TestOpenSearchFilterValidation):
|
||||
"""Custom (non-identity) filter keys must be honored, not silently dropped."""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.client_mock.search.return_value = {"hits": {"hits": []}}
|
||||
|
||||
def _search_body(self):
|
||||
return self.client_mock.search.call_args[1]["body"]
|
||||
|
||||
def test_search_honors_custom_filter_key(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.user_id.keyword": "alice"}}, clauses)
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_keyword_search_honors_custom_filter_key(self):
|
||||
self.os_db.keyword_search(query="report", filters={"category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_list_honors_custom_filter_key(self):
|
||||
self.os_db.list(filters={"category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_non_string_filter_values_use_plain_field(self):
|
||||
"""Non-string identity/scalar values must match against the plain payload field, not .keyword."""
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"age": 30, "archived": False})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.age": 30}}, clauses)
|
||||
self.assertIn({"term": {"payload.archived": False}}, clauses)
|
||||
|
||||
def test_custom_filter_keys_still_validated(self):
|
||||
with self.assertRaises(ValueError):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"bad key!": "x"})
|
||||
|
||||
def test_search_ignores_or_operator_filter(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "$or": [{"a": 1}]})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
def test_search_ignores_operator_shaped_filter_value(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "score": {"gte": 5}})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
def test_search_ignores_wildcard_filter_value(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "category": "*"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
def test_list_ignores_or_operator_filter(self):
|
||||
self.os_db.list(filters={"user_id": "alice", "$or": [{"a": 1}]})
|
||||
self.client_mock.search.assert_called_once()
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
Reference in New Issue
Block a user