diff --git a/spring-vault-core/src/main/java/org/springframework/vault/client/ReactiveVaultClients.java b/spring-vault-core/src/main/java/org/springframework/vault/client/ReactiveVaultClients.java index 68a632cf..aae05465 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/client/ReactiveVaultClients.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/client/ReactiveVaultClients.java @@ -15,6 +15,8 @@ */ package org.springframework.vault.client; +import reactor.core.publisher.Mono; + import org.springframework.core.codec.ByteArrayDecoder; import org.springframework.core.codec.ByteArrayEncoder; import org.springframework.core.codec.StringDecoder; @@ -23,6 +25,8 @@ import org.springframework.http.codec.CodecConfigurer.CustomCodecs; import org.springframework.http.codec.json.Jackson2JsonDecoder; import org.springframework.http.codec.json.Jackson2JsonEncoder; import org.springframework.util.Assert; +import org.springframework.web.reactive.function.client.ClientRequest; +import org.springframework.web.reactive.function.client.ExchangeFilterFunction; import org.springframework.web.reactive.function.client.ExchangeStrategies; import org.springframework.web.reactive.function.client.WebClient; import org.springframework.web.util.UriBuilderFactory; @@ -89,4 +93,31 @@ public class ReactiveVaultClients { return WebClient.builder().uriBuilderFactory(uriBuilderFactory) .exchangeStrategies(strategies).clientConnector(connector).build(); } + + /** + * Create a {@link ExchangeFilterFunction} that associates each request with a + * {@code X-Vault-Namespace} header if the header is not present. + * + * @param namespace the Vault namespace to use. Must not be {@literal null} or empty. + * @return the {@link ExchangeFilterFunction} to register with {@link WebClient}. + * @see VaultHttpHeaders#VAULT_NAMESPACE + * @since 2.2 + */ + public static ExchangeFilterFunction namespace(String namespace) { + + Assert.hasText(namespace, "Vault Namespace must not be empty!"); + + return ExchangeFilterFunction.ofRequestProcessor(request -> { + + return Mono.fromSupplier(() -> { + + return ClientRequest.from(request).headers(headers -> { + + if (!headers.containsKey(VaultHttpHeaders.VAULT_NAMESPACE)) { + headers.add(VaultHttpHeaders.VAULT_NAMESPACE, namespace); + } + }).build(); + }); + }); + } } diff --git a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClients.java b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClients.java index 38ce1ea1..4fa10cce 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClients.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultClients.java @@ -20,6 +20,7 @@ import java.util.ArrayList; import java.util.List; import java.util.Map; +import org.springframework.http.HttpHeaders; import org.springframework.http.client.ClientHttpRequestFactory; import org.springframework.http.client.ClientHttpRequestInterceptor; import org.springframework.http.converter.ByteArrayHttpMessageConverter; @@ -27,6 +28,7 @@ import org.springframework.http.converter.HttpMessageConverter; import org.springframework.http.converter.StringHttpMessageConverter; import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter; import org.springframework.lang.Nullable; +import org.springframework.util.Assert; import org.springframework.web.client.RestTemplate; import org.springframework.web.util.DefaultUriBuilderFactory; import org.springframework.web.util.DefaultUriTemplateHandler; @@ -128,6 +130,32 @@ public class VaultClients { return restTemplate; } + /** + * Create a {@link ClientHttpRequestInterceptor} that associates each request with a + * {@code X-Vault-Namespace} header if the header is not present. + * + * @param namespace the Vault namespace to use. Must not be {@literal null} or empty. + * @return the {@link ClientHttpRequestInterceptor} to register with + * {@link RestTemplate}. + * @see VaultHttpHeaders#VAULT_NAMESPACE + * @since 2.2 + */ + public static ClientHttpRequestInterceptor createNamespaceInterceptor(String namespace) { + + Assert.hasText(namespace, "Vault Namespace must not be empty!"); + + return (request, body, execution) -> { + + HttpHeaders headers = request.getHeaders(); + + if (!headers.containsKey(VaultHttpHeaders.VAULT_NAMESPACE)) { + headers.add(VaultHttpHeaders.VAULT_NAMESPACE, namespace); + } + + return execution.execute(request, body); + }; + } + private static DefaultUriTemplateHandler createUriTemplateHandler( VaultEndpointProvider endpointProvider) { diff --git a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultHttpHeaders.java b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultHttpHeaders.java index 766a229a..ca09f54c 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/client/VaultHttpHeaders.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/client/VaultHttpHeaders.java @@ -31,6 +31,13 @@ public abstract class VaultHttpHeaders { */ public static final String VAULT_TOKEN = "X-Vault-Token"; + /** + * The HTTP {@code X-Vault-Namespace} header field name. + * + * @since 2.2 + */ + public static final String VAULT_NAMESPACE = "X-Vault-Namespace"; + private VaultHttpHeaders() { } diff --git a/spring-vault-core/src/test/java/org/springframework/vault/client/ReactiveVaultClientsUnitTests.java b/spring-vault-core/src/test/java/org/springframework/vault/client/ReactiveVaultClientsUnitTests.java new file mode 100644 index 00000000..24fa4ec5 --- /dev/null +++ b/spring-vault-core/src/test/java/org/springframework/vault/client/ReactiveVaultClientsUnitTests.java @@ -0,0 +1,81 @@ +/* + * Copyright 2018 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 + * + * http://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.vault.client; + +import java.util.Collections; + +import org.junit.Test; +import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; + +import org.springframework.http.HttpMethod; +import org.springframework.http.HttpStatus; +import org.springframework.http.client.reactive.ClientHttpConnector; +import org.springframework.http.client.reactive.ClientHttpRequest; +import org.springframework.mock.http.client.reactive.MockClientHttpRequest; +import org.springframework.mock.http.client.reactive.MockClientHttpResponse; +import org.springframework.web.reactive.function.client.WebClient; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Unit tests for {@link ReactiveVaultClients}. + * + * @author Mark Paluch + */ +public class ReactiveVaultClientsUnitTests { + + @Test + public void shouldApplyNamespace() { + + ClientHttpRequest request = new MockClientHttpRequest(HttpMethod.POST, + "/auth/foo"); + MockClientHttpResponse response = new MockClientHttpResponse(HttpStatus.OK); + + ClientHttpConnector connector = (method, uri, fn) -> fn.apply(request).then( + Mono.just(response)); + + WebClient webClient = WebClient.builder().clientConnector(connector) + .filter(ReactiveVaultClients.namespace("foo/bar")).build(); + + webClient.get().uri("/auth/foo").retrieve().bodyToMono(String.class) + .as(StepVerifier::create).verifyComplete(); + + assertThat(request.getHeaders()).containsEntry(VaultHttpHeaders.VAULT_NAMESPACE, + Collections.singletonList("foo/bar")); + } + + @Test + public void shouldAllowNamespaceOverride() { + + ClientHttpRequest request = new MockClientHttpRequest(HttpMethod.POST, + "/auth/foo"); + MockClientHttpResponse response = new MockClientHttpResponse(HttpStatus.OK); + + ClientHttpConnector connector = (method, uri, fn) -> fn.apply(request).then( + Mono.just(response)); + + WebClient webClient = WebClient.builder().clientConnector(connector) + .filter(ReactiveVaultClients.namespace("foo/bar")).build(); + + webClient.get().uri("/auth/foo").header(VaultHttpHeaders.VAULT_NAMESPACE, "baz") + .retrieve().bodyToMono(String.class).as(StepVerifier::create) + .verifyComplete(); + + assertThat(request.getHeaders()).containsEntry(VaultHttpHeaders.VAULT_NAMESPACE, + Collections.singletonList("baz")); + } +} diff --git a/spring-vault-core/src/test/java/org/springframework/vault/client/VaultClientsUnitTests.java b/spring-vault-core/src/test/java/org/springframework/vault/client/VaultClientsUnitTests.java index b81361aa..9706cca6 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/client/VaultClientsUnitTests.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/client/VaultClientsUnitTests.java @@ -19,13 +19,21 @@ import java.net.URI; import org.junit.Test; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.test.web.client.MockRestServiceServer; import org.springframework.vault.client.VaultClients.PrefixAwareUriTemplateHandler; +import org.springframework.web.client.RestTemplate; import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.header; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.method; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo; +import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess; /** - * Unit tests for - * {@link org.springframework.vault.client.VaultClients.PrefixAwareUriTemplateHandler}. + * Unit tests for {@link org.springframework.vault.client.VaultClients}. * * @author Mark Paluch */ @@ -55,4 +63,42 @@ public class VaultClientsUnitTests { assertThat(uri).hasScheme("https").hasHost("foo").hasPort(-1) .hasPath("/path/bar"); } + + @Test + public void shouldApplyNamespace() { + + RestTemplate restTemplate = VaultClients.createRestTemplate(); + restTemplate.getInterceptors().add( + VaultClients.createNamespaceInterceptor("foo/bar")); + restTemplate.setUriTemplateHandler(new PrefixAwareUriTemplateHandler()); + + MockRestServiceServer mockRest = MockRestServiceServer.createServer(restTemplate); + + mockRest.expect(requestTo("/auth/foo")).andExpect(method(HttpMethod.GET)) + .andExpect(header(VaultHttpHeaders.VAULT_NAMESPACE, "foo/bar")) + .andRespond(withSuccess()); + + restTemplate.getForEntity("/auth/foo", String.class); + } + + @Test + public void shouldAllowNamespaceOverride() { + + RestTemplate restTemplate = VaultClients.createRestTemplate(); + restTemplate.getInterceptors().add( + VaultClients.createNamespaceInterceptor("foo/bar")); + restTemplate.setUriTemplateHandler(new PrefixAwareUriTemplateHandler()); + + MockRestServiceServer mockRest = MockRestServiceServer.createServer(restTemplate); + + mockRest.expect(requestTo("/auth/foo")).andExpect(method(HttpMethod.GET)) + .andExpect(header(VaultHttpHeaders.VAULT_NAMESPACE, "baz")) + .andRespond(withSuccess()); + + HttpHeaders headers = new HttpHeaders(); + headers.add(VaultHttpHeaders.VAULT_NAMESPACE, "baz"); + + restTemplate.exchange("/auth/foo", HttpMethod.GET, new HttpEntity<>(headers), + String.class); + } }