diff --git a/src/modules/Elsa.Sql.MySql/MySqlClient.cs b/src/modules/Elsa.Sql.MySql/MySqlClient.cs index 60d84fd3e..60669f7d1 100644 --- a/src/modules/Elsa.Sql.MySql/MySqlClient.cs +++ b/src/modules/Elsa.Sql.MySql/MySqlClient.cs @@ -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 +/// +/// MySql client implementation. +/// +/// +public class MySqlClient(string connectionString) : BaseSqlClient(connectionString) { - private string? _connectionString; + protected override DbConnection CreateConnection() => new MySqlConnection(_connectionString); - /// - /// MySql client implimentation. - /// - /// - public MySqlClient(string? connectionString) => _connectionString = connectionString; - - /// - /// - /// - public async Task ExecuteCommandAsync(string sqlCommand) - { - using var connection = new MySqlConnection(_connectionString); - connection.Open(); - var command = new MySqlCommand(sqlCommand, connection); - - var result = await command.ExecuteNonQueryAsync(); - return result; - } - - /// - /// - /// - public async Task ExecuteScalarAsync(string sqlQuery) - { - using var connection = new MySqlConnection(_connectionString); - connection.Open(); - var command = new MySqlCommand(sqlQuery, connection); - - var result = await command.ExecuteScalarAsync(); - return result; - } - - /// - /// - /// - public async Task 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); } \ No newline at end of file diff --git a/src/modules/Elsa.Sql.PostgreSql/PostgreSqlClient.cs b/src/modules/Elsa.Sql.PostgreSql/PostgreSqlClient.cs index 495439f30..5eb5982ba 100644 --- a/src/modules/Elsa.Sql.PostgreSql/PostgreSqlClient.cs +++ b/src/modules/Elsa.Sql.PostgreSql/PostgreSqlClient.cs @@ -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 +/// +/// PostgreSQL client implementation. +/// +/// +public class PostgreSqlClient(string connectionString) : BaseSqlClient(connectionString) { - private string? _connectionString; + protected override DbConnection CreateConnection() => new NpgsqlConnection(_connectionString); - /// - /// PostgreSQL client implimentation. - /// - /// - public PostgreSqlClient(string? connectionString) => _connectionString = connectionString; - - /// - /// - /// - public async Task ExecuteCommandAsync(string sqlCommand) - { - using var connection = new NpgsqlConnection(_connectionString); - connection.Open(); - var command = new NpgsqlCommand(sqlCommand, connection); - - var result = await command.ExecuteNonQueryAsync(); - return result; - } - - /// - /// - /// - public async Task ExecuteScalarAsync(string sqlQuery) - { - using var connection = new NpgsqlConnection(_connectionString); - connection.Open(); - var command = new NpgsqlCommand(sqlQuery, connection); - - var result = await command.ExecuteScalarAsync(); - return result; - } - - /// - /// - /// - public async Task 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); } \ No newline at end of file diff --git a/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs b/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs index f5b57c263..e8780c359 100644 --- a/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs +++ b/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs @@ -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 +/// +/// Microsoft SQL server client implementation. +/// +/// +public class SqlServerClient(string connectionString) : BaseSqlClient(connectionString) { - private string? _connectionString; + protected override DbConnection CreateConnection() => new SqlConnection(_connectionString); - /// - /// Microsoft SQL server client implimentation. - /// - /// - public SqlServerClient(string? connectionString) => _connectionString = connectionString; - - /// - /// - /// - public async Task ExecuteCommandAsync(string sqlCommand) - { - using var connection = new SqlConnection(_connectionString); - connection.Open(); - var command = new SqlCommand(sqlCommand, connection); - - var result = await command.ExecuteNonQueryAsync(); - return result; - } - - /// - /// - /// - public async Task ExecuteScalarAsync(string sqlQuery) - { - using var connection = new SqlConnection(_connectionString); - connection.Open(); - var command = new SqlCommand(sqlQuery, connection); - - var result = await command.ExecuteScalarAsync(); - return result; - } - - /// - /// - /// - public async Task 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); } \ No newline at end of file diff --git a/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs b/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs index d67e69342..ec43f5245 100644 --- a/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs +++ b/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs @@ -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 +/// +/// Sqlite client implementation. +/// +/// +public class SqliteClient(string connectionString) : BaseSqlClient(connectionString) { - private string? _connectionString; + protected override DbConnection CreateConnection() => new SqliteConnection(_connectionString); - /// - /// Sqlite client implimentation. - /// - /// - public SqliteClient(string? connectionString) => _connectionString = connectionString; - - /// - /// - /// - public async Task ExecuteCommandAsync(string sqlCommand) - { - using var connection = new SqliteConnection(_connectionString); - connection.Open(); - var command = new SqliteCommand(sqlCommand, connection); - - var result = await command.ExecuteNonQueryAsync(); - return result; - } - - /// - /// - /// - public async Task ExecuteScalarAsync(string sqlQuery) - { - using var connection = new SqliteConnection(_connectionString); - connection.Open(); - var command = new SqliteCommand(sqlQuery, connection); - - var result = await command.ExecuteScalarAsync(); - return result; - } - - /// - /// - /// - public async Task 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); } \ No newline at end of file diff --git a/src/modules/Elsa.Sql/Client/BaseSqlClient.cs b/src/modules/Elsa.Sql/Client/BaseSqlClient.cs index 6ed7dbfd9..9f7cca113 100644 --- a/src/modules/Elsa.Sql/Client/BaseSqlClient.cs +++ b/src/modules/Elsa.Sql/Client/BaseSqlClient.cs @@ -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 { + /// + /// The connection string used to connect with the database. + /// + protected readonly string _connectionString; + + /// + /// The marker used when injecting parameters into a query. + /// Default: "@" + /// + public virtual string ParameterMarker { get; set; } = "@"; + + /// + /// The text following the ParameterMarker when injecting parameters into a query. + /// Default: "param" + /// + public virtual string ParameterText { get; set; } = "p"; + + /// + /// Set to true to add a counter to the end of the parameter string. + /// Default: true + /// + public virtual bool IncrementParameter { get; set; } = true; + + /// + /// Default base implementation for an SQL client. + /// + /// + protected BaseSqlClient(string connectionString) => _connectionString = connectionString; + + /// + /// Create a connection using the client specific connection. + /// + /// + protected abstract DbConnection CreateConnection(); + + /// + /// Create a command using the client specific connection. + /// + /// + /// + /// + protected abstract DbCommand CreateCommand(string query, DbConnection connection); + + /// + /// + /// + public async Task 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; + } + + /// + /// + /// + public async Task 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; + } + + /// + /// + /// + public async Task 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)); + } + + /// + /// Add parameters into the query to prevent SQL injection. + /// + /// Command to add the parameters to + /// Parameters to add + /// + private DbCommand AddCommandParameters(DbCommand command, Dictionary 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; + } + /// /// Returns data as a . /// /// Reader to return data from. /// of data. - 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 /// /// Reader to return data from. /// of data. - protected static DataTable ReadAsDataTable(IDataReader reader) + private DataTable ReadAsDataTable(IDataReader reader) { var data = new DataTable(); var schemaTable =reader.GetSchemaTable(); diff --git a/src/modules/Elsa.Sql/Client/ISqlClient.cs b/src/modules/Elsa.Sql/Client/ISqlClient.cs index 2312097b8..2cc3eb5b9 100644 --- a/src/modules/Elsa.Sql/Client/ISqlClient.cs +++ b/src/modules/Elsa.Sql/Client/ISqlClient.cs @@ -1,27 +1,43 @@ using System.Data; +using Elsa.Sql.Models; namespace Elsa.Sql.Client; public interface ISqlClient { + /// + /// The marker used when injecting parameters into a query. + /// + public string ParameterMarker { get; set; } + + /// + /// The text following the ParameterMarker when injecting parameters into a query. + /// + public string ParameterText { get; set; } + + /// + /// Set to true to add a counter to the end of the parameter string. + /// + public bool IncrementParameter { get; set; } + /// /// Asynchronously executes a Transact-SQL statement against the connection and returns the number of rows affected. /// - /// The command to execute + /// The evaluated query to execute. /// The number of rows affected. - public Task ExecuteCommandAsync(string sqlCommand); + public Task ExecuteCommandAsync(EvaluatedQuery evaluatedQuery); /// /// 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. /// - /// The query to execute + /// The evaluated query to execute. /// 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. - public Task ExecuteScalarAsync(string sqlQuery); + public Task ExecuteScalarAsync(EvaluatedQuery evaluatedQuery); /// /// Asynchronously executes the query, and returns a dataset of data returned by the query. /// - /// Query to execute + /// The evaluated query to execute. /// DataSet of the queried data - public Task ExecuteQueryAsync(string sqlQuery); + public Task ExecuteQueryAsync(EvaluatedQuery evaluatedQuery); } \ No newline at end of file diff --git a/src/modules/Elsa.Sql/Contracts/ISqlEvaluator.cs b/src/modules/Elsa.Sql/Contracts/ISqlEvaluator.cs index 512d9fbd9..0817f409b 100644 --- a/src/modules/Elsa.Sql/Contracts/ISqlEvaluator.cs +++ b/src/modules/Elsa.Sql/Contracts/ISqlEvaluator.cs @@ -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 /// The context in which the expression is evaluated. /// A set of options. /// An optional cancellation token. - /// The result of the evaluation. - Task EvaluateAsync( + /// The result. + Task EvaluateAsync( string expression, ExpressionExecutionContext context, ExpressionEvaluatorOptions options, diff --git a/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs b/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs new file mode 100644 index 000000000..1570fab0e --- /dev/null +++ b/src/modules/Elsa.Sql/Models/EvaluatedQuery.cs @@ -0,0 +1,35 @@ +namespace Elsa.Sql.Models +{ + /// + /// Represents a safely formatted SQL expression result. + /// + public class EvaluatedQuery + { + /// + /// Query with parameterized values + /// + public string Query { get; set; } + + /// + /// Parameters to inject into the query at execution + /// + public Dictionary Parameters { get; set; } = new Dictionary(); + + /// + /// An evaluated query response. + /// + /// The evaluated query + public EvaluatedQuery(string query) => Query = query; + + /// + /// An evaluated query response. + /// + /// The evaluated query + /// Parameters to pass into the parameterized query + public EvaluatedQuery(string query, Dictionary parameters) + { + Query = query; + Parameters = parameters; + } + } +} \ No newline at end of file diff --git a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs index 88bde44d4..10ee5a0df 100644 --- a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs +++ b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs @@ -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; /// public class SqlEvaluator() : ISqlEvaluator { + private WorkflowExecutionContext executionContext; + private ActivityExecutionContext activityContext; + private ExpressionExecutionContext expressionContext; + /// - public async Task EvaluateAsync( + public async Task 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(); + var client = factory.CreateClient(activityContext.ActivityState["Client"].ToString(), activityContext.ActivityState["ConnectionString"].ToString()); var sb = new StringBuilder(); 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); - 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}}}}}.") }; } } \ No newline at end of file