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
};
}