using System.Data.Common; using System.Reflection; using Cleanuparr.Persistence; using Cleanuparr.Persistence.Providers; using Cleanuparr.Shared.Configuration; using Cleanuparr.Shared.Helpers; using Microsoft.EntityFrameworkCore; using Microsoft.EntityFrameworkCore.Infrastructure; using Microsoft.EntityFrameworkCore.Metadata; using Npgsql; namespace Cleanuparr.Infrastructure.Features.DatabaseMigration; public sealed record MigrationResult(bool Success, string? Error, IReadOnlyDictionary TableCounts); public sealed class SqliteToPostgresMigrator { private readonly ModelDataCopier _copier = new(); public async Task RunAsync(bool force, IProgress? progress, CancellationToken cancellationToken) { Dictionary counts = new(); try { string connectionString = BuildPostgresConnectionString(); if (!force) { IReadOnlyList provisioned = await GetProvisionedSchemasAsync(connectionString, cancellationToken); if (provisioned.Count > 0) { return new MigrationResult( false, $"Target PostgreSQL already contains applied migrations in schema(s): {string.Join(", ", provisioned)}. Re-run with --force to wipe and re-import.", new Dictionary()); } } await MigrateSourceSchemaAsync(progress, cancellationToken); try { await MigrateTargetSchemaAsync(connectionString, cancellationToken); await using NpgsqlConnection connection = new(connectionString); await connection.OpenAsync(cancellationToken); await using DbTransaction transaction = await connection.BeginTransactionAsync(cancellationToken); try { await CopyContextAsync(DbContextKind.Data, connection, transaction, counts, progress, cancellationToken); await CopyContextAsync(DbContextKind.Events, connection, transaction, counts, progress, cancellationToken); await CopyContextAsync(DbContextKind.Users, connection, transaction, counts, progress, cancellationToken); await transaction.CommitAsync(cancellationToken); } catch { await transaction.RollbackAsync(cancellationToken); throw; } } catch { await DropCleanuparrSchemasAsync(connectionString, progress, cancellationToken); throw; } return new MigrationResult(true, null, counts); } catch (Exception exception) { return new MigrationResult(false, exception.ToString(), counts); } } private static string BuildPostgresConnectionString() => PostgresDatabaseProvider.BuildConnectionString( DatabaseConfigProvider.GetRequired(ConfigurationKeys.PostgresHost), DatabaseConfigProvider.GetOptional(ConfigurationKeys.PostgresPort), DatabaseConfigProvider.GetRequired(ConfigurationKeys.PostgresUser), DatabaseConfigProvider.GetRequired(ConfigurationKeys.PostgresPassword), DatabaseConfigProvider.GetRequired(ConfigurationKeys.PostgresDatabase), DatabaseConfigProvider.GetOptional(ConfigurationKeys.PostgresExtraParams)); private async Task MigrateSourceSchemaAsync(IProgress? progress, CancellationToken cancellationToken) { progress?.Report("Applying pending SQLite migrations to the source databases..."); await using (DataContext data = BuildSourceContext(DbContextKind.Data)) { await data.Database.MigrateAsync(cancellationToken); } await using (EventsContext events = BuildSourceContext(DbContextKind.Events)) { await events.Database.MigrateAsync(cancellationToken); } await using (UsersContext users = BuildSourceContext(DbContextKind.Users)) { await users.Database.MigrateAsync(cancellationToken); } } private async Task MigrateTargetSchemaAsync(string connectionString, CancellationToken cancellationToken) { await using DataContext data = BuildTargetContext(connectionString, DbContextKind.Data, keepKeys: false); await data.Database.MigrateAsync(cancellationToken); await using EventsContext events = BuildTargetContext(connectionString, DbContextKind.Events, keepKeys: false); await events.Database.MigrateAsync(cancellationToken); await using UsersContext users = BuildTargetContext(connectionString, DbContextKind.Users, keepKeys: false); await users.Database.MigrateAsync(cancellationToken); } internal static async Task DropCleanuparrSchemasAsync(string connectionString, IProgress? progress, CancellationToken cancellationToken) { try { PostgresDatabaseProvider provider = new(); string[] schemas = { provider.GetSchema(DbContextKind.Data)!, provider.GetSchema(DbContextKind.Events)!, provider.GetSchema(DbContextKind.Users)!, }; await using NpgsqlConnection connection = new(connectionString); await connection.OpenAsync(cancellationToken); foreach (string schema in schemas) { await using NpgsqlCommand command = new($"DROP SCHEMA IF EXISTS \"{schema}\" CASCADE;", connection); await command.ExecuteNonQueryAsync(cancellationToken); } } catch (Exception exception) { progress?.Report($"Warning: failed to clean up target schemas after error: {exception.Message}"); } } private static async Task> GetProvisionedSchemasAsync(string connectionString, CancellationToken cancellationToken) { PostgresDatabaseProvider provider = new(); List provisioned = new(); if (await HasAppliedMigrationsAsync(connectionString, DbContextKind.Data, cancellationToken)) { provisioned.Add(provider.GetSchema(DbContextKind.Data)!); } if (await HasAppliedMigrationsAsync(connectionString, DbContextKind.Events, cancellationToken)) { provisioned.Add(provider.GetSchema(DbContextKind.Events)!); } if (await HasAppliedMigrationsAsync(connectionString, DbContextKind.Users, cancellationToken)) { provisioned.Add(provider.GetSchema(DbContextKind.Users)!); } return provisioned; } private static async Task HasAppliedMigrationsAsync(string connectionString, DbContextKind kind, CancellationToken cancellationToken) where TContext : DbContext { await using TContext context = BuildTargetContext(connectionString, kind, keepKeys: false); IEnumerable applied = await context.Database.GetAppliedMigrationsAsync(cancellationToken); return applied.Any(); } private async Task CopyContextAsync( DbContextKind kind, NpgsqlConnection connection, DbTransaction transaction, Dictionary counts, IProgress? progress, CancellationToken cancellationToken) where TContext : DbContext { await using TContext source = BuildSourceContext(kind); await using TContext target = BuildTargetContextOnConnection(connection, kind); await target.Database.UseTransactionAsync(transaction, cancellationToken); await TruncateAsync(target, kind, cancellationToken); await _copier.CopyAsync(source, target, progress, cancellationToken); await VerifyAndRecordCountsAsync(source, target, counts, cancellationToken); } public static TContext Instantiate(DbContextOptionsBuilder builder, IDatabaseProvider provider) where TContext : DbContext => (TContext)Activator.CreateInstance(typeof(TContext), builder.Options, provider)!; private static TContext BuildSourceContext(DbContextKind kind) where TContext : DbContext { SqliteDatabaseProvider provider = new(); DbContextOptionsBuilder builder = new(); provider.ConfigureContext(builder, kind); return Instantiate(builder, provider); } private static TContext BuildTargetContext(string connectionString, DbContextKind kind, bool keepKeys) where TContext : DbContext { PostgresDatabaseProvider provider = new(); string schema = provider.GetSchema(kind)!; DbContextOptionsBuilder builder = new(); builder.UseNpgsql(connectionString, options => options .MigrationsAssembly(PostgresDatabaseProvider.MigrationsAssembly) .MigrationsHistoryTable("__ef_migrations_history", schema)) .UseLowerCaseNamingConvention() .UseSnakeCaseNamingConvention(); if (keepKeys) { builder.ReplaceService(); } return Instantiate(builder, provider); } private static TContext BuildTargetContextOnConnection(NpgsqlConnection connection, DbContextKind kind) where TContext : DbContext { PostgresDatabaseProvider provider = new(); string schema = provider.GetSchema(kind)!; DbContextOptionsBuilder builder = new(); builder.UseNpgsql(connection, options => options .MigrationsAssembly(PostgresDatabaseProvider.MigrationsAssembly) .MigrationsHistoryTable("__ef_migrations_history", schema)) .UseLowerCaseNamingConvention() .UseSnakeCaseNamingConvention() .ReplaceService(); return Instantiate(builder, provider); } private static async Task TruncateAsync(DbContext target, DbContextKind kind, CancellationToken cancellationToken) { string schema = new PostgresDatabaseProvider().GetSchema(kind)!; List tables = target.Model.GetEntityTypes() .Where(entityType => !entityType.IsOwned()) .Select(entityType => entityType.GetTableName()) .Where(name => name is not null) .Distinct() .Select(name => $"\"{schema}\".\"{name}\"") .ToList(); if (tables.Count == 0) { return; } string sql = $"TRUNCATE TABLE {string.Join(", ", tables)} CASCADE;"; await target.Database.ExecuteSqlRawAsync(sql, cancellationToken); } internal static async Task VerifyAndRecordCountsAsync( DbContext source, DbContext target, Dictionary counts, CancellationToken cancellationToken) { foreach (IEntityType entityType in ModelDataCopier.OrderByDependencies(target.Model.GetEntityTypes())) { string? table = entityType.GetTableName(); if (table is null) { continue; } int sourceCount = await CountAsync(source, entityType.ClrType, cancellationToken); int targetCount = await CountAsync(target, entityType.ClrType, cancellationToken); if (sourceCount != targetCount) { throw new InvalidOperationException( $"Row count mismatch for table '{table}': source has {sourceCount} rows, target has {targetCount} rows."); } counts[table] = targetCount; } } private static async Task CountAsync(DbContext context, Type entityClrType, CancellationToken cancellationToken) { MethodInfo setMethod = typeof(DbContext) .GetMethods() .Single(method => method.Name == nameof(DbContext.Set) && method.IsGenericMethod && method.GetParameters().Length == 0) .MakeGenericMethod(entityClrType); object dbSet = setMethod.Invoke(context, null)!; MethodInfo countMethod = typeof(EntityFrameworkQueryableExtensions) .GetMethods() .Single(method => method.Name == nameof(EntityFrameworkQueryableExtensions.CountAsync) && method.GetParameters().Length == 2) .MakeGenericMethod(entityClrType); object task = countMethod.Invoke(null, new object[] { dbSet, cancellationToken })!; return await (Task)task; } }