/* * Copyright (c) 2023 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 System; using System.Collections.Generic; using System.Net.Http; using System.Net.Security; using System.Security.Cryptography.X509Certificates; using FluentAssertions; using Microsoft.VisualStudio.TestTools.UnitTesting; using NSubstitute; using ProtonVPN.Api.Handlers.TlsPinning; using ProtonVPN.Configurations.Contracts; using ProtonVPN.Configurations.Contracts.Entities; using ProtonVPN.Configurations.Entities; namespace ProtonVPN.Api.Tests.Handlers.TlsPinning; [TestClass] public class TlsPinnedCertificateHandlerTest { private IReportClient _reportClient; private X509Certificate _apiCert; private X509Certificate _alternativeHostCert; private readonly string _unknownHost = "unknown.host.com"; private readonly string _apiHost = "api.protonvpn.ch"; private readonly string _alternativeHost = "alternative.host.com"; [TestInitialize] public void TestInitialize() { _reportClient = Substitute.For(); _apiCert = X509Certificate.CreateFromCertFile("TestData\\api.protonvpn.ch.cer"); _alternativeHostCert = X509Certificate.CreateFromCertFile("TestData\\alternative.host.cer"); } #region Unknown domain tests [TestMethod] public void ItShouldReturnTrue() { IConfiguration config = GetApiTlsPinningConfig(false); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_unknownHost, _apiCert, SslPolicyErrors.None).Should().BeTrue(); } [TestMethod] public void ItShouldReturnFalseWhenEnforceIsOn() { IConfiguration config = GetApiTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_unknownHost, _apiCert, SslPolicyErrors.None).Should().BeFalse(); } [TestMethod] public void ItShouldReturnFalseWhenEnforceIsOnAndSslError() { IConfiguration config = GetApiTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_unknownHost, _apiCert, SslPolicyErrors.RemoteCertificateNameMismatch).Should().BeFalse(); } #endregion #region Known domain tests [TestMethod] public void ItShouldReturnFalseWhenSslError() { IConfiguration config = GetApiTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_apiHost, _apiCert, SslPolicyErrors.RemoteCertificateNameMismatch).Should().BeFalse(); } [TestMethod] public void ItShouldReturnTrueWhenEnforceIsOff() { IConfiguration config = GetApiTlsPinningConfig(false); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_apiHost, _apiCert, SslPolicyErrors.None).Should().BeTrue(); } [TestMethod] public void ItShouldReturnTrueWhenPinIsValid() { IConfiguration config = GetApiTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_apiHost, _apiCert, SslPolicyErrors.None).Should().BeTrue(); } [TestMethod] public void ItShouldReturnFalseWhenPinIsNotValid() { IConfiguration config = GetApiTlsPinningConfig(true); config.TlsPinning.PinnedDomains.Returns(new List()); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_apiHost, _apiCert, SslPolicyErrors.None).Should().BeFalse(); } #endregion #region Alternative api tests [TestMethod] public void ItShouldReturnTrueWhenAlternativeHostPinIsValid() { IConfiguration config = GetAlternativeApiTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_alternativeHost, _alternativeHostCert, SslPolicyErrors.None).Should().BeTrue(); } [TestMethod] public void ItShouldReturnFalseWhenAlternativeHostPinIsInvalid() { IConfiguration config = GetAlternativeApiTlsPinningConfig(true); config.TlsPinning.PinnedDomains.Returns(new List()); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_alternativeHost, _alternativeHostCert, SslPolicyErrors.None).Should().BeFalse(); } #endregion [TestMethod] public void ItShouldSendTlsPinReportWhenPinIsNotValid() { IConfiguration config = GetIncorrectTlsPinningConfig(true); TestCertificateHandler certificateHandler = CreateTestCertificateHandler(config); certificateHandler.GetValidationResult(_apiHost, _apiCert, SslPolicyErrors.None); _reportClient.ReceivedWithAnyArgs().Send(null); } private IConfiguration GetApiTlsPinningConfig(bool enforce) { IConfiguration config = Substitute.For(); PinConfigBuilder builder = new(enforce); builder.AddDomain(_apiHost, enforce, new List { "IEwk65VSaxv3s1/88vF/rM8PauJoIun3rzVCX5mLS3M=", "drtmcR2kFkM8qJClsuWgUzxgBkePfRCkRpqUesyDmeE=", "YRGlaY0jyJ4Jw2/4M8FIftwbDIQfh8Sdro96CeEel54=", "AfMENBVvOS8MnISprtvyPsjKlPooqh8nMB/pvCrpJpw=", }); ITlsPinningConfiguration tlsPinningConfig = builder.Config(); config.TlsPinning.Returns(tlsPinningConfig); config.IsCertificateValidationEnabled.Returns(true); return config; } private IConfiguration GetIncorrectTlsPinningConfig(bool enforce) { IConfiguration config = Substitute.For(); PinConfigBuilder builder = new(enforce); builder.AddDomain(_apiHost, enforce, new List { "wrong pin", "another wrong pin" }); ITlsPinningConfiguration tlsPinningConfig = builder.Config(); config.TlsPinning.Returns(tlsPinningConfig); config.IsCertificateValidationEnabled.Returns(true); return config; } private IConfiguration GetAlternativeApiTlsPinningConfig(bool enforce) { IConfiguration config = Substitute.For(); PinConfigBuilder builder = new(enforce); builder.AddDomain(_alternativeHost, enforce, new List { "EU6TS9MO0L/GsDHvVc9D5fChYLNy5JdGYpJw0ccgetM=", "iKPIHPnDNqdkvOnTClQ8zQAIKG0XavaPkcEo0LBAABA=", "MSlVrBCdL0hKyczvgYVSRNm88RicyY04Q2y5qrBt0xA=", "C2UxW0T1Ckl9s+8cXfjXxlEqwAfPM4HiW2y3UdtBeCw=", }); ITlsPinningConfiguration tlsPinningConfig = builder.Config(); config.TlsPinning.Returns(tlsPinningConfig); config.IsCertificateValidationEnabled.Returns(true); return config; } private TestCertificateHandler CreateTestCertificateHandler(IConfiguration config) { return new(new CertificateValidator(_reportClient, config), config); } } internal class PinConfigBuilder { private readonly ITlsPinningConfiguration _config; private readonly List _domains = new(); public PinConfigBuilder(bool enforce) { _config = Substitute.For(); _config.PinnedDomains.Returns(new List()); _config.Enforce.Returns(enforce); } public PinConfigBuilder AddDomain(string domain, bool enforce, List pins) { _domains.Add(new TlsPinnedDomain { Name = domain, PublicKeyHashes = pins, Enforce = enforce, SendReport = true }); return this; } public ITlsPinningConfiguration Config() { _config.PinnedDomains.Returns(_domains); return _config; } } internal class TestCertificateHandler : TlsPinnedCertificateHandler { public TestCertificateHandler(ICertificateValidator certificateValidator, IConfiguration config) : base(certificateValidator, config) { } public bool GetValidationResult(string host, X509Certificate cert, SslPolicyErrors sslPolicyErrors) { return CertificateCustomValidationCallback( new HttpRequestMessage { Headers = { Host = host }, RequestUri = new UriBuilder(new Uri("https://host.com")).Uri }, cert, new X509Chain(), sslPolicyErrors); } }