fix: sql injection, prompt injection (#4997)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
committed by
GitHub
parent
b66cf0f272
commit
1b95c99db4
@@ -127,10 +127,10 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
|
||||
# Verify pgvector instance properties
|
||||
self.assertEqual(pgvector.collection_name, "test_collection")
|
||||
self.assertEqual(pgvector.embedding_model_dims, 3)
|
||||
@@ -179,8 +179,8 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -233,7 +233,7 @@ class TestPGVector(unittest.TestCase):
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -277,7 +277,7 @@ class TestPGVector(unittest.TestCase):
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -319,7 +319,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify insert query was executed (psycopg3 uses executemany)
|
||||
insert_calls = [call for call in self.mock_cursor.executemany.call_args_list
|
||||
if "INSERT INTO test_collection" in str(call)]
|
||||
if "INSERT INTO" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(insert_calls) > 0)
|
||||
|
||||
# Verify data format
|
||||
@@ -392,7 +392,9 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_execute_values.assert_called_once()
|
||||
call_args = mock_execute_values.call_args
|
||||
|
||||
self.assertIn("INSERT INTO test_collection", call_args[0][1])
|
||||
mock_psycopg2.sql.SQL.assert_any_call(
|
||||
"INSERT INTO {} (id, vector, payload) VALUES %s"
|
||||
)
|
||||
|
||||
# The data argument should be a list of tuples, one per vector
|
||||
data_arg = call_args[0][2]
|
||||
@@ -400,6 +402,9 @@ class TestPGVector(unittest.TestCase):
|
||||
self.assertEqual(data_arg[0][0], self.test_ids[0])
|
||||
self.assertEqual(data_arg[1][0], self.test_ids[1])
|
||||
|
||||
# Restore the module after the sys.modules patch reverts
|
||||
importlib.reload(sys.modules['mem0.vector_stores.pgvector'])
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@patch('mem0.vector_stores.pgvector.ConnectionPool')
|
||||
@patch.object(PGVector, '_get_cursor')
|
||||
@@ -534,7 +539,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DELETE FROM test_collection" in str(call)]
|
||||
if "DELETE FROM" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -573,7 +578,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DELETE FROM test_collection" in str(call)]
|
||||
if "DELETE FROM" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -615,7 +620,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify update queries were executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -657,7 +662,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify update queries were executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -867,7 +872,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete_col query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DROP TABLE IF EXISTS test_collection" in str(call)]
|
||||
if "DROP TABLE IF EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -906,7 +911,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete_col query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DROP TABLE IF EXISTS test_collection" in str(call)]
|
||||
if "DROP TABLE IF EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -1802,7 +1807,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify the update query was executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -1843,7 +1848,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify the update query was executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -1861,7 +1866,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Only raise exception on the delete operation, not during setup
|
||||
def execute_side_effect(*args, **kwargs):
|
||||
if args and "DELETE FROM" in str(args[0]):
|
||||
if args and ("DELETE FROM" in str(args[0]) or "DELETE" in repr(args[0])):
|
||||
raise Exception("Database error")
|
||||
return MagicMock()
|
||||
mock_cursor.execute.side_effect = execute_side_effect
|
||||
@@ -2003,9 +2008,9 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify only vector update query was executed (not payload)
|
||||
vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET vector" in str(call) and "payload" not in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET vector" in str(call) and "payload" not in str(call)]
|
||||
payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
|
||||
self.assertTrue(len(vector_update_calls) > 0)
|
||||
self.assertEqual(len(payload_update_calls), 0)
|
||||
@@ -2045,9 +2050,9 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify both vector and payload update queries were executed
|
||||
vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET vector" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET vector" in str(call)]
|
||||
payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
|
||||
self.assertTrue(len(vector_update_calls) > 0)
|
||||
self.assertTrue(len(payload_update_calls) > 0)
|
||||
|
||||
Reference in New Issue
Block a user