fix(upstash): escape quotes in filter values and validate filter types (#5981)

This commit is contained in:
Hrushikesh Yadav
2026-08-04 15:16:24 +05:30
committed by GitHub
parent 965140eb19
commit 3ac9ba4b50
2 changed files with 89 additions and 9 deletions
+37 -8
View File
@@ -1,5 +1,6 @@
import logging
from typing import Dict, List, Optional
import re
from typing import Any, Dict, List, Optional
from pydantic import BaseModel
@@ -13,6 +14,23 @@ except ImportError:
logger = logging.getLogger(__name__)
_SAFE_FILTER_KEY = re.compile(r"[a-zA-Z_][a-zA-Z0-9_]*\Z")
def _validate_filter(key: str, value: Any) -> None:
if not isinstance(key, str) or not _SAFE_FILTER_KEY.fullmatch(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__}"
)
if isinstance(value, str) and ('"' in value or "\\" in value):
raise ValueError(
f"Filter value for {key!r} contains prohibited characters "
f"(double quote or backslash): {value!r}"
)
class OutputData(BaseModel):
id: Optional[str] # memory id
@@ -92,7 +110,9 @@ class UpstashVector(VectorStoreBase):
)
def _stringify(self, x):
return f'"{x}"' if isinstance(x, str) else x
if isinstance(x, str):
return f'"{x}"'
return x
def search(
self,
@@ -113,6 +133,9 @@ class UpstashVector(VectorStoreBase):
List[OutputData]: Search results.
"""
if filters:
for k, v in filters.items():
_validate_filter(k, v)
filters_str = " AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()]) if filters else None
response = []
@@ -160,13 +183,16 @@ class UpstashVector(VectorStoreBase):
Returns:
List[OutputData]: Search results, or None if sparse/BM25 search is not supported.
"""
try:
filters_str = (
" AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()])
if filters
else None
)
if filters:
for k, v in filters.items():
_validate_filter(k, v)
filters_str = (
" AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()])
if filters
else None
)
try:
response = self.client.query(
data=query,
top_k=top_k,
@@ -252,6 +278,9 @@ class UpstashVector(VectorStoreBase):
Returns:
List[OutputData]: Search results.
"""
if filters:
for k, v in filters.items():
_validate_filter(k, v)
filters_str = " AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()]) if filters else None
info = self.client.info()
+52 -1
View File
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock, call, patch
import pytest
from mem0.configs.vector_stores.upstash_vector import UpstashVectorConfig
from mem0.vector_stores.upstash_vector import UpstashVector
from mem0.vector_stores.upstash_vector import UpstashVector, _validate_filter
@dataclass
@@ -228,6 +228,57 @@ def test_update_vector_with_embeddings(upstash_instance_with_embeddings):
)
def test_filter_rejects_dict_value():
with pytest.raises(ValueError):
_validate_filter("user_id", {"$ne": ""})
def test_filter_rejects_list_value():
with pytest.raises(ValueError):
_validate_filter("user_id", ["alice", "bob"])
def test_filter_rejects_invalid_key():
with pytest.raises(ValueError):
_validate_filter("user_id; DROP", "alice")
def test_filter_accepts_scalars():
_validate_filter("user_id", "alice")
_validate_filter("count", 42)
_validate_filter("score", 0.95)
_validate_filter("active", True)
def test_filter_rejects_double_quote_in_value():
with pytest.raises(ValueError, match="prohibited characters"):
_validate_filter("user_id", 'alice" OR 1=1 --')
def test_filter_rejects_backslash_in_value():
with pytest.raises(ValueError, match="prohibited characters"):
_validate_filter("user_id", "alice\\bob")
def test_search_rejects_dict_filter(upstash_instance):
with pytest.raises(ValueError):
upstash_instance.search(
query="test", vectors=[[0.1]], filters={"user_id": {"$ne": ""}}
)
def test_keyword_search_raises_on_invalid_filter(upstash_instance):
with pytest.raises(ValueError):
upstash_instance.keyword_search(
query="test", filters={"user_id": 'alice" OR 1=1'}
)
def test_filter_rejects_key_with_trailing_newline():
with pytest.raises(ValueError, match="Invalid filter key"):
_validate_filter("user_id\n", "alice")
def test_insert_vectors_with_embeddings_missing_data(upstash_instance_with_embeddings):
vectors = [[0.1, 0.2, 0.3]]
payloads = [{"name": "vector1"}] # Missing data field