Compare commits

...

3 Commits

Author SHA1 Message Date
utkarsh240799 5355a226a6 fix: remove unused StatementParameterListItem import in tests
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-26 19:14:21 +05:30
utkarsh240799 1eb6bd8303 fix: ensure Databricks SQL compatibility for parameterized queries
- Add explicit type='TIMESTAMP' on StatementParameterListItem for
  created_at/updated_at columns instead of relying on implicit
  STRING->TIMESTAMP casting
- Fix pre-existing bug: update() used Python list repr [0.1, 0.2]
  for embedding which is invalid Databricks SQL, now uses
  _format_sql_value() to produce array(0.1, 0.2) syntax

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-26 19:11:13 +05:30
utkarsh240799 c6cc8f14b8 fix: prevent SQL injection in Databricks vector store (#4073)
Replace f-string interpolation with Databricks parameterized queries
using StatementParameterListItem in delete(), update(), and insert()
methods. Add column name validation in update() to reject invalid SQL
identifiers from payload keys. Embedding vectors remain inlined as
they are numeric arrays not supported by the parameterization API.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-26 19:03:10 +05:30
2 changed files with 286 additions and 26 deletions
+51 -12
View File
@@ -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":
+235 -14
View File
@@ -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}"