From 24da6ab288dc6cd5bdfe07bf6c98d08de940788e Mon Sep 17 00:00:00 2001 From: Matt Date: Sun, 16 Mar 2025 00:44:35 +0000 Subject: [PATCH] Simplify SqlEvaluator implementation and tidy code. --- src/modules/Elsa.Sql.MySql/MySqlClient.cs | 6 --- .../Elsa.Sql.SqlServer/SqlServerClient.cs | 2 - src/modules/Elsa.Sql.Sqlite/SqliteClient.cs | 2 - src/modules/Elsa.Sql/Client/BaseSqlClient.cs | 43 ++++--------------- src/modules/Elsa.Sql/Client/ISqlClient.cs | 15 +++++++ src/modules/Elsa.Sql/Services/SqlEvaluator.cs | 40 +++++++++++------ 6 files changed, 51 insertions(+), 57 deletions(-) diff --git a/src/modules/Elsa.Sql.MySql/MySqlClient.cs b/src/modules/Elsa.Sql.MySql/MySqlClient.cs index 425f701de..60669f7d1 100644 --- a/src/modules/Elsa.Sql.MySql/MySqlClient.cs +++ b/src/modules/Elsa.Sql.MySql/MySqlClient.cs @@ -10,12 +10,6 @@ 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 96efacb63..e8780c359 100644 --- a/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs +++ b/src/modules/Elsa.Sql.SqlServer/SqlServerClient.cs @@ -10,8 +10,6 @@ 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 e6ad5afcb..ec43f5245 100644 --- a/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs +++ b/src/modules/Elsa.Sql.Sqlite/SqliteClient.cs @@ -10,8 +10,6 @@ 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 f8714bd26..314f103ba 100644 --- a/src/modules/Elsa.Sql/Client/BaseSqlClient.cs +++ b/src/modules/Elsa.Sql/Client/BaseSqlClient.cs @@ -1,6 +1,5 @@ using System.Data; using System.Data.Common; -using System.Text; using Elsa.Sql.Models; namespace Elsa.Sql.Client; @@ -19,14 +18,14 @@ public abstract class BaseSqlClient : ISqlClient public virtual string ParameterMarker { get; set; } = "@"; /// - /// The text following the ParameterMarkerwhen injecting parameters into a query - /// Default: + /// The text following the ParameterMarker when injecting parameters into a query. + /// Default: "param" /// - public virtual string ParameterText { get; set; } = string.Empty; + public virtual string ParameterText { get; set; } = "p"; /// - /// Set to true to add a counter to the end of the parameter string - /// Default: false + /// Set to true to add a counter to the end of the parameter string. + /// Default: true /// public virtual bool IncrementParameter { get; set; } = true; @@ -57,8 +56,7 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var query = ReplaceQueryParameters(evaluatedQuery); - var command = CreateCommand(query, connection); + var command = CreateCommand(evaluatedQuery.Query, connection); AddCommandParameters(command, evaluatedQuery.Parameters); var result = await command.ExecuteNonQueryAsync(); @@ -72,8 +70,7 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var query = ReplaceQueryParameters(evaluatedQuery); - var command = CreateCommand(query, connection); + var command = CreateCommand(evaluatedQuery.Query, connection); AddCommandParameters(command, evaluatedQuery.Parameters); var result = await command.ExecuteScalarAsync(); @@ -87,8 +84,7 @@ public abstract class BaseSqlClient : ISqlClient { using var connection = CreateConnection(); connection.Open(); - var query = ReplaceQueryParameters(evaluatedQuery); - var command = CreateCommand(query, connection); + var command = CreateCommand(evaluatedQuery.Query, connection); AddCommandParameters(command, evaluatedQuery.Parameters); using var reader = await command.ExecuteReaderAsync(); @@ -96,28 +92,7 @@ public abstract class BaseSqlClient : ISqlClient } /// - /// 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. + /// Add parameters into the query to prevent SQL injection. /// /// Command to add the parameters to /// Parameters to add diff --git a/src/modules/Elsa.Sql/Client/ISqlClient.cs b/src/modules/Elsa.Sql/Client/ISqlClient.cs index e5b743b3b..2cc3eb5b9 100644 --- a/src/modules/Elsa.Sql/Client/ISqlClient.cs +++ b/src/modules/Elsa.Sql/Client/ISqlClient.cs @@ -5,6 +5,21 @@ 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. /// diff --git a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs index 7bfdc42ea..c0947bc2e 100644 --- a/src/modules/Elsa.Sql/Services/SqlEvaluator.cs +++ b/src/modules/Elsa.Sql/Services/SqlEvaluator.cs @@ -3,6 +3,7 @@ using Elsa.Expressions.Models; using Elsa.Extensions; using Elsa.Sql.Contracts; using Elsa.Sql.Models; +using Elsa.Workflows; namespace Elsa.Sql.Services; @@ -14,6 +15,10 @@ namespace Elsa.Sql.Services; /// public class SqlEvaluator() : ISqlEvaluator { + private WorkflowExecutionContext executionContext; + private ActivityExecutionContext activityContext; + private ExpressionExecutionContext expressionContext; + /// public async Task EvaluateAsync( string expression, @@ -23,6 +28,14 @@ public class SqlEvaluator() : ISqlEvaluator { 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(); @@ -49,11 +62,12 @@ public class SqlEvaluator() : ISqlEvaluator 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}}}}}'."); + object? value = ResolveValue(key); + if (value is null) throw new NullReferenceException($"No value found for {{{{{key}}}}}."); // Replace with parameterized name - string paramName = $"{{{{p{paramIndex++}}}}}"; + var counterValue = client.IncrementParameter ? paramIndex++.ToString() : string.Empty; + string paramName = $"{client.ParameterMarker}{client.ParameterText}{counterValue}"; parameters[paramName] = value; sb.Append(paramName); @@ -63,19 +77,19 @@ public class SqlEvaluator() : ISqlEvaluator 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, + "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, _ => null }; }