fix: sql injection, prompt injection (#4997)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Harsh Vardhan Gupta
2026-04-29 00:51:16 +05:30
committed by GitHub
parent b66cf0f272
commit 1b95c99db4
8 changed files with 233 additions and 150 deletions
+27 -22
View File
@@ -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)