diff --git a/OpenSSH_GUI.Core/Extensions/SshConfigFilesExtension.cs b/OpenSSH_GUI.Core/Extensions/SshConfigFilesExtension.cs index 1e87942..958e9f8 100644 --- a/OpenSSH_GUI.Core/Extensions/SshConfigFilesExtension.cs +++ b/OpenSSH_GUI.Core/Extensions/SshConfigFilesExtension.cs @@ -35,7 +35,16 @@ public static void ValidateDirectories(ILogger? logger = null) try { if (!Directory.Exists(GetRootSshPath())) Directory.CreateDirectory(GetRootSshPath()); - if (!Directory.Exists(GetBaseSshPath())) Directory.CreateDirectory(GetBaseSshPath()); + if (!Directory.Exists(GetBaseSshPath())) + { + // OpenSSH expects the user's .ssh directory to be accessible by the owner only. + if (OperatingSystem.IsWindows()) + Directory.CreateDirectory(GetBaseSshPath()); + else + Directory.CreateDirectory( + GetBaseSshPath(), + UnixFileMode.UserRead | UnixFileMode.UserWrite | UnixFileMode.UserExecute); + } } catch (Exception e) { diff --git a/OpenSSH_GUI.Core/Lib/Misc/ServerConnection.cs b/OpenSSH_GUI.Core/Lib/Misc/ServerConnection.cs index 3af05d0..0749309 100644 --- a/OpenSSH_GUI.Core/Lib/Misc/ServerConnection.cs +++ b/OpenSSH_GUI.Core/Lib/Misc/ServerConnection.cs @@ -1,6 +1,7 @@ using System.Reactive.Disposables; using System.Reactive.Disposables.Fluent; using System.Reactive.Linq; +using System.Text; using OpenSSH_GUI.Core.Enums; using OpenSSH_GUI.Core.Extensions; using OpenSSH_GUI.Core.Lib.AuthorizedKeys; @@ -13,6 +14,13 @@ namespace OpenSSH_GUI.Core.Lib.Misc; public sealed partial class ServerConnection : ReactiveObject, IDisposable { + private const string RemoteSshDirectory = ".ssh"; + + // SftpClient.ChangePermissions expects the octal digits written as a decimal number (e.g. 700 => rwx------). + private const short RemoteSshDirectoryMode = 700; + + private const short RemoteSshFileMode = 600; + private readonly CompositeDisposable _disposables = new(); [ObservableAsProperty(ReadOnly = true)] @@ -21,44 +29,30 @@ public sealed partial class ServerConnection : ReactiveObject, IDisposable [Reactive(SetModifier = AccessModifier.Private)] private DateTime _connectionTime = DateTime.Now; - [ObservableAsProperty(ReadOnly = true)] - private string _createEmptyFileCommand = string.Empty; - [Reactive(SetModifier = AccessModifier.Private)] private bool _isConnected; [ObservableAsProperty(ReadOnly = true)] private string _lineSeparator = string.Empty; - [ObservableAsProperty(ReadOnly = true)] - private string _readContentsCommand = string.Empty; - [Reactive(SetModifier = AccessModifier.Private)] private PlatformID _serverOs = PlatformID.Other; private ServerConnection(ConnectionCredentials? credentials = null) { ConnectionCredentials = credentials ?? ConnectionCredentials.Empty; - ClientConnection = new SshClient(ConnectionCredentials.GetConnectionInfo()) + var connectionInfo = ConnectionCredentials.GetConnectionInfo(); + ClientConnection = new SshClient(connectionInfo) { KeepAliveInterval = TimeSpan.FromSeconds(10) }; + FileTransferConnection = new SftpClient(connectionInfo); _connectionStringHelper = this.WhenAnyValue(obj => obj.IsConnected) .Select(c => c ? $"{ConnectionCredentials.Username}@{ConnectionCredentials.Hostname}" : string.Empty) .ToProperty(this, obj => obj.ConnectionString) .DisposeWith(_disposables); - _readContentsCommandHelper = this.WhenAnyValue(obj => obj.ServerOs) - .Select(c => c == PlatformID.Win32NT ? "type" : "cat") - .ToProperty(this, obj => obj.ReadContentsCommand) - .DisposeWith(_disposables); - - _createEmptyFileCommandHelper = this.WhenAnyValue(obj => obj.ServerOs) - .Select(c => c == PlatformID.Win32NT ? "echo. >" : "touch") - .ToProperty(this, obj => obj.CreateEmptyFileCommand) - .DisposeWith(_disposables); - _lineSeparatorHelper = this.WhenAnyValue(obj => obj.ServerOs) .Select(e => e.GetLineSeparator()) .ToProperty(this, obj => obj.LineSeparator) @@ -79,18 +73,34 @@ private SshClient ClientConnection init => this.RaiseAndSetIfChanged(ref field, value); } + /// + /// SFTP channel used for all remote file access. Transferring file contents over SFTP avoids + /// building shell command lines from file paths or file contents. + /// + private SftpClient FileTransferConnection + { + get; + init => this.RaiseAndSetIfChanged(ref field, value); + } + /// - public void Dispose() { _disposables.Dispose(); } + public void Dispose() + { + _disposables.Dispose(); + FileTransferConnection.Dispose(); + ClientConnection.Dispose(); + } public static ServerConnection WithCredentials(ConnectionCredentials credentials) => new(credentials); public async ValueTask ConnectToServerAsync(CancellationToken token = default) { await ClientConnection.ConnectAsync(token); - IsConnected = ClientConnection.IsConnected; + if (ClientConnection.IsConnected) + await FileTransferConnection.ConnectAsync(token); + IsConnected = ClientConnection.IsConnected && FileTransferConnection.IsConnected; if (!IsConnected) return ServerOs != PlatformID.Other && IsConnected; ServerOs = await GetServerOsAsync(token); - await CheckForFilesAndCreateThemIfTheyNotExistAsync(token); ConnectionTime = DateTime.Now; return ServerOs != PlatformID.Other && IsConnected; } @@ -99,6 +109,7 @@ public ValueTask DisconnectFromServerAsync(CancellationToken token = defau { try { + FileTransferConnection.Disconnect(); ClientConnection.Disconnect(); IsConnected = ClientConnection.IsConnected; return ValueTask.FromResult(true); @@ -113,12 +124,8 @@ public async ValueTask GetKnownHostsFromServerAsync(Cancellation { if (!IsConnected) throw new InvalidOperationException("No connection to get known hosts from"); - var path = await ResolveRemoteEnvVariablesAsync( - SshConfigFiles.Known_Hosts.GetPathOfFile(false, ServerOs), - token); - using var command = ClientConnection.CreateCommand($"{ReadContentsCommand} {path}"); - await command.ExecuteAsync(token); - return await KnownHostsFile.InitializeAsync(command.OutputStream, true, false, token); + var content = await ReadRemoteFileAsync(SshConfigFiles.Known_Hosts, token); + return await KnownHostsFile.InitializeAsync(content, true, true, token); } public async ValueTask WriteKnownHostsToServerAsync(KnownHostsFile knownHostsFile, @@ -127,12 +134,9 @@ public async ValueTask WriteKnownHostsToServerAsync(KnownHostsFile knownHo if (!knownHostsFile.KnownHosts.Any(e => e.ChangesMade)) return false; if (!IsConnected) return false; - var path = await ResolveRemoteEnvVariablesAsync( - SshConfigFiles.Known_Hosts.GetPathOfFile(false, ServerOs), token); var content = await knownHostsFile.GetUpdatedContentsAsync(ServerOs); - using var command = ClientConnection.CreateCommand(BuildRemoteWriteCommand(ServerOs, content, path)); - await command.ExecuteAsync(token); - return command.ExitStatus == 0; + await WriteRemoteFileAsync(SshConfigFiles.Known_Hosts, content, token); + return true; } public async ValueTask GetAuthorizedKeysFromServerAsync(CancellationToken token = default) @@ -140,11 +144,8 @@ public async ValueTask GetAuthorizedKeysFromServerAsync(Canc if (!IsConnected) throw new InvalidOperationException("No connection to get authorized keys from"); - var path = await ResolveRemoteEnvVariablesAsync( - SshConfigFiles.Authorized_Keys.GetPathOfFile(false, ServerOs), token); - using var command = ClientConnection.CreateCommand($"{ReadContentsCommand} {path}"); - await command.ExecuteAsync(token); - return await AuthorizedKeysFile.ParseAsync(command.OutputStream, token); + await using var content = await ReadRemoteFileAsync(SshConfigFiles.Authorized_Keys, token); + return await AuthorizedKeysFile.ParseAsync(content, token); } public async ValueTask WriteAuthorizedKeysChangesToServerAsync(AuthorizedKeysFile authorizedKeysFile, @@ -153,62 +154,57 @@ public async ValueTask WriteAuthorizedKeysChangesToServerAsync(AuthorizedK if (!authorizedKeysFile.ChangesMade) return false; if (!IsConnected) return false; - var path = await ResolveRemoteEnvVariablesAsync( - SshConfigFiles.Authorized_Keys.GetPathOfFile(false, ServerOs), token); var content = authorizedKeysFile.ExportFileContent(ServerOs); - using var command = ClientConnection.CreateCommand(BuildRemoteWriteCommand(ServerOs, content, path)); - await command.ExecuteAsync(token); - return command.ExitStatus == 0; + await WriteRemoteFileAsync(SshConfigFiles.Authorized_Keys, content, token); + return true; } - private async ValueTask ResolveRemoteEnvVariablesAsync(string originalPath, - CancellationToken token = default) + /// + /// Returns the SFTP path of the given file inside the remote user's SSH directory. + /// The path is relative to the SFTP working directory, which is the user's home directory + /// on both OpenSSH for Unix and OpenSSH for Windows. + /// + private static string GetRemotePath(SshConfigFiles file) => + $"{RemoteSshDirectory}/{Enum.GetName(file)!.ToLowerInvariant()}"; + + /// + /// Reads the given remote file via SFTP. A missing file yields an empty stream. + /// + private async ValueTask ReadRemoteFileAsync(SshConfigFiles file, CancellationToken token) { - if (!IsConnected) return originalPath; - var parts = originalPath.Split('%', StringSplitOptions.RemoveEmptyEntries); - var result = string.Empty; - foreach (var part in parts) - if (part.Contains('\\') || part.Contains('/')) - { - result += part.Trim(); - } - else - { - var cmdText = ServerOs is PlatformID.Unix or PlatformID.MacOSX - ? $"echo ${part}" - : $"echo %{part}%"; - using var command = ClientConnection.CreateCommand(cmdText); - await command.ExecuteAsync(token); - result += command.Result.Trim(); - } - - return result; + var path = GetRemotePath(file); + var content = new MemoryStream(); + if (await FileTransferConnection.ExistsAsync(path, token)) + await FileTransferConnection.DownloadFileAsync(path, content, token); + content.Seek(0, SeekOrigin.Begin); + return content; } - private async ValueTask CheckForFilesAndCreateThemIfTheyNotExistAsync(CancellationToken token = default) + /// + /// Replaces the contents of the given remote file via SFTP. Missing files and the + /// SSH directory are created with owner-only permissions on Unix hosts. + /// + private async ValueTask WriteRemoteFileAsync(SshConfigFiles file, string content, CancellationToken token) { - if (!ClientConnection.IsConnected) return; - - var authKeyPath = SshConfigFiles.Authorized_Keys.GetPathOfFile(false); - var knownHostPath = SshConfigFiles.Known_Hosts.GetPathOfFile(false); - - using var authorizedKeysFileCheck = ClientConnection.CreateCommand($"{ReadContentsCommand} {authKeyPath}"); - await authorizedKeysFileCheck.ExecuteAsync(token); - - using var knownHostsFileCheck = ClientConnection.CreateCommand($"{ReadContentsCommand} {knownHostPath}"); - await knownHostsFileCheck.ExecuteAsync(token); + var path = GetRemotePath(file); + var isUnix = ServerOs is PlatformID.Unix or PlatformID.MacOSX; - if (authorizedKeysFileCheck.ExitStatus != 0) + if (!await FileTransferConnection.ExistsAsync(RemoteSshDirectory, token)) { - using var createAuthCmd = ClientConnection.CreateCommand($"{CreateEmptyFileCommand} {authKeyPath}"); - await createAuthCmd.ExecuteAsync(token); + await FileTransferConnection.CreateDirectoryAsync(RemoteSshDirectory, token); + if (isUnix) FileTransferConnection.ChangePermissions(RemoteSshDirectory, RemoteSshDirectoryMode); } - if (knownHostsFileCheck.ExitStatus != 0) + var isNewFile = !await FileTransferConnection.ExistsAsync(path, token); + + await using (var remoteFile = + await FileTransferConnection.OpenAsync(path, FileMode.Create, FileAccess.Write, token)) { - using var createKnownCmd = ClientConnection.CreateCommand($"{CreateEmptyFileCommand} {knownHostPath}"); - await createKnownCmd.ExecuteAsync(token); + var bytes = new UTF8Encoding(false).GetBytes(content); + await remoteFile.WriteAsync(bytes, token); } + + if (isNewFile && isUnix) FileTransferConnection.ChangePermissions(path, RemoteSshFileMode); } private async ValueTask GetServerOsAsync(CancellationToken token = default) @@ -228,53 +224,4 @@ private async ValueTask GetServerOsAsync(CancellationToken token = d return PlatformID.Other; } - - /// - /// Builds a platform-appropriate shell command to write the given content to a file on the remote host. - /// - /// The of the remote host. - /// The content to write into the file. - /// The full remote path of the target file. - /// If true, appends to the file instead of overwriting it. - /// A shell command string ready to be executed on the remote host. - /// - /// Thrown when no write command can be constructed for the given . - /// - private static string BuildRemoteWriteCommand(PlatformID platformId, string content, string filePath, - bool append = false) - { - var redirectOperator = append ? ">>" : ">"; - - return platformId is PlatformID.Unix or PlatformID.MacOSX - ? BuildUnixCommand(content, filePath, redirectOperator) - : BuildWindowsCommand(content, filePath, redirectOperator); - } - - /// - /// Builds a Unix shell write command using printf for reliable, escape-safe output. - /// - /// The content to write. - /// The target file path on the remote host. - /// Shell redirect operator (> or >>). - /// A Unix shell command string. - private static string BuildUnixCommand(string content, string filePath, string redirectOperator) - { - var escaped = content.Replace("'", "'\\''"); - return $"printf '%s' '{escaped}' {redirectOperator} '{filePath}'"; - } - - /// - /// Builds a Windows shell write command using PowerShell's Set-Content or Add-Content - /// for reliable Unicode-safe file writing. - /// - /// The content to write. - /// The target file path on the remote host. - /// Shell redirect operator (> or >>), used to determine append mode. - /// A PowerShell command string. - private static string BuildWindowsCommand(string content, string filePath, string redirectOperator) - { - var escaped = content.Replace("'", "''"); - var cmdlet = redirectOperator == ">>" ? "Add-Content" : "Set-Content"; - return $"powershell -Command \"{cmdlet} -Path '{filePath}' -Value '{escaped}' -NoNewline -Encoding UTF8\""; - } -} \ No newline at end of file +} diff --git a/OpenSSH_GUI.Core/Services/KeyFileBackupService.cs b/OpenSSH_GUI.Core/Services/KeyFileBackupService.cs index d1485af..72e209c 100644 --- a/OpenSSH_GUI.Core/Services/KeyFileBackupService.cs +++ b/OpenSSH_GUI.Core/Services/KeyFileBackupService.cs @@ -19,6 +19,9 @@ public sealed class KeyFileBackupService : IKeyFileBackupService, IDisposable { private const string BackupFileExtension = "bak"; + private const UnixFileMode OwnerOnlyDirectoryMode = + UnixFileMode.UserRead | UnixFileMode.UserWrite | UnixFileMode.UserExecute; + private static readonly string BackupDirectory = Path.Combine(SshConfigFilesExtension.GetBaseSshPath(), AppDomain.CurrentDomain.FriendlyName); @@ -76,8 +79,16 @@ public void BeginOperationLog() { if (_operationLogger is not null) return; - if (!Directory.Exists(BackupDirectory)) + // The directory holds copies of private keys - restrict it to the owner. + if (OperatingSystem.IsWindows()) + { Directory.CreateDirectory(BackupDirectory); + } + else + { + Directory.CreateDirectory(BackupDirectory, OwnerOnlyDirectoryMode); + File.SetUnixFileMode(BackupDirectory, OwnerOnlyDirectoryMode); + } var operationLogFile = Path.Combine(BackupDirectory, Path.ChangeExtension("operation_log", "log")); _loggerFactory = new SerilogLoggerFactory( diff --git a/OpenSSH_GUI.Core/Services/KeyFileWriterService.cs b/OpenSSH_GUI.Core/Services/KeyFileWriterService.cs index cd088e4..6f5d3a9 100644 --- a/OpenSSH_GUI.Core/Services/KeyFileWriterService.cs +++ b/OpenSSH_GUI.Core/Services/KeyFileWriterService.cs @@ -15,6 +15,8 @@ namespace OpenSSH_GUI.Core.Services; /// public class KeyFileWriterService(ILogger logger) : IKeyFileWriterService { + private const UnixFileMode OwnerOnlyFileMode = UnixFileMode.UserRead | UnixFileMode.UserWrite; + /// public async ValueTask WriteToFile(string filePath, string content, bool overwrite = false, Encoding? encoding = null) @@ -36,22 +38,29 @@ public async ValueTask WriteToFile(string filePath, string content, throw new IOException("File already exists"); } + // Create truncates an existing file, CreateNew closes the race between the existence check and the open. var options = new FileStreamOptions { BufferSize = 0, - Access = FileAccess.ReadWrite, - Mode = FileMode.OpenOrCreate, - Share = FileShare.ReadWrite + Access = FileAccess.Write, + Mode = overwrite ? FileMode.Create : FileMode.CreateNew, + Share = FileShare.None }; if (!OperatingSystem.IsWindows()) { - options.UnixCreateMode = UnixFileMode.UserRead | UnixFileMode.UserWrite; + options.UnixCreateMode = OwnerOnlyFileMode; } await using var fileStream = fileInfo.Open(options); logger.LogDebug("Opened file {filePath}", filePath); + if (!OperatingSystem.IsWindows()) + { + // UnixCreateMode only applies to newly created files - enforce 0600 on overwritten files as well. + File.SetUnixFileMode(fileStream.SafeFileHandle, OwnerOnlyFileMode); + } + byte[]? rented = null; var maxByteCount = encoding.GetMaxByteCount(content.Length); var buffer = maxByteCount <= 256 diff --git a/OpenSSH_GUI.Tests/Core/Services/KeyFileWriterServiceTests.cs b/OpenSSH_GUI.Tests/Core/Services/KeyFileWriterServiceTests.cs new file mode 100644 index 0000000..4988f1e --- /dev/null +++ b/OpenSSH_GUI.Tests/Core/Services/KeyFileWriterServiceTests.cs @@ -0,0 +1,66 @@ +using Microsoft.Extensions.Logging.Abstractions; +using OpenSSH_GUI.Core.Services; +using Shouldly; +using Xunit; + +namespace OpenSSH_GUI.Tests.Core.Services; + +public sealed class KeyFileWriterServiceTests : IDisposable +{ + private readonly string _directory = Directory.CreateTempSubdirectory("keyfilewriter").FullName; + private readonly KeyFileWriterService _service = new(NullLogger.Instance); + + public void Dispose() { Directory.Delete(_directory, true); } + + [Fact] + public async Task WriteToFile_Overwrite_TruncatesLongerExistingContent() + { + var path = Path.Combine(_directory, "id_test"); + await File.WriteAllTextAsync(path, new string('x', 4096), TestContext.Current.CancellationToken); + + await _service.WriteToFile(path, "short", true); + + (await File.ReadAllTextAsync(path, TestContext.Current.CancellationToken)).ShouldBe("short"); + } + + [Fact] + public async Task WriteToFile_WithoutOverwrite_ThrowsAndKeepsExistingContent() + { + var path = Path.Combine(_directory, "id_test"); + await File.WriteAllTextAsync(path, "original", TestContext.Current.CancellationToken); + + await Should.ThrowAsync(async () => await _service.WriteToFile(path, "new")); + + (await File.ReadAllTextAsync(path, TestContext.Current.CancellationToken)).ShouldBe("original"); + } + + [Fact] + public async Task WriteToFile_Overwrite_RestrictsExistingFileToOwner() + { + Assert.SkipWhen(OperatingSystem.IsWindows(), "Unix file modes only"); + var path = Path.Combine(_directory, "id_test"); + await File.WriteAllTextAsync(path, "public", TestContext.Current.CancellationToken); +#pragma warning disable CA1416 + File.SetUnixFileMode( + path, + UnixFileMode.UserRead | UnixFileMode.UserWrite | UnixFileMode.GroupRead | UnixFileMode.OtherRead); + + await _service.WriteToFile(path, "private", true); + + File.GetUnixFileMode(path).ShouldBe(UnixFileMode.UserRead | UnixFileMode.UserWrite); +#pragma warning restore CA1416 + } + + [Fact] + public async Task WriteToFile_NewFile_IsCreatedOwnerOnly() + { + Assert.SkipWhen(OperatingSystem.IsWindows(), "Unix file modes only"); + var path = Path.Combine(_directory, "id_new"); + + await _service.WriteToFile(path, "private"); + +#pragma warning disable CA1416 + File.GetUnixFileMode(path).ShouldBe(UnixFileMode.UserRead | UnixFileMode.UserWrite); +#pragma warning restore CA1416 + } +}