From c08e7fc9a92d2ae1590fe64a5a29df98a2d990f1 Mon Sep 17 00:00:00 2001 From: Sipke Schoorstra Date: Mon, 17 Feb 2025 16:56:56 +0100 Subject: [PATCH] Add BulkUpsertExtensions for enhanced bulk upsert operations This commit introduces a dedicated `BulkUpsertExtensions` class to streamline bulk upsert operations in Entity Framework Core. It supports multiple database providers and replaces the previous implementation in `QueryableExtensions` for better modularity and maintainability. --- .../Extensions/BulkUpsertExtensions.cs | 384 ++++++++++++++++++ .../Extensions/QueryableExtensions.cs | 25 -- 2 files changed, 384 insertions(+), 25 deletions(-) create mode 100644 src/modules/Elsa.EntityFrameworkCore.Common/Extensions/BulkUpsertExtensions.cs diff --git a/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/BulkUpsertExtensions.cs b/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/BulkUpsertExtensions.cs new file mode 100644 index 000000000..212edb898 --- /dev/null +++ b/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/BulkUpsertExtensions.cs @@ -0,0 +1,384 @@ + + +using System.Text; +using Microsoft.EntityFrameworkCore; +using Microsoft.EntityFrameworkCore.Infrastructure; +using Microsoft.EntityFrameworkCore.Metadata; +using System.Linq.Expressions; + +// ReSharper disable once CheckNamespace +namespace Elsa.EntityFrameworkCore.Extensions; + +/// +/// Provides extension methods to perform bulk upsert operations for entities +/// in an Entity Framework Core context, supporting multiple database providers. +/// +public static class BulkUpsertExtensions +{ + /// + /// Performs a bulk upsert operation on a list of entities in the specified database context using a key selector. + /// + /// The type of the database context. + /// The type of the entity being upserted. + /// The database context where the bulk upsert operation will be executed. + /// The list of entities to be upserted. + /// An expression used to determine the key for upsert operations. + /// A token to observe while waiting for the operation to complete. + public static async Task BulkUpsertAsync( + this TDbContext dbContext, + IList entities, + Expression> keySelector, + CancellationToken cancellationToken = default) + where TDbContext : DbContext + where TEntity : class, new() + { + await BulkUpsertAsync(dbContext, entities, keySelector, 50, cancellationToken); + } + + /// + /// Performs a bulk upsert operation on a list of entities in the specified database context using a key selector and optional batch size. + /// + /// The type of the database context. + /// The type of the entity being upserted. + /// The database context where the bulk upsert operation will be executed. + /// The list of entities to be upserted. + /// An expression used to determine the key for upsert operations. + /// The size of each batch for processing the upsert operation. Defaults to 50. + /// A token to observe while waiting for the operation to complete. + /// Thrown if the database provider for the context is not supported. + public static async Task BulkUpsertAsync( + this TDbContext dbContext, + IList entities, + Expression> keySelector, + int batchSize = 50, + CancellationToken cancellationToken = default) + where TDbContext : DbContext + where TEntity : class, new() + { + if (entities.Count == 0) + return; + + // Identify the current provider (e.g., "Microsoft.EntityFrameworkCore.SqlServer") + var providerName = dbContext.Database.ProviderName?.ToLowerInvariant() ?? string.Empty; + + // Determine the method for generating SQL based on the provider + Func, Expression>, (string, object[])> generateSql = providerName switch + { + var pn when pn.Contains("sqlserver") => GenerateSqlServerUpsert, + var pn when pn.Contains("sqlite") => GenerateSqliteUpsert, + var pn when pn.Contains("postgres") => GeneratePostgresUpsert, + var pn when pn.Contains("mysql") => GenerateMySqlUpsert, + var pn when pn.Contains("oracle") => GenerateOracleUpsert, + _ => throw new NotSupportedException($"Provider '{providerName}' is not supported.") + }; + + // Loop through batched entities + foreach (var batch in entities.Chunk(batchSize)) + { + // Generate SQL and parameters + var (sql, parameters) = generateSql(dbContext, batch, keySelector); + + await dbContext.Database.ExecuteSqlRawAsync(sql, parameters, cancellationToken); + } + } + + private static (string, object[]) GenerateSqlServerUpsert( + DbContext dbContext, + IList entities, + Expression> keySelector) + where TEntity : class + { + var entityType = dbContext.Model.FindEntityType(typeof(TEntity))!; + var tableName = $"[{entityType.GetSchema()}].[{entityType.GetTableName()}]"; + var storeObject = StoreObjectIdentifier.Table(entityType.GetTableName()!, entityType.GetSchema()); + + // Include shadow properties + var props = entityType.GetProperties().ToList(); + + var keyProp = entityType.FindProperty(keySelector.GetMemberAccess().Name)!; + var keyColumnName = $"[{keyProp.GetColumnName(storeObject)}]"; + var columnNames = props + .Select(p => $"[{p.GetColumnName(storeObject)}]") + .ToList(); + + var mergeSql = new StringBuilder(); + mergeSql.AppendLine($"MERGE {tableName} AS Target"); + mergeSql.AppendLine("USING (VALUES"); + + var parameters = new List(); + for (var i = 0; i < entities.Count; i++) + { + var entity = entities[i]; + var values = new List(); + + for (var j = 0; j < props.Count; j++) + { + var property = props[j]; + var paramName = $"@p{i}_{j}"; + + // If it's a shadow property, retrieve value via Entry(..).Property(..) + object? value = property.IsShadowProperty() + ? dbContext.Entry(entity).Property(property.Name).CurrentValue + : property.PropertyInfo?.GetValue(entity); + + values.Add(paramName); + parameters.Add(value); + } + + var line = $"({string.Join(", ", values)}){(i < entities.Count - 1 ? "," : string.Empty)}"; + mergeSql.AppendLine(line); + } + + mergeSql.AppendLine($") AS Source ({string.Join(", ", columnNames)})"); + mergeSql.AppendLine($"ON Target.{keyColumnName} = Source.{keyColumnName}"); + mergeSql.AppendLine("WHEN MATCHED THEN"); + mergeSql.AppendLine($" UPDATE SET {string.Join(", ", columnNames.Where(c => c != keyColumnName).Select(c => $"Target.{c} = Source.{c}"))}"); + mergeSql.AppendLine("WHEN NOT MATCHED THEN"); + mergeSql.AppendLine($" INSERT ({string.Join(", ", columnNames)})"); + mergeSql.AppendLine($" VALUES ({string.Join(", ", columnNames.Select(c => $"Source.{c}"))});"); + + return (mergeSql.ToString(), parameters.ToArray()); + } + + private static (string, object[]) GenerateSqliteUpsert( + DbContext dbContext, + IList entities, + Expression> keySelector) + where TEntity : class + { + var entityType = dbContext.Model.FindEntityType(typeof(TEntity))!; + var tableName = entityType.GetTableName(); + var storeObject = StoreObjectIdentifier.Table(tableName!, entityType.GetSchema()); + + var props = entityType.GetProperties().ToList(); + + var keyProp = entityType.FindProperty(keySelector.GetMemberAccess().Name)!; + var keyColumnName = keyProp.GetColumnName(storeObject); + var columnNames = props + .Select(p => p.GetColumnName(storeObject)!) + .ToList(); + + var sb = new StringBuilder(); + var parameters = new List(); + + sb.Append($"INSERT INTO \"{tableName}\" ({string.Join(", ", columnNames.Select(c => $"\"{c}\""))}) VALUES "); + + for (var i = 0; i < entities.Count; i++) + { + var entity = entities[i]; + var placeholders = new List(); + + for (var j = 0; j < props.Count; j++) + { + var property = props[j]; + var paramName = $"@p{i}_{j}"; + + object? value = property.IsShadowProperty() + ? dbContext.Entry(entity).Property(property.Name).CurrentValue + : property.PropertyInfo?.GetValue(entity); + + placeholders.Add(paramName); + parameters.Add(value); + } + + sb.Append($"({string.Join(", ", placeholders)})"); + if (i < entities.Count - 1) + sb.Append(", "); + } + + sb.AppendLine(); + sb.AppendLine($"ON CONFLICT(\"{keyColumnName}\") DO UPDATE SET"); + + var updateAssignments = columnNames + .Where(c => c != keyColumnName) + .Select(c => $"\"{c}\"=excluded.\"{c}\""); + + sb.AppendLine(string.Join(", ", updateAssignments) + ";"); + + return (sb.ToString(), parameters.ToArray()); + } + + private static (string, object[]) GeneratePostgresUpsert( + DbContext dbContext, + IList entities, + Expression> keySelector) + where TEntity : class + { + var entityType = dbContext.Model.FindEntityType(typeof(TEntity))!; + var tableName = entityType.GetTableName(); + var storeObject = StoreObjectIdentifier.Table(tableName!, entityType.GetSchema()); + + var props = entityType.GetProperties().ToList(); + + var keyProp = entityType.FindProperty(keySelector.GetMemberAccess().Name)!; + var keyColumnName = keyProp.GetColumnName(storeObject); + var columnNames = props + .Select(p => p.GetColumnName(storeObject)!) + .ToList(); + + var sb = new StringBuilder(); + var parameters = new List(); + var parameterCount = 0; + + sb.Append($"INSERT INTO \"{storeObject.Schema}\".\"{storeObject.Name}\" ({string.Join(", ", columnNames.Select(c => $"\"{c}\""))}) VALUES "); + + for (var i = 0; i < entities.Count; i++) + { + var entity = entities[i]; + var placeholders = new List(); + + foreach (var property in props) + { + var paramName = $"{{{parameterCount++}}}"; + + object? value = property.IsShadowProperty() + ? dbContext.Entry(entity).Property(property.Name).CurrentValue + : property.PropertyInfo?.GetValue(entity); + + placeholders.Add(paramName); + parameters.Add(value); + } + + sb.Append($"({string.Join(", ", placeholders)})"); + if (i < entities.Count - 1) + sb.Append(", "); + } + + sb.AppendLine(); + sb.AppendLine($"ON CONFLICT (\"{keyColumnName}\") DO UPDATE SET"); + + var updateAssignments = columnNames + .Where(c => c != keyColumnName) + .Select(c => $"\"{c}\" = EXCLUDED.\"{c}\""); + + sb.AppendLine(string.Join(", ", updateAssignments) + ";"); + + return (sb.ToString(), parameters.ToArray()); + } + + private static (string, object[]) GenerateMySqlUpsert( + DbContext dbContext, + IList entities, + Expression> keySelector) + where TEntity : class + { + var entityType = dbContext.Model.FindEntityType(typeof(TEntity))!; + var tableName = entityType.GetTableName(); + var storeObject = StoreObjectIdentifier.Table(tableName!, entityType.GetSchema()); + + var props = entityType.GetProperties().ToList(); + + var keyProp = entityType.FindProperty(keySelector.GetMemberAccess().Name)!; + var keyColumnName = keyProp.GetColumnName(storeObject); + var columnNames = props + .Select(p => p.GetColumnName(storeObject)!) + .ToList(); + + var sb = new StringBuilder(); + var parameters = new List(); + + sb.Append($"INSERT INTO `{tableName}` ({string.Join(", ", columnNames.Select(c => $"`{c}`"))}) VALUES "); + + for (var i = 0; i < entities.Count; i++) + { + var entity = entities[i]; + var placeholders = new List(); + + for (var j = 0; j < props.Count; j++) + { + var property = props[j]; + var paramName = $"@p{i}_{j}"; + + object? value = property.IsShadowProperty() + ? dbContext.Entry(entity).Property(property.Name).CurrentValue + : property.PropertyInfo?.GetValue(entity); + + placeholders.Add(paramName); + parameters.Add(value); + } + + sb.Append($"({string.Join(", ", placeholders)})"); + if (i < entities.Count - 1) + sb.Append(", "); + } + + sb.AppendLine(); + sb.AppendLine("ON DUPLICATE KEY UPDATE"); + + var updateAssignments = columnNames + .Where(c => c != keyColumnName) + .Select(c => $"`{c}` = VALUES(`{c}`)"); + + sb.AppendLine(string.Join(", ", updateAssignments) + ";"); + + return (sb.ToString(), parameters.ToArray()); + } + + private static (string, object[]) GenerateOracleUpsert( + DbContext dbContext, + IList entities, + Expression> keySelector) + where TEntity : class + { + var entityType = dbContext.Model.FindEntityType(typeof(TEntity))!; + var schema = entityType.GetSchema(); + var tableName = entityType.GetTableName(); + var storeObject = StoreObjectIdentifier.Table(tableName!, schema); + var fullName = !string.IsNullOrEmpty(schema) ? $"{schema}.{tableName}" : tableName; + + var props = entityType.GetProperties().ToList(); + + var keyProp = entityType.FindProperty(keySelector.GetMemberAccess().Name)!; + var keyColumnName = keyProp.GetColumnName(storeObject); + + var columnNames = props + .Select(p => p.GetColumnName(storeObject)!) + .ToList(); + + var sb = new StringBuilder(); + var parameters = new List(); + + sb.AppendLine($"MERGE INTO {fullName} Target"); + sb.AppendLine("USING (SELECT"); + + for (var i = 0; i < entities.Count; i++) + { + var entity = entities[i]; + var lineParts = new List(); + + for (var j = 0; j < props.Count; j++) + { + var property = props[j]; + var paramName = $":p{i}_{j}"; + + object? value = property.IsShadowProperty() + ? dbContext.Entry(entity).Property(property.Name).CurrentValue + : property.PropertyInfo?.GetValue(entity); + + parameters.Add(value); + + // Oracle aliases must match the column name + var alias = property.GetColumnName(storeObject); + lineParts.Add($"{paramName} AS {alias}"); + } + + // Comma if not last + var suffix = (i < entities.Count - 1) ? " FROM DUAL UNION ALL SELECT" : " FROM DUAL"; + sb.AppendLine(string.Join(", ", lineParts) + suffix); + } + + sb.AppendLine($") Source ON (Target.{keyColumnName} = Source.{keyColumnName})"); + sb.AppendLine("WHEN MATCHED THEN UPDATE SET"); + + var updateSetClauses = columnNames + .Where(c => c != keyColumnName) + .Select(c => $"Target.{c} = Source.{c}"); + + sb.AppendLine(string.Join(", ", updateSetClauses)); + sb.AppendLine("WHEN NOT MATCHED THEN"); + sb.AppendLine($"INSERT ({string.Join(", ", columnNames)})"); + sb.AppendLine($"VALUES ({string.Join(", ", columnNames.Select(c => $"Source.{c}"))});"); + + return (sb.ToString(), parameters.ToArray()); + } +} \ No newline at end of file diff --git a/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/QueryableExtensions.cs b/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/QueryableExtensions.cs index 191378a4b..11171ace7 100644 --- a/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/QueryableExtensions.cs +++ b/src/modules/Elsa.EntityFrameworkCore.Common/Extensions/QueryableExtensions.cs @@ -12,31 +12,6 @@ namespace Elsa.EntityFrameworkCore.Extensions; [PublicAPI] public static class QueryableExtensions { - /// - /// Inserts or updates a list of entities in bulk. - /// - public static async Task BulkUpsertAsync(this TDbContext dbContext, IList entities, Expression> keySelector, CancellationToken cancellationToken = default) where TDbContext : DbContext where TEntity : class, new() - { - var set = dbContext.Set(); - var compiledKeySelector = keySelector.Compile(); - var containsLambda = entities.Any() ? keySelector.BuildContainsExpression(entities) : default; - var existingEntitiesQuery = set.AsNoTracking(); - - if (containsLambda != null) - existingEntitiesQuery = existingEntitiesQuery.Where(containsLambda); - - var existingEntities = await existingEntitiesQuery.ToListAsync(cancellationToken); - var entitiesToUpdate = entities.IntersectBy(existingEntities.Select(compiledKeySelector), compiledKeySelector).ToList(); - var entitiesToInsert = entities.Except(entitiesToUpdate).ToList(); - - if (entitiesToUpdate.Any()) - set.UpdateRange(entitiesToUpdate); - if (entitiesToInsert.Any()) - await set.AddRangeAsync(entitiesToInsert, cancellationToken); - - await dbContext.SaveChangesAsync(cancellationToken); - } - /// /// Inserts a list of entities in bulk. ///