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:
microbluey
2026-07-23 02:54:16 +08:00
committed by GitHub
parent a9cb4bb644
commit a58e0586ad
2 changed files with 78 additions and 21 deletions
+59
View File
@@ -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"}}])