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
+9 -7
View File
@@ -79,8 +79,8 @@ class TestRerank:
reranker = LLMReranker({"provider": "openai"})
reranker.rerank("query", [{"text": "some text"}])
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
assert "some text" in prompt_sent
user_msg = mock_llm_instance.generate_response.call_args[1]["messages"][1]["content"]
assert "some text" in user_msg
def test_content_field_extraction(self, mock_llm):
_, mock_llm_instance = mock_llm
@@ -89,8 +89,8 @@ class TestRerank:
reranker = LLMReranker({"provider": "openai"})
reranker.rerank("query", [{"content": "some content"}])
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
assert "some content" in prompt_sent
user_msg = mock_llm_instance.generate_response.call_args[1]["messages"][1]["content"]
assert "some content" in user_msg
def test_fallback_score_on_llm_error(self, mock_llm):
_, mock_llm_instance = mock_llm
@@ -106,12 +106,14 @@ class TestRerank:
_, mock_llm_instance = mock_llm
mock_llm_instance.generate_response.return_value = "0.7"
custom_prompt = "Rate this: query={query} doc={document}"
custom_prompt = "Rate relevance on a scale of 0.0 to 1.0."
reranker = LLMReranker({"provider": "openai", "scoring_prompt": custom_prompt})
reranker.rerank("my query", [{"memory": "my doc"}])
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
assert prompt_sent == "Rate this: query=my query doc=my doc"
messages = mock_llm_instance.generate_response.call_args[1]["messages"]
assert messages[0]["content"] == custom_prompt
assert "my query" in messages[1]["content"]
assert "my doc" in messages[1]["content"]
def test_original_doc_not_mutated(self, mock_llm):
_, mock_llm_instance = mock_llm
+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)