fix(opensearch): validate filter values to prevent term query injection (#5986)

This commit is contained in:
Hrushikesh Yadav
2026-07-01 20:47:45 +05:30
committed by GitHub
parent 152d1e66f7
commit a36a392cd3
2 changed files with 65 additions and 0 deletions
+16
View File
@@ -1,4 +1,5 @@
import logging
import re
import time
from typing import Any, Dict, List, Optional
@@ -14,6 +15,18 @@ from mem0.vector_stores.base import VectorStoreBase
logger = logging.getLogger(__name__)
_SAFE_FILTER_KEY = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_.]*$")
def _validate_filter(key: str, value) -> None:
if not isinstance(key, str) or not _SAFE_FILTER_KEY.match(key):
raise ValueError(f"Invalid filter key: {key!r}")
if not isinstance(value, (str, int, float, bool)):
raise ValueError(
f"Filter value for {key!r} must be str, int, float, or bool, "
f"got {type(value).__name__}"
)
class OutputData(BaseModel):
id: str
@@ -196,6 +209,7 @@ class OpenSearchDB(VectorStoreBase):
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}})
# Combine knn with filters if needed
@@ -246,6 +260,7 @@ class OpenSearchDB(VectorStoreBase):
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}})
if filter_clauses:
@@ -360,6 +375,7 @@ class OpenSearchDB(VectorStoreBase):
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}})
if filter_clauses:
+49
View File
@@ -545,3 +545,52 @@ def test_memory_initialization_opensearch_aws_auth(
assert memory.config.vector_store.provider == "opensearch"
assert mock_vector_factory.call_count >= 2
class TestOpenSearchFilterValidation(unittest.TestCase):
"""Validate that non-scalar filter values are rejected to prevent term injection."""
def setUp(self):
self.client_mock = MagicMock(spec=OpenSearch)
self.client_mock.indices = MagicMock()
self.client_mock.indices.exists = MagicMock(return_value=False)
self.client_mock.indices.create = MagicMock()
self.client_mock.search = MagicMock()
patcher = patch("mem0.vector_stores.opensearch.OpenSearch", return_value=self.client_mock)
self.mock_os = patcher.start()
self.addCleanup(patcher.stop)
self.os_db = OpenSearchDB(
host="localhost",
port=9200,
collection_name="test_collection",
embedding_model_dims=1536,
verify_certs=False,
use_ssl=False,
)
self.client_mock.reset_mock()
def test_search_rejects_dict_filter_value(self):
with self.assertRaises(ValueError):
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": {"$ne": ""}})
def test_search_rejects_list_filter_value(self):
with self.assertRaises(ValueError):
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": ["alice", "bob"]})
def test_list_rejects_dict_filter_value(self):
result = self.os_db.list(filters={"user_id": {"$ne": ""}})
self.assertEqual(result, [[]])
self.client_mock.search.assert_not_called()
def test_keyword_search_rejects_dict_filter_value(self):
with self.assertRaises(ValueError):
self.os_db.keyword_search(query="test", filters={"user_id": {"$ne": ""}})
def test_search_accepts_string_filter(self):
mock_response = {"hits": {"hits": []}}
self.client_mock.search.return_value = mock_response
results = self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice"})
self.assertEqual(results, [])
self.client_mock.search.assert_called_once()