From a2ffca32660537b7b304305f36774ae2a13452cb Mon Sep 17 00:00:00 2001 From: Utkarsh Date: Thu, 26 Mar 2026 20:22:34 +0530 Subject: [PATCH] fix: prevent SQL injection in Databricks vector store (#4558) Co-authored-by: utkarsh240799 Co-authored-by: Claude Opus 4.6 (1M context) --- mem0/vector_stores/databricks.py | 63 +++++-- tests/vector_stores/test_databricks.py | 249 +++++++++++++++++++++++-- 2 files changed, 286 insertions(+), 26 deletions(-) diff --git a/mem0/vector_stores/databricks.py b/mem0/vector_stores/databricks.py index b8ec82efd..a8cd8ec86 100644 --- a/mem0/vector_stores/databricks.py +++ b/mem0/vector_stores/databricks.py @@ -1,11 +1,13 @@ import json import logging +import re import uuid from typing import Optional, List from datetime import datetime, date from databricks.sdk.service.catalog import ColumnInfo, ColumnTypeName, TableType, DataSourceFormat from databricks.sdk.service.catalog import TableConstraint, PrimaryKeyConstraint from databricks.sdk import WorkspaceClient +from databricks.sdk.service.sql import StatementParameterListItem from databricks.sdk.service.vectorsearch import ( VectorIndexType, DeltaSyncVectorIndexSpecRequest, @@ -28,6 +30,9 @@ class MemoryResult(BaseModel): excluded_keys = {"user_id", "agent_id", "run_id", "hash", "data", "created_at", "updated_at"} +# Pattern for valid SQL identifiers to prevent column name injection +_VALID_SQL_IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + class Databricks(VectorStoreBase): def __init__( @@ -388,28 +393,46 @@ class Databricks(VectorStoreBase): # Determine the number of items to process num_items = len(payloads) if payloads else len(vectors) if vectors else 0 + params = [] value_tuples = [] for i in range(num_items): - values = [] + placeholders = [] for col in self.columns: + param_name = f"{col.name}_{i}" if col.name == "memory_id": val = ids[i] if ids and i < len(ids) else str(uuid.uuid4()) elif col.name == "embedding": val = vectors[i] if vectors and i < len(vectors) else [] + # Vectors are numeric arrays — ARRAY type not supported by StatementParameterListItem, + # so we inline using _format_sql_value (values are floats from the embedding model). + placeholders.append(self._format_sql_value(val)) + continue elif col.name == "memory": val = payloads[i].get("data") if payloads and i < len(payloads) else None else: val = payloads[i].get(col.name) if payloads and i < len(payloads) else None - values.append(val) - formatted = [self._format_sql_value(v) for v in values] - value_tuples.append(f"({', '.join(formatted)})") + + if val is None: + placeholders.append("NULL") + else: + placeholders.append(f":{param_name}") + if isinstance(val, dict): + val = json.dumps(val) + # Use explicit type for TIMESTAMP columns so Databricks doesn't + # rely on implicit STRING→TIMESTAMP casting. + param_type = "TIMESTAMP" if col.type_name == ColumnTypeName.TIMESTAMP else None + params.append(StatementParameterListItem(name=param_name, value=str(val), type=param_type)) + value_tuples.append(f"({', '.join(placeholders)})") insert_sql = f"INSERT INTO {self.fully_qualified_table_name} ({', '.join(self.column_names)}) VALUES {', '.join(value_tuples)}" # Execute the insert try: response = self.client.statement_execution.execute_statement( - statement=insert_sql, warehouse_id=self.warehouse_id, wait_timeout="30s" + statement=insert_sql, + warehouse_id=self.warehouse_id, + wait_timeout="30s", + parameters=params, ) if response.status.state.value == "SUCCEEDED": logger.info( @@ -494,10 +517,13 @@ class Databricks(VectorStoreBase): try: logger.info(f"Deleting vector with ID {vector_id} from Delta table {self.fully_qualified_table_name}") - delete_sql = f"DELETE FROM {self.fully_qualified_table_name} WHERE memory_id = '{vector_id}'" + delete_sql = f"DELETE FROM {self.fully_qualified_table_name} WHERE memory_id = :vector_id" response = self.client.statement_execution.execute_statement( - statement=delete_sql, warehouse_id=self.warehouse_id, wait_timeout="30s" + statement=delete_sql, + warehouse_id=self.warehouse_id, + wait_timeout="30s", + parameters=[StatementParameterListItem(name="vector_id", value=str(vector_id))], ) if response.status.state.value == "SUCCEEDED": @@ -519,8 +545,8 @@ class Databricks(VectorStoreBase): payload (dict, optional): New payload data. """ - update_sql = f"UPDATE {self.fully_qualified_table_name} SET " set_clauses = [] + params = [] if not vector_id: logger.error("vector_id is required for update operation") return @@ -528,25 +554,38 @@ class Databricks(VectorStoreBase): if not isinstance(vector, list): logger.error("vector must be a list of float values") return - set_clauses.append(f"embedding = {vector}") + # Vectors are numeric arrays — safe to inline since StatementParameterListItem + # doesn't support ARRAY types, and values are validated as list of floats above. + # Use array() SQL syntax, not Python list repr which is invalid Databricks SQL. + set_clauses.append(f"embedding = {self._format_sql_value(vector)}") if payload: if not isinstance(payload, dict): logger.error("payload must be a dictionary") return for key, value in payload.items(): if key not in excluded_keys: - set_clauses.append(f"{key} = '{value}'") + if not _VALID_SQL_IDENTIFIER.match(key): + logger.warning(f"Skipping invalid column name in payload: {key!r}") + continue + param_name = f"payload_{key}" + set_clauses.append(f"{key} = :{param_name}") + params.append(StatementParameterListItem(name=param_name, value=str(value))) if not set_clauses: logger.error("No fields to update") return + update_sql = f"UPDATE {self.fully_qualified_table_name} SET " update_sql += ", ".join(set_clauses) - update_sql += f" WHERE memory_id = '{vector_id}'" + update_sql += " WHERE memory_id = :vector_id" + params.append(StatementParameterListItem(name="vector_id", value=str(vector_id))) try: logger.info(f"Updating vector with ID {vector_id} in Delta table {self.fully_qualified_table_name}") response = self.client.statement_execution.execute_statement( - statement=update_sql, warehouse_id=self.warehouse_id, wait_timeout="30s" + statement=update_sql, + warehouse_id=self.warehouse_id, + wait_timeout="30s", + parameters=params, ) if response.status.state.value == "SUCCEEDED": diff --git a/tests/vector_stores/test_databricks.py b/tests/vector_stores/test_databricks.py index b1ea2c94a..93c0324d5 100644 --- a/tests/vector_stores/test_databricks.py +++ b/tests/vector_stores/test_databricks.py @@ -153,9 +153,16 @@ def test_insert_generates_sql(db_instance_direct, mock_workspace_client): sql = kwargs["statement"] if "statement" in kwargs else args[0] assert "INSERT INTO" in sql assert "catalog.schema.table" in sql - assert "id1" in sql - # Embedding list rendered + # Embedding list rendered inline (ARRAY type not supported by parameterized queries) assert "array(0.1, 0.2, 0.3, 0.4)" in sql + # User-supplied values should use parameterized queries, not inline values + assert ":memory_id_0" in sql + assert "id1" not in sql # id should be in params, not in SQL + params = kwargs["parameters"] + param_names = {p.name for p in params} + assert "memory_id_0" in param_names + id_param = next(p for p in params if p.name == "memory_id_0") + assert id_param.value == "id1" # ---------------------- Search Tests ---------------------- # @@ -246,7 +253,13 @@ def test_delete_vector(db_instance_delta, mock_workspace_client): db_instance_delta.delete("id-delete") args, kwargs = mock_workspace_client.statement_execution.execute_statement.call_args sql = kwargs.get("statement") or args[0] - assert "DELETE FROM" in sql and "id-delete" in sql + assert "DELETE FROM" in sql + assert ":vector_id" in sql + assert "id-delete" not in sql # value should be in params, not in SQL + params = kwargs["parameters"] + assert len(params) == 1 + assert params[0].name == "vector_id" + assert params[0].value == "id-delete" # ---------------------- Update Tests ---------------------- # @@ -260,10 +273,17 @@ def test_update_vector(db_instance_direct, mock_workspace_client): ) args, kwargs = mock_workspace_client.statement_execution.execute_statement.call_args sql = kwargs.get("statement") or args[0] - assert "UPDATE" in sql and "id-upd" in sql - assert "embedding = [0.4, 0.5, 0.6, 0.7]" in sql - assert "custom = 'val'" in sql + assert "UPDATE" in sql + assert "embedding = array(0.4, 0.5, 0.6, 0.7)" in sql + assert "custom = :payload_custom" in sql + assert ":vector_id" in sql + assert "id-upd" not in sql # value should be in params, not in SQL + assert "'val'" not in sql # value should be in params, not in SQL assert "user_id" not in sql # excluded + params = kwargs["parameters"] + param_map = {p.name: p.value for p in params} + assert param_map["payload_custom"] == "val" + assert param_map["vector_id"] == "id-upd" # ---------------------- Get Tests ---------------------- # @@ -703,9 +723,12 @@ def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): payloads=[{"data": "test memory", "user_id": "u1", "hash": "h1"}], ids=["mem-001"], ) - insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + insert_call = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + insert_sql = insert_call["statement"] assert "INSERT INTO cat.sch.tbl" in insert_sql - assert "mem-001" in insert_sql + assert ":memory_id_0" in insert_sql # parameterized, not inline + insert_params = {p.name: p.value for p in insert_call["parameters"]} + assert insert_params["memory_id_0"] == "mem-001" # SEARCH mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( @@ -753,16 +776,23 @@ def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): # UPDATE db.update(vector_id="mem-001", payload={"memory": "updated memory"}) - update_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + update_call = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + update_sql = update_call["statement"] assert "UPDATE cat.sch.tbl" in update_sql - assert "mem-001" in update_sql - assert "updated memory" in update_sql + assert ":vector_id" in update_sql + assert ":payload_memory" in update_sql + update_params = {p.name: p.value for p in update_call["parameters"]} + assert update_params["vector_id"] == "mem-001" + assert update_params["payload_memory"] == "updated memory" # DELETE db.delete("mem-001") - delete_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + delete_call = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + delete_sql = delete_call["statement"] assert "DELETE FROM cat.sch.tbl" in delete_sql - assert "mem-001" in delete_sql + assert ":vector_id" in delete_sql + delete_params = {p.name: p.value for p in delete_call["parameters"]} + assert delete_params["vector_id"] == "mem-001" def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): @@ -796,8 +826,11 @@ def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): payloads=[{"data": "direct memory", "user_id": "u1", "hash": "h1"}], ids=["mem-da-001"], ) - insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + insert_call = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + insert_sql = insert_call["statement"] assert "array(0.1, 0.2, 0.3, 0.4)" in insert_sql + insert_params = {p.name: p.value for p in insert_call["parameters"]} + assert insert_params["memory_id_0"] == "mem-da-001" # SEARCH with vector mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( @@ -838,3 +871,191 @@ def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert "query_vector" in list_kwargs assert "query_text" not in list_kwargs + + +# ---------------------- SQL Injection Prevention Tests (Issue #4073) ---------------------- # + + +SQLI_PAYLOADS = [ + "'; DELETE FROM table; --", + "' OR '1'='1", + "'; DROP TABLE memories; --", + "1' UNION SELECT * FROM secrets --", +] + + +@pytest.mark.parametrize("malicious_id", SQLI_PAYLOADS) +def test_delete_prevents_sql_injection(db_instance_delta, mock_workspace_client, malicious_id): + """Verify delete() uses parameterized queries so SQL injection payloads are never in the SQL string.""" + db_instance_delta.delete(malicious_id) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + # The malicious payload must NOT appear in the SQL statement itself + assert malicious_id not in sql + assert ":vector_id" in sql + # It should be safely passed as a parameter + assert kwargs["parameters"][0].value == malicious_id + + +@pytest.mark.parametrize("malicious_id", SQLI_PAYLOADS) +def test_update_prevents_sql_injection_in_vector_id(db_instance_delta, mock_workspace_client, malicious_id): + """Verify update() parameterizes vector_id.""" + db_instance_delta.update(vector_id=malicious_id, payload={"memory": "safe value"}) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert malicious_id not in sql + assert ":vector_id" in sql + param_map = {p.name: p.value for p in kwargs["parameters"]} + assert param_map["vector_id"] == malicious_id + + +@pytest.mark.parametrize("malicious_value", SQLI_PAYLOADS) +def test_update_prevents_sql_injection_in_payload(db_instance_delta, mock_workspace_client, malicious_value): + """Verify update() parameterizes payload values.""" + db_instance_delta.update(vector_id="safe-id", payload={"memory": malicious_value}) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert malicious_value not in sql + assert ":payload_memory" in sql + param_map = {p.name: p.value for p in kwargs["parameters"]} + assert param_map["payload_memory"] == malicious_value + + +@pytest.mark.parametrize("malicious_id", SQLI_PAYLOADS) +def test_insert_prevents_sql_injection_in_ids(db_instance_delta, mock_workspace_client, malicious_id): + """Verify insert() parameterizes memory IDs.""" + db_instance_delta.insert( + vectors=[[0.1, 0.2]], + payloads=[{"data": "test", "hash": "h1"}], + ids=[malicious_id], + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert malicious_id not in sql + assert ":memory_id_0" in sql + param_map = {p.name: p.value for p in kwargs["parameters"]} + assert param_map["memory_id_0"] == malicious_id + + +@pytest.mark.parametrize("malicious_data", SQLI_PAYLOADS) +def test_insert_prevents_sql_injection_in_payload_data(db_instance_delta, mock_workspace_client, malicious_data): + """Verify insert() parameterizes payload data values.""" + db_instance_delta.insert( + vectors=[[0.1, 0.2]], + payloads=[{"data": malicious_data, "hash": "h1"}], + ids=["safe-id"], + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert malicious_data not in sql + param_map = {p.name: p.value for p in kwargs["parameters"]} + assert param_map["memory_0"] == malicious_data + + +@pytest.mark.parametrize("malicious_key", [ + "memory; DROP TABLE x--", + "col' OR '1'='1", + "valid_col; DELETE FROM t", + "col\nname", +]) +def test_update_rejects_malicious_column_names(db_instance_delta, mock_workspace_client, malicious_key): + """Verify update() skips payload keys that are not valid SQL identifiers.""" + db_instance_delta.update( + vector_id="safe-id", + payload={malicious_key: "some value", "memory": "legit"}, + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + # The malicious key must NOT appear as a column name in the SQL + assert malicious_key not in sql + # The legitimate key should still be present + assert ":payload_memory" in sql + + +def test_insert_multi_row(db_instance_delta, mock_workspace_client): + """Verify multi-row insert generates unique parameter names per row.""" + db_instance_delta.insert( + vectors=[[0.1, 0.2], [0.3, 0.4]], + payloads=[ + {"data": "first memory", "hash": "h1"}, + {"data": "second memory", "hash": "h2"}, + ], + ids=["id-1", "id-2"], + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + params = kwargs["parameters"] + param_map = {p.name: p.value for p in params} + # Each row should have distinct parameter names + assert ":memory_id_0" in sql + assert ":memory_id_1" in sql + assert param_map["memory_id_0"] == "id-1" + assert param_map["memory_id_1"] == "id-2" + assert param_map["memory_0"] == "first memory" + assert param_map["memory_1"] == "second memory" + + +def test_insert_with_none_values(db_instance_delta, mock_workspace_client): + """Verify insert() uses NULL literal for None values, not parameters.""" + db_instance_delta.insert( + vectors=[[0.1, 0.2]], + payloads=[{"data": "test", "hash": "h1"}], # agent_id, run_id etc will be None + ids=["id-1"], + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert "NULL" in sql # None values should be literal NULL + # Only non-None values should have parameters + param_names = {p.name for p in kwargs["parameters"]} + assert "memory_id_0" in param_names + assert "memory_0" in param_names + assert "hash_0" in param_names + + +def test_update_payload_only(db_instance_delta, mock_workspace_client): + """Verify update() works with payload only (no vector).""" + db_instance_delta.update(vector_id="id-1", payload={"memory": "new text"}) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert "UPDATE" in sql + assert "embedding" not in sql + assert ":payload_memory" in sql + assert ":vector_id" in sql + + +def test_update_vector_only(db_instance_direct, mock_workspace_client): + """Verify update() works with vector only (no payload).""" + db_instance_direct.update(vector_id="id-1", vector=[0.1, 0.2, 0.3, 0.4]) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + sql = kwargs["statement"] + assert "UPDATE" in sql + assert "embedding = array(0.1, 0.2, 0.3, 0.4)" in sql + assert ":vector_id" in sql + param_map = {p.name: p.value for p in kwargs["parameters"]} + assert param_map["vector_id"] == "id-1" + assert len(kwargs["parameters"]) == 1 # only vector_id param + + +def test_insert_timestamp_params_have_explicit_type(db_instance_delta, mock_workspace_client): + """Verify insert() sets type='TIMESTAMP' on created_at/updated_at parameters + so Databricks doesn't rely on implicit STRING->TIMESTAMP casting.""" + db_instance_delta.insert( + vectors=[[0.1, 0.2]], + payloads=[{ + "data": "test", + "hash": "h1", + "created_at": "2024-01-01T00:00:00+00:00", + "updated_at": "2024-01-02T00:00:00+00:00", + }], + ids=["id-1"], + ) + kwargs = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs + params = kwargs["parameters"] + ts_params = [p for p in params if "created_at" in p.name or "updated_at" in p.name] + assert len(ts_params) == 2 + for p in ts_params: + assert p.type == "TIMESTAMP", f"Parameter {p.name} should have type=TIMESTAMP, got {p.type}" + # Non-timestamp params should not have a type set (defaults to STRING) + non_ts_params = [p for p in params if "created_at" not in p.name and "updated_at" not in p.name] + for p in non_ts_params: + assert p.type is None, f"Parameter {p.name} should not have explicit type, got {p.type}"