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