diff --git a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultRestTemplateFactory.java b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultRestTemplateFactory.java new file mode 100644 index 00000000..5511dccc --- /dev/null +++ b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultRestTemplateFactory.java @@ -0,0 +1,58 @@ +/* + * Copyright 2019-2020 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.cloud.vault.config; + +import java.util.function.Consumer; +import java.util.function.Function; + +import org.springframework.http.client.ClientHttpRequestFactory; +import org.springframework.lang.Nullable; +import org.springframework.vault.client.RestTemplateBuilder; +import org.springframework.vault.client.RestTemplateFactory; +import org.springframework.web.client.RestTemplate; + +/** + * Default {@link RestTemplateFactory} implementation. + * + * @author Mark Paluch + * @since 3.0 + */ +class DefaultRestTemplateFactory implements RestTemplateFactory { + + private final ClientHttpRequestFactory requestFactory; + + private final Function builderFunction; + + DefaultRestTemplateFactory(ClientHttpRequestFactory requestFactory, + Function builderFunction) { + this.requestFactory = requestFactory; + this.builderFunction = builderFunction; + } + + @Override + public RestTemplate create(@Nullable Consumer customizer) { + + RestTemplateBuilder builder = this.builderFunction.apply(this.requestFactory); + + if (customizer != null) { + customizer.accept(builder); + } + + return builder.build(); + } + +} diff --git a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultWebClientFactory.java b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultWebClientFactory.java new file mode 100644 index 00000000..c6b58194 --- /dev/null +++ b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/DefaultWebClientFactory.java @@ -0,0 +1,58 @@ +/* + * Copyright 2019-2020 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.cloud.vault.config; + +import java.util.function.Consumer; +import java.util.function.Function; + +import org.springframework.http.client.reactive.ClientHttpConnector; +import org.springframework.lang.Nullable; +import org.springframework.vault.client.WebClientBuilder; +import org.springframework.vault.client.WebClientFactory; +import org.springframework.web.reactive.function.client.WebClient; + +/** + * Default implementation of {@link WebClientFactory}. + * + * @author Mark Paluch + * @since 3.0 + */ +class DefaultWebClientFactory implements WebClientFactory { + + private final ClientHttpConnector connector; + + private final Function builderFunction; + + DefaultWebClientFactory(ClientHttpConnector connector, + Function builderFunction) { + this.connector = connector; + this.builderFunction = builderFunction; + } + + @Override + public WebClient create(@Nullable Consumer customizer) { + + WebClientBuilder builder = this.builderFunction.apply(this.connector); + + if (customizer != null) { + customizer.accept(builder); + } + + return builder.build(); + } + +} diff --git a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultBootstrapConfiguration.java b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultBootstrapConfiguration.java index 001cd2ab..5bf6f6ff 100644 --- a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultBootstrapConfiguration.java +++ b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultBootstrapConfiguration.java @@ -49,6 +49,7 @@ import org.springframework.vault.authentication.SimpleSessionManager; import org.springframework.vault.client.ClientHttpRequestFactoryFactory; import org.springframework.vault.client.RestTemplateBuilder; import org.springframework.vault.client.RestTemplateCustomizer; +import org.springframework.vault.client.RestTemplateFactory; import org.springframework.vault.client.RestTemplateRequestCustomizer; import org.springframework.vault.client.SimpleVaultEndpointProvider; import org.springframework.vault.client.VaultEndpointProvider; @@ -58,7 +59,6 @@ import org.springframework.vault.core.VaultOperations; import org.springframework.vault.core.VaultTemplate; import org.springframework.vault.support.ClientOptions; import org.springframework.vault.support.SslConfiguration; -import org.springframework.web.client.RestOperations; import org.springframework.web.client.RestTemplate; /** @@ -83,15 +83,22 @@ public class VaultBootstrapConfiguration implements InitializingBean { private final List> requestCustomizers; + private ClientFactoryWrapper clientFactoryWrapper; + /** * Used for Vault communication. */ private RestTemplateBuilder restTemplateBuilder; + /** + * Used for Vault communication. + */ + private RestTemplateFactory restTemplateFactory; + /** * Used for external (AWS, GCP) communication. */ - private RestOperations externalRestOperations; + private RestTemplate externalRestOperations; public VaultBootstrapConfiguration(ConfigurableApplicationContext applicationContext, VaultProperties vaultProperties, ObjectProvider endpointProvider, @@ -118,21 +125,39 @@ public class VaultBootstrapConfiguration implements InitializingBean { @Override public void afterPropertiesSet() { - ClientHttpRequestFactory clientHttpRequestFactory = clientHttpRequestFactoryWrapper() - .getClientHttpRequestFactory(); + this.clientFactoryWrapper = createClientFactoryWrapper(); - this.restTemplateBuilder = RestTemplateBuilder.builder().requestFactory(clientHttpRequestFactory) + this.restTemplateBuilder = restTemplateBuilder(this.clientFactoryWrapper.getClientHttpRequestFactory()); + + this.externalRestOperations = new RestTemplate(this.clientFactoryWrapper.getClientHttpRequestFactory()); + + this.restTemplateFactory = new DefaultRestTemplateFactory( + this.clientFactoryWrapper.getClientHttpRequestFactory(), this::restTemplateBuilder); + + this.customizers.forEach(customizer -> customizer.customize(this.externalRestOperations)); + } + + /** + * Create a {@link RestTemplateBuilder} initialized with {@link VaultEndpointProvider} + * and {@link ClientHttpRequestFactory}. May be overridden by subclasses. + * @return the {@link RestTemplateBuilder}. + * @since 2.3 + * @see #clientHttpRequestFactoryWrapper() + */ + protected RestTemplateBuilder restTemplateBuilder(ClientHttpRequestFactory requestFactory) { + + RestTemplateBuilder builder = RestTemplateBuilder.builder().requestFactory(requestFactory) .endpointProvider(this.endpointProvider); - this.customizers.forEach(this.restTemplateBuilder::customizers); - this.requestCustomizers.forEach(this.restTemplateBuilder::requestCustomizers); + this.customizers.forEach(builder::customizers); + this.requestCustomizers.forEach(builder::requestCustomizers); if (StringUtils.hasText(this.vaultProperties.getNamespace())) { this.restTemplateBuilder.defaultHeader(VaultHttpHeaders.VAULT_NAMESPACE, this.vaultProperties.getNamespace()); } - this.externalRestOperations = new RestTemplate(clientHttpRequestFactory); + return builder; } /** @@ -147,14 +172,19 @@ public class VaultBootstrapConfiguration implements InitializingBean { @Bean @ConditionalOnMissingBean public ClientFactoryWrapper clientHttpRequestFactoryWrapper() { + return this.clientFactoryWrapper; + } - ClientOptions clientOptions = new ClientOptions(Duration.ofMillis(this.vaultProperties.getConnectionTimeout()), - Duration.ofMillis(this.vaultProperties.getReadTimeout())); - - SslConfiguration sslConfiguration = VaultConfigurationUtil - .createSslConfiguration(this.vaultProperties.getSsl()); - - return new ClientFactoryWrapper(ClientHttpRequestFactoryFactory.create(clientOptions, sslConfiguration)); + /** + * Create a {@link RestTemplateFactory} bean that is used to produce + * {@link RestTemplate}. + * @return the {@link RestTemplateFactory}. + * @see #clientHttpRequestFactoryWrapper() + * @since 3.0 + */ + @Bean + public RestTemplateFactory vaulRestTemplateFactory() { + return this.restTemplateFactory; } /** @@ -215,7 +245,7 @@ public class VaultBootstrapConfiguration implements InitializingBean { VaultProperties.SessionLifecycle lifecycle = this.vaultProperties.getSession().getLifecycle(); if (lifecycle.isEnabled()) { - RestTemplate restTemplate = this.restTemplateBuilder.build(); + RestTemplate restTemplate = this.restTemplateFactory.create(); LifecycleAwareSessionManagerSupport.RefreshTrigger trigger = new LifecycleAwareSessionManagerSupport.FixedTimeoutRefreshTrigger( lifecycle.getRefreshBeforeExpiry(), lifecycle.getExpiryThreshold()); return new LifecycleAwareSessionManager(clientAuthentication, @@ -236,13 +266,24 @@ public class VaultBootstrapConfiguration implements InitializingBean { @ConditionalOnAuthentication public ClientAuthentication clientAuthentication() { - RestTemplate restTemplate = this.restTemplateBuilder.build(); + RestTemplate restTemplate = this.restTemplateFactory.create(); ClientAuthenticationFactory factory = new ClientAuthenticationFactory(this.vaultProperties, restTemplate, this.externalRestOperations); return factory.createClientAuthentication(); } + protected ClientFactoryWrapper createClientFactoryWrapper() { + + ClientOptions clientOptions = new ClientOptions(Duration.ofMillis(this.vaultProperties.getConnectionTimeout()), + Duration.ofMillis(this.vaultProperties.getReadTimeout())); + + SslConfiguration sslConfiguration = VaultConfigurationUtil + .createSslConfiguration(this.vaultProperties.getSsl()); + + return new ClientFactoryWrapper(ClientHttpRequestFactoryFactory.create(clientOptions, sslConfiguration)); + } + /** * Wrapper to keep {@link TaskScheduler} local to Spring Cloud Vault. */ diff --git a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfiguration.java b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfiguration.java index 05eccf77..d27ce57f 100644 --- a/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfiguration.java +++ b/spring-cloud-vault-config/src/main/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfiguration.java @@ -26,6 +26,7 @@ import reactor.core.publisher.Mono; import reactor.netty.http.client.HttpClient; import org.springframework.beans.factory.BeanFactory; +import org.springframework.beans.factory.InitializingBean; import org.springframework.beans.factory.ListableBeanFactory; import org.springframework.beans.factory.ObjectFactory; import org.springframework.beans.factory.ObjectProvider; @@ -60,6 +61,7 @@ import org.springframework.vault.client.VaultEndpointProvider; import org.springframework.vault.client.VaultHttpHeaders; import org.springframework.vault.client.WebClientBuilder; import org.springframework.vault.client.WebClientCustomizer; +import org.springframework.vault.client.WebClientFactory; import org.springframework.vault.core.ReactiveVaultOperations; import org.springframework.vault.core.ReactiveVaultTemplate; import org.springframework.vault.support.ClientOptions; @@ -83,41 +85,60 @@ import org.springframework.web.reactive.function.client.WebClient; @ConditionalOnClass({ Flux.class, WebClient.class, ReactiveVaultOperations.class, HttpClient.class }) @EnableConfigurationProperties({ VaultProperties.class }) @Order(Ordered.LOWEST_PRECEDENCE - 10) -public class VaultReactiveBootstrapConfiguration { - - private final BeanFactory beanFactory; +public class VaultReactiveBootstrapConfiguration implements InitializingBean { private final VaultProperties vaultProperties; + private final VaultEndpointProvider endpointProvider; + + private final List customizers; + + private ClientHttpConnector clientHttpConnector; + /** * Used for Vault communication. */ - private final WebClientBuilder webClientBuilder; + private WebClientBuilder webClientBuilder; - public VaultReactiveBootstrapConfiguration(BeanFactory beanFactory, VaultProperties vaultProperties, + /** + * Used for Vault communication. + */ + private WebClientFactory webClientFactory; + + public VaultReactiveBootstrapConfiguration(VaultProperties vaultProperties, ObjectProvider endpointProvider, ObjectProvider> webClientCustomizers) { - this.beanFactory = beanFactory; this.vaultProperties = vaultProperties; - VaultEndpointProvider provider = endpointProvider.getIfAvailable(); + this.endpointProvider = endpointProvider.getIfAvailable( + () -> SimpleVaultEndpointProvider.of(VaultConfigurationUtil.createVaultEndpoint(vaultProperties))); + this.customizers = new ArrayList<>(webClientCustomizers.getIfAvailable(Collections::emptyList)); + AnnotationAwareOrderComparator.sort(this.customizers); + } - if (provider == null) { - provider = SimpleVaultEndpointProvider.of(VaultConfigurationUtil.createVaultEndpoint(vaultProperties)); - } + @Override + public void afterPropertiesSet() { - this.webClientBuilder = WebClientBuilder.builder().httpConnector(createConnector(this.vaultProperties)) - .endpointProvider(provider); - List customizers = new ArrayList<>( - webClientCustomizers.getIfAvailable(Collections::emptyList)); - AnnotationAwareOrderComparator.sort(customizers); + this.clientHttpConnector = createConnector(this.vaultProperties); - customizers.forEach(this.webClientBuilder::customizers); + this.webClientBuilder = webClientBuilder(this.clientHttpConnector); + + this.webClientFactory = new DefaultWebClientFactory(this.clientHttpConnector, this::webClientBuilder); + } + + protected WebClientBuilder webClientBuilder(ClientHttpConnector connector) { + + WebClientBuilder builder = WebClientBuilder.builder().httpConnector(connector) + .endpointProvider(this.endpointProvider); + + this.customizers.forEach(builder::customizers); if (StringUtils.hasText(this.vaultProperties.getNamespace())) { - this.webClientBuilder.defaultHeader(VaultHttpHeaders.VAULT_NAMESPACE, this.vaultProperties.getNamespace()); + builder.defaultHeader(VaultHttpHeaders.VAULT_NAMESPACE, this.vaultProperties.getNamespace()); } + + return builder; } /** @@ -127,7 +148,7 @@ public class VaultReactiveBootstrapConfiguration { * @param vaultProperties the Vault properties. * @return the {@link ClientHttpConnector}. */ - private static ClientHttpConnector createConnector(VaultProperties vaultProperties) { + protected ClientHttpConnector createConnector(VaultProperties vaultProperties) { ClientOptions clientOptions = new ClientOptions(Duration.ofMillis(vaultProperties.getConnectionTimeout()), Duration.ofMillis(vaultProperties.getReadTimeout())); @@ -137,6 +158,16 @@ public class VaultReactiveBootstrapConfiguration { return ClientHttpConnectorFactory.create(clientOptions, sslConfiguration); } + /** + * Create a {@link WebClientFactory} bean that is used to produce {@link WebClient}. + * @return the {@link WebClientFactory}. + * @since 3.0 + */ + @Bean + public WebClientFactory vaultWebClientFactory() { + return this.webClientFactory; + } + /** * Creates a {@link ReactiveVaultTemplate}. * @return the {@link ReactiveVaultTemplate} bean. @@ -144,13 +175,13 @@ public class VaultReactiveBootstrapConfiguration { */ @Bean @ConditionalOnMissingBean(ReactiveVaultOperations.class) - public ReactiveVaultTemplate reactiveVaultTemplate() { + public ReactiveVaultTemplate reactiveVaultTemplate(ObjectProvider sessionManager) { if (this.vaultProperties.getAuthentication() == VaultProperties.AuthenticationMethod.NONE) { return new ReactiveVaultTemplate(this.webClientBuilder); } - return new ReactiveVaultTemplate(this.webClientBuilder, beanFactory.getBean(ReactiveSessionManager.class)); + return new ReactiveVaultTemplate(this.webClientBuilder, sessionManager.getObject()); } /** @@ -171,7 +202,7 @@ public class VaultReactiveBootstrapConfiguration { VaultProperties.SessionLifecycle lifecycle = this.vaultProperties.getSession().getLifecycle(); if (lifecycle.isEnabled()) { - WebClient webClient = this.webClientBuilder.build(); + WebClient webClient = this.webClientFactory.create(); ReactiveLifecycleAwareSessionManager.RefreshTrigger trigger = new ReactiveLifecycleAwareSessionManager.FixedTimeoutRefreshTrigger( lifecycle.getRefreshBeforeExpiry(), lifecycle.getExpiryThreshold()); return new ReactiveLifecycleAwareSessionManager(vaultTokenSupplier, @@ -245,7 +276,7 @@ public class VaultReactiveBootstrapConfiguration { } private VaultTokenSupplier createAuthenticationStepsOperator(AuthenticationStepsFactory factory) { - WebClient webClient = this.webClientBuilder.build(); + WebClient webClient = this.webClientFactory.create(); return new AuthenticationStepsOperator(factory.getAuthenticationSteps(), webClient); } diff --git a/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultBootstrapConfigurationTests.java b/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultBootstrapConfigurationTests.java index 349885bf..47788833 100644 --- a/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultBootstrapConfigurationTests.java +++ b/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultBootstrapConfigurationTests.java @@ -26,6 +26,7 @@ import org.springframework.test.util.ReflectionTestUtils; import org.springframework.vault.authentication.ClientAuthentication; import org.springframework.vault.authentication.SessionManager; import org.springframework.vault.authentication.SimpleSessionManager; +import org.springframework.vault.client.RestTemplateFactory; import org.springframework.vault.core.VaultTemplate; import static org.assertj.core.api.Assertions.assertThat; @@ -50,6 +51,7 @@ public class VaultBootstrapConfigurationTests { assertThat(context).doesNotHaveBean(SessionManager.class); assertThat(context).doesNotHaveBean(ClientAuthentication.class); assertThat(context).hasSingleBean(VaultTemplate.class); + assertThat(context).hasSingleBean(RestTemplateFactory.class); }); } diff --git a/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfigurationTests.java b/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfigurationTests.java index 7545d8ff..1fccb729 100644 --- a/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfigurationTests.java +++ b/spring-cloud-vault-config/src/test/java/org/springframework/cloud/vault/config/VaultReactiveBootstrapConfigurationTests.java @@ -37,6 +37,7 @@ import org.springframework.vault.authentication.ReactiveSessionManager; import org.springframework.vault.authentication.SessionManager; import org.springframework.vault.authentication.SimpleSessionManager; import org.springframework.vault.authentication.VaultTokenSupplier; +import org.springframework.vault.client.WebClientFactory; import org.springframework.vault.core.ReactiveVaultOperations; import org.springframework.vault.support.VaultToken; import org.springframework.web.reactive.function.client.WebClient; @@ -66,6 +67,7 @@ public class VaultReactiveBootstrapConfigurationTests { .isNotInstanceOf(LifecycleAwareSessionManager.class) .isNotInstanceOf(SimpleSessionManager.class); assertThat(context.getBeanNamesForType(WebClient.class)).isEmpty(); + assertThat(context).hasSingleBean(WebClientFactory.class); }); }