diff --git a/src/modules/Elsa.Sql.MySql/MySqlClient.cs b/src/modules/Elsa.Sql.MySql/MySqlClient.cs index 60669f7d1..425f701de 100644 --- a/src/modules/Elsa.Sql.MySql/MySqlClient.cs +++ b/src/modules/Elsa.Sql.MySql/MySqlClient.cs @@ -10,6 +10,12 @@ namespace Elsa.Sql.MySql; /// public class MySqlClient(string connectionString) : BaseSqlClient(connectionString) { + public override string ParameterMarker { get; set; } = "@"; + + public override string ParameterText { get; set; } = ""; + + public override bool IncrementParameter { get; set; } = false; + protected override DbConnection CreateConnection() => new MySqlConnection(_connectionString); protected override DbCommand CreateCommand(string query, DbConnection connection) => new MySqlCommand(query, (MySqlConnection)connection); diff --git a/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs b/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs index e8780c359..96efacb63 100644 --- a/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs +++ b/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs @@ -10,6 +10,8 @@ namespace Elsa.Sql.SqlServer; /// public class SqlServerClient(string connectionString) : BaseSqlClient(connectionString) { + public override string ParameterText { get; set; } = "p"; + protected override DbConnection CreateConnection() => new SqlConnection(_connectionString); protected override DbCommand CreateCommand(string query, DbConnection connection) => new SqlCommand(query, (SqlConnection)connection); diff --git a/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs b/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs index ec43f5245..e6ad5afcb 100644 --- a/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs +++ b/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs @@ -10,6 +10,8 @@ namespace Elsa.Sql.Sqlite; /// public class SqliteClient(string connectionString) : BaseSqlClient(connectionString) { + public override string ParameterText { get; set; } = "p"; + protected override DbConnection CreateConnection() => new SqliteConnection(_connectionString); protected override DbCommand CreateCommand(string query, DbConnection connection) => new SqliteCommand(query, (SqliteConnection)connection); diff --git a/src/modules/Elsa.Sql/Client/BaseSqlClient.cs b/src/modules/Elsa.Sql/Client/BaseSqlClient.cs index 872d4cfcc..f8714bd26 100644 --- a/src/modules/Elsa.Sql/Client/BaseSqlClient.cs +++ b/src/modules/Elsa.Sql/Client/BaseSqlClient.cs @@ -1,5 +1,6 @@ using System.Data; using System.Data.Common; +using System.Text; using Elsa.Sql.Models; namespace Elsa.Sql.Client; @@ -11,6 +12,24 @@ public abstract class BaseSqlClient : ISqlClient /// protected readonly string _connectionString; + /// + /// The marker used when injecting parameters into a query. + /// Default: "@" + /// + public virtual string ParameterMarker { get; set; } = "@"; + + /// + /// The text following the ParameterMarkerwhen injecting parameters into a query + /// Default: + /// + public virtual string ParameterText { get; set; } = string.Empty; + + /// + /// Set to true to add a counter to the end of the parameter string + /// Default: false + /// + public virtual bool IncrementParameter { get; set; } = true; + /// /// Create a connection using the client specific connection. /// @@ -38,8 +57,9 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var command = CreateCommand(evaluatedQuery.Query, connection); - AddParameters(command, evaluatedQuery.Parameters); + var query = ReplaceQueryParameters(evaluatedQuery); + var command = CreateCommand(query, connection); + AddCommandParameters(command, evaluatedQuery.Parameters); var result = await command.ExecuteNonQueryAsync(); return result; @@ -52,8 +72,9 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var command = CreateCommand(evaluatedQuery.Query, connection); - AddParameters(command, evaluatedQuery.Parameters); + var query = ReplaceQueryParameters(evaluatedQuery); + var command = CreateCommand(query, connection); + AddCommandParameters(command, evaluatedQuery.Parameters); var result = await command.ExecuteScalarAsync(); return result; @@ -66,20 +87,42 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var command = CreateCommand(evaluatedQuery.Query, connection); - AddParameters(command, evaluatedQuery.Parameters); + var query = ReplaceQueryParameters(evaluatedQuery); + var command = CreateCommand(query, connection); + AddCommandParameters(command, evaluatedQuery.Parameters); using var reader = await command.ExecuteReaderAsync(); return await Task.FromResult(ReadAsDataSet(reader)); } + /// + /// Replace the evaluated parameters with client specific parameters. + /// + /// Query to replace parameters for. + /// + private string ReplaceQueryParameters(EvaluatedQuery evaluatedQuery) + { + var count = 1; + var clientUpdatedParams = new Dictionary(); + var queryBuilder = new StringBuilder(evaluatedQuery.Query); + foreach (var param in evaluatedQuery.Parameters) + { + var counterValue = IncrementParameter ? count++.ToString() : string.Empty; + var newKey = $"{ParameterMarker}{ParameterText}{counterValue}"; + queryBuilder.Replace(param.Key, newKey); + clientUpdatedParams[newKey] = param.Value; + } + evaluatedQuery.Parameters = clientUpdatedParams; + return queryBuilder.ToString(); + } + /// /// Inject parameters into the query to prevent SQL injection. /// /// Command to add the parameters to /// Parameters to add /// - private DbCommand AddParameters(DbCommand command, Dictionary parameters) + private DbCommand AddCommandParameters(DbCommand command, Dictionary parameters) { // Add parameters dynamically foreach (var param in parameters) diff --git a/src/modules/Elsa.Sql/Elsa.Sql.csproj b/src/modules/Elsa.Sql/Elsa.Sql.csproj index 670705651..bfb24be1e 100644 --- a/src/modules/Elsa.Sql/Elsa.Sql.csproj +++ b/src/modules/Elsa.Sql/Elsa.Sql.csproj @@ -9,7 +9,6 @@ - diff --git a/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs b/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs index 36aae3131..1570fab0e 100644 --- a/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs +++ b/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs @@ -8,12 +8,12 @@ /// /// Query with parameterized values /// - public string Query { get; } + public string Query { get; set; } /// /// Parameters to inject into the query at execution /// - public Dictionary Parameters { get; } = new Dictionary(); + public Dictionary Parameters { get; set; } = new Dictionary(); /// /// An evaluated query response. diff --git a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs index e8a06998d..7bfdc42ea 100644 --- a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs +++ b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs @@ -21,39 +21,43 @@ public class SqlEvaluator() : ISqlEvaluator ExpressionEvaluatorOptions options, CancellationToken cancellationToken = default) { - if (!expression.Contains("@")) return new EvaluatedQuery(expression); + if (!expression.Contains("{{")) return new EvaluatedQuery(expression); var sb = new StringBuilder(); - var parameters = new Dictionary(); int start = 0; + var parameters = new Dictionary(); 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); + // Extract key + string key = expression.Substring(openIndex + 2, closeIndex - openIndex - 2).Trim(); + if (string.IsNullOrEmpty(key)) throw new FormatException("Empty placeholder '{{}}' is not allowed."); + + // Resolve value object? value = ResolveValue(key, context); - if (value is null) throw new NullReferenceException($"No value found for '{key}'."); + if (value is null) throw new NullReferenceException($"No value found for '{{{{{key}}}}}'."); - string paramName = $"@param{paramIndex++}"; + // Replace with parameterized name + string paramName = $"{{{{p{paramIndex++}}}}}"; parameters[paramName] = value; sb.Append(paramName); - start = endIndex; + start = closeIndex + 2; } return new EvaluatedQuery(sb.ToString(), parameters);