diff --git a/mem0/vector_stores/opensearch.py b/mem0/vector_stores/opensearch.py index 8b8b85377..ad06a2419 100644 --- a/mem0/vector_stores/opensearch.py +++ b/mem0/vector_stores/opensearch.py @@ -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}} diff --git a/tests/vector_stores/test_opensearch.py b/tests/vector_stores/test_opensearch.py index 6b924e782..3b93aadbd 100644 --- a/tests/vector_stores/test_opensearch.py +++ b/tests/vector_stores/test_opensearch.py @@ -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"}}])