More work on getting the protocol fixed

This commit is contained in:
Kenneth Skovhede
2024-09-23 23:18:36 +02:00
parent efab444d94
commit 3bafe74c0a
8 changed files with 200 additions and 43 deletions
+3 -3
View File
@@ -159,12 +159,12 @@ public static class Program
var settings = Settings.Load();
if (!string.IsNullOrWhiteSpace(keydata.JWT) && settings.JWT != keydata.JWT)
settings = settings with { JWT = keydata.JWT };
if (keydata.ServerCertificates.Any())
settings = settings with { ServerCertificates = MiniServerCertificate.MergeCertificates(keydata.ServerCertificates, settings.ServerCertificates) };
if (keydata.ServerCertificates != null && keydata.ServerCertificates.Any())
settings = settings with { ServerCertificates = keydata.ServerCertificates };
if (!string.IsNullOrWhiteSpace(keydata.LocalEncryptionKey) && settings.SettingsEncryptionKey != keydata.LocalEncryptionKey)
{
Log.WriteMessage(LogMessageType.Information, LogTag, "ReKey", "Changing the local settings encryption key");
// Log.WriteMessage(LogMessageType.Information, LogTag, "ReKey", "Changing the local settings encryption key");
// TODO: Implement changing the database encryption key
// FIXMEGlobal.Provider.GetRequiredService<Connection>().ChangeDbKey(keydata.LocalEncryptionKey);
// settings = settings with { SettingsEncryptionKey = keydata.LocalEncryptionKey };
@@ -110,6 +110,10 @@ internal sealed record EnvelopedMessage
/// </summary>
public string? Payload { get; init; }
/// <summary>
/// The public key hash
/// </summary>
public string? PublicKeyHash { get; init; }
/// <summary>
/// The signature of the payload
/// </summary>
public string? Signature { get; init; }
@@ -185,6 +189,21 @@ internal sealed record EnvelopedMessage
).Replace("-", "").ToLower();
}
/// <summary>
/// Computes the signature of the payload
/// </summary>
/// <param name="key">The private key to use</param>
/// <returns>The computed signature</returns>
public string? ComputePayloadSignature(RSA key)
{
if (key is null || Payload is null && MessageId is null)
return null;
return BitConverter.ToString(
key.SignData(Encoding.UTF8.GetBytes($"{Payload}::{MessageId}"), HashAlgorithmName.SHA256, RSASignaturePadding.Pss)
).Replace("-", "").ToLower();
}
/// <summary>
/// Validates the message signature
/// </summary>
@@ -203,10 +222,10 @@ internal sealed record EnvelopedMessage
/// <summary>
/// Creates the signature on the returned message
/// </summary>
/// <param name="pemPrivatekey">The private key to use</param>
/// <param name="key">The private key to use</param>
/// <returns>The signed message</returns>
public EnvelopedMessage WithSignature(string? pemPrivatekey)
=> this with { Signature = ComputePayloadSignature(pemPrivatekey) };
public EnvelopedMessage WithSignature(RSA key)
=> this with { Signature = ComputePayloadSignature(key) };
/// <summary>
/// Creates a new message with a payload
@@ -22,6 +22,7 @@
using System.Net;
using System.Security.Cryptography;
using System.Text;
using System.Text.Json;
using Duplicati.Library.Logging;
namespace Duplicati.Library.RemoteControl;
@@ -36,20 +37,32 @@ public class KeepRemoteConnection : IDisposable
/// </summary>
private static readonly string LogTag = Log.LogTagFromType<KeepRemoteConnection>();
/// <summary>
/// The interval between reconnect attempts
/// </summary>
private static readonly TimeSpan ReconnectInterval = TimeSpan.FromSeconds(30);
/// <summary>
/// The interval between heartbeats
/// </summary>
private static readonly TimeSpan HeartbeatInterval = TimeSpan.FromSeconds(5);
/// <summary>
/// The interval between certificate refreshes
/// </summary>
private static readonly TimeSpan CertificateRefreshInterval = TimeSpan.FromDays(7);
/// <summary>
/// The client key to use for signing messages
/// </summary>
private static readonly string? ClientKey = null; // TODO: Fill in
private static readonly RSA ClientKey = RSA.Create(2048);
/// <summary>
/// The client ID to use for identifying the client
/// </summary>
private static readonly string ClientId = AutoUpdater.UpdaterManager.MachineID;
private static readonly string ClientId = string.IsNullOrWhiteSpace(AutoUpdater.UpdaterManager.MachineID)
? Guid.NewGuid().ToString()
: AutoUpdater.UpdaterManager.MachineID;
/// <summary>
/// The stats the connection can be in
@@ -94,6 +107,43 @@ public class KeepRemoteConnection : IDisposable
/// The task that runs the connection
/// </summary>
private Task _runnerTask;
/// <summary>
/// The currently negotiated server certificate
/// </summary>
private MiniServerCertificate? _serverCertificate;
/// <summary>
/// The time the certificate was last refreshed
/// </summary>
private DateTime _lastCertificateRefresh = DateTime.UnixEpoch;
/// <summary>
/// Task for requesting certificate refresh
/// </summary>
private TaskCompletionSource<bool> _refreshCertificates = new TaskCompletionSource<bool>();
/// <summary>
/// The callback to call when rekeying
/// </summary>
private readonly Func<ClaimedClientData, Task> _onReKey;
/// <summary>
/// The callback to call when a message is received
/// </summary>
private readonly Func<CommandMessage, Task> _onMessage;
/// <summary>
/// The current JWT token
/// </summary>
private string _token;
/// <summary>
/// The server URL
/// </summary>
private string _serverUrl;
/// <summary>
/// The certificate URL
/// </summary>
private string _certificateUrl;
/// <summary>
/// The server keys
/// </summary>
private IEnumerable<MiniServerCertificate> _serverKeys;
/// <summary>
/// Creates a new connection to the remote server
@@ -104,8 +154,24 @@ public class KeepRemoteConnection : IDisposable
/// <param name="cancellationToken">The token to cancel the connection</param>
private KeepRemoteConnection(string serverUrl, string JWT, string certificateUrl, IEnumerable<MiniServerCertificate> serverKeys, CancellationToken cancellationToken, Func<ClaimedClientData, Task> onReKey, Func<CommandMessage, Task> onMessage)
{
_serverUrl = serverUrl;
_certificateUrl = certificateUrl;
_token = JWT;
_serverKeys = serverKeys;
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_onReKey = onReKey;
_onMessage = onMessage;
_client = new Websocket.Client.WebsocketClient(new Uri(serverUrl));
_client.ReconnectTimeout = TimeSpan.FromSeconds(30);
_runnerTask = RunMainLoop();
}
/// <summary>
/// Runs the inner loop of the connection
/// </summary>
private Task RunMainLoop()
{
_client.ReconnectTimeout = ReconnectInterval;
_client.ReconnectionHappened.Subscribe(info =>
{
@@ -115,8 +181,8 @@ public class KeepRemoteConnection : IDisposable
_client.DisconnectionHappened.Subscribe(info =>
{
// TODO: If disconnected due to certifiate error, we should try to get fresh certificates
_state = ConnectionState.NotConnected;
_serverCertificate = null;
Log.WriteMessage(LogMessageType.Warning, LogTag, "WebsocketDisconnect", "Disconnected from the server");
});
@@ -127,8 +193,28 @@ public class KeepRemoteConnection : IDisposable
try
{
var envelope = EnvelopedMessage.ForceParse(msg.Text);
var machineKey = serverKeys.FirstOrDefault(x => x.Identifier == envelope.From && x.Expiry > DateTimeOffset.Now)?.PublicKey;
envelope.ValidateSignature(machineKey);
if (_serverCertificate == null || _state == ConnectionState.Connected)
{
if (envelope.GetMessageType() != MessageType.Welcome)
throw new ProtocolViolationException("Expected welcome message");
if (string.IsNullOrWhiteSpace(envelope.PublicKeyHash))
throw new ProtocolViolationException("No public key hash in welcome message");
_serverCertificate = _serverKeys.FirstOrDefault(x => x.PublicKeyHash == envelope.PublicKeyHash && x.Expiry > DateTimeOffset.Now);
if (_serverCertificate == null)
{
_refreshCertificates.TrySetResult(true);
throw new ProtocolViolationException("No valid server certificate");
}
}
if (_serverCertificate == null)
{
_refreshCertificates.TrySetResult(true);
throw new ProtocolViolationException("No valid server certificate");
}
envelope.ValidateSignature(_serverCertificate?.PublicKey);
if (_state == ConnectionState.Connected)
{
@@ -140,7 +226,7 @@ public class KeepRemoteConnection : IDisposable
Log.WriteMessage(LogMessageType.Information, LogTag, "WebsocketAuthenticated", "Connected with the server");
_challenge = RandomNumberGenerator.GetHexString(64);
SendEnvelope(envelope.RespondWith(new AuthMessage(JWT, _challenge)));
SendEnvelope(envelope.RespondWith(new AuthMessage(_token, _challenge)));
}
else if (_state == ConnectionState.WelcomeReceived)
{
@@ -155,12 +241,15 @@ public class KeepRemoteConnection : IDisposable
throw new ProtocolViolationException("Invalid Json message");
using RSA rsa = RSA.Create();
rsa.ImportFromPem(machineKey);
rsa.ImportFromPem(_serverCertificate?.PublicKey);
if (!rsa.VerifyData(Encoding.UTF8.GetBytes(_challenge!), Convert.FromHexString(authMessage.SignedChallenge), HashAlgorithmName.SHA256, RSASignaturePadding.Pss))
throw new EnvelopeJsonParsingException("Invalid Json message");
if ((authMessage.WillReplaceToken ?? false) && authMessage.NewToken != null)
await onReKey(new ClaimedClientData(authMessage.NewToken, serverUrl, certificateUrl, serverKeys, null));
{
_token = authMessage.NewToken;
await InvokeReKey();
}
_state = ConnectionState.Authenticated;
}
@@ -172,7 +261,7 @@ public class KeepRemoteConnection : IDisposable
break;
case MessageType.Command:
await onMessage(new CommandMessage(envelope.GetPayload<CommandRequestMessage>(), response =>
await _onMessage(new CommandMessage(envelope.GetPayload<CommandRequestMessage>(), response =>
{
SendEnvelope(envelope.RespondWith(response));
return true;
@@ -198,13 +287,16 @@ public class KeepRemoteConnection : IDisposable
}
});
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_runnerTask = Task.WhenAny(
return Task.WhenAny(
_client.Start(),
RunHeartbeatLoop()
RunHeartbeatLoop(),
RunCertificateRefreshLoop()
);
}
private Task InvokeReKey()
=> _onReKey(new ClaimedClientData(_token, _serverUrl, _certificateUrl, _serverKeys, null));
/// <summary>
/// Creates a new connection to the remote server
/// </summary>
@@ -317,6 +409,40 @@ public class KeepRemoteConnection : IDisposable
}
}
/// <summary>
/// Runs a loop that refreshes the server certificates
/// </summary>
/// <returns>An awaitable task</returns>
private async Task RunCertificateRefreshLoop()
{
while (!_cancellationTokenSource.Token.IsCancellationRequested)
{
var t = await Task.WhenAny(_refreshCertificates.Task, Task.Delay(CertificateRefreshInterval, _cancellationTokenSource.Token));
if (_cancellationTokenSource.Token.IsCancellationRequested)
return;
if (t == _refreshCertificates.Task)
Interlocked.Exchange(ref _refreshCertificates, new TaskCompletionSource<bool>());
if (_lastCertificateRefresh.AddMinutes(5) < DateTime.Now)
{
using var client = new HttpClient();
var response = await client.GetAsync(_certificateUrl);
if (response.IsSuccessStatusCode)
{
using var stream = await response.Content.ReadAsStreamAsync(_cancellationTokenSource.Token);
var serverKeys = await JsonSerializer.DeserializeAsync<IEnumerable<MiniServerCertificate>>(stream, cancellationToken: _cancellationTokenSource.Token);
if (serverKeys != null && serverKeys.Any())
{
_lastCertificateRefresh = DateTime.Now;
_serverKeys = serverKeys;
await InvokeReKey();
}
}
}
}
}
/// </inheritdoc>
public void Dispose()
{
@@ -84,6 +84,26 @@ public class RegisterForRemote : IDisposable
Disposed
}
/// <summary>
/// Data returned when the machine is claimed
/// </summary>
/// <param name="Success">True if the claim was successful</param>
/// <param name="StatusMessage">The status message for the claim</param>
/// <param name="JWT">The JWT token for the machine</param>
/// <param name="ServerUrl">The URL for the remote server</param>
/// <param name="CertificateUrl">The URL for getting new server certificates</param>
/// <param name="ServerCertificates">The certificates for the remote server</param>
/// <param name="LocalEncryptionKey">The encryption key for the local settings</param>
private sealed record EnvelopedClaimedClientData(
bool Success,
string StatusMessage,
string JWT,
string ServerUrl,
string CertificateUrl,
IEnumerable<MiniServerCertificate> ServerCertificates,
string? LocalEncryptionKey
);
/// <summary>
/// The current state of the registration process
/// </summary>
@@ -230,8 +250,13 @@ public class RegisterForRemote : IDisposable
var response = await _httpClient.PostAsync(_registerClientData!.StatusLink, CreateMachineData(), _cancellationTokenSource.Token);
response.EnsureSuccessStatusCode();
return await response.Content.ReadFromJsonAsync<ClaimedClientData>()
var result = await response.Content.ReadFromJsonAsync<EnvelopedClaimedClientData>()
?? throw new Exception("Failed to read machine claim data");
if (!result.Success)
throw new Exception($"Failed to claim machine: {result.StatusMessage}");
return new ClaimedClientData(result.JWT, result.ServerUrl, result.CertificateUrl, result.ServerCertificates, result.LocalEncryptionKey);
}
/// </inheritdoc>
+6 -21
View File
@@ -40,6 +40,8 @@ public sealed record RegisterClientData(
/// <summary>
/// Data returned when the machine is claimed
/// </summary>
/// <param name="Success">True if the claim was successful</param>
/// <param name="StatusMessage">The status message for the claim</param>
/// <param name="JWT">The JWT token for the machine</param>
/// <param name="ServerUrl">The URL for the remote server</param>
/// <param name="CertificateUrl">The URL for getting new server certificates</param>
@@ -58,30 +60,13 @@ public sealed record ClaimedClientData(
/// This data is serialized to various files, and chosen instead of X509 certificates.
/// If we find a non-complex certificate format, we should switch to that.
/// </summary>
/// <param name="Identifier">The machine identifier the key is valid for</param>
/// <param name="PublicKeyHash">The hash of the certificate public key</param>
/// <param name="PublicKey">The certificate public key</param>
/// <param name="Obtained">The date the certificate was obtained</param>
/// <param name="Expiry">The expiry date of the certificate</param>
/// <param name="Revoked">The date the certificate was revoked, or null if not revoked</param>
public sealed record MiniServerCertificate(
string Identifier,
string PublicKeyHash,
string PublicKey,
DateTimeOffset Obtained,
DateTimeOffset Expiry,
DateTimeOffset? Revoked
)
{
/// <summary>
/// Merges two sets of certificates, keeping the newest
/// </summary>
/// <param name="newcerts">The new certificates</param>
/// <param name="oldcerts">The old certificates</param>
/// <returns>The merged certificates</returns>
public static IEnumerable<MiniServerCertificate> MergeCertificates(IEnumerable<MiniServerCertificate>? newcerts, IEnumerable<MiniServerCertificate>? oldcerts)
=> (newcerts ?? []).Concat(oldcerts ?? [])
.DistinctBy(x => x.Identifier)
.Where(x => x.Revoked == null)
.Where(x => x.Expiry > DateTimeOffset.UtcNow)
.ToArray();
}
DateTimeOffset Expiry
);
@@ -146,7 +146,7 @@ backupApp.controller('SystemSettingsController', function($rootScope, $scope, $r
AppService.get('/remotecontrol/status').then(function(data) {
mapRemoteControlStatus(data.data);
}, () => { });
}, () => { });
}
@@ -112,6 +112,7 @@ public class RemoteControllerRegistrationService(Connection connection, IHttpCli
{
Token = claimData.JWT,
ServerCertificates = claimData.ServerCertificates,
CertificateUrl = claimData.CertificateUrl,
ServerUrl = claimData.ServerUrl
});
@@ -79,6 +79,7 @@ public class RemoteControllerService(Connection connection, IHttpClientFactory h
_keepRemoteConnection = KeepRemoteConnection.CreateRemoteListener(
config.ServerUrl,
config.Token,
config.CertificateUrl,
config.ServerCertificates,
CancellationToken.None,
ReKey,
@@ -100,7 +101,7 @@ public class RemoteControllerService(Connection connection, IHttpClientFactory h
connection.ApplicationSettings.RemoteControlConfig = JsonConvert.SerializeObject(new RemoteControlConfig
{
Token = data.JWT,
ServerCertificates = MiniServerCertificate.MergeCertificates(data.ServerCertificates, oldCerts),
ServerCertificates = data.ServerCertificates ?? oldCerts ?? [],
ServerUrl = data.ServerUrl,
CertificateUrl = data.CertificateUrl
});