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
+19 -21
View File
@@ -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}}
+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"}}])