fix(neptune): escape filter values in openCypher queries to prevent injection (#5982)

Signed-off-by: Hrushikesh Yadav <yadavhrushikesh65@gmail.com>
This commit is contained in:
Hrushikesh Yadav
2026-07-01 18:48:04 +05:30
committed by GitHub
parent ad7e09851c
commit bc05fd9623
2 changed files with 80 additions and 6 deletions
+27 -6
View File
@@ -1,7 +1,8 @@
import logging
import re
import time
import uuid
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel
@@ -14,6 +15,22 @@ 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: Any) -> 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__}"
)
def _escape_cypher(value: str) -> str:
return value.replace("\\", "\\\\").replace("'", "\\'")
class OutputData(BaseModel):
id: Optional[str] # memory id
score: Optional[float] # distance
@@ -403,19 +420,21 @@ class NeptuneAnalyticsVector(VectorStoreBase):
def _get_where_clause(filters: dict):
"""
Build WHERE clause for Cypher queries from filters.
Args:
filters (dict): Filter conditions as key-value pairs.
Returns:
str: Formatted WHERE clause for Cypher query.
"""
where_clause = ""
for i, (k, v) in enumerate(filters.items()):
_validate_filter(k, v)
escaped_v = _escape_cypher(str(v))
if i == 0:
where_clause += f"WHERE n.{k} = '{v}' "
where_clause += f"WHERE n.{k} = '{escaped_v}' "
else:
where_clause += f"AND n.{k} = '{v}' "
where_clause += f"AND n.{k} = '{escaped_v}' "
return where_clause
@staticmethod
@@ -434,7 +453,9 @@ class NeptuneAnalyticsVector(VectorStoreBase):
"""
conditions = []
for k, v in filters.items():
conditions.append(f"{{equals:{{property: '{k}', value: '{v}'}}}}")
_validate_filter(k, v)
escaped_v = _escape_cypher(str(v))
conditions.append(f"{{equals:{{property: '{k}', value: '{escaped_v}'}}}}")
if len(conditions) == 1:
filter_clause = f", nodeFilter: {conditions[0]}"
@@ -6,6 +6,11 @@ import pytest
from dotenv import load_dotenv
from mem0.utils.factory import VectorStoreFactory
from mem0.vector_stores.neptune_analytics import (
NeptuneAnalyticsVector,
_escape_cypher,
_validate_filter,
)
load_dotenv()
@@ -184,3 +189,51 @@ class TestNeptuneAnalyticsOperations:
with pytest.raises(ValueError):
VectorStoreFactory.create("neptune", config)
class TestNeptuneFilterValidation:
def test_filter_rejects_dict_value(self):
with pytest.raises(ValueError):
_validate_filter("user_id", {"$ne": ""})
def test_filter_rejects_list_value(self):
with pytest.raises(ValueError):
_validate_filter("user_id", ["alice"])
def test_filter_rejects_invalid_key(self):
with pytest.raises(ValueError):
_validate_filter("user_id'; DROP", "alice")
def test_filter_accepts_scalars(self):
_validate_filter("user_id", "alice")
_validate_filter("count", 42)
_validate_filter("label", "MEM0_VECTOR_test")
def test_escape_cypher_quotes(self):
assert _escape_cypher("alice") == "alice"
assert _escape_cypher("it's") == "it\\'s"
assert _escape_cypher("a\\b") == "a\\\\b"
def test_where_clause_escapes_values(self):
clause = NeptuneAnalyticsVector._get_where_clause(
{"user_id": "it's a test"}
)
assert "it\\'s a test" in clause
def test_where_clause_rejects_dict(self):
with pytest.raises(ValueError):
NeptuneAnalyticsVector._get_where_clause(
{"user_id": {"$ne": ""}}
)
def test_node_filter_escapes_values(self):
clause = NeptuneAnalyticsVector._get_node_filter_clause(
{"label": "it's"}
)
assert "it\\'s" in clause
def test_node_filter_rejects_dict(self):
with pytest.raises(ValueError):
NeptuneAnalyticsVector._get_node_filter_clause(
{"user_id": {"$ne": ""}}
)