// Copyright (C) 2025, The Duplicati Team // https://duplicati.com, hello@duplicati.com // // Permission is hereby granted, free of charge, to any person obtaining a // copy of this software and associated documentation files (the "Software"), // to deal in the Software without restriction, including without limitation // the rights to use, copy, modify, merge, publish, distribute, sublicense, // and/or sell copies of the Software, and to permit persons to whom the // Software is furnished to do so, subject to the following conditions: // // The above copyright notice and this permission notice shall be included in // all copies or substantial portions of the Software. // // THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS // OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, // FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE // AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER // LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING // FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER // DEALINGS IN THE SOFTWARE. using System.Diagnostics; using System.Net; using System.Net.Http.Headers; using System.Security.Cryptography; using System.Text.Json; using Duplicati.Library.AutoUpdater; using Duplicati.Library.Logging; using Duplicati.Library.Utility; using Uri = System.Uri; namespace Duplicati.Library.RemoteControl; /// /// Support class for keeping a connection to a remote server /// public class KeepRemoteConnection : IDisposable { /// /// The protocol version to use /// private const int PROTOCOL_VERSION = 1; /// /// The log tag for messages from this class /// private static readonly string LogTag = Log.LogTagFromType(); /// /// The interval between reconnect attempts /// private static readonly TimeSpan ReconnectInterval = TimeSpan.FromSeconds(30); /// /// The interval between heartbeats /// private static readonly TimeSpan HeartbeatInterval = TimeSpan.FromSeconds(60); /// /// The time between reconnect attempts if no response is received /// private static readonly TimeSpan NoResponseTimeout = HeartbeatInterval * 2; /// /// The minimum time between certificate refreshes /// private static readonly TimeSpan MinimumCertificateRefreshInterval = TimeSpan.FromMinutes(5); /// /// The interval between certificate refreshes /// private static readonly TimeSpan CertificateRefreshInterval = TimeSpan.FromDays(7); /// /// The client key to use for signing messages /// private static readonly RSA ClientKey = RSA.Create(2048); /// /// The client ID to use for identifying the client /// private static readonly string ClientId = Guid.NewGuid().ToString(); /// /// The JSON options to use for deserialization /// internal static readonly JsonSerializerOptions JsonOptions = new JsonSerializerOptions { PropertyNamingPolicy = JsonNamingPolicy.CamelCase, PropertyNameCaseInsensitive = true }; /// /// The stats the connection can be in /// public enum ConnectionState { /// /// The connection is not established /// NotConnected, /// /// We received a welcome message /// WelcomeReceived, /// /// The connection is authenticated /// Authenticated, /// /// The connection is in an error state /// Error } /// /// The websocket client /// private readonly Websocket.Client.WebsocketClient _client; /// /// The cancellation token source /// private readonly CancellationTokenSource _cancellationTokenSource; /// /// The current state of the connection /// private ConnectionState _state = ConnectionState.NotConnected; /// /// The task that runs the connection /// private Task _runnerTask; /// /// The currently negotiated server certificate /// private MiniServerCertificate? _serverCertificate; /// /// The public key of the server /// private RSA? _serverPublicKey; /// /// The callback to call when connecting /// private readonly Func, Task>> _onConnect; /// /// The callback to call when rekeying /// private readonly Func _onReKey; /// /// The callback to call when a control message is received /// private readonly Func _onControl; /// /// The callback to call when a command message is received /// private readonly Func _onMessage; /// /// The current JWT token /// private string _token; /// /// The server URL /// private string _serverUrl; /// /// The certificate URL /// private string _certificateUrl; /// /// The server keys /// private IEnumerable _serverKeys; /// /// The last time a message was received /// private DateTimeOffset _lastMessageReceived = DateTimeOffset.MinValue; /// /// Creates a new connection to the remote server /// /// The url to use /// The JWT token to use /// The certificate url to use /// The server keys to use /// The token to cancel the connection /// The callback to call when connecting /// The callback to call when rekeying /// The callback to call when a control message is received /// The callback to call when a command message is received private KeepRemoteConnection( string serverUrl, string JWT, string certificateUrl, IEnumerable serverKeys, CancellationToken cancellationToken, Func, Task>> onConnect, Func onReKey, Func onControl, Func onMessage) { _serverUrl = serverUrl; _certificateUrl = certificateUrl; _token = JWT; _serverKeys = serverKeys; _cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); _onConnect = onConnect; _onReKey = onReKey; _onControl = onControl; _onMessage = onMessage; _client = new Websocket.Client.WebsocketClient(new Uri(serverUrl)); _runnerTask = RunMainLoop(); } /// /// Runs the inner loop of the connection /// private async Task RunMainLoop() { _client.ReconnectTimeout = NoResponseTimeout; _client.LostReconnectTimeout = ReconnectInterval; _client.IsReconnectionEnabled = true; // Set up the periodic refreshers using var reconnectHelper = new PeriodicRefresher( Timeout.InfiniteTimeSpan, ReconnectInterval, async token => { await _client.Start(); }, _cancellationTokenSource.Token); using var heartbeatHelper = new PeriodicRefresher( HeartbeatInterval, TimeSpan.FromSeconds(1), _ => { // Reconnect if we have disconnected if (!_client.IsRunning) { reconnectHelper.Signal(); } // If we do not get any response from the server, we should reconnect else if ((_state == ConnectionState.Authenticated || _state == ConnectionState.WelcomeReceived) && _lastMessageReceived.Add(NoResponseTimeout) < DateTimeOffset.Now) { _state = ConnectionState.Error; Log.WriteMessage(LogMessageType.Warning, LogTag, "WebsocketDisconnect", "No response from server"); _client.Stop(System.Net.WebSockets.WebSocketCloseStatus.NormalClosure, "No response"); } SendEnvelope(new EnvelopedMessage() { From = ClientId, To = "server", Type = "ping", ErrorMessage = null, Payload = null, MessageId = Guid.NewGuid().ToString() }); return Task.CompletedTask; }, _cancellationTokenSource.Token); using var certificateRfreshHelper = new PeriodicRefresher( CertificateRefreshInterval, MinimumCertificateRefreshInterval, RefreshCertificates, _cancellationTokenSource.Token); _client.DisconnectionHappened.Subscribe(info => { _state = ConnectionState.NotConnected; _serverCertificate = null; _serverPublicKey = null; reconnectHelper.Signal(); Log.WriteMessage(LogMessageType.Warning, LogTag, "WebsocketDisconnect", "Disconnected from the server"); }); _client.MessageReceived.Subscribe(async msg => { // Ignore messages if we are in an error state if (_state == ConnectionState.Error) return; _lastMessageReceived = DateTimeOffset.Now; if (_state == ConnectionState.NotConnected) Log.WriteMessage(LogMessageType.Verbose, LogTag, "WebsocketMessage", "Received message from server: {0}", msg); else // Encrypted messages are not logged, as the content has no meaning before being decrypted Log.WriteMessage(LogMessageType.Verbose, LogTag, "WebsocketMessage", "Received encrypted message from server"); try { if (string.IsNullOrWhiteSpace(msg.Text)) throw new ProtocolViolationException("Empty message"); if (_serverCertificate == null || _serverPublicKey == null || _state == ConnectionState.NotConnected) { // Should be safe from replay, as the response is encrypted with the server public key // So even a replay attack would not let the attacker know the client's token var welcomeEnvelope = EnvelopedMessage.ForceParse(msg.Text); if (welcomeEnvelope.GetMessageType() != MessageType.Welcome) throw new ProtocolViolationException("Expected welcome message"); if (string.IsNullOrWhiteSpace(welcomeEnvelope.Payload)) throw new ProtocolViolationException("No payload in welcome message"); var welcomeMessage = welcomeEnvelope.GetPayload() ?? throw new ProtocolViolationException("Invalid welcome message"); if (string.IsNullOrWhiteSpace(welcomeMessage.PublicKeyHash)) throw new ProtocolViolationException("No public key hash in welcome message"); _serverCertificate = _serverKeys.FirstOrDefault(x => x.PublicKeyHash == welcomeMessage.PublicKeyHash && x.Expiry > DateTimeOffset.Now); if (_serverCertificate == null) { certificateRfreshHelper.Signal(); throw new ProtocolViolationException("No valid server certificate"); } try { var tmp = RSA.Create(); tmp.ImportFromPem(_serverCertificate.PublicKey); _serverPublicKey = tmp; } catch { certificateRfreshHelper.Signal(); throw new ProtocolViolationException("Invalid server certificate"); } _state = ConnectionState.WelcomeReceived; // Prepare basic metadata and allow additional metadata to be added var metadata = await _onConnect(new Dictionary() { { "client-version", UpdaterManager.SelfVersion?.Version ?? "0.0.0.0" }, { "client-id", ClientId }, { "client-uptime", (DateTime.Now - Process.GetCurrentProcess().StartTime).ToString() }, { "machine-name", DataFolderManager.MachineName }, { "machine-id", DataFolderManager.MachineID }, { "install-id", DataFolderManager.InstallID }, { "machine-os", UpdaterManager.OperatingSystemName }, { "package-id", UpdaterManager.PackageTypeId }, { "update-channel", UpdaterManager.CurrentChannel.ToString() } }); SendEnvelope( welcomeEnvelope.RespondWith( new AuthMessage( _token, ClientKey.ExportRSAPublicKeyPem(), UpdaterManager.SelfVersion?.Version ?? "0.0.0.0", PROTOCOL_VERSION, metadata ), "auth" ), force: true); return; } if (_serverCertificate == null || _serverPublicKey == null || _serverCertificate.HasExpired()) { certificateRfreshHelper.Signal(); throw new ProtocolViolationException("No valid server certificate"); } var envelope = TransportHelper.ParseFromEncryptedMessage(msg.Text, ClientKey); if (_state == ConnectionState.WelcomeReceived) { if (envelope.GetMessageType() != MessageType.Auth) throw new ProtocolViolationException("Expected welcome message"); var authMessage = envelope.GetPayload(); if (!authMessage.Accepted ?? false) throw new ProtocolViolationException("Authentication failed"); _state = ConnectionState.Authenticated; if ((authMessage.WillReplaceToken ?? false) && authMessage.NewToken != null) { _token = authMessage.NewToken; await InvokeReKey(); } } else if (_state == ConnectionState.Authenticated) { Log.WriteVerboseMessage(LogTag, "WebsocketMessage", "Processing message of type {0}", envelope.GetMessageType()); switch (envelope.GetMessageType()) { case MessageType.Pong: break; case MessageType.Command: await _onMessage(new CommandMessage( envelope.GetPayload(), response => SendEnvelope(envelope.RespondWith(response)) )); break; case MessageType.Control: await _onControl(new ControlMessage( envelope.GetPayload(), response => SendEnvelope(envelope.RespondWith(response)) )); break; default: throw new ProtocolViolationException("Unexpected message"); } } else { throw new ProtocolViolationException("Unexpected message"); } } catch (Exception ex) { _state = ConnectionState.Error; Log.WriteMessage(LogMessageType.Error, LogTag, "WebsocketMessage", ex, "Failed to process message: {0}", msg); await _client.Stop(System.Net.WebSockets.WebSocketCloseStatus.NormalClosure, "Error"); reconnectHelper.Signal(); } }); // Start the connection reconnectHelper.Signal(); var t = await Task.WhenAny( heartbeatHelper.RunLoopAsync(), reconnectHelper.RunLoopAsync(), certificateRfreshHelper.RunLoopAsync() ); _cancellationTokenSource.Cancel(); // Re-throw any exceptions await t; } /// /// Helper method to invoke the rekey callback /// /// An awaitable task private Task InvokeReKey() => _onReKey(new ClaimedClientData(_token, _serverUrl, _certificateUrl, _serverKeys, null)); /// /// Creates a new connection to the remote server /// /// The url to use /// The JWT to use /// The certificate url to use /// The server keys to use /// The token to cancel the connection /// The callback to call when connecting /// The callback to call when rekeying /// The callback to call when a control message is received /// The callback to call when a command message is received /// public static Task Start( string serverUrl, string JWT, string certificateUrl, IEnumerable serverKeys, CancellationToken cancellationToken, Func, Task>> onConnect, Func onReKey, Func onControl, Func onMessage) => Task.Run(async () => { using var connection = new KeepRemoteConnection(serverUrl, JWT, certificateUrl, serverKeys, cancellationToken, onConnect, onReKey, onControl, onMessage); await connection._runnerTask; }); /// /// Gets the task representing the connection /// /// The task public Task Run() => _runnerTask; /// /// Stops the connection /// /// An awaitable task public Task Stop() { _cancellationTokenSource.Cancel(); return _runnerTask; } /// /// Sends an enveloped message to the remote server /// /// The envelope to send /// True if the message was sent private bool SendEnvelope(EnvelopedMessage envelope, bool force = true) { if ((_state != ConnectionState.Authenticated && !force) || _serverPublicKey == null) return false; _client.Send(TransportHelper.CreateEncryptedMessage(envelope with { From = ClientId }, _serverPublicKey)); return true; } /// /// Sends a new command to the server /// /// The message to send /// True if the message was sent public bool SendCommand(CommandRequestMessage message) { if (_state != ConnectionState.Authenticated || _serverPublicKey == null) return false; _client.Send(TransportHelper.CreateEncryptedMessage(new EnvelopedMessage() { From = ClientId, To = "server", Type = "command", MessageId = Guid.NewGuid().ToString(), Payload = JsonSerializer.Serialize(message, options: JsonOptions), ErrorMessage = null }, _serverPublicKey)); return true; } /// /// The current state of the connection /// public ConnectionState State => _state; /// /// Creates a new connection to the remote server /// /// The url to use /// The JWT token to use /// The certificate url to use /// The server keys to use /// The callback to call when connecting /// The callback to call when rekeying /// The callback to call when a control message is received /// The callback to call when a message is received /// The connection object public static KeepRemoteConnection CreateRemoteListener( string serverUrl, string JWT, string certificateUrl, IEnumerable serverKeys, CancellationToken cancellationToken, Func, Task>> onConnect, Func onReKey, Func onControl, Func onMessage) => new KeepRemoteConnection(serverUrl, JWT, certificateUrl, serverKeys, cancellationToken, onConnect, onReKey, onControl, onMessage); /// /// Requests a certificate refresh /// /// The cancellation token /// An awaitable task private async Task RefreshCertificates(CancellationToken cancelToken) { using var client = HttpClientHelper.CreateClient(); // We won't set infiniteTimeout and keep the default 100s timeout var response = await client.GetAsync(_certificateUrl, cancelToken); if (response.IsSuccessStatusCode) { await using var stream = await response.Content.ReadAsStreamAsync(cancelToken); var serverKeys = await JsonSerializer.DeserializeAsync>(stream, options: RegisterForRemote.JsonOptions, cancellationToken: cancelToken); if (serverKeys != null && serverKeys.Any()) { _serverKeys = serverKeys .Where(x => !x.HasExpired() && !string.IsNullOrWhiteSpace(x.PublicKeyHash) && !string.IsNullOrWhiteSpace(x.PublicKey)) .ToList(); await InvokeReKey(); } } } /// public void Dispose() { _cancellationTokenSource.Cancel(); _client.Dispose(); _cancellationTokenSource.Dispose(); } /// /// A wrapper for allowing external code to handle a command message /// public sealed class CommandMessage { /// /// The callback method that will receive the response /// private readonly Func _respondCommand; /// /// The command request message /// public CommandRequestMessage CommandRequestMessage { get; } /// /// Creates a new command message /// /// The command request message /// The callback method that will receive the response public CommandMessage(CommandRequestMessage commandRequestMessage, Func respondCommand) { CommandRequestMessage = commandRequestMessage; _respondCommand = respondCommand; } /// /// Responds to the command message /// /// The response to send /// True if the response was sent public bool Respond(CommandResponseMessage response) => _respondCommand(response); /// /// Handles the command message with a configured http client. /// The client must be configured with the correct base address and authorization headers. /// /// The pre-configured http client /// An awaitable task public async Task Handle(HttpClient client) { try { Log.WriteVerboseMessage(LogTag, "WebsocketCommand", "Handling command {0} {1}", CommandRequestMessage.Method, CommandRequestMessage.Path); var request = new HttpRequestMessage(new HttpMethod(CommandRequestMessage.Method), CommandRequestMessage.Path); if (!string.IsNullOrWhiteSpace(CommandRequestMessage.Body)) request.Content = new ByteArrayContent(Convert.FromBase64String(CommandRequestMessage.Body)); if (CommandRequestMessage.Headers != null) { foreach (var header in CommandRequestMessage.Headers) { if (header.Key == "Content-Type" && request.Content != null) request.Content.Headers.ContentType = new MediaTypeHeaderValue(header.Value); else request.Headers.Add(header.Key, header.Value); } } var response = await client.SendAsync(request); var responseBody = await response.Content.ReadAsByteArrayAsync(); var responseHeaders = response.Headers.ToDictionary(x => x.Key, x => x.Value.First()); Respond(new CommandResponseMessage((int)response.StatusCode, responseBody == null ? null : Convert.ToBase64String(responseBody), responseHeaders)); } catch (Exception ex) { Respond(new CommandResponseMessage(500, ex.Message, null)); } } } /// /// A wrapper for allowing external code to handle a control message /// public sealed class ControlMessage { /// /// The callback method that will receive the response /// private readonly Func _respondCommand; /// /// The command request message /// public ControlRequestMessage ControlRequestMessage { get; } /// /// Creates a new command message /// /// The command request message /// The callback method that will receive the response public ControlMessage(ControlRequestMessage controlRequestMessage, Func respondCommand) { ControlRequestMessage = controlRequestMessage; _respondCommand = respondCommand; } /// /// Responds to the command message /// /// The response to send /// True if the response was sent public bool Respond(ControlResponseMessage response) => _respondCommand(response); } }