From 49b740b4b67ffddf71a30d55986dbc396cd50a3f Mon Sep 17 00:00:00 2001 From: Jordan Brown Date: Sun, 22 Mar 2020 16:59:21 -0400 Subject: [PATCH] Add SQLite Setup Wizard support --- src/Tgstation.Server.Host/Core/SetupWizard.cs | 315 ++++++++++++------ 1 file changed, 212 insertions(+), 103 deletions(-) diff --git a/src/Tgstation.Server.Host/Core/SetupWizard.cs b/src/Tgstation.Server.Host/Core/SetupWizard.cs index d685044ec7..6fa087c27a 100644 --- a/src/Tgstation.Server.Host/Core/SetupWizard.cs +++ b/src/Tgstation.Server.Host/Core/SetupWizard.cs @@ -1,4 +1,5 @@ using Microsoft.AspNetCore.Hosting; +using Microsoft.Data.Sqlite; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Options; using MySql.Data.MySqlClient; @@ -8,6 +9,7 @@ using System.Collections.Generic; using System.Data.Common; using System.Data.SqlClient; using System.Globalization; +using System.IO; using System.Linq; using System.Text; using System.Text.RegularExpressions; @@ -148,6 +150,129 @@ namespace Tgstation.Server.Host.Core while (true); } + /// + /// Ensure a given works. + /// + /// The test . + /// The may have derived data populated. + /// The database name (or path in the case of a database). + /// Whether or not the database exists. + /// The for the operation. + /// A representing the running operation. + async Task TestDatabaseConnection( + DbConnection testConnection, + DatabaseConfiguration databaseConfiguration, + string databaseName, + bool dbExists, + CancellationToken cancellationToken) + { + bool isSqliteDB = databaseConfiguration.DatabaseType == DatabaseType.Sqlite; + using (testConnection) + { + await console.WriteAsync("Testing connection...", true, cancellationToken).ConfigureAwait(false); + await testConnection.OpenAsync(cancellationToken).ConfigureAwait(false); + await console.WriteAsync("Connection successful!", true, cancellationToken).ConfigureAwait(false); + + if (databaseConfiguration.DatabaseType == DatabaseType.MariaDB + || databaseConfiguration.DatabaseType == DatabaseType.MySql) + { + await console.WriteAsync("Checking MySQL/MariaDB version...", true, cancellationToken).ConfigureAwait(false); + using (var command = testConnection.CreateCommand()) + { + command.CommandText = "SELECT VERSION()"; + var fullVersion = (string)await command.ExecuteScalarAsync(cancellationToken).ConfigureAwait(false); + await console.WriteAsync(String.Format(CultureInfo.InvariantCulture, "Found {0}", fullVersion), true, cancellationToken).ConfigureAwait(false); + var splits = fullVersion.Split('-'); + databaseConfiguration.MySqlServerVersion = splits.First(); + } + } + + if (!isSqliteDB && !dbExists) + { + await console.WriteAsync("Testing create DB permission...", true, cancellationToken).ConfigureAwait(false); + using (var command = testConnection.CreateCommand()) + { + // I really don't care about user sanitization here, they want to fuck their own DB? so be it +#pragma warning disable CA2100 // Review SQL queries for security vulnerabilities + command.CommandText = $"CREATE DATABASE {databaseName}"; +#pragma warning restore CA2100 // Review SQL queries for security vulnerabilities + await command.ExecuteNonQueryAsync(cancellationToken).ConfigureAwait(false); + } + + await console.WriteAsync("Success!", true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("Dropping test database...", true, cancellationToken).ConfigureAwait(false); + using (var command = testConnection.CreateCommand()) + { +#pragma warning disable CA2100 // Review SQL queries for security vulnerabilities + command.CommandText = $"DROP DATABASE {databaseName}"; +#pragma warning restore CA2100 // Review SQL queries for security vulnerabilities + try + { + await command.ExecuteNonQueryAsync(cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException) + { + throw; + } + catch (Exception e) + { + await console.WriteAsync(e.Message, true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("This should be okay, but you may want to manually drop the database before continuing!", true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("Press any key to continue...", true, cancellationToken).ConfigureAwait(false); + await console.PressAnyKeyAsync(cancellationToken).ConfigureAwait(false); + } + } + } + } + + if (isSqliteDB && !dbExists) + await Task.WhenAll( + console.WriteAsync("Deleting test database file...", true, cancellationToken), + ioManager.DeleteFile(databaseName, cancellationToken)).ConfigureAwait(false); + } + + async Task ValidateNonExistantSqliteDBName(string databaseName, CancellationToken cancellationToken) + { + var resolvedPath = ioManager.ResolvePath(databaseName); + try + { + var directoryName = ioManager.GetDirectoryName(resolvedPath); + bool directoryExisted = await ioManager.DirectoryExists(directoryName, cancellationToken).ConfigureAwait(false); + await ioManager.CreateDirectory(directoryName, cancellationToken).ConfigureAwait(false); + try + { + await ioManager.WriteAllBytes(resolvedPath, Array.Empty(), cancellationToken).ConfigureAwait(false); + } + catch + { + if (!directoryExisted) + await ioManager.DeleteDirectory(directoryName, cancellationToken).ConfigureAwait(false); + throw; + } + } + catch (IOException) + { + return null; + } + + if (!Path.IsPathRooted(databaseName)) + { + await console.WriteAsync("Note, this relative path (currently) resolves to the following:", true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync(resolvedPath, true, cancellationToken).ConfigureAwait(false); + bool writeResolved = await PromptYesNo( + "Would you like to save the relative path in the configuration? If not, the full path will be saved. (y/n): ", + cancellationToken) + .ConfigureAwait(false); + + if (writeResolved) + databaseName = resolvedPath; + } + + await ioManager.DeleteFile(databaseName, cancellationToken).ConfigureAwait(false); + return databaseName; + } + /// /// Prompts the user to create a /// @@ -163,7 +288,17 @@ namespace Tgstation.Server.Host.Core var databaseConfiguration = new DatabaseConfiguration(); do { - await console.WriteAsync(String.Format(CultureInfo.InvariantCulture, "Please enter one of {0}, {1}, or {2}: ", DatabaseType.MariaDB, DatabaseType.SqlServer, DatabaseType.MySql), false, cancellationToken).ConfigureAwait(false); + await console.WriteAsync( + String.Format( + CultureInfo.InvariantCulture, + "Please enter one of {0}, {1}, {2}, or {3}: ", + DatabaseType.Sqlite, + DatabaseType.MariaDB, + DatabaseType.MySql, + DatabaseType.SqlServer), + false, + cancellationToken) + .ConfigureAwait(false); var databaseTypeString = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); if (Enum.TryParse(databaseTypeString, out var databaseType)) { @@ -177,52 +312,68 @@ namespace Tgstation.Server.Host.Core string serverAddress; uint? mySQLServerPort = null; - do - { - await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Enter the server's address and port [: or ] (blank for local): ", false, cancellationToken).ConfigureAwait(false); - serverAddress = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); - if (String.IsNullOrWhiteSpace(serverAddress)) + + bool isSqliteDB = databaseConfiguration.DatabaseType == DatabaseType.Sqlite; + if (isSqliteDB) + serverAddress = null; + else + do { - serverAddress = null; - break; - } - else if (databaseConfiguration.DatabaseType != DatabaseType.SqlServer) - { - var m = Regex.Match(serverAddress, @"^(?.+):(?.+)$"); - if (m.Success) + await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("Enter the server's address and port [: or ] (blank for local): ", false, cancellationToken).ConfigureAwait(false); + serverAddress = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); + if (String.IsNullOrWhiteSpace(serverAddress)) { - serverAddress = m.Groups["server"].Value; - if (uint.TryParse(m.Groups["port"].Value, out uint port)) + serverAddress = null; + break; + } + else if (databaseConfiguration.DatabaseType == DatabaseType.SqlServer) + { + var m = Regex.Match(serverAddress, @"^(?.+):(?.+)$"); + if (m.Success) { - mySQLServerPort = port; - break; - } - else - { - await console.WriteAsync($@"Failed to parse port ""{m.Groups["port"].Value}"", please try again.", true, cancellationToken).ConfigureAwait(false); + serverAddress = m.Groups["server"].Value; + if (uint.TryParse(m.Groups["port"].Value, out uint port)) + { + mySQLServerPort = port; + break; + } + else + { + await console.WriteAsync($@"Failed to parse port ""{m.Groups["port"].Value}"", please try again.", true, cancellationToken).ConfigureAwait(false); + } } + else break; } else break; } - else break; - } - while (true); + while (true); await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Enter the database name (Can be from previous installation. Otherwise, should not exist): ", false, cancellationToken).ConfigureAwait(false); + await console.WriteAsync($"Enter the database {(isSqliteDB ? "file path" : "name")} (Can be from previous installation. Otherwise, should not exist): ", false, cancellationToken).ConfigureAwait(false); + string databaseName; + bool dbExists = false; do { databaseName = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); if (!String.IsNullOrWhiteSpace(databaseName)) + { + dbExists = isSqliteDB + ? await ioManager.FileExists(databaseName, cancellationToken).ConfigureAwait(false) + : await PromptYesNo("Does this database already exist? (y/n): ", cancellationToken).ConfigureAwait(false); + + if (!dbExists && isSqliteDB) + databaseName = await ValidateNonExistantSqliteDBName(databaseName, cancellationToken).ConfigureAwait(false); + } + + if (String.IsNullOrWhiteSpace(databaseName)) + await console.WriteAsync("Invalid database name!", true, cancellationToken).ConfigureAwait(false); + else break; - await console.WriteAsync("Invalid database name!", true, cancellationToken).ConfigureAwait(false); } while (true); - var dbExists = await PromptYesNo("Does this database already exist? (y/n): ", cancellationToken).ConfigureAwait(false); - bool useWinAuth; if (databaseConfiguration.DatabaseType == DatabaseType.SqlServer && platformIdentifier.IsWindows) useWinAuth = await PromptYesNo("Use Windows Authentication? (y/n): ", cancellationToken).ConfigureAwait(false); @@ -233,27 +384,28 @@ namespace Tgstation.Server.Host.Core string username = null; string password = null; - if (!useWinAuth) - { - await console.WriteAsync("Enter username: ", false, cancellationToken).ConfigureAwait(false); - username = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Enter password: ", false, cancellationToken).ConfigureAwait(false); - password = await console.ReadLineAsync(true, cancellationToken).ConfigureAwait(false); - } - else - { - await console.WriteAsync("IMPORTANT: If using the service runner, ensure this computer's LocalSystem account has CREATE DATABASE permissions on the target server!", true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("The account it uses in MSSQL is usually \"NT AUTHORITY\\SYSTEM\" and the role it needs is usually \"dbcreator\".", true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("We'll run a sanity test here, but it won't be indicative of the service's permissions if that is the case", true, cancellationToken).ConfigureAwait(false); - } + if (!isSqliteDB) + if (!useWinAuth) + { + await console.WriteAsync("Enter username: ", false, cancellationToken).ConfigureAwait(false); + username = await console.ReadLineAsync(false, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("Enter password: ", false, cancellationToken).ConfigureAwait(false); + password = await console.ReadLineAsync(true, cancellationToken).ConfigureAwait(false); + } + else + { + await console.WriteAsync("IMPORTANT: If using the service runner, ensure this computer's LocalSystem account has CREATE DATABASE permissions on the target server!", true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("The account it uses in MSSQL is usually \"NT AUTHORITY\\SYSTEM\" and the role it needs is usually \"dbcreator\".", true, cancellationToken).ConfigureAwait(false); + await console.WriteAsync("We'll run a sanity test here, but it won't be indicative of the service's permissions if that is the case", true, cancellationToken).ConfigureAwait(false); + } await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); DbConnection testConnection; - void CreateTestConnection(string connectionString) - { - testConnection = dbConnectionFactory.CreateConnection(connectionString, databaseConfiguration.DatabaseType); - } + void CreateTestConnection(string connectionString) => + testConnection = dbConnectionFactory.CreateConnection( + connectionString, + databaseConfiguration.DatabaseType); if (databaseConfiguration.DatabaseType == DatabaseType.SqlServer) { @@ -262,6 +414,7 @@ namespace Tgstation.Server.Host.Core ApplicationName = application.VersionPrefix, DataSource = serverAddress ?? "(local)" }; + if (useWinAuth) csb.IntegratedSecurity = true; else @@ -274,8 +427,20 @@ namespace Tgstation.Server.Host.Core csb.InitialCatalog = databaseName; databaseConfiguration.ConnectionString = csb.ConnectionString; } + else if(databaseConfiguration.DatabaseType == DatabaseType.Sqlite) + { + var csb = new SqliteConnectionStringBuilder + { + DataSource = databaseName, + Mode = dbExists ? SqliteOpenMode.ReadOnly : SqliteOpenMode.ReadWriteCreate + }; + + CreateTestConnection(csb.ConnectionString); + databaseConfiguration.ConnectionString = csb.ConnectionString; + } else { + // MySQL/MariaDB var csb = new MySqlConnectionStringBuilder { Server = serverAddress ?? "127.0.0.1", @@ -293,63 +458,7 @@ namespace Tgstation.Server.Host.Core try { - using (testConnection) - { - await console.WriteAsync("Testing connection...", true, cancellationToken).ConfigureAwait(false); - await testConnection.OpenAsync(cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Connection successful!", true, cancellationToken).ConfigureAwait(false); - - if (databaseConfiguration.DatabaseType != DatabaseType.SqlServer) - { - await console.WriteAsync("Checking MySQL/MariaDB version...", true, cancellationToken).ConfigureAwait(false); - using (var command = testConnection.CreateCommand()) - { - command.CommandText = "SELECT VERSION()"; - var fullVersion = (string)await command.ExecuteScalarAsync(cancellationToken).ConfigureAwait(false); - await console.WriteAsync(String.Format(CultureInfo.InvariantCulture, "Found {0}", fullVersion), true, cancellationToken).ConfigureAwait(false); - var splits = fullVersion.Split('-'); - databaseConfiguration.MySqlServerVersion = splits[0]; - } - } - - if (!dbExists) - { - await console.WriteAsync("Testing create DB permission...", true, cancellationToken).ConfigureAwait(false); - using (var command = testConnection.CreateCommand()) - { - // I really don't care about user sanitization here, they want to fuck their own DB? so be it -#pragma warning disable CA2100 // Review SQL queries for security vulnerabilities - command.CommandText = $"CREATE DATABASE {databaseName}"; -#pragma warning restore CA2100 // Review SQL queries for security vulnerabilities - await command.ExecuteNonQueryAsync(cancellationToken).ConfigureAwait(false); - } - - await console.WriteAsync("Success!", true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Dropping test database...", true, cancellationToken).ConfigureAwait(false); - using (var command = testConnection.CreateCommand()) - { -#pragma warning disable CA2100 // Review SQL queries for security vulnerabilities - command.CommandText = $"DROP DATABASE {databaseName}"; -#pragma warning restore CA2100 // Review SQL queries for security vulnerabilities - try - { - await command.ExecuteNonQueryAsync(cancellationToken).ConfigureAwait(false); - } - catch (OperationCanceledException) - { - throw; - } - catch (Exception e) - { - await console.WriteAsync(e.Message, true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync(null, true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("This should be okay, but you may want to manually drop the database before continuing!", true, cancellationToken).ConfigureAwait(false); - await console.WriteAsync("Press any key to continue...", true, cancellationToken).ConfigureAwait(false); - await console.PressAnyKeyAsync(cancellationToken).ConfigureAwait(false); - } - } - } - } + await TestDatabaseConnection(testConnection, databaseConfiguration, databaseName, dbExists, cancellationToken).ConfigureAwait(false); return databaseConfiguration; }