Simplify SqlEvaluator implementation and tidy code.
This commit is contained in:
parent
d3dc6f3597
commit
24da6ab288
|
|
@ -10,12 +10,6 @@ namespace Elsa.Sql.MySql;
|
|||
/// <param name="connectionString"></param>
|
||||
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);
|
||||
|
|
|
|||
|
|
@ -10,8 +10,6 @@ namespace Elsa.Sql.SqlServer;
|
|||
/// <param name="connectionString"></param>
|
||||
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);
|
||||
|
|
|
|||
|
|
@ -10,8 +10,6 @@ namespace Elsa.Sql.Sqlite;
|
|||
/// <param name="connectionString"></param>
|
||||
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);
|
||||
|
|
|
|||
|
|
@ -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; } = "@";
|
||||
|
||||
/// <summary>
|
||||
/// The text following the <c>ParameterMarker</c>when injecting parameters into a query
|
||||
/// Default: <see cref="string.Empty"/>
|
||||
/// The text following the <c>ParameterMarker</c> when injecting parameters into a query.
|
||||
/// Default: "param"
|
||||
/// </summary>
|
||||
public virtual string ParameterText { get; set; } = string.Empty;
|
||||
public virtual string ParameterText { get; set; } = "p";
|
||||
|
||||
/// <summary>
|
||||
/// 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
|
||||
/// </summary>
|
||||
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
|
|||
}
|
||||
|
||||
/// <summary>
|
||||
/// Replace the evaluated parameters with client specific parameters.
|
||||
/// </summary>
|
||||
/// <param name="evaluatedQuery">Query to replace parameters for.</param>
|
||||
/// <returns></returns>
|
||||
private string ReplaceQueryParameters(EvaluatedQuery evaluatedQuery)
|
||||
{
|
||||
var count = 1;
|
||||
var clientUpdatedParams = new Dictionary<string, object>();
|
||||
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();
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Inject parameters into the query to prevent SQL injection.
|
||||
/// 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>
|
||||
|
|
|
|||
|
|
@ -5,6 +5,21 @@ 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>
|
||||
|
|
|
|||
|
|
@ -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;
|
|||
/// </remarks>
|
||||
public class SqlEvaluator() : ISqlEvaluator
|
||||
{
|
||||
private WorkflowExecutionContext executionContext;
|
||||
private ActivityExecutionContext activityContext;
|
||||
private ExpressionExecutionContext expressionContext;
|
||||
|
||||
/// <inheritdoc />
|
||||
public async Task<EvaluatedQuery> 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<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?>();
|
||||
|
|
@ -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
|
||||
};
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue