diff --git a/src/Migrator.Tests/ConstraintCatalogParsingTests.cs b/src/Migrator.Tests/ConstraintCatalogParsingTests.cs new file mode 100644 index 00000000..132a679c --- /dev/null +++ b/src/Migrator.Tests/ConstraintCatalogParsingTests.cs @@ -0,0 +1,94 @@ +using System; +using System.Data; +using System.Linq; +using DotNetProjects.Migrator.Framework; +using DotNetProjects.Migrator.Providers; +using DotNetProjects.Migrator.Providers.Impl.Mysql; +using NSubstitute; +using NSubstitute.Extensions; +using NUnit.Framework; +using UniqueConstraint = DotNetProjects.Migrator.Framework.UniqueConstraint; +using ForeignKeyConstraint = DotNetProjects.Migrator.Framework.ForeignKeyConstraint; + +namespace Migrator.Tests; + +public class ConstraintCatalogParsingTests +{ + private MySqlTransformationProvider provider; + private IDbCommand command; + private DataTable rows; + private DataTableReader reader; + + [SetUp] + public void SetUp() + { + var connection = Substitute.For(); + connection.State.Returns(ConnectionState.Open); + command = Substitute.For(); + connection.CreateCommand().Returns(command); + command.CreateParameter().Returns(_ => Substitute.For()); + command.Parameters.Returns(Substitute.For()); + provider = Substitute.ForPartsOf(new MysqlDialect(), connection, "default", null); + rows = new DataTable(); + foreach (var name in new[] { "name", "type", "column", "ordinal", "expression" }) rows.Columns.Add(name, typeof(object)); + provider.Configure().GetForeignKeyConstraints(Arg.Any()).Returns(_ => + { + Assert.That(reader.IsClosed, Is.True, "Close the catalog reader before querying foreign keys."); + return new[] { new ForeignKeyConstraint("FK_Parent", "Parent", new[] { "Id" }, "Items", new[] { "ParentId" }) }; + }); + } + + [TearDown] + public void TearDown() { reader?.Dispose(); rows.Dispose(); provider.Dispose(); } + + private TableConstraint[] Read() + { + reader = rows.CreateDataReader(); + provider.Configure().ExecuteQuery(Arg.Any(), Arg.Any()).Returns(reader); + return ConstraintMetadataReader.Read(provider, "Items"); + } + + [TestCase("P")][TestCase("PK")][TestCase(" primary key ")][TestCase("PN")] + [TestCase("U")][TestCase("UQ")][TestCase(" unique ")] + public void CompositeKeysKeepCatalogOrderAndConstraintBoundaries(string type) + { + rows.Rows.Add("First", type, "Z", 1, DBNull.Value); + rows.Rows.Add("First", type, "A", 2, DBNull.Value); + rows.Rows.Add("Second", "U", "Other", 1, DBNull.Value); + var result = Read(); + var first = result[0]; + var columns = first is PrimaryKeyConstraint pk ? pk.KeyColumns : ((UniqueConstraint)first).KeyColumns; + Assert.That(columns, Is.EqualTo(new[] { "Z", "A" })); + Assert.That(first.Name, Is.EqualTo("First")); + if (type.Trim().StartsWith("P", StringComparison.OrdinalIgnoreCase)) + Assert.That(((PrimaryKeyConstraint)first).NonClustered, Is.EqualTo(type == "PN")); + else Assert.That(first, Is.TypeOf()); + Assert.That(((UniqueConstraint)result[1]).KeyColumns, Is.EqualTo(new[] { "Other" })); + Assert.That(result[2].Name, Is.EqualTo("FK_Parent")); + Assert.That(result, Has.Length.EqualTo(3)); + command.Received(1).Dispose(); + } + + [TestCase("C", "CHECK (Amount > 0)", "Amount > 0")] + [TestCase("K", "Amount > 0", "Amount > 0")] + [TestCase(" check ", null, null)] + public void ChecksPreserveExpressionsIncludingMissingCatalogText(string type, string expression, string expected) + { + rows.Rows.Add("Positive", type, DBNull.Value, 0, (object)expression ?? DBNull.Value); + Assert.That(((CheckConstraint)Read()[0]).CheckConstraintString, Is.EqualTo(expected)); + } + + [Test] + public void EmptyCatalogStillIncludesForeignKeys() + => Assert.That(Read().Select(c => c.Name), Is.EqualTo(new[] { "FK_Parent" })); + + [Test] + public void UnknownCatalogTypeFailsAndDisposesResources() + { + rows.Rows.Add("Unknown", "unexpected", "Id", 1, DBNull.Value); + Assert.Throws(() => Read()); + Assert.That(reader.IsClosed, Is.True); + command.Received(1).Dispose(); + provider.DidNotReceive().GetForeignKeyConstraints(Arg.Any()); + } +} diff --git a/src/Migrator.Tests/SQLiteFallbackAlterationTests.cs b/src/Migrator.Tests/SQLiteFallbackAlterationTests.cs index 59789cfb..81659f5e 100644 --- a/src/Migrator.Tests/SQLiteFallbackAlterationTests.cs +++ b/src/Migrator.Tests/SQLiteFallbackAlterationTests.cs @@ -15,7 +15,14 @@ public class SQLiteFallbackAlterationTests private sealed class LegacyProvider(IDbConnection connection) : SQLiteTransformationProvider(new SQLiteDialect(), connection, "default", null) { + public bool FailParentRebuild { get; set; } public override object ExecuteScalar(string sql) => sql == "SELECT sqlite_version()" ? "3.25.0" : base.ExecuteScalar(sql); + public override int ExecuteNonQuery(string sql) + { + if (FailParentRebuild && sql.StartsWith("CREATE TABLE", StringComparison.OrdinalIgnoreCase) && sql.Contains("OriginalTemp")) + throw new InvalidOperationException("Injected rebuild failure."); + return base.ExecuteNonQuery(sql); + } } private SqliteConnection connection; @@ -165,4 +172,57 @@ public void RenameUpdatesSelfReferencingParentColumns() Assert.That(provider.ExecuteScalar("SELECT ParentId FROM Nodes WHERE Id=2"), Is.EqualTo(1L)); Assert.That(provider.CheckForeignKeyIntegrity(), Is.True); } + + [TestCase("Obsolete")] + [TestCase("obsolete")] + public void RemoveColumnUpdatesIncomingKeysCaseInsensitively(string name) + { + provider.ExecuteNonQuery("CREATE TABLE Original (Id INTEGER PRIMARY KEY, Obsolete INTEGER UNIQUE); CREATE TABLE Child (Value INTEGER REFERENCES Original(Obsolete)); INSERT INTO Original VALUES (1, 7); INSERT INTO Child VALUES (7)"); + provider.RemoveColumn("Original", name); + Assert.That(provider.ColumnExists("Original", "Obsolete"), Is.False); + Assert.That(provider.GetForeignKeyConstraints("Child"), Is.Empty); + Assert.That(provider.ExecuteScalar("SELECT Value FROM Child"), Is.EqualTo(7L)); + Assert.That(provider.ExecuteScalar("SELECT Id FROM Original"), Is.EqualTo(1L)); + Assert.That(provider.CheckForeignKeyIntegrity(), Is.True); + } + + [TestCase(false)] + [TestCase(true)] + public void FailureRemovingColumnPreservesEarlierDependentTables(bool callerTransaction) + { + provider.ExecuteNonQuery("CREATE TABLE Original (Id INTEGER PRIMARY KEY, Obsolete INTEGER UNIQUE); CREATE TABLE Child (Value INTEGER REFERENCES Original(Obsolete)); INSERT INTO Original VALUES (1, 7); INSERT INTO Child VALUES (7)"); + if (callerTransaction) provider.BeginTransaction(); + provider.FailParentRebuild = true; + Assert.That(Assert.Throws(() => provider.RemoveColumn("Original", "Obsolete")).Message, Is.EqualTo("Injected rebuild failure.")); + Assert.That(provider.HasActiveTransaction, Is.EqualTo(callerTransaction)); + if (callerTransaction) provider.Rollback(); + Assert.That(provider.GetForeignKeyConstraints("Child"), Has.Length.EqualTo(1)); + Assert.That(provider.ExecuteScalar("SELECT Obsolete FROM Original"), Is.EqualTo(7L)); + Assert.That(provider.ExecuteScalar("SELECT Value FROM Child"), Is.EqualTo(7L)); + Assert.That(provider.GetTables(), Is.EquivalentTo(new[] { "Original", "Child" })); + Assert.That(provider.CheckForeignKeyIntegrity(), Is.True); + } + + [Test] + public void RemovingSelfReferencedColumnPreservesOtherValues() + { + provider.ExecuteNonQuery("CREATE TABLE Nodes (Id INTEGER PRIMARY KEY, Obsolete INTEGER UNIQUE, ParentValue INTEGER REFERENCES Nodes(Obsolete)); INSERT INTO Nodes VALUES (1, 7, NULL), (2, 8, 7)"); + provider.RemoveColumn("Nodes", "obsolete"); + Assert.That(provider.GetForeignKeyConstraints("Nodes"), Is.Empty); + Assert.That(provider.ExecuteScalar("SELECT ParentValue FROM Nodes WHERE Id=2"), Is.EqualTo(7L)); + Assert.That(provider.CheckForeignKeyIntegrity(), Is.True); + } + + [Test] + public void SuccessfulRemovalRemainsInTheCallerTransaction() + { + provider.ExecuteNonQuery("CREATE TABLE Original (Id INTEGER PRIMARY KEY, Obsolete INTEGER UNIQUE); CREATE TABLE Child (Value INTEGER REFERENCES Original(Obsolete)); INSERT INTO Original VALUES (1, 7); INSERT INTO Child VALUES (7)"); + provider.BeginTransaction(); + provider.RemoveColumn("Original", "Obsolete"); + Assert.That(provider.HasActiveTransaction, Is.True); + Assert.That(provider.GetForeignKeyConstraints("Child"), Is.Empty); + provider.Rollback(); + Assert.That(provider.GetForeignKeyConstraints("Child"), Has.Length.EqualTo(1)); + Assert.That(provider.ExecuteScalar("SELECT Obsolete FROM Original"), Is.EqualTo(7L)); + } } diff --git a/src/Migrator.Tests/SchemaCatalogContractTests.cs b/src/Migrator.Tests/SchemaCatalogContractTests.cs new file mode 100644 index 00000000..664631fe --- /dev/null +++ b/src/Migrator.Tests/SchemaCatalogContractTests.cs @@ -0,0 +1,36 @@ +using System.Data; +using System.Data.Common; +using System.Linq; +using DotNetProjects.Migrator.Providers.Impl.SQLite; +using NSubstitute; +using NUnit.Framework; + +namespace Migrator.Tests; + +public class SchemaCatalogContractTests +{ + [Test] + public void SchemaColumnEnumerationReturnsColumnNamesAndScopesTheRequest() + { + using var table = new DataTable(); + table.Columns.Add("TABLE_NAME"); table.Columns.Add("COLUMN_NAME"); + table.Rows.Add("Orders", "Id"); table.Rows.Add("Orders", "Total"); + var connection = Substitute.For(); + connection.GetSchema("Columns", Arg.Any()).Returns(table); + using var provider = new SQLiteTransformationProvider(new SQLiteDialect(), connection, "default", null); + Assert.That(provider.GetColumns("sales", "Orders").ToArray(), Is.EqualTo(new[] { "Id", "Total" })); + connection.Received(1).GetSchema("Columns", Arg.Is(x => x.Length == 4 && x[0] == null && x[1] == "sales" && x[2] == "Orders" && x[3] == null)); + } + + [Test] + public void SchemaTableEnumerationPreservesReturnedNamesAndSchemaRestriction() + { + using var table = new DataTable(); table.Columns.Add("TABLE_NAME"); + table.Rows.Add("Orders"); table.Rows.Add("Order Details"); + var connection = Substitute.For(); + connection.GetSchema("Tables", Arg.Any()).Returns(table); + using var provider = new SQLiteTransformationProvider(new SQLiteDialect(), connection, "default", null); + Assert.That(provider.GetTables("sales").ToArray(), Is.EqualTo(new[] { "Orders", "Order Details" })); + connection.Received(1).GetSchema("Tables", Arg.Is(x => x.Length == 4 && x[1] == "sales" && x[2] == null)); + } +} diff --git a/src/Migrator.Tests/SqlPreviewBoundaryTests.cs b/src/Migrator.Tests/SqlPreviewBoundaryTests.cs new file mode 100644 index 00000000..ce6a4246 --- /dev/null +++ b/src/Migrator.Tests/SqlPreviewBoundaryTests.cs @@ -0,0 +1,86 @@ +using System; +using System.Collections.Generic; +using System.Data; +using DotNetProjects.Migrator; +using DotNetProjects.Migrator.Framework; +using DotNetProjects.Migrator.Framework.Fluent; +using DotNetProjects.Migrator.Providers; +using NUnit.Framework; + +namespace Migrator.Tests; + +public class SqlPreviewBoundaryTests +{ + private static IEnumerable Literals() + { + yield return new TestCaseData(null, "NULL"); + yield return new TestCaseData(DBNull.Value, "NULL"); + yield return new TestCaseData("O'Brien", "'O''Brien'"); + yield return new TestCaseData("", "''"); + yield return new TestCaseData((byte)255, "255"); + yield return new TestCaseData((sbyte)-128, "-128"); + yield return new TestCaseData(short.MinValue, "-32768"); + yield return new TestCaseData(ushort.MaxValue, "65535"); + yield return new TestCaseData(int.MinValue, "-2147483648"); + yield return new TestCaseData(uint.MaxValue, "4294967295"); + yield return new TestCaseData(long.MinValue, "-9223372036854775808"); + yield return new TestCaseData(ulong.MaxValue, "18446744073709551615"); + yield return new TestCaseData(12.5m, "12.5"); + yield return new TestCaseData(12.5f, "12.5"); + yield return new TestCaseData(12.5d, "12.5"); + yield return new TestCaseData(new DateTime(2024, 2, 29, 12, 34, 56).AddTicks(1234567), "'2024-02-29 12:34:56.1234567'"); + yield return new TestCaseData(Guid.Parse("00112233-4455-6677-8899-aabbccddeeff"), "'00112233-4455-6677-8899-aabbccddeeff'"); + } + + [TestCaseSource(nameof(Literals)), SetCulture("de-DE")] + public void LiteralsPreserveValuesWithoutUsingCurrentCulture(object value, string expected) + => Assert.That(new SqlGenerationContext(ProviderTypes.SQLite).Literal(value), Is.EqualTo(expected)); + + [TestCase(ProviderTypes.SQLite, "1", "0")] + [TestCase(ProviderTypes.SqlServer, "1", "0")] + [TestCase(ProviderTypes.PostgreSQL, "TRUE", "FALSE")] + [TestCase(ProviderTypes.PostgreSQL82, "TRUE", "FALSE")] + public void BooleanLiteralsUseTheTargetDialect(ProviderTypes provider, string yes, string no) + { + var context = new SqlGenerationContext(provider); + Assert.That(context.Literal(true), Is.EqualTo(yes)); + Assert.That(context.Literal(false), Is.EqualTo(no)); + } + + [Test] + public void UnsupportedValuesCannotMasqueradeAsSqlLiterals() + { + var context = new SqlGenerationContext(ProviderTypes.SQLite); + foreach (var value in new object[] { DayOfWeek.Monday, RawSql.Insert("DROP TABLE Items"), new Version(1, 2), new byte[] { 1 }, TimeSpan.FromDays(1) }) + Assert.Throws(() => context.Literal(value)); + } + + [TestCase(ProviderTypes.SQLite, "\"two words\"", "\"a.b\"")] + [TestCase(ProviderTypes.PostgreSQL, "\"two words\"", "\"a.b\"")] + [TestCase(ProviderTypes.SqlServer, "[two words]", "[a.b]")] + public void IdentifierAtomsAreQuotedWithoutSplittingDots(ProviderTypes type, string spaced, string dotted) + { + var context = new SqlGenerationContext(type); + Assert.That(context.Quote("two words"), Is.EqualTo(spaced)); + Assert.That(context.Quote("a.b"), Is.EqualTo(dotted)); + } + + [TestCase(ProviderTypes.SQLite, "\"sales.region\".\"Order Lines\"", "\"sales.region\".\"Order Lines\"")] + [TestCase(ProviderTypes.PostgreSQL, "\"sales.region\".\"Order Lines\"", "\"sales.region\".\"Order Lines\"")] + [TestCase(ProviderTypes.SqlServer, "[sales.region].[Order]]Lines]", "[sales.region].[Order]]Lines]")] + public void QualifiedTableNamesPreserveQuotedComponents(ProviderTypes type, string name, string expected) + => Assert.That(new SqlGenerationContext(type).Table(name), Is.EqualTo(expected)); + + [Test, Category("SQLite")] + public void PreviewWithSpacesAndDottedColumnNamesExecutesLikeTheMigration() + { + var builder = new MigrationBuilder(); + builder.Create.Table("Order Details").WithColumn("Id").AsInt32().WithColumn("Line.Item").AsString(); + builder.Insert.IntoTable("Order Details").Row(new[] { "Id", "Line.Item" }, new object[] { 1, "O'Brien" }); + builder.Create.Index("Index with spaces").OnTable("Order Details").WithColumns("Line.Item"); + using var provider = ProviderFactory.Create(ProviderTypes.SQLite, "Data Source=:memory:", null); + foreach (var sql in builder.Preview(new SqlGenerationContext(ProviderTypes.SQLite))) provider.ExecuteNonQuery(sql); + Assert.That(provider.ExecuteScalar("SELECT \"Line.Item\" FROM \"Order Details\" WHERE Id=1"), Is.EqualTo("O'Brien")); + Assert.That(provider.IndexExists("Order Details", "Index with spaces"), Is.True); + } +} diff --git a/src/Migrator/Framework/Fluent/SqlGenerationContext.cs b/src/Migrator/Framework/Fluent/SqlGenerationContext.cs index 645d8ef4..141754fc 100644 --- a/src/Migrator/Framework/Fluent/SqlGenerationContext.cs +++ b/src/Migrator/Framework/Fluent/SqlGenerationContext.cs @@ -19,15 +19,13 @@ private void RequireKnownSchema() public Dialect Dialect { get; } public SqlGenerationContext(ProviderTypes provider, Func existingTable = null) { Provider = provider; Dialect = ProviderFactory.DialectForProvider(provider) ?? throw new ArgumentException("Unknown provider."); this.existingTable = existingTable; } - private string AlwaysQuote(string name) + public string Quote(string name) => Dialect.QuoteColumnNameIfRequired(name); + public string Table(string name) => Dialect.QuoteTableNameIfRequired(name); + private static readonly HashSet NumericTypes = new() { - if (string.IsNullOrWhiteSpace(name)) throw new ArgumentException("Identifier cannot be empty."); - var template = Dialect.QuoteTemplate; - var closing = template[^1].ToString(); - return string.Format(CultureInfo.InvariantCulture, template, name.Replace(closing, closing + closing)); - } - public string Quote(string name) => Dialect.ColumnNameNeedsQuote || Dialect.IsReservedWord(name) ? AlwaysQuote(name) : name; - public string Table(string name) => string.Join(".", name.Split('.').Select(part => Dialect.TableNameNeedsQuote || Dialect.IsReservedWord(part) ? AlwaysQuote(part) : part)); + typeof(byte), typeof(sbyte), typeof(short), typeof(ushort), typeof(int), typeof(uint), + typeof(long), typeof(ulong), typeof(decimal), typeof(float), typeof(double) + }; public string Column(Column column) => Dialect.GetAndMapColumnProperties(Definitions.CopyColumn(column)).ColumnSql; public string Literal(object value) => value switch { @@ -37,7 +35,7 @@ private string AlwaysQuote(string name) DateTime d => "'" + d.ToString("yyyy-MM-dd HH:mm:ss.fffffff", CultureInfo.InvariantCulture) + "'", Guid g when Dialect is DotNetProjects.Migrator.Providers.Impl.Oracle.OracleDialect => Dialect.Default(g)[8..], Guid g => "'" + g + "'", - byte or sbyte or short or ushort or int or uint or long or ulong or decimal or float or double => Convert.ToString(value, CultureInfo.InvariantCulture), + object number when NumericTypes.Contains(number.GetType()) => Convert.ToString(value, CultureInfo.InvariantCulture), _ => throw new NotSupportedException("No portable SQL literal for " + value.GetType().Name) }; public void RequireTable(string table) diff --git a/src/Migrator/Providers/ConstraintMetadataReader.cs b/src/Migrator/Providers/ConstraintMetadataReader.cs index 59fd8feb..017f9faf 100644 --- a/src/Migrator/Providers/ConstraintMetadataReader.cs +++ b/src/Migrator/Providers/ConstraintMetadataReader.cs @@ -16,6 +16,22 @@ namespace DotNetProjects.Migrator.Providers; internal static class ConstraintMetadataReader { public static TableConstraint[] Read(TransformationProvider provider, string table) + { + var query = CatalogQuery(provider, table); + List constraints; + using (var command = provider.CreateCommand()) + { + AddParameter(command, "lookup_table", query.Table); + if (query.IncludeSchema) AddParameter(command, "lookup_schema", query.Schema); + using var reader = provider.ExecuteQuery(command, query.Sql); + constraints = ReadConstraints(reader); + } + // Some drivers allow only one active reader on a connection. + constraints.AddRange(provider.GetForeignKeyConstraints(table)); + return constraints.ToArray(); + } + + private static (string Sql, string Table, string Schema, bool IncludeSchema) CatalogQuery(TransformationProvider provider, string table) { string sql; var parameterTable = provider.QuoteTableNameIfRequired(table); @@ -73,42 +89,55 @@ FROM information_schema.TABLE_CONSTRAINTS c LEFT JOIN information_schema.KEY_COL } else throw new NotSupportedException("Structured constraint inspection is not implemented for " + provider.Dialect.GetType().Name + "."); + return (sql, parameterTable, schema, oracle || provider.Dialect is MysqlDialect); + } + + private enum ConstraintKind { Primary, NonClusteredPrimary, Unique, Check } + private static readonly Dictionary CatalogKinds = new(StringComparer.OrdinalIgnoreCase) + { + ["P"] = ConstraintKind.Primary, ["PK"] = ConstraintKind.Primary, ["PRIMARY KEY"] = ConstraintKind.Primary, + ["PN"] = ConstraintKind.NonClusteredPrimary, + ["U"] = ConstraintKind.Unique, ["UQ"] = ConstraintKind.Unique, ["UNIQUE"] = ConstraintKind.Unique, + ["C"] = ConstraintKind.Check, ["K"] = ConstraintKind.Check, ["CHECK"] = ConstraintKind.Check + }; + + private static TableConstraint CreateConstraint(IDataRecord row, string name) + { + if (!CatalogKinds.TryGetValue(row.GetString(1).Trim(), out var kind)) + throw new MigrationException("Unknown catalog constraint type."); + return kind switch + { + ConstraintKind.Primary => new PrimaryKeyConstraint { Name = name }, + ConstraintKind.NonClusteredPrimary => new PrimaryKeyConstraint { Name = name, NonClustered = true }, + ConstraintKind.Unique => new UniqueConstraint { Name = name }, + _ => new CheckConstraint(name, row.IsDBNull(4) ? null : CheckExpression(row.GetString(4))) + }; + } + + private static List ReadConstraints(IDataReader reader) + { var constraints = new List(); - using (var command = provider.CreateCommand()) + TableConstraint current = null; + var keys = new List(); + void Complete() { - AddParameter(command, "lookup_table", parameterTable); - if (oracle || provider.Dialect is MysqlDialect) AddParameter(command, "lookup_schema", schema); - using var reader = provider.ExecuteQuery(command, sql); - string lastName = null; - TableConstraint current = null; - var keys = new List(); - void Complete() + if (current is PrimaryKeyConstraint pk) pk.KeyColumns = keys.ToArray(); + if (current is UniqueConstraint unique) unique.KeyColumns = keys.ToArray(); + if (current != null) constraints.Add(current); + } + while (reader.Read()) + { + var name = reader.GetString(0); + if (current?.Name != name) { - if (current is PrimaryKeyConstraint pk) pk.KeyColumns = keys.ToArray(); - if (current is UniqueConstraint unique) unique.KeyColumns = keys.ToArray(); - if (current != null) constraints.Add(current); + Complete(); + keys.Clear(); + current = CreateConstraint(reader, name); } - while (reader.Read()) - { - var name = reader.GetString(0); - if (name != lastName) - { - Complete(); keys.Clear(); lastName = name; - current = reader.GetString(1).Trim().ToUpperInvariant() switch - { - "P" or "PK" or "PRIMARY KEY" => new PrimaryKeyConstraint { Name = name }, - "PN" => new PrimaryKeyConstraint { Name = name, NonClustered = true }, - "U" or "UQ" or "UNIQUE" => new UniqueConstraint { Name = name }, - "C" or "K" or "CHECK" => new CheckConstraint(name, reader.IsDBNull(4) ? null : CheckExpression(reader.GetString(4))), - _ => throw new MigrationException("Unknown catalog constraint type.") - }; - } - if (!reader.IsDBNull(2)) keys.Add(reader.GetString(2)); - } - Complete(); + if (!reader.IsDBNull(2)) keys.Add(reader.GetString(2)); } - constraints.AddRange(provider.GetForeignKeyConstraints(table)); - return constraints.ToArray(); + Complete(); + return constraints; } internal static string CheckExpression(string source) diff --git a/src/Migrator/Providers/Impl/SQLite/SQLiteTransformationProvider.cs b/src/Migrator/Providers/Impl/SQLite/SQLiteTransformationProvider.cs index 283c846c..c1b1f4fd 100644 --- a/src/Migrator/Providers/Impl/SQLite/SQLiteTransformationProvider.cs +++ b/src/Migrator/Providers/Impl/SQLite/SQLiteTransformationProvider.cs @@ -474,48 +474,57 @@ public void MoveIndexesFromOriginalTable(string origTable, string newTable) public override void RemoveColumn(string tableName, string column) { if (Version.Parse(Convert.ToString(ExecuteScalar("SELECT sqlite_version()"))) >= new Version(3, 35, 0) - && TableExists(tableName)) - { - var info = GetSQLiteTableInfo(tableName); - var definition = info.Columns.SingleOrDefault(c => c.Name.Equals(column, StringComparison.OrdinalIgnoreCase)); - bool Matches(string name) => string.Equals(name, column, StringComparison.OrdinalIgnoreCase); - var dependent = definition == null || info.PrimaryKey?.KeyColumns.Contains(column, StringComparer.OrdinalIgnoreCase) == true || info.Uniques.Any(u => u.KeyColumns.Contains(column, StringComparer.OrdinalIgnoreCase)) - || info.CheckConstraints.Count != 0 - || info.Uniques.Any(u => u.KeyColumns.Any(Matches)) - || info.Indexes.Any(i => i.KeyColumns.Any(Matches) || i.FilterItems.Count != 0) - || info.ForeignKeys.Any(f => f.ChildColumns.Any(Matches)) - || GetTables().Any(t => GetForeignKeyConstraints(t).Any(f => f.ParentTable.Equals(tableName, StringComparison.OrdinalIgnoreCase) && f.ParentColumns.Any(Matches))); - if (!dependent) - { - // SQLite itself validates trigger/view dependencies atomically. A rejection is - // surfaced rather than retrying with a potentially lossy reconstruction. - ExecuteNonQuery($"ALTER TABLE {Dialect.Quote(tableName)} DROP COLUMN {Dialect.QuoteIdentifier(definition.Name)}"); - return; - } - } - // In SQLite we need to recreate the table even if we only want to add, alter or drop a foreign key. So we not only recreate the table given - // as parameter but also the tables with FKs pointing to the column you want to remove. - // In order to perform it smoothly, the PRAGMA foreign keys should be set off. - - var isPragmaForeignKeysOn = IsPragmaForeignKeysOn(); - - if (isPragmaForeignKeysOn) - { - throw new Exception($"{nameof(RemoveColumn)} requires foreign keys off."); - } - - if (!TableExists(tableName)) - { - throw new MigrationException($"The table '{tableName}' does not exist"); - } - - if (!ColumnExists(tableName, column)) + && TableExists(tableName) && CanDropColumnNatively(tableName, column)) { - throw new MigrationException($"The table '{tableName}' does not have a column named '{column}'"); + // Native SQLite validates trigger and view dependencies atomically. + ExecuteNonQuery($"ALTER TABLE {Dialect.Quote(tableName)} DROP COLUMN {Dialect.QuoteIdentifier(column)}"); + return; } + if (IsPragmaForeignKeysOn()) throw new Exception($"{nameof(RemoveColumn)} requires foreign keys off."); + if (!TableExists(tableName)) throw new MigrationException($"The table '{tableName}' does not exist"); + if (!ColumnExists(tableName, column)) throw new MigrationException($"The table '{tableName}' does not have a column named '{column}'"); var sqliteInfoMainTable = GetSQLiteTableInfo(tableName); - + ValidateColumnRemoval(sqliteInfoMainTable, column); + var affected = new List(); + foreach (var name in GetTables()) + { + var info = string.Equals(name, tableName, StringComparison.OrdinalIgnoreCase) + ? sqliteInfoMainTable : GetSQLiteTableInfo(name); + var references = info.ForeignKeys.Where(f => + string.Equals(f.ParentTable, tableName, StringComparison.OrdinalIgnoreCase) + && f.ParentColumns.Contains(column, StringComparer.OrdinalIgnoreCase)).ToArray(); + if (references.Any(f => f.ParentColumns.Length > 1)) + throw new MigrationException($"You need to delete/adjust the FK in table {name} pointing to {tableName}."); + foreach (var reference in references) info.ForeignKeys.Remove(reference); + if (references.Length != 0 && info != sqliteInfoMainTable) affected.Add(info); + } + sqliteInfoMainTable.Uniques.RemoveAll(x => x.KeyColumns.Length == 1 && x.KeyColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); + sqliteInfoMainTable.ColumnMappings.RemoveAll(x => x.OldName.Equals(column, StringComparison.OrdinalIgnoreCase)); + sqliteInfoMainTable.Columns.RemoveAll(x => x.Name.Equals(column, StringComparison.OrdinalIgnoreCase)); + sqliteInfoMainTable.Indexes.RemoveAll(x => x.KeyColumns.Length == 1 && x.KeyColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); + sqliteInfoMainTable.ForeignKeys.RemoveAll(x => x.ChildColumns.Length == 1 && x.ChildColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); + + affected.Add(sqliteInfoMainTable); + RecreateTablesAtomically(affected); + } + + private bool CanDropColumnNatively(string tableName, string column) + { + var info = GetSQLiteTableInfo(tableName); + var definition = info.Columns.SingleOrDefault(c => c.Name.Equals(column, StringComparison.OrdinalIgnoreCase)); + bool Matches(string name) => string.Equals(name, column, StringComparison.OrdinalIgnoreCase); + var dependent = definition == null || info.PrimaryKey?.KeyColumns.Contains(column, StringComparer.OrdinalIgnoreCase) == true + || info.CheckConstraints.Count != 0 + || info.Uniques.Any(u => u.KeyColumns.Any(Matches)) + || info.Indexes.Any(i => i.KeyColumns.Any(Matches) || i.FilterItems.Count != 0) + || info.ForeignKeys.Any(f => f.ChildColumns.Any(Matches)) + || GetTables().Any(t => GetForeignKeyConstraints(t).Any(f => f.ParentTable.Equals(tableName, StringComparison.OrdinalIgnoreCase) && f.ParentColumns.Any(Matches))); + return !dependent; + } + + private static void ValidateColumnRemoval(SQLiteTableInfo sqliteInfoMainTable, string column) + { if (sqliteInfoMainTable.PrimaryKey?.KeyColumns.Any(x => x.Equals(column, StringComparison.OrdinalIgnoreCase)) == true) throw new MigrationException("Remove the named primary-key constraint before removing one of its columns."); @@ -526,7 +535,7 @@ public override void RemoveColumn(string tableName, string column) throw new MigrationException("A check constraint contains the column you want to remove. Remove the check constraint first"); } - if (!sqliteInfoMainTable.ColumnMappings.Any(x => x.OldName == column)) + if (!sqliteInfoMainTable.ColumnMappings.Any(x => x.OldName.Equals(column, StringComparison.OrdinalIgnoreCase))) { throw new MigrationException("Column not found"); } @@ -580,55 +589,27 @@ public override void RemoveColumn(string tableName, string column) throw new Exception(stringBuilder.ToString()); } - var allTableNames = GetTables(); + } - // Remove foreign keys with single parent column pointing to the column to be removed. - foreach (var allTableName in allTableNames) + private void RecreateTablesAtomically(IEnumerable tables) + { + var ownsTransaction = !HasActiveTransaction; + if (ownsTransaction) BeginTransaction(); + try { - if (allTableName == tableName) - { - continue; - } - - var sqliteTableInfoOther = GetSQLiteTableInfo(allTableName); - var recreateOtherTable = false; - - for (var i = sqliteTableInfoOther.ForeignKeys.Count - 1; i >= 0; i--) - { - if (!sqliteTableInfoOther.ForeignKeys[i].ParentTable.Equals(tableName, StringComparison.OrdinalIgnoreCase)) - { - continue; - } - - if (sqliteTableInfoOther.ForeignKeys[i].ParentColumns.Contains(column) && sqliteTableInfoOther.ForeignKeys[i].ParentColumns.Length > 1) - { - StringBuilder stringBuilder = new(); - stringBuilder.Append($"You need to delete/adjust the FK in table {allTableName} pointing to {tableName}."); - stringBuilder.Append("Other foreign key if exists with just one parent column we adjust silently."); - - throw new Exception(stringBuilder.ToString()); - } - - if (sqliteTableInfoOther.ForeignKeys[i].ParentColumns.Contains(column) && sqliteTableInfoOther.ForeignKeys[i].ParentColumns.Length == 1) - { - recreateOtherTable = true; - sqliteTableInfoOther.ForeignKeys.RemoveAt(i); - } - } - - if (recreateOtherTable) + foreach (var info in tables) RecreateTable(info); + if (ownsTransaction) { - RecreateTable(sqliteTableInfoOther); + if (!CheckForeignKeyIntegrity()) throw new MigrationException("SQLite column removal would leave invalid foreign keys."); + Commit(); } } - - sqliteInfoMainTable.Uniques.RemoveAll(x => x.KeyColumns.Length == 1 && x.KeyColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); - sqliteInfoMainTable.ColumnMappings.RemoveAll(x => x.OldName.Equals(column, StringComparison.OrdinalIgnoreCase)); - sqliteInfoMainTable.Columns.RemoveAll(x => x.Name.Equals(column, StringComparison.OrdinalIgnoreCase)); - sqliteInfoMainTable.Indexes.RemoveAll(x => x.KeyColumns.Length == 1 && x.KeyColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); - sqliteInfoMainTable.ForeignKeys.RemoveAll(x => x.ChildColumns.Length == 1 && x.ChildColumns[0].Equals(column, StringComparison.OrdinalIgnoreCase)); - - RecreateTable(sqliteInfoMainTable); + catch (Exception ex) + { + if (ownsTransaction) + try { Rollback(); } catch (Exception rollback) { ex.Data["RollbackException"] = rollback; } + throw; + } } public override void RenameColumn(string tableName, string oldColumnName, string newColumnName) diff --git a/src/Migrator/Providers/TransformationProvider.cs b/src/Migrator/Providers/TransformationProvider.cs index 2a433e81..8e0b36b0 100644 --- a/src/Migrator/Providers/TransformationProvider.cs +++ b/src/Migrator/Providers/TransformationProvider.cs @@ -2002,7 +2002,7 @@ public IEnumerable GetColumns(string schema, string table) var c = _connection as DbConnection; var tables = c.GetSchema("Columns", tableRestrictions); - return from DataRow row in tables.Rows select (row["TABLE_NAME"] as string); + return from DataRow row in tables.Rows select (row["COLUMN_NAME"] as string); } protected void ValidateIndex(string tableName, Index index)