Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion OpenSSH_GUI.Core/Extensions/SshConfigFilesExtension.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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)
{
Expand Down
205 changes: 76 additions & 129 deletions OpenSSH_GUI.Core/Lib/Misc/ServerConnection.cs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
锘縰sing 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;
Expand All @@ -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)]
Expand All @@ -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)
Expand All @@ -79,18 +73,34 @@ private SshClient ClientConnection
init => this.RaiseAndSetIfChanged(ref field, value);
}

/// <summary>
/// SFTP channel used for all remote file access. Transferring file contents over SFTP avoids
/// building shell command lines from file paths or file contents.
/// </summary>
private SftpClient FileTransferConnection
{
get;
init => this.RaiseAndSetIfChanged(ref field, value);
}

/// <inheritdoc />
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<bool> 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;
}
Expand All @@ -99,6 +109,7 @@ public ValueTask<bool> DisconnectFromServerAsync(CancellationToken token = defau
{
try
{
FileTransferConnection.Disconnect();
ClientConnection.Disconnect();
IsConnected = ClientConnection.IsConnected;
return ValueTask.FromResult(true);
Expand All @@ -113,12 +124,8 @@ public async ValueTask<KnownHostsFile> 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<bool> WriteKnownHostsToServerAsync(KnownHostsFile knownHostsFile,
Expand All @@ -127,24 +134,18 @@ public async ValueTask<bool> 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<AuthorizedKeysFile> GetAuthorizedKeysFromServerAsync(CancellationToken token = default)
{
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<bool> WriteAuthorizedKeysChangesToServerAsync(AuthorizedKeysFile authorizedKeysFile,
Expand All @@ -153,62 +154,57 @@ public async ValueTask<bool> 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<string> ResolveRemoteEnvVariablesAsync(string originalPath,
CancellationToken token = default)
/// <summary>
/// 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.
/// </summary>
private static string GetRemotePath(SshConfigFiles file) =>
$"{RemoteSshDirectory}/{Enum.GetName(file)!.ToLowerInvariant()}";

/// <summary>
/// Reads the given remote file via SFTP. A missing file yields an empty stream.
/// </summary>
private async ValueTask<Stream> 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)
/// <summary>
/// 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.
/// </summary>
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<PlatformID> GetServerOsAsync(CancellationToken token = default)
Expand All @@ -228,53 +224,4 @@ private async ValueTask<PlatformID> GetServerOsAsync(CancellationToken token = d

return PlatformID.Other;
}

/// <summary>
/// Builds a platform-appropriate shell command to write the given content to a file on the remote host.
/// </summary>
/// <param name="platformId">The <see cref="PlatformID" /> of the remote host.</param>
/// <param name="content">The content to write into the file.</param>
/// <param name="filePath">The full remote path of the target file.</param>
/// <param name="append">If <c>true</c>, appends to the file instead of overwriting it.</param>
/// <returns>A shell command string ready to be executed on the remote host.</returns>
/// <exception cref="PlatformNotSupportedException">
/// Thrown when no write command can be constructed for the given <paramref name="platformId" />.
/// </exception>
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);
}

/// <summary>
/// Builds a Unix shell write command using <c>printf</c> for reliable, escape-safe output.
/// </summary>
/// <param name="content">The content to write.</param>
/// <param name="filePath">The target file path on the remote host.</param>
/// <param name="redirectOperator">Shell redirect operator (<c>&gt;</c> or <c>&gt;&gt;</c>).</param>
/// <returns>A Unix shell command string.</returns>
private static string BuildUnixCommand(string content, string filePath, string redirectOperator)
{
var escaped = content.Replace("'", "'\\''");
return $"printf '%s' '{escaped}' {redirectOperator} '{filePath}'";
}

/// <summary>
/// Builds a Windows shell write command using PowerShell's <c>Set-Content</c> or <c>Add-Content</c>
/// for reliable Unicode-safe file writing.
/// </summary>
/// <param name="content">The content to write.</param>
/// <param name="filePath">The target file path on the remote host.</param>
/// <param name="redirectOperator">Shell redirect operator (<c>&gt;</c> or <c>&gt;&gt;</c>), used to determine append mode.</param>
/// <returns>A PowerShell command string.</returns>
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\"";
}
}
}
13 changes: 12 additions & 1 deletion OpenSSH_GUI.Core/Services/KeyFileBackupService.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand Down Expand Up @@ -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(
Expand Down
Loading
Loading