fix: prevent arbitrary code execution via pickle in FAISS vector store (#4833)
This commit is contained in:
@@ -1,4 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
import tempfile
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
@@ -6,7 +8,13 @@ import faiss
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from mem0.vector_stores.faiss import FAISS, OutputData
|
||||
from mem0.vector_stores.faiss import (
|
||||
FAISS,
|
||||
OutputData,
|
||||
SafeUnpickler,
|
||||
_safe_pickle_load,
|
||||
_validate_docstore_structure,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -273,8 +281,8 @@ def test_delete_col(faiss_instance):
|
||||
# Call delete_col
|
||||
faiss_instance.delete_col()
|
||||
|
||||
# Verify os.remove was called twice (for index and docstore files)
|
||||
assert mock_remove.call_count == 2
|
||||
# Verify os.remove was called for index, json docstore, and legacy pkl files
|
||||
assert mock_remove.call_count == 3
|
||||
|
||||
# Verify the internal state was reset
|
||||
assert faiss_instance.index is None
|
||||
@@ -299,3 +307,360 @@ def test_normalize_L2(faiss_instance, mock_faiss_index):
|
||||
|
||||
# Verify faiss.normalize_L2 was called
|
||||
mock_normalize.assert_called_once()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Security Tests for Pickle Deserialization Vulnerability Fix
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestSafeUnpickler:
|
||||
"""Tests for the SafeUnpickler class that prevents arbitrary code execution."""
|
||||
|
||||
def test_safe_unpickler_allows_basic_types(self):
|
||||
"""SafeUnpickler should allow basic Python types."""
|
||||
# Create a legitimate pickle with basic types
|
||||
data = (
|
||||
{"key1": "value1", "key2": {"nested": "dict"}},
|
||||
{0: "id1", 1: "id2"},
|
||||
)
|
||||
pickled = pickle.dumps(data)
|
||||
|
||||
# Should load successfully
|
||||
import io
|
||||
|
||||
result = SafeUnpickler(io.BytesIO(pickled)).load()
|
||||
assert result == data
|
||||
|
||||
def test_safe_unpickler_blocks_os_system(self):
|
||||
"""SafeUnpickler should block os.system execution attempts."""
|
||||
# Generate the malicious payload dynamically to ensure correct format
|
||||
import io
|
||||
|
||||
class Evil:
|
||||
def __reduce__(self):
|
||||
return (os.system, ("echo pwned",))
|
||||
|
||||
malicious_payload = pickle.dumps(Evil())
|
||||
|
||||
with pytest.raises(pickle.UnpicklingError) as exc_info:
|
||||
SafeUnpickler(io.BytesIO(malicious_payload)).load()
|
||||
|
||||
assert "Unsafe pickle" in str(exc_info.value)
|
||||
assert "posix.system" in str(exc_info.value)
|
||||
|
||||
def test_safe_unpickler_blocks_subprocess(self):
|
||||
"""SafeUnpickler should block subprocess execution attempts."""
|
||||
import subprocess
|
||||
|
||||
# Create a malicious pickle that tries to use subprocess
|
||||
class MaliciousSubprocess:
|
||||
def __reduce__(self):
|
||||
return (subprocess.call, (["echo", "pwned"],))
|
||||
|
||||
malicious_payload = pickle.dumps(MaliciousSubprocess())
|
||||
|
||||
import io
|
||||
|
||||
with pytest.raises(pickle.UnpicklingError) as exc_info:
|
||||
SafeUnpickler(io.BytesIO(malicious_payload)).load()
|
||||
|
||||
assert "Unsafe pickle" in str(exc_info.value)
|
||||
|
||||
def test_safe_unpickler_blocks_eval(self):
|
||||
"""SafeUnpickler should block eval/exec attempts."""
|
||||
|
||||
# Create a malicious pickle that tries to use eval
|
||||
class MaliciousEval:
|
||||
def __reduce__(self):
|
||||
return (eval, ("__import__('os').system('touch pwned')",))
|
||||
|
||||
malicious_payload = pickle.dumps(MaliciousEval())
|
||||
|
||||
import io
|
||||
|
||||
with pytest.raises(pickle.UnpicklingError) as exc_info:
|
||||
SafeUnpickler(io.BytesIO(malicious_payload)).load()
|
||||
|
||||
assert "Unsafe pickle" in str(exc_info.value)
|
||||
|
||||
def test_safe_unpickler_blocks_arbitrary_modules(self):
|
||||
"""SafeUnpickler should block imports from arbitrary modules."""
|
||||
|
||||
# Create a pickle that tries to load a class from a non-builtins module
|
||||
class ArbitraryClass:
|
||||
def __reduce__(self):
|
||||
return (type, ("Evil", (), {}))
|
||||
|
||||
malicious_payload = pickle.dumps(ArbitraryClass())
|
||||
|
||||
import io
|
||||
|
||||
# This should either work (type is a builtin) or fail safely
|
||||
# The key is it shouldn't execute arbitrary code
|
||||
try:
|
||||
result = SafeUnpickler(io.BytesIO(malicious_payload)).load()
|
||||
# If it loads, verify it's just a benign type object
|
||||
assert isinstance(result, type)
|
||||
except pickle.UnpicklingError:
|
||||
# This is also acceptable - blocking unknown patterns
|
||||
pass
|
||||
|
||||
|
||||
class TestSafePickleLoad:
|
||||
"""Tests for the _safe_pickle_load function."""
|
||||
|
||||
def test_safe_pickle_load_with_valid_file(self):
|
||||
"""_safe_pickle_load should load valid pickle files."""
|
||||
with tempfile.NamedTemporaryFile(mode="wb", suffix=".pkl", delete=False) as f:
|
||||
data = ({"id1": {"data": "test"}}, {0: "id1"})
|
||||
pickle.dump(data, f)
|
||||
temp_path = f.name
|
||||
|
||||
try:
|
||||
result = _safe_pickle_load(temp_path)
|
||||
assert result == data
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
def test_safe_pickle_load_blocks_malicious_file(self):
|
||||
"""_safe_pickle_load should block malicious pickle files."""
|
||||
|
||||
# Generate the malicious payload dynamically
|
||||
class Evil:
|
||||
def __reduce__(self):
|
||||
return (os.system, ("echo pwned",))
|
||||
|
||||
malicious_payload = pickle.dumps(Evil())
|
||||
|
||||
with tempfile.NamedTemporaryFile(mode="wb", suffix=".pkl", delete=False) as f:
|
||||
f.write(malicious_payload)
|
||||
temp_path = f.name
|
||||
|
||||
try:
|
||||
with pytest.raises(pickle.UnpicklingError) as exc_info:
|
||||
_safe_pickle_load(temp_path)
|
||||
assert "Unsafe pickle" in str(exc_info.value)
|
||||
finally:
|
||||
os.unlink(temp_path)
|
||||
|
||||
|
||||
class TestValidateDocstoreStructure:
|
||||
"""Tests for the _validate_docstore_structure function."""
|
||||
|
||||
def test_valid_structure(self):
|
||||
"""Should accept valid docstore structure."""
|
||||
data = ({"id1": {"data": "test"}}, {0: "id1"})
|
||||
docstore, index_to_id = _validate_docstore_structure(data)
|
||||
assert docstore == {"id1": {"data": "test"}}
|
||||
assert index_to_id == {0: "id1"}
|
||||
|
||||
def test_invalid_tuple_length(self):
|
||||
"""Should reject tuples with wrong length."""
|
||||
with pytest.raises(ValueError, match="expected tuple"):
|
||||
_validate_docstore_structure(({}, {}, {}))
|
||||
|
||||
def test_invalid_docstore_type(self):
|
||||
"""Should reject non-dict docstore."""
|
||||
with pytest.raises(ValueError, match="docstore must be a dict"):
|
||||
_validate_docstore_structure(("not a dict", {}))
|
||||
|
||||
def test_invalid_index_to_id_type(self):
|
||||
"""Should reject non-dict index_to_id."""
|
||||
with pytest.raises(ValueError, match="index_to_id must be a dict"):
|
||||
_validate_docstore_structure(({}, "not a dict"))
|
||||
|
||||
def test_invalid_docstore_key_type(self):
|
||||
"""Should reject non-string docstore keys."""
|
||||
with pytest.raises(ValueError, match="Invalid docstore key type"):
|
||||
_validate_docstore_structure(({123: {"data": "test"}}, {0: "id1"}))
|
||||
|
||||
def test_invalid_index_to_id_key_type(self):
|
||||
"""Should reject non-int index_to_id keys."""
|
||||
with pytest.raises(ValueError, match="Invalid index_to_id key type"):
|
||||
_validate_docstore_structure(({"id1": {"data": "test"}}, {"0": "id1"}))
|
||||
|
||||
|
||||
class TestFAISSSecurityIntegration:
|
||||
"""Integration tests for FAISS security fixes."""
|
||||
|
||||
def test_faiss_saves_as_json(self):
|
||||
"""FAISS should save docstore as JSON, not pickle."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
mock_index = Mock()
|
||||
mock_index.d = 128
|
||||
mock_index.ntotal = 0
|
||||
|
||||
with patch("mem0.vector_stores.faiss.faiss.IndexFlatL2", return_value=mock_index):
|
||||
with patch("mem0.vector_stores.faiss.faiss.write_index"):
|
||||
faiss_store = FAISS(
|
||||
collection_name="test_security",
|
||||
path=os.path.join(temp_dir, "test_faiss"),
|
||||
distance_strategy="euclidean",
|
||||
)
|
||||
faiss_store.index = mock_index
|
||||
|
||||
# Insert some data
|
||||
faiss_store.docstore = {"id1": {"data": "test"}}
|
||||
faiss_store.index_to_id = {0: "id1"}
|
||||
faiss_store._save()
|
||||
|
||||
# Verify JSON file was created
|
||||
json_path = os.path.join(temp_dir, "test_faiss", "test_security.json")
|
||||
pkl_path = os.path.join(temp_dir, "test_faiss", "test_security.pkl")
|
||||
|
||||
assert os.path.exists(json_path), "JSON docstore file should be created"
|
||||
assert not os.path.exists(pkl_path), "Pickle file should NOT be created"
|
||||
|
||||
# Verify JSON content
|
||||
with open(json_path, "r") as f:
|
||||
data = json.load(f)
|
||||
assert data["docstore"] == {"id1": {"data": "test"}}
|
||||
assert data["index_to_id"] == {"0": "id1"}
|
||||
|
||||
def test_faiss_loads_json_preferentially(self):
|
||||
"""FAISS should prefer JSON over pickle when both exist."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
faiss_path = os.path.join(temp_dir, "test_faiss")
|
||||
os.makedirs(faiss_path)
|
||||
|
||||
# Create both JSON and pickle files with different data
|
||||
json_data = {"docstore": {"id1": {"source": "json"}}, "index_to_id": {"0": "id1"}}
|
||||
pkl_data = ({"id1": {"source": "pickle"}}, {0: "id1"})
|
||||
|
||||
with open(os.path.join(faiss_path, "test_pref.json"), "w") as f:
|
||||
json.dump(json_data, f)
|
||||
|
||||
with open(os.path.join(faiss_path, "test_pref.pkl"), "wb") as f:
|
||||
pickle.dump(pkl_data, f)
|
||||
|
||||
mock_index = Mock()
|
||||
mock_index.d = 128
|
||||
mock_index.ntotal = 1
|
||||
|
||||
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
|
||||
with patch("mem0.vector_stores.faiss.faiss.write_index"):
|
||||
faiss_store = FAISS.__new__(FAISS)
|
||||
faiss_store.collection_name = "test_pref"
|
||||
faiss_store.path = faiss_path
|
||||
faiss_store.index = None
|
||||
faiss_store.docstore = {}
|
||||
faiss_store.index_to_id = {}
|
||||
|
||||
faiss_store._load(
|
||||
os.path.join(faiss_path, "test_pref.faiss"),
|
||||
os.path.join(faiss_path, "test_pref.pkl"),
|
||||
)
|
||||
|
||||
# Should have loaded from JSON, not pickle
|
||||
assert faiss_store.docstore == {"id1": {"source": "json"}}
|
||||
|
||||
def test_faiss_blocks_malicious_pickle_on_load(self):
|
||||
"""FAISS should block loading of malicious pickle files."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
faiss_path = os.path.join(temp_dir, "test_faiss")
|
||||
os.makedirs(faiss_path)
|
||||
|
||||
# Create a malicious pickle file (RCE payload)
|
||||
class Evil:
|
||||
def __reduce__(self):
|
||||
return (os.system, (f"touch {temp_dir}/pwned",))
|
||||
|
||||
malicious_payload = pickle.dumps(Evil())
|
||||
|
||||
with open(os.path.join(faiss_path, "malicious.pkl"), "wb") as f:
|
||||
f.write(malicious_payload)
|
||||
|
||||
mock_index = Mock()
|
||||
mock_index.ntotal = 1
|
||||
|
||||
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
|
||||
faiss_store = FAISS.__new__(FAISS)
|
||||
faiss_store.collection_name = "malicious"
|
||||
faiss_store.path = faiss_path
|
||||
faiss_store.index = None
|
||||
faiss_store.docstore = {}
|
||||
faiss_store.index_to_id = {}
|
||||
|
||||
# Should raise an error, not execute the malicious payload
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
faiss_store._load(
|
||||
os.path.join(faiss_path, "malicious.faiss"),
|
||||
os.path.join(faiss_path, "malicious.pkl"),
|
||||
)
|
||||
|
||||
assert "malicious pickle" in str(exc_info.value).lower() or "unsafe" in str(exc_info.value).lower()
|
||||
|
||||
# Verify the malicious command was NOT executed
|
||||
pwned_file = os.path.join(temp_dir, "pwned")
|
||||
assert not os.path.exists(pwned_file), "Malicious payload should NOT have been executed!"
|
||||
|
||||
def test_faiss_migrates_legacy_pickle_to_json(self):
|
||||
"""FAISS should auto-migrate valid pickle files to JSON format."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
faiss_path = os.path.join(temp_dir, "test_faiss")
|
||||
os.makedirs(faiss_path)
|
||||
|
||||
# Create a legitimate legacy pickle file
|
||||
pkl_data = ({"id1": {"data": "legacy"}}, {0: "id1"})
|
||||
with open(os.path.join(faiss_path, "legacy.pkl"), "wb") as f:
|
||||
pickle.dump(pkl_data, f)
|
||||
|
||||
mock_index = Mock()
|
||||
mock_index.d = 128
|
||||
mock_index.ntotal = 1
|
||||
|
||||
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
|
||||
with patch("mem0.vector_stores.faiss.faiss.write_index"):
|
||||
faiss_store = FAISS.__new__(FAISS)
|
||||
faiss_store.collection_name = "legacy"
|
||||
faiss_store.path = faiss_path
|
||||
faiss_store.index = None
|
||||
faiss_store.docstore = {}
|
||||
faiss_store.index_to_id = {}
|
||||
|
||||
faiss_store._load(
|
||||
os.path.join(faiss_path, "legacy.faiss"),
|
||||
os.path.join(faiss_path, "legacy.pkl"),
|
||||
)
|
||||
|
||||
# Data should be loaded correctly
|
||||
assert faiss_store.docstore == {"id1": {"data": "legacy"}}
|
||||
assert faiss_store.index_to_id == {0: "id1"}
|
||||
|
||||
# JSON file should now exist (auto-migrated)
|
||||
json_path = os.path.join(faiss_path, "legacy.json")
|
||||
assert os.path.exists(json_path), "JSON file should be created during migration"
|
||||
|
||||
def test_delete_col_removes_json_and_pkl(self):
|
||||
"""delete_col should remove both JSON and legacy pickle files."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
faiss_path = os.path.join(temp_dir, "test_faiss")
|
||||
os.makedirs(faiss_path)
|
||||
|
||||
# Create both file types
|
||||
json_path = os.path.join(faiss_path, "test_del.json")
|
||||
pkl_path = os.path.join(faiss_path, "test_del.pkl")
|
||||
faiss_index_path = os.path.join(faiss_path, "test_del.faiss")
|
||||
|
||||
with open(json_path, "w") as f:
|
||||
json.dump({"docstore": {}, "index_to_id": {}}, f)
|
||||
with open(pkl_path, "wb") as f:
|
||||
pickle.dump(({}, {}), f)
|
||||
with open(faiss_index_path, "w") as f:
|
||||
f.write("dummy")
|
||||
|
||||
with patch("faiss.IndexFlatL2"):
|
||||
faiss_store = FAISS.__new__(FAISS)
|
||||
faiss_store.collection_name = "test_del"
|
||||
faiss_store.path = faiss_path
|
||||
faiss_store.index = Mock()
|
||||
faiss_store.docstore = {}
|
||||
faiss_store.index_to_id = {}
|
||||
|
||||
faiss_store.delete_col()
|
||||
|
||||
# Both files should be deleted
|
||||
assert not os.path.exists(json_path), "JSON file should be deleted"
|
||||
assert not os.path.exists(pkl_path), "PKL file should be deleted"
|
||||
assert not os.path.exists(faiss_index_path), "FAISS index should be deleted"
|
||||
|
||||
Reference in New Issue
Block a user