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;
}