using System; using Microsoft.Data.SqlClient; using SyncEngine.Access; namespace SyncEngine.SqlServer { public static class SqlIdentityHelper { public static string GetIdentityStagingTableName(string tableName) { return SqlIdentifierHelper.SanitizeConstraintName(tableName) + "__identity_src"; } public static AccessColumnSchema GetIdentityColumn(AccessTableSchema table) { if (table?.Columns == null) { return null; } foreach (var col in table.Columns) { if (col.IsAutoIncrement) { return col; } } return null; } public static bool TableExists(SqlConnection connection, string tableName) { using (var cmd = new SqlCommand("SELECT CASE WHEN OBJECT_ID(@table, 'U') IS NULL THEN 0 ELSE 1 END", connection)) { cmd.Parameters.AddWithValue("@table", tableName); return Convert.ToInt32(cmd.ExecuteScalar()) == 1; } } public static bool ColumnIsIdentity(SqlConnection connection, string tableName, string columnName) { using (var cmd = new SqlCommand(@" SELECT COLUMNPROPERTY(OBJECT_ID(@table), @column, 'IsIdentity')", connection)) { cmd.Parameters.AddWithValue("@table", tableName); cmd.Parameters.AddWithValue("@column", columnName); var result = cmd.ExecuteScalar(); return result != null && result != DBNull.Value && Convert.ToInt32(result) == 1; } } public static bool NeedsIdentityRebuild(SqlConnection connection, AccessTableSchema table) { var identityCol = GetIdentityColumn(table); if (identityCol == null) { return false; } var stagingName = GetIdentityStagingTableName(table.TableName); if (TableExists(connection, stagingName) && !TableExists(connection, table.TableName)) { return true; } if (!TableExists(connection, table.TableName)) { return false; } if (!ColumnExists(connection, table.TableName, identityCol.Name)) { return false; } return !ColumnIsIdentity(connection, table.TableName, identityCol.Name); } public static void SetIdentityInsert(SqlConnection connection, SqlTransaction tx, string qualifiedTable, bool on) { var sql = string.Format("SET IDENTITY_INSERT {0} {1}", qualifiedTable, on ? "ON" : "OFF"); using (var cmd = tx == null ? new SqlCommand(sql, connection) : new SqlCommand(sql, connection, tx)) { cmd.ExecuteNonQuery(); } } public static void ReseedIdentity(SqlConnection connection, string qualifiedTable, object seedValue) { if (seedValue == null || seedValue == DBNull.Value) { return; } var seed = Convert.ToInt64(seedValue); var tableLiteral = ToCheckIdentTableLiteral(qualifiedTable); using (var cmd = new SqlCommand( string.Format("DBCC CHECKIDENT ({0}, RESEED, {1})", tableLiteral, seed), connection)) { cmd.ExecuteNonQuery(); } } /// DBCC CHECKIDENT requires a quoted 'schema.table' string, not bracketed identifiers. internal static string ToCheckIdentTableLiteral(string qualifiedTable) { if (string.IsNullOrWhiteSpace(qualifiedTable)) { throw new ArgumentException("Table name is required", nameof(qualifiedTable)); } var inner = qualifiedTable.Trim(); var schemaTable = inner; if (inner.StartsWith("[", StringComparison.Ordinal)) { var closeSchema = inner.IndexOf("].[", StringComparison.Ordinal); if (closeSchema > 0) { var schema = inner.Substring(1, closeSchema - 1); var table = inner.Substring(closeSchema + 3, inner.Length - closeSchema - 4); schemaTable = schema + "." + table; } else { schemaTable = inner.Trim('[', ']'); } } return "N'" + schemaTable.Replace("'", "''") + "'"; } public static object GetMaxColumnValue(SqlConnection connection, string qualifiedTable, string columnName) { using (var cmd = new SqlCommand(string.Format("SELECT MAX([{0}]) FROM {1}", columnName, qualifiedTable), connection)) { return cmd.ExecuteScalar(); } } private static bool ColumnExists(SqlConnection connection, string tableName, string columnName) { using (var cmd = new SqlCommand("SELECT COL_LENGTH(@table, @column)", connection)) { cmd.Parameters.AddWithValue("@table", tableName); cmd.Parameters.AddWithValue("@column", columnName); var result = cmd.ExecuteScalar(); return result != null && result != DBNull.Value; } } } }