fix(upstash): escape quotes in filter values and validate filter types (#5981)
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user