From dfac33609d45d0ad8679e8c47d45c69970e8f8ff Mon Sep 17 00:00:00 2001 From: Nico Espinosa Date: Thu, 1 Oct 2026 11:21:30 +0200 Subject: [PATCH] xds: Keep saved certs and roots after SslContext update Fixes #13058. CertProviderSslContextProvider cleared the saved key, cert chain and trusted roots after every SslContext build. When the identity cert and the CA roots come from separate certificate provider instances, the roots provider usually does not send another update, so the next identity cert rotation found no roots and did not rebuild the SslContext. New connections kept the old, possibly expired, identity cert. The same problem occurred with a single shared provider instance when only the identity cert changed. The saved identity credentials and trust roots are now kept as the latest known values. An update from either provider rebuilds the SslContext with the latest values from both. This generalizes the fix in #12340, which kept the roots only when using system root certs. As a result, a root-only update now also rebuilds the SslContext. When a shared provider instance updates the cert and the roots in the same refresh, the SslContext is built twice, and the first build briefly uses the new identity cert with the old roots. --- .../CertProviderSslContextProvider.java | 14 +- ...tProviderClientSslContextProviderTest.java | 168 +++++++++++++++--- ...tProviderServerSslContextProviderTest.java | 164 ++++++++++++++--- 3 files changed, 284 insertions(+), 62 deletions(-) diff --git a/xds/src/main/java/io/grpc/xds/internal/security/certprovider/CertProviderSslContextProvider.java b/xds/src/main/java/io/grpc/xds/internal/security/certprovider/CertProviderSslContextProvider.java index cb99ca6ad95..5b270c187db 100644 --- a/xds/src/main/java/io/grpc/xds/internal/security/certprovider/CertProviderSslContextProvider.java +++ b/xds/src/main/java/io/grpc/xds/internal/security/certprovider/CertProviderSslContextProvider.java @@ -167,34 +167,24 @@ public final void updateSpiffeTrustMap(Map> spiffe updateSslContextWhenReady(); } + // Saved credentials and roots persist across builds, so an update from either provider + // rebuilds using the latest values from both. private void updateSslContextWhenReady() { if (isMtls()) { if (savedKey != null && (savedTrustedRoots != null || savedSpiffeTrustMap != null)) { updateSslContext(); - clearKeysAndCerts(); } } else if (isRegularTlsAndClientSide()) { if (savedTrustedRoots != null || savedSpiffeTrustMap != null) { updateSslContext(); - clearKeysAndCerts(); } } else if (isRegularTlsAndServerSide()) { if (savedKey != null) { updateSslContext(); - clearKeysAndCerts(); } } } - private void clearKeysAndCerts() { - savedKey = null; - if (!isUsingSystemRootCerts) { - savedTrustedRoots = null; - savedSpiffeTrustMap = null; - } - savedCertChain = null; - } - protected final boolean isMtls() { return certInstance != null && (rootCertInstance != null || isUsingSystemRootCerts); } diff --git a/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderClientSslContextProviderTest.java b/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderClientSslContextProviderTest.java index 37aa2c65b08..5147a1b5def 100644 --- a/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderClientSslContextProviderTest.java +++ b/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderClientSslContextProviderTest.java @@ -154,9 +154,6 @@ public void testProviderForClient_mtls() throws Exception { // now generate root cert update watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -166,27 +163,157 @@ public void testProviderForClient_mtls() throws Exception { CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); - // just do root cert update: sslContext should still be the same + // just do root cert update: sslContext should be updated watcherCaptor[0].updateTrustedRoots( ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); - assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + testCallback = testCallback1; // now update id cert: sslContext should be updated i.e.different from the previous one watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); } + @Test + public void testProviderForClient_mtls_sharedInstance_certUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[1]; + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "testca", 0); + CertProviderClientSslContextProvider provider = + getSslContextProvider( + "gcp_id", + "gcp_id", + CommonBootstrapperTestUtils.getTestBootstrapInfo(), + /* alpnProtocols= */ null, + /* staticCertValidationContext= */ null, false); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do id cert update: sslContext should be updated + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another id cert update: sslContext should be updated again + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + /** + * Helper method to build CertProviderClientSslContextProvider with separate cert and root + * instances. watcherCaptor[0] is the cert watcher and watcherCaptor[1] is the root watcher. + */ + private CertProviderClientSslContextProvider getSslContextProviderWithSeparateInstances( + CertificateProvider.DistributorWatcher[] watcherCaptor) { + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "testca", 0); + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "file_watcher", 1); + return getSslContextProvider( + "gcp_id", + "file_provider", + CommonBootstrapperTestUtils.getTestBootstrapInfo(), + /* alpnProtocols= */ null, + /* staticCertValidationContext= */ null, false); + } + + @Test + public void testProviderForClient_mtls_separateInstances_certUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderClientSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do id cert update: sslContext should be updated + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another id cert update: sslContext should be updated again + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + @Test + public void testProviderForClient_mtls_separateInstances_rootUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderClientSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do root cert update: sslContext should be updated + watcherCaptor[1].updateTrustedRoots( + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another root cert update: sslContext should be updated again + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + @Test + public void testProviderForClient_mtls_separateInstances_ignoresOtherInstanceUpdates() + throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderClientSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // root cert update from the cert instance: sslContext should still be the same + watcherCaptor[0].updateTrustedRoots( + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + // id cert update from the root instance: sslContext should still be the same + watcherCaptor[1].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + } + @Test public void testProviderForClient_systemRootCerts_mtls() throws Exception { final CertificateProvider.DistributorWatcher[] watcherCaptor = @@ -214,8 +341,6 @@ public void testProviderForClient_systemRootCerts_mtls() throws Exception { watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(CLIENT_KEY_FILE), ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); @@ -231,8 +356,6 @@ public void testProviderForClient_systemRootCerts_mtls() throws Exception { watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -298,9 +421,6 @@ public void testProviderForClient_mtls_newXds() throws Exception { // now generate root cert update watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -310,22 +430,18 @@ public void testProviderForClient_mtls_newXds() throws Exception { CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); - // just do root cert update: sslContext should still be the same + // just do root cert update: sslContext should be updated watcherCaptor[0].updateTrustedRoots( ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); - assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + testCallback = testCallback1; // now update id cert: sslContext should be updated i.e.different from the previous one watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); @@ -388,7 +504,6 @@ public void testProviderForClient_tls() throws Exception { assertThat(provider.getSslContextAndTrustManager()).isNotNull(); assertThat(provider.savedKey).isNull(); assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -517,9 +632,6 @@ public void testProviderForClient_deprecatedCertProviderField() throws Exception // Generate root cert update watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); diff --git a/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderServerSslContextProviderTest.java b/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderServerSslContextProviderTest.java index 93559f47245..39206df9393 100644 --- a/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderServerSslContextProviderTest.java +++ b/xds/src/test/java/io/grpc/xds/internal/security/certprovider/CertProviderServerSslContextProviderTest.java @@ -140,9 +140,6 @@ public void testProviderForServer_mtls() throws Exception { // now generate root cert update watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -152,22 +149,18 @@ public void testProviderForServer_mtls() throws Exception { CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); - // just do root cert update: sslContext should still be the same + // just do root cert update: sslContext should be updated watcherCaptor[0].updateTrustedRoots( ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); - assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + testCallback = testCallback1; // now update id cert: sslContext should be updated i.e.different from the previous one watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); @@ -209,9 +202,6 @@ public void testProviderForServer_mtls_newXds() throws Exception { // now generate root cert update watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); @@ -221,27 +211,159 @@ public void testProviderForServer_mtls_newXds() throws Exception { CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); - // just do root cert update: sslContext should still be the same + // just do root cert update: sslContext should be updated watcherCaptor[0].updateTrustedRoots( ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); - assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + testCallback = testCallback1; // now update id cert: sslContext should be updated i.e.different from the previous one watcherCaptor[0].updateCertificate( CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); - assertThat(provider.savedTrustedRoots).isNull(); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); } + @Test + public void testProviderForServer_mtls_sharedInstance_certUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[1]; + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "testca", 0); + CertProviderServerSslContextProvider provider = + getSslContextProvider( + "gcp_id", + "gcp_id", + CommonBootstrapperTestUtils.getTestBootstrapInfo(), + /* alpnProtocols= */ null, + /* staticCertValidationContext= */ null, + /* requireClientCert= */ true); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + watcherCaptor[0].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do id cert update: sslContext should be updated + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another id cert update: sslContext should be updated again + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + /** + * Helper method to build CertProviderServerSslContextProvider with separate cert and root + * instances. watcherCaptor[0] is the cert watcher and watcherCaptor[1] is the root watcher. + */ + private CertProviderServerSslContextProvider getSslContextProviderWithSeparateInstances( + CertificateProvider.DistributorWatcher[] watcherCaptor) { + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "testca", 0); + TestCertificateProvider.createAndRegisterProviderProvider( + certificateProviderRegistry, watcherCaptor, "file_watcher", 1); + return getSslContextProvider( + "gcp_id", + "file_provider", + CommonBootstrapperTestUtils.getTestBootstrapInfo(), + /* alpnProtocols= */ null, + /* staticCertValidationContext= */ null, + /* requireClientCert= */ true); + } + + @Test + public void testProviderForServer_mtls_separateInstances_certUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderServerSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do id cert update: sslContext should be updated + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another id cert update: sslContext should be updated again + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + @Test + public void testProviderForServer_mtls_separateInstances_rootUpdateOnly() throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderServerSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // just do root cert update: sslContext should be updated + watcherCaptor[1].updateTrustedRoots( + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isNotSameInstanceAs(testCallback.updatedSslContext); + + // another root cert update: sslContext should be updated again + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback2 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback2.updatedSslContext) + .isNotSameInstanceAs(testCallback1.updatedSslContext); + } + + @Test + public void testProviderForServer_mtls_separateInstances_ignoresOtherInstanceUpdates() + throws Exception { + final CertificateProvider.DistributorWatcher[] watcherCaptor = + new CertificateProvider.DistributorWatcher[2]; + CertProviderServerSslContextProvider provider = + getSslContextProviderWithSeparateInstances(watcherCaptor); + watcherCaptor[0].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_0_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); + watcherCaptor[1].updateTrustedRoots(ImmutableList.of(getCertFromResourceName(CA_PEM_FILE))); + TestCallback testCallback = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback.updatedSslContext).isNotNull(); + + // root cert update from the cert instance: sslContext should still be the same + watcherCaptor[0].updateTrustedRoots( + ImmutableList.of(getCertFromResourceName(CLIENT_PEM_FILE))); + // id cert update from the root instance: sslContext should still be the same + watcherCaptor[1].updateCertificate( + CommonCertProviderTestUtils.getPrivateKey(SERVER_1_KEY_FILE), + ImmutableList.of(getCertFromResourceName(SERVER_1_PEM_FILE))); + TestCallback testCallback1 = CommonTlsContextTestsUtil.getValueThruCallback(provider); + assertThat(testCallback1.updatedSslContext).isSameInstanceAs(testCallback.updatedSslContext); + } + @Test public void testProviderForServer_queueExecutor() throws Exception { final CertificateProvider.DistributorWatcher[] watcherCaptor = @@ -302,8 +424,6 @@ public void testProviderForServer_tls() throws Exception { ImmutableList.of(getCertFromResourceName(SERVER_0_PEM_FILE))); assertThat(provider.getSslContextAndTrustManager()).isNotNull(); - assertThat(provider.savedKey).isNull(); - assertThat(provider.savedCertChain).isNull(); assertThat(provider.savedTrustedRoots).isNull(); TestCallback testCallback =