diff --git a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClient.java b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClient.java index 3894e883..d5b4c7e9 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClient.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClient.java @@ -26,6 +26,11 @@ import org.springframework.http.HttpEntity; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpInputMessage; import org.springframework.http.HttpMethod; +import org.springframework.http.HttpRequest; +import org.springframework.http.client.ClientHttpRequestExecution; +import org.springframework.http.client.ClientHttpRequestFactory; +import org.springframework.http.client.ClientHttpRequestInterceptor; +import org.springframework.http.client.ClientHttpResponse; import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter; import org.springframework.util.Assert; import org.springframework.vault.core.VaultTemplate; @@ -62,6 +67,45 @@ public class VaultClient extends VaultAccessor { this(new RestTemplate(), new VaultEndpoint()); } + /** + * Creates a new {@link VaultClient} for a {@link ClientHttpRequestFactory} and {@link VaultEndpoint}. + * + * @param requestFactory must not be {@literal null}. + * @param endpoint must not be {@literal null}. + */ + public VaultClient(ClientHttpRequestFactory requestFactory, VaultEndpoint endpoint) { + + super(newRestTemplate(requestFactory)); + + Assert.notNull(endpoint, "VaultEndpoint must not be null"); + this.endpoint = endpoint; + } + + /** + * Create a {@link RestTemplate} using an interceptor given a {@link ClientHttpRequestFactory}. This forces + * {@link RestTemplate} to create the body representation instead of streaming the body to the TCP channel. Streaming + * the body without knowing the size in advance will skip the {@link HttpHeaders#CONTENT_LENGTH} makes Vault upset. + * + * @param requestFactory must not be {@literal null}. + * @return the {@link RestTemplate} + */ + private static RestTemplate newRestTemplate(ClientHttpRequestFactory requestFactory) { + + Assert.notNull(requestFactory, "ClientHttpRequestFactory must not be null"); + + RestTemplate restTemplate = new RestTemplate(requestFactory); + restTemplate.getInterceptors().add(new ClientHttpRequestInterceptor() { + + @Override + public ClientHttpResponse intercept(HttpRequest request, byte[] body, ClientHttpRequestExecution execution) + throws IOException { + return execution.execute(request, body); + } + }); + + return restTemplate; + } + /** * Creates a new {@link VaultClient} for a {@link RestTemplate} and {@link VaultEndpoint}. * @@ -193,6 +237,9 @@ public class VaultClient extends VaultAccessor { public VaultResponseEntity exchange(String pathTemplate, HttpMethod method, HttpEntity requestEntity, Class responseType, Map uriVariables) throws RestClientException { + Assert.hasText(pathTemplate, "Path template must not be null or empty"); + Assert.isTrue(!pathTemplate.startsWith("/"), "Path template must not start with a slash (/)"); + URI uri = uriVariables != null ? buildUri(pathTemplate, uriVariables) : getEndpoint().createUri(pathTemplate); return exchange(uri, method, requestEntity, responseType); @@ -219,6 +266,9 @@ public class VaultClient extends VaultAccessor { HttpEntity requestEntity, ParameterizedTypeReference responseType, Map uriVariables) throws RestClientException { + Assert.hasText(pathTemplate, "Path template must not be null or empty"); + Assert.isTrue(!pathTemplate.startsWith("/"), "Path template must not start with a slash (/)"); + URI uri = uriVariables != null ? buildUri(pathTemplate, uriVariables) : getEndpoint().createUri(pathTemplate); return exchange(uri, method, requestEntity, responseType); @@ -235,6 +285,9 @@ public class VaultClient extends VaultAccessor { */ public T doWithRestTemplate(String pathTemplate, Map uriVariables, RestTemplateCallback callback) { + Assert.hasText(pathTemplate, "Path template must not be null or empty"); + Assert.isTrue(!pathTemplate.startsWith("/"), "Path template must not start with a slash (/)"); + URI uri = uriVariables != null ? buildUri(pathTemplate, uriVariables) : getEndpoint().createUri(pathTemplate); return super.doWithRestTemplate(uri, callback); diff --git a/spring-vault-core/src/main/java/org/springframework/vault/config/AbstractVaultConfiguration.java b/spring-vault-core/src/main/java/org/springframework/vault/config/AbstractVaultConfiguration.java index fc78fb1d..8445db6f 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/config/AbstractVaultConfiguration.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/config/AbstractVaultConfiguration.java @@ -32,7 +32,6 @@ import org.springframework.vault.core.VaultClientFactory; import org.springframework.vault.core.VaultTemplate; import org.springframework.vault.support.ClientOptions; import org.springframework.vault.support.SslConfiguration; -import org.springframework.web.client.RestTemplate; /** * Base class for Spring Vault configuration using JavaConfig. @@ -111,10 +110,7 @@ public abstract class AbstractVaultConfiguration { */ @Bean public VaultClient vaultClient() { - - RestTemplate restTemplate = new RestTemplate(clientHttpRequestFactoryWrapper().getClientHttpRequestFactory()); - - return new VaultClient(restTemplate, vaultEndpoint()); + return new VaultClient(clientHttpRequestFactoryWrapper().getClientHttpRequestFactory(), vaultEndpoint()); } /** diff --git a/spring-vault-core/src/test/java/org/springframework/vault/authentication/ClientCertificateAuthenticationIntegrationTests.java b/spring-vault-core/src/test/java/org/springframework/vault/authentication/ClientCertificateAuthenticationIntegrationTests.java index 98470f00..dffa246c 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/authentication/ClientCertificateAuthenticationIntegrationTests.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/authentication/ClientCertificateAuthenticationIntegrationTests.java @@ -38,7 +38,6 @@ import org.springframework.vault.support.SslConfiguration; import org.springframework.vault.support.VaultToken; import org.springframework.vault.util.IntegrationTestSupport; import org.springframework.vault.util.Settings; -import org.springframework.web.client.RestTemplate; /** * Integration tests for {@link ClientCertificateAuthentication}. @@ -76,7 +75,7 @@ public class ClientCertificateAuthenticationIntegrationTests extends Integration ClientHttpRequestFactory clientHttpRequestFactory = ClientHttpRequestFactoryFactory.create(new ClientOptions(), prepareCertAuthenticationMethod()); - VaultClient vaultClient = new VaultClient(new RestTemplate(clientHttpRequestFactory), new VaultEndpoint()); + VaultClient vaultClient = new VaultClient(clientHttpRequestFactory, new VaultEndpoint()); ClientCertificateAuthentication authentication = new ClientCertificateAuthentication(vaultClient); VaultToken login = authentication.login(); @@ -90,7 +89,7 @@ public class ClientCertificateAuthenticationIntegrationTests extends Integration ClientHttpRequestFactory clientHttpRequestFactory = ClientHttpRequestFactoryFactory.create(new ClientOptions(), Settings.createSslConfiguration()); - VaultClient vaultClient = new VaultClient(new RestTemplate(clientHttpRequestFactory), new VaultEndpoint()); + VaultClient vaultClient = new VaultClient(clientHttpRequestFactory, new VaultEndpoint()); new ClientCertificateAuthentication(vaultClient).login(); } diff --git a/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTokenTemplateIntegrationTests.java b/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTokenTemplateIntegrationTests.java index 61458d8a..7ec86fbd 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTokenTemplateIntegrationTests.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTokenTemplateIntegrationTests.java @@ -147,7 +147,7 @@ public class VaultTokenTemplateIntegrationTests extends IntegrationTestSupport { return vaultOperations.doWithVault(new VaultOperations.ClientCallback>() { @Override public VaultResponseEntity doWithVault(VaultClient client) { - return client.getForEntity("/auth/token/lookup-self", tokenResponse.getToken(), String.class); + return client.getForEntity("auth/token/lookup-self", tokenResponse.getToken(), String.class); } }); } diff --git a/spring-vault-core/src/test/java/org/springframework/vault/util/TestRestTemplateFactory.java b/spring-vault-core/src/test/java/org/springframework/vault/util/TestRestTemplateFactory.java index 2bfecc8e..07200c63 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/util/TestRestTemplateFactory.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/util/TestRestTemplateFactory.java @@ -15,16 +15,20 @@ */ package org.springframework.vault.util; +import java.io.IOException; import java.util.concurrent.atomic.AtomicReference; import org.springframework.beans.factory.DisposableBean; import org.springframework.beans.factory.InitializingBean; +import org.springframework.http.HttpRequest; +import org.springframework.http.client.ClientHttpRequestExecution; import org.springframework.http.client.ClientHttpRequestFactory; +import org.springframework.http.client.ClientHttpRequestInterceptor; +import org.springframework.http.client.ClientHttpResponse; import org.springframework.util.Assert; import org.springframework.vault.config.ClientHttpRequestFactoryFactory; import org.springframework.vault.support.ClientOptions; import org.springframework.vault.support.SslConfiguration; -import org.springframework.web.client.DefaultResponseErrorHandler; import org.springframework.web.client.RestTemplate; /** @@ -70,11 +74,17 @@ public class TestRestTemplateFactory { Assert.notNull(requestFactory, "ClientHttpRequestFactory must not be null!"); - RestTemplate RestTemplate = new RestTemplate(); - RestTemplate.setErrorHandler(new DefaultResponseErrorHandler()); - RestTemplate.setRequestFactory(requestFactory); + RestTemplate template = new RestTemplate(); + template.setRequestFactory(requestFactory); + template.getInterceptors().add(new ClientHttpRequestInterceptor() { + @Override + public ClientHttpResponse intercept(HttpRequest request, byte[] body, ClientHttpRequestExecution execution) + throws IOException { + return execution.execute(request, body); + } + }); - return RestTemplate; + return template; } private static void initializeClientHttpRequestFactory(SslConfiguration sslConfiguration) throws Exception {