Merge pull request #6503 from KnibbsyMan/feat/sql-injection-prevention

FEAT - Automatic SQL Expression Parameterization
This commit is contained in:
Sipke Schoorstra 2025-03-17 10:05:22 +01:00 committed by GitHub
commit a04aaf69e8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 248 additions and 227 deletions

View file

@ -1,55 +1,16 @@
using MySql.Data.MySqlClient;
using Elsa.Sql.Client;
using System.Data;
using System.Data.Common;
namespace Elsa.Sql.MySql;
public class MySqlClient : BaseSqlClient, ISqlClient
/// <summary>
/// MySql client implementation.
/// </summary>
/// <param name="connectionString"></param>
public class MySqlClient(string connectionString) : BaseSqlClient(connectionString)
{
private string? _connectionString;
protected override DbConnection CreateConnection() => new MySqlConnection(_connectionString);
/// <summary>
/// MySql client implimentation.
/// </summary>
/// <param name="connectionString"></param>
public MySqlClient(string? connectionString) => _connectionString = connectionString;
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<int?> ExecuteCommandAsync(string sqlCommand)
{
using var connection = new MySqlConnection(_connectionString);
connection.Open();
var command = new MySqlCommand(sqlCommand, connection);
var result = await command.ExecuteNonQueryAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<object?> ExecuteScalarAsync(string sqlQuery)
{
using var connection = new MySqlConnection(_connectionString);
connection.Open();
var command = new MySqlCommand(sqlQuery, connection);
var result = await command.ExecuteScalarAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<DataSet?> ExecuteQueryAsync(string sqlQuery)
{
using var connection = new MySqlConnection(_connectionString);
connection.Open();
var command = new MySqlCommand(sqlQuery, connection);
using var reader = await command.ExecuteReaderAsync();
return await Task.FromResult(ReadAsDataSet(reader));
}
protected override DbCommand CreateCommand(string query, DbConnection connection) => new MySqlCommand(query, (MySqlConnection)connection);
}

View file

@ -1,55 +1,16 @@
using Npgsql;
using Elsa.Sql.Client;
using System.Data;
using System.Data.Common;
namespace Elsa.Sql.PostgreSql;
public class PostgreSqlClient : BaseSqlClient, ISqlClient
/// <summary>
/// PostgreSQL client implementation.
/// </summary>
/// <param name="connectionString"></param>
public class PostgreSqlClient(string connectionString) : BaseSqlClient(connectionString)
{
private string? _connectionString;
protected override DbConnection CreateConnection() => new NpgsqlConnection(_connectionString);
/// <summary>
/// PostgreSQL client implimentation.
/// </summary>
/// <param name="connectionString"></param>
public PostgreSqlClient(string? connectionString) => _connectionString = connectionString;
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<int?> ExecuteCommandAsync(string sqlCommand)
{
using var connection = new NpgsqlConnection(_connectionString);
connection.Open();
var command = new NpgsqlCommand(sqlCommand, connection);
var result = await command.ExecuteNonQueryAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<object?> ExecuteScalarAsync(string sqlQuery)
{
using var connection = new NpgsqlConnection(_connectionString);
connection.Open();
var command = new NpgsqlCommand(sqlQuery, connection);
var result = await command.ExecuteScalarAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<DataSet?> ExecuteQueryAsync(string sqlQuery)
{
using var connection = new NpgsqlConnection(_connectionString);
connection.Open();
var command = new NpgsqlCommand(sqlQuery, connection);
using var reader = await command.ExecuteReaderAsync();
return await Task.FromResult(ReadAsDataSet(reader));
}
protected override DbCommand CreateCommand(string query, DbConnection connection) => new NpgsqlCommand(query, (NpgsqlConnection)connection);
}

View file

@ -1,55 +1,16 @@
using System.Data;
using System.Data.Common;
using Elsa.Sql.Client;
using Microsoft.Data.SqlClient;
namespace Elsa.Sql.SqlServer;
public class SqlServerClient : BaseSqlClient, ISqlClient
/// <summary>
/// Microsoft SQL server client implementation.
/// </summary>
/// <param name="connectionString"></param>
public class SqlServerClient(string connectionString) : BaseSqlClient(connectionString)
{
private string? _connectionString;
protected override DbConnection CreateConnection() => new SqlConnection(_connectionString);
/// <summary>
/// Microsoft SQL server client implimentation.
/// </summary>
/// <param name="connectionString"></param>
public SqlServerClient(string? connectionString) => _connectionString = connectionString;
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<int?> ExecuteCommandAsync(string sqlCommand)
{
using var connection = new SqlConnection(_connectionString);
connection.Open();
var command = new SqlCommand(sqlCommand, connection);
var result = await command.ExecuteNonQueryAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<object?> ExecuteScalarAsync(string sqlQuery)
{
using var connection = new SqlConnection(_connectionString);
connection.Open();
var command = new SqlCommand(sqlQuery, connection);
var result = await command.ExecuteScalarAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<DataSet?> ExecuteQueryAsync(string sqlQuery)
{
using var connection = new SqlConnection(_connectionString);
connection.Open();
var command = new SqlCommand(sqlQuery, connection);
using var reader = await command.ExecuteReaderAsync();
return await Task.FromResult(ReadAsDataSet(reader));
}
protected override DbCommand CreateCommand(string query, DbConnection connection) => new SqlCommand(query, (SqlConnection)connection);
}

View file

@ -1,55 +1,16 @@
using System.Data;
using System.Data.Common;
using Elsa.Sql.Client;
using Microsoft.Data.Sqlite;
namespace Elsa.Sql.Sqlite;
public class SqliteClient : BaseSqlClient, ISqlClient
/// <summary>
/// Sqlite client implementation.
/// </summary>
/// <param name="connectionString"></param>
public class SqliteClient(string connectionString) : BaseSqlClient(connectionString)
{
private string? _connectionString;
protected override DbConnection CreateConnection() => new SqliteConnection(_connectionString);
/// <summary>
/// Sqlite client implimentation.
/// </summary>
/// <param name="connectionString"></param>
public SqliteClient(string? connectionString) => _connectionString = connectionString;
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<int?> ExecuteCommandAsync(string sqlCommand)
{
using var connection = new SqliteConnection(_connectionString);
connection.Open();
var command = new SqliteCommand(sqlCommand, connection);
var result = await command.ExecuteNonQueryAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<object?> ExecuteScalarAsync(string sqlQuery)
{
using var connection = new SqliteConnection(_connectionString);
connection.Open();
var command = new SqliteCommand(sqlQuery, connection);
var result = await command.ExecuteScalarAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<DataSet?> ExecuteQueryAsync(string sqlQuery)
{
using var connection = new SqliteConnection(_connectionString);
connection.Open();
var command = new SqliteCommand(sqlQuery, connection);
using var reader = await command.ExecuteReaderAsync();
return await Task.FromResult(ReadAsDataSet(reader));
}
protected override DbCommand CreateCommand(string query, DbConnection connection) => new SqliteCommand(query, (SqliteConnection)connection);
}

View file

@ -1,15 +1,120 @@
using System.Data;
using System.Data.Common;
using Elsa.Sql.Models;
namespace Elsa.Sql.Client;
public abstract class BaseSqlClient
public abstract class BaseSqlClient : ISqlClient
{
/// <summary>
/// The connection string used to connect with the database.
/// </summary>
protected readonly string _connectionString;
/// <summary>
/// The marker used when injecting parameters into a query.
/// Default: "@"
/// </summary>
public virtual string ParameterMarker { get; set; } = "@";
/// <summary>
/// The text following the <c>ParameterMarker</c> when injecting parameters into a query.
/// Default: "param"
/// </summary>
public virtual string ParameterText { get; set; } = "p";
/// <summary>
/// Set to true to add a counter to the end of the parameter string.
/// Default: true
/// </summary>
public virtual bool IncrementParameter { get; set; } = true;
/// <summary>
/// Default base implementation for an SQL client.
/// </summary>
/// <param name="connectionString"></param>
protected BaseSqlClient(string connectionString) => _connectionString = connectionString;
/// <summary>
/// Create a connection using the client specific connection.
/// </summary>
/// <returns></returns>
protected abstract DbConnection CreateConnection();
/// <summary>
/// Create a command using the client specific connection.
/// </summary>
/// <param name="query"></param>
/// <param name="connection"></param>
/// <returns></returns>
protected abstract DbCommand CreateCommand(string query, DbConnection connection);
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<int?> ExecuteCommandAsync(EvaluatedQuery evaluatedQuery)
{
using var connection = CreateConnection();
connection.Open();
var command = CreateCommand(evaluatedQuery.Query, connection);
AddCommandParameters(command, evaluatedQuery.Parameters);
var result = await command.ExecuteNonQueryAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<object?> ExecuteScalarAsync(EvaluatedQuery evaluatedQuery)
{
using var connection = CreateConnection();
connection.Open();
var command = CreateCommand(evaluatedQuery.Query, connection);
AddCommandParameters(command, evaluatedQuery.Parameters);
var result = await command.ExecuteScalarAsync();
return result;
}
/// <summary>
/// <inheritdoc/>
/// </summary>
public async Task<DataSet?> ExecuteQueryAsync(EvaluatedQuery evaluatedQuery)
{
using var connection = CreateConnection();
connection.Open();
var command = CreateCommand(evaluatedQuery.Query, connection);
AddCommandParameters(command, evaluatedQuery.Parameters);
using var reader = await command.ExecuteReaderAsync();
return await Task.FromResult(ReadAsDataSet(reader));
}
/// <summary>
/// Add parameters into the query to prevent SQL injection.
/// </summary>
/// <param name="command">Command to add the parameters to</param>
/// <param name="parameters">Parameters to add</param>
/// <returns></returns>
private DbCommand AddCommandParameters(DbCommand command, Dictionary<string, object?> parameters)
{
foreach (var param in parameters)
{
var dbParam = command.CreateParameter();
dbParam.ParameterName = param.Key;
dbParam.Value = param.Value ?? DBNull.Value;
command.Parameters.Add(dbParam);
}
return command;
}
/// <summary>
/// Returns <see cref="IDataReader"/> data as a <see cref="DataSet"/>.
/// </summary>
/// <param name="reader">Reader to return data from.</param>
/// <returns><see cref="DataSet"/> of data.</returns>
protected static DataSet ReadAsDataSet(IDataReader reader)
private DataSet ReadAsDataSet(IDataReader reader)
{
var dataSet = new DataSet("dataset");
dataSet.Tables.Add(ReadAsDataTable(reader));
@ -21,7 +126,7 @@ public abstract class BaseSqlClient
/// </summary>
/// <param name="reader">Reader to return data from.</param>
/// <returns><see cref="DataTable"/> of data.</returns>
protected static DataTable ReadAsDataTable(IDataReader reader)
private DataTable ReadAsDataTable(IDataReader reader)
{
var data = new DataTable();
var schemaTable =reader.GetSchemaTable();

View file

@ -1,27 +1,43 @@
using System.Data;
using Elsa.Sql.Models;
namespace Elsa.Sql.Client;
public interface ISqlClient
{
/// <summary>
/// The marker used when injecting parameters into a query.
/// </summary>
public string ParameterMarker { get; set; }
/// <summary>
/// The text following the <c>ParameterMarker</c> when injecting parameters into a query.
/// </summary>
public string ParameterText { get; set; }
/// <summary>
/// Set to true to add a counter to the end of the parameter string.
/// </summary>
public bool IncrementParameter { get; set; }
/// <summary>
/// Asynchronously executes a Transact-SQL statement against the connection and returns the number of rows affected.
/// </summary>
/// <param name="sqlCommand">The command to execute</param>
/// <param name="evaluatedQuery">The evaluated query to execute.</param>
/// <returns>The number of rows affected.</returns>
public Task<int?> ExecuteCommandAsync(string sqlCommand);
public Task<int?> ExecuteCommandAsync(EvaluatedQuery evaluatedQuery);
/// <summary>
/// Asynchronously executes the query, and returns the first column of the first row in the result set returned by the query. Additional columns or rows are ignored.
/// </summary>
/// <param name="sqlQuery">The query to execute</param>
/// <param name="evaluatedQuery">The evaluated query to execute.</param>
/// <returns>The first column of the first row in the result set, or a null reference if the result set is empty. Returns a maximum of 2033 characters.</returns>
public Task<object?> ExecuteScalarAsync(string sqlQuery);
public Task<object?> ExecuteScalarAsync(EvaluatedQuery evaluatedQuery);
/// <summary>
/// Asynchronously executes the query, and returns a dataset of data returned by the query.
/// </summary>
/// <param name="sqlQuery">Query to execute</param>
/// <param name="evaluatedQuery">The evaluated query to execute.</param>
/// <returns>DataSet of the queried data</returns>
public Task<DataSet?> ExecuteQueryAsync(string sqlQuery);
public Task<DataSet?> ExecuteQueryAsync(EvaluatedQuery evaluatedQuery);
}

View file

@ -1,4 +1,5 @@
using Elsa.Expressions.Models;
using Elsa.Sql.Models;
using JetBrains.Annotations;
namespace Elsa.Sql.Contracts;
@ -16,8 +17,8 @@ public interface ISqlEvaluator
/// <param name="context">The context in which the expression is evaluated.</param>
/// <param name="options">A set of options.</param>
/// <param name="cancellationToken">An optional cancellation token.</param>
/// <returns>The result of the evaluation.</returns>
Task<string?> EvaluateAsync(
/// <returns>The <see cref="EvaluatedQuery"/> result.</returns>
Task<EvaluatedQuery> EvaluateAsync(
string expression,
ExpressionExecutionContext context,
ExpressionEvaluatorOptions options,

View file

@ -0,0 +1,35 @@
namespace Elsa.Sql.Models
{
/// <summary>
/// Represents a safely formatted SQL expression result.
/// </summary>
public class EvaluatedQuery
{
/// <summary>
/// Query with parameterized values
/// </summary>
public string Query { get; set; }
/// <summary>
/// Parameters to inject into the query at execution
/// </summary>
public Dictionary<string, object?> Parameters { get; set; } = new Dictionary<string, object?>();
/// <summary>
/// An evaluated query response.
/// </summary>
/// <param name="query">The evaluated query</param>
public EvaluatedQuery(string query) => Query = query;
/// <summary>
/// An evaluated query response.
/// </summary>
/// <param name="query">The evaluated query</param>
/// <param name="parameters">Parameters to pass into the parameterized query</param>
public EvaluatedQuery(string query, Dictionary<string, object?> parameters)
{
Query = query;
Parameters = parameters;
}
}
}

View file

@ -2,6 +2,8 @@
using Elsa.Expressions.Models;
using Elsa.Extensions;
using Elsa.Sql.Contracts;
using Elsa.Sql.Models;
using Elsa.Workflows;
namespace Elsa.Sql.Services;
@ -13,60 +15,78 @@ namespace Elsa.Sql.Services;
/// </remarks>
public class SqlEvaluator() : ISqlEvaluator
{
private WorkflowExecutionContext executionContext;
private ActivityExecutionContext activityContext;
private ExpressionExecutionContext expressionContext;
/// <inheritdoc />
public async Task<string?> EvaluateAsync(
public async Task<EvaluatedQuery> EvaluateAsync(
string expression,
ExpressionExecutionContext context,
ExpressionEvaluatorOptions options,
CancellationToken cancellationToken = default)
{
if (!expression.Contains("@")) return expression;
if (!expression.Contains("{{")) return new EvaluatedQuery(expression);
expressionContext = context;
executionContext = context.GetWorkflowExecutionContext();
activityContext = context.GetActivityExecutionContext();
// Create client
var factory = context.GetRequiredService<ISqlClientFactory>();
var client = factory.CreateClient(activityContext.ActivityState["Client"].ToString(), activityContext.ActivityState["ConnectionString"].ToString());
var sb = new StringBuilder();
int start = 0;
var parameters = new Dictionary<string, object?>();
int paramIndex = 0;
while (start < expression.Length)
{
int atIndex = expression.IndexOf('@', start);
if (atIndex == -1)
int openIndex = expression.IndexOf("{{", start);
if (openIndex == -1)
{
sb.Append(expression.Substring(start));
sb.Append(expression.AsSpan(start));
break;
}
sb.Append(expression.Substring(start, atIndex - start));
// Append everything before {{
sb.Append(expression.AsSpan(start, openIndex - start));
int endIndex = atIndex + 1;
while (endIndex < expression.Length && (char.IsLetterOrDigit((char)expression[endIndex]) || expression[endIndex] == '.' || expression[endIndex] == '_'))
{
endIndex++;
}
// Find the closing }}
int closeIndex = expression.IndexOf("}}", openIndex + 2);
if (closeIndex == -1) throw new FormatException("Unmatched '{{' found in SQL expression.");
string key = expression.Substring(atIndex + 1, endIndex - atIndex - 1);
object? value = ResolveValue(key, context);
if (value is null) throw new NullReferenceException($"No value found for '{key}'.");
// Extract key
string key = expression.Substring(openIndex + 2, closeIndex - openIndex - 2).Trim();
if (string.IsNullOrEmpty(key)) throw new FormatException("Empty placeholder '{{}}' is not allowed.");
sb.Append(value?.ToString() ?? $"<{key}>");
start = endIndex;
// Resolve value and replace with parameterized name
var counterValue = client.IncrementParameter ? paramIndex++.ToString() : string.Empty;
string paramName = $"{client.ParameterMarker}{client.ParameterText}{counterValue}";
parameters[paramName] = ResolveValue(key);
sb.Append(paramName);
start = closeIndex + 2;
}
return sb.ToString();
return new EvaluatedQuery(sb.ToString(), parameters);
}
private object? ResolveValue(string key, ExpressionExecutionContext context)
private object? ResolveValue(string key)
{
return key switch
{
"Workflow.Definition.Id" => context.GetWorkflowExecutionContext().Workflow.Identity.DefinitionId,
"Workflow.Definition.Version.Id" => context.GetWorkflowExecutionContext().Workflow.Identity.Id,
"Workflow.Definition.Version" => context.GetWorkflowExecutionContext().Workflow.Identity.Version,
"Workflow.Instance.Id" => context.GetActivityExecutionContext().WorkflowExecutionContext.Id,
"Correlation.Id" => context.GetActivityExecutionContext().WorkflowExecutionContext.CorrelationId,
"LastResult" => context.GetLastResult(),
var i when i.StartsWith("Input.") => context.GetWorkflowExecutionContext().Input.TryGetValue(i.Substring(6), out var v) ? v : null,
var o when o.StartsWith("Output.") => context.GetWorkflowExecutionContext().Output.TryGetValue(o.Substring(7), out var v) ? v : null,
var v when v.StartsWith("Variables.") => context.GetWorkflowExecutionContext().Variables.FirstOrDefault(x => x.Name == v.Substring(10), null)?.Value ?? null,
_ => null
"Workflow.Definition.Id" => executionContext.Workflow.Identity.DefinitionId,
"Workflow.Definition.Version.Id" => executionContext.Workflow.Identity.Id,
"Workflow.Definition.Version" => executionContext.Workflow.Identity.Version,
"Workflow.Instance.Id" => activityContext.WorkflowExecutionContext.Id,
"Correlation.Id" => activityContext.WorkflowExecutionContext.CorrelationId,
"LastResult" => expressionContext.GetLastResult(),
var i when i.StartsWith("Input.") => executionContext.Input.TryGetValue(i.Substring(6), out var v) ? v : null,
var o when o.StartsWith("Output.") => executionContext.Output.TryGetValue(o.Substring(7), out var v) ? v : null,
var v when v.StartsWith("Variables.") => executionContext.Variables.FirstOrDefault(x => x.Name == v.Substring(10), null)?.Value ?? null,
_ => throw new NullReferenceException($"No matching property found for {{{{{key}}}}}.")
};
}
}