fix(milvus,baidu): sanitize filter values to prevent expression injection (#5746)
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import Dict, Optional
|
||||
|
||||
@@ -47,6 +48,8 @@ class OutputData(BaseModel):
|
||||
|
||||
|
||||
class BaiduDB(VectorStoreBase):
|
||||
_SAFE_FILTER_KEY = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
endpoint: str,
|
||||
@@ -404,8 +407,16 @@ class BaiduDB(VectorStoreBase):
|
||||
"""
|
||||
conditions = []
|
||||
for key, value in filters.items():
|
||||
if not self._SAFE_FILTER_KEY.match(key):
|
||||
raise ValueError(f"Invalid filter key: {key!r}")
|
||||
if isinstance(value, str):
|
||||
conditions.append(f'metadata["{key}"] = "{value}"')
|
||||
else:
|
||||
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
|
||||
conditions.append(f'metadata["{key}"] = "{escaped}"')
|
||||
elif isinstance(value, (int, float, bool)):
|
||||
conditions.append(f'metadata["{key}"] = {value}')
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Filter value for {key!r} must be str, int, float, or bool, "
|
||||
f"got {type(value).__name__}"
|
||||
)
|
||||
return " AND ".join(conditions)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
@@ -143,6 +144,8 @@ class MilvusDB(VectorStoreBase):
|
||||
data = [_build_record(idx, embedding, metadata) for idx, embedding, metadata in zip(ids, vectors, payloads)]
|
||||
self.client.insert(collection_name=self.collection_name, data=data, **kwargs)
|
||||
|
||||
_SAFE_FILTER_KEY = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
|
||||
|
||||
def _create_filter(self, filters: dict):
|
||||
"""Prepare filters for efficient query.
|
||||
|
||||
@@ -154,10 +157,18 @@ class MilvusDB(VectorStoreBase):
|
||||
"""
|
||||
operands = []
|
||||
for key, value in filters.items():
|
||||
if not self._SAFE_FILTER_KEY.match(key):
|
||||
raise ValueError(f"Invalid filter key: {key!r}")
|
||||
if isinstance(value, str):
|
||||
operands.append(f'(metadata["{key}"] == "{value}")')
|
||||
else:
|
||||
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
|
||||
operands.append(f'(metadata["{key}"] == "{escaped}")')
|
||||
elif isinstance(value, (int, float, bool)):
|
||||
operands.append(f'(metadata["{key}"] == {value})')
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Filter value for {key!r} must be str, int, float, or bool, "
|
||||
f"got {type(value).__name__}"
|
||||
)
|
||||
|
||||
return " and ".join(operands)
|
||||
|
||||
|
||||
@@ -235,3 +235,35 @@ def test_col_info(mochow_instance, mock_mochow_client):
|
||||
result = mochow_instance.col_info()
|
||||
|
||||
assert result == mock_table_info
|
||||
|
||||
|
||||
def test_create_filter_rejects_dict_value(mochow_instance):
|
||||
"""Dict filter values could contain expression injection payloads."""
|
||||
with pytest.raises(ValueError, match="must be str, int, float, or bool"):
|
||||
mochow_instance._create_filter({"user_id": {"$ne": ""}})
|
||||
|
||||
|
||||
def test_create_filter_rejects_malicious_key(mochow_instance):
|
||||
"""Keys with special characters must be rejected."""
|
||||
with pytest.raises(ValueError, match="Invalid filter key"):
|
||||
mochow_instance._create_filter({'"] = "") or true or ("': "x"})
|
||||
|
||||
|
||||
def test_create_filter_escapes_quotes_in_value(mochow_instance):
|
||||
"""Double-quotes inside string values must be escaped."""
|
||||
result = mochow_instance._create_filter({"user_id": 'alice"}'})
|
||||
assert '\\"' in result
|
||||
assert 'alice\\"' in result
|
||||
|
||||
|
||||
def test_create_filter_escapes_backslash_and_quote(mochow_instance):
|
||||
"""Backslashes and double-quotes in the same value must both be escaped."""
|
||||
result = mochow_instance._create_filter({"user_id": r'alice\path"beta'})
|
||||
assert result == r'metadata["user_id"] = "alice\\path\"beta"'
|
||||
|
||||
|
||||
def test_create_filter_renders_boolean(mochow_instance):
|
||||
"""Boolean values must be rendered unquoted in the backend's expected format."""
|
||||
result = mochow_instance._create_filter({"active": True, "deleted": False})
|
||||
assert 'metadata["active"] = True' in result
|
||||
assert 'metadata["deleted"] = False' in result
|
||||
|
||||
@@ -310,6 +310,33 @@ class TestMilvusDB:
|
||||
with pytest.raises(ValueError, match="no vector data"):
|
||||
milvus_db.update(vector_id="test_id", vector=None, payload={"data": "test"})
|
||||
|
||||
def test_create_filter_rejects_expression_injection(self, milvus_db):
|
||||
"""Crafted string value must not break out of the quoted expression."""
|
||||
with pytest.raises(ValueError, match="must be str, int, float, or bool"):
|
||||
milvus_db._create_filter({"user_id": {"$ne": ""}})
|
||||
|
||||
def test_create_filter_rejects_malicious_key(self, milvus_db):
|
||||
"""Keys with special characters must be rejected."""
|
||||
with pytest.raises(ValueError, match="Invalid filter key"):
|
||||
milvus_db._create_filter({'"] == "") or true or ("': "x"})
|
||||
|
||||
def test_create_filter_escapes_quotes_in_value(self, milvus_db):
|
||||
"""Double-quotes inside string values must be escaped."""
|
||||
result = milvus_db._create_filter({"user_id": 'alice"}'})
|
||||
assert '\\"' in result
|
||||
assert 'alice\\"' in result
|
||||
|
||||
def test_create_filter_escapes_backslash_and_quote(self, milvus_db):
|
||||
"""Backslashes and double-quotes in the same value must both be escaped."""
|
||||
result = milvus_db._create_filter({"user_id": r'alice\path"beta'})
|
||||
assert result == r'(metadata["user_id"] == "alice\\path\"beta")'
|
||||
|
||||
def test_create_filter_renders_boolean(self, milvus_db):
|
||||
"""Boolean values must be rendered unquoted in the backend's expected format."""
|
||||
result = milvus_db._create_filter({"active": True, "deleted": False})
|
||||
assert '(metadata["active"] == True)' in result
|
||||
assert '(metadata["deleted"] == False)' in result
|
||||
|
||||
def test_collection_already_exists(self, mock_milvus_client):
|
||||
"""Test that existing collection is not recreated."""
|
||||
mock_milvus_client.has_collection.return_value = True
|
||||
|
||||
Reference in New Issue
Block a user