/*
* Copyright (c) 2026 Proton AG
*
* This file is part of ProtonVPN.
*
* ProtonVPN is free software: you can redistribute it and/or modify
* it under the terms of the GNU General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* ProtonVPN is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU General Public License for more details.
*
* You should have received a copy of the GNU General Public License
* along with ProtonVPN. If not, see .
*/
using ProtonVPN.Api.Contracts;
using ProtonVPN.Api.Contracts.Certificates;
using ProtonVPN.Client.EventMessaging.Contracts;
using ProtonVPN.Client.Logic.Auth.Contracts;
using ProtonVPN.Client.Logic.Auth.Contracts.Messages;
using ProtonVPN.Client.Logic.Auth.Contracts.Models;
using ProtonVPN.Client.Logic.Users.Contracts.Messages;
using ProtonVPN.Client.Settings.Contracts;
using ProtonVPN.Crypto.Contracts;
using ProtonVPN.Logging.Contracts;
using ProtonVPN.Logging.Contracts.Events.UserCertificateLogs;
namespace ProtonVPN.Client.Logic.Auth;
public class ConnectionCertificateManager : IConnectionCertificateManager
{
private const string USER_GROUP_EXTENSION_OID = "1.3.6.1.4.1.56809.1.0.0.2";
private const string FREE_USER_GROUP = "vpn-free";
private const string PAID_USER_GROUP = "vpn-paid";
private readonly ISettings _settings;
private readonly IConnectionKeyManager _connectionKeyManager;
private readonly IApiClient _apiClient;
private readonly ILogger _logger;
private readonly IEventMessageSender _eventMessageSender;
private readonly ICertificateParser _certificateParser;
private readonly SemaphoreSlim _semaphore = new(1, 1);
public ConnectionCertificateManager(
ISettings settings,
IConnectionKeyManager connectionKeyManager,
IApiClient apiClient,
ILogger logger,
IEventMessageSender eventMessageSender,
ICertificateParser certificateParser)
{
_settings = settings;
_connectionKeyManager = connectionKeyManager;
_apiClient = apiClient;
_logger = logger;
_eventMessageSender = eventMessageSender;
_certificateParser = certificateParser;
}
public void DeleteKeyPairAndCertificate()
{
_connectionKeyManager.DeleteKeyPair();
_settings.ConnectionCertificate = null;
SendUpdateMessage(null);
_logger.Info("Connection certificate deleted.");
}
public void DeleteKeyPairAndCertificateIfMatches(string expiredCertificatePem)
{
if (expiredCertificatePem == _settings.ConnectionCertificate?.Pem)
{
DeleteKeyPairAndCertificate();
}
}
private enum NewCertificateRequestParameter
{
NewCertificateIfCurrentIsOld = 0,
ForceNewCertificate = 1,
ForceNewKeyPairAndCertificate = 2
}
public async Task RequestNewCertificateAsync(CancellationToken cancellationToken = default, string? expiredCertificatePem = null)
{
await EnqueueRequestAsync(NewCertificateRequestParameter.NewCertificateIfCurrentIsOld,
cancellationToken,
expiredCertificatePem);
}
public bool IsCertificateOutOfSyncWithPlan()
{
ConnectionCertificate? connectionCertificate = _settings.ConnectionCertificate;
if (connectionCertificate is null || string.IsNullOrWhiteSpace(connectionCertificate.Value.Pem))
{
return true;
}
VpnPlan vpnPlan = _settings.VpnPlan;
const string outOfSyncLog = "User plan and connection certificate are out of sync.";
List userGroups = _certificateParser.GetExtensionStrings(connectionCertificate.Value.Pem, USER_GROUP_EXTENSION_OID);
if (vpnPlan.IsPaid && userGroups.Contains(FREE_USER_GROUP))
{
_logger.Warn($"{outOfSyncLog} Paid plan, but free certificate.");
return true;
}
if (!vpnPlan.IsPaid && userGroups.Contains(PAID_USER_GROUP))
{
_logger.Warn($"{outOfSyncLog} Free plan, but paid certificate.");
return true;
}
_logger.Info("User plan and connection certificate are in sync.");
return false;
}
public async Task ForceRequestNewCertificateAsync(CancellationToken cancellationToken = default)
{
await EnqueueRequestAsync(NewCertificateRequestParameter.ForceNewCertificate, cancellationToken);
}
public async Task ForceRequestNewKeyPairAndCertificateAsync()
{
await EnqueueRequestAsync(NewCertificateRequestParameter.ForceNewKeyPairAndCertificate);
}
private async Task EnqueueRequestAsync(
NewCertificateRequestParameter parameter,
CancellationToken cancellationToken = default,
string? expiredCertificatePem = null)
{
await _semaphore.WaitAsync(cancellationToken);
try
{
if (parameter != NewCertificateRequestParameter.NewCertificateIfCurrentIsOld ||
IsToRequest(expiredCertificatePem))
{
LogNewCertificateRequest(parameter);
RegenerateKeyPairIfRequested(parameter);
ApiResponseResult response = await RequestAsync(cancellationToken);
if (response.Failure)
{
_logger.Error("Connection certificate request failed with " +
$"Status Code {response.ResponseMessage.StatusCode}, " +
$"Internal Code {response.Value.Code}, " +
$"Error '{response.Value.Error}'.");
}
}
else
{
SendMessageWithCurrentCertificate();
}
}
catch (Exception e)
{
_logger.Error("Connection certificate request failed.", e);
}
finally
{
_semaphore.Release();
}
}
private void LogNewCertificateRequest(NewCertificateRequestParameter parameter)
{
switch (parameter)
{
case NewCertificateRequestParameter.NewCertificateIfCurrentIsOld:
_logger.Info("Requesting a new connection certificate since the current one is considered old.");
break;
case NewCertificateRequestParameter.ForceNewCertificate:
_logger.Info("Forcing a new connection certificate request.");
break;
case NewCertificateRequestParameter.ForceNewKeyPairAndCertificate:
_logger.Info("Generating new connection key pair and forcing a new connection certificate request.");
break;
}
}
private bool IsToRequest(string? expiredCertificatePem)
{
ConnectionCertificate? connectionCertificate = _settings.ConnectionCertificate;
DateTimeOffset utcNow = DateTimeOffset.UtcNow;
return connectionCertificate is null ||
string.IsNullOrWhiteSpace(connectionCertificate.Value.Pem) ||
utcNow >= connectionCertificate.Value.RefreshUtcDate ||
utcNow >= connectionCertificate.Value.ExpirationUtcDate ||
expiredCertificatePem == connectionCertificate.Value.Pem;
}
private void RegenerateKeyPairIfRequested(NewCertificateRequestParameter parameter)
{
if (parameter == NewCertificateRequestParameter.ForceNewKeyPairAndCertificate)
{
_connectionKeyManager.RegenerateKeyPair();
}
}
private async Task> RequestAsync(CancellationToken cancellationToken)
{
ApiResponseResult certificateResponseData =
await RequestConnectionCertificateAsync(cancellationToken);
if (certificateResponseData.Failure && certificateResponseData.Value.Code == ResponseCodes.CLIENT_PUBLIC_KEY_CONFLICT)
{
_logger.Warn("New connection certificate failed because the " +
"client public key is already in use. Generating a new key pair and retrying.");
_connectionKeyManager.RegenerateKeyPair();
certificateResponseData = await RequestConnectionCertificateAsync(cancellationToken);
}
return certificateResponseData;
}
private async Task> RequestConnectionCertificateAsync(CancellationToken cancellationToken)
{
CertificateRequest certificateRequest = CreateCertificateRequestData();
ApiResponseResult certificateResponseData =
await _apiClient.RequestConnectionCertificateAsync(certificateRequest, cancellationToken);
if (certificateResponseData.Success)
{
ConnectionCertificate connectionCertificate = new()
{
Pem = certificateResponseData.Value.Certificate,
RequestUtcDate = DateTimeOffset.UtcNow,
RefreshUtcDate = DateTimeOffset.FromUnixTimeSeconds(certificateResponseData.Value.RefreshTime),
ExpirationUtcDate = DateTimeOffset.FromUnixTimeSeconds(certificateResponseData.Value.ExpirationTime),
};
_settings.ConnectionCertificate = connectionCertificate;
string userGroups = string.Join(", ", _certificateParser.GetExtensionStrings(connectionCertificate.Pem, USER_GROUP_EXTENSION_OID));
_logger.Info("New connection certificate successfully saved. " +
$"User groups: [{userGroups}]. " +
$"Expires at {connectionCertificate.ExpirationUtcDate}.");
SendUpdateMessage(connectionCertificate);
}
return certificateResponseData;
}
private CertificateRequest CreateCertificateRequestData()
{
return new()
{
ClientPublicKey = GetOrCreateClientPublicKeyPem(),
Features = [],
};
}
private string GetOrCreateClientPublicKeyPem()
{
string? clientPublicKey = _connectionKeyManager.GetPublicKey()?.Pem;
if (string.IsNullOrEmpty(clientPublicKey))
{
_connectionKeyManager.RegenerateKeyPair();
clientPublicKey = _connectionKeyManager.GetPublicKey()?.Pem;
}
return clientPublicKey ?? string.Empty;
}
private void SendUpdateMessage(ConnectionCertificate? connectionCertificate)
{
ConnectionCertificateUpdatedMessage message = new()
{
Certificate = connectionCertificate
};
_eventMessageSender.Send(message);
}
private void SendMessageWithCurrentCertificate()
{
SendUpdateMessage(_settings.ConnectionCertificate);
}
}