From 196d3c7470867b8273b19c96db26bad3948763f2 Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Mon, 16 Oct 2023 14:24:59 +0200 Subject: [PATCH] Polishing. Extract method to avoid code duplications. Add author and since tags. Add reactive batch rewrap support. See: gh-687 Original pull request: gh-819 --- .../core/ReactiveVaultTransitOperations.java | 9 +++ .../core/ReactiveVaultTransitTemplate.java | 22 +++++- .../vault/core/VaultTransitOperations.java | 2 + .../vault/core/VaultTransitTemplate.java | 74 ++++++------------- .../ReactiveVaultTransitIntegrationTests.java | 55 ++++++++++++-- .../VaultTransitTemplateIntegrationTests.java | 11 ++- 6 files changed, 104 insertions(+), 69 deletions(-) diff --git a/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitOperations.java b/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitOperations.java index d8f36dc5..c3799686 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitOperations.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitOperations.java @@ -214,6 +214,15 @@ public interface ReactiveVaultTransitOperations { */ Mono rewrap(String keyName, String ciphertext, VaultTransitContext transitContext); + /** + * Rewrap the provided batch of cipher text using the latest version of the named key. + * @param batchRequest a list of {@link Ciphertext} which includes cipher text and a + * context + * @return the rewrapped result in the order of {@code batchRequest} ciphertexts. + * @see #rewrap(String, String) + */ + Flux rewrap(String keyName, List batchRequest); + /** * Create a HMAC using {@code keyName} of given {@link Plaintext} using the default * hash algorithm. The key can be of any type supported by transit; the raw key will diff --git a/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitTemplate.java b/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitTemplate.java index 35ec5b29..cadf83cf 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitTemplate.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/core/ReactiveVaultTransitTemplate.java @@ -38,6 +38,7 @@ import org.springframework.vault.support.VaultTransitKeyCreationRequest; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import java.util.ArrayList; import java.util.Base64; import java.util.Collections; import java.util.LinkedHashMap; @@ -223,9 +224,7 @@ public class ReactiveVaultTransitTemplate implements ReactiveVaultTransitOperati Assert.hasText(ciphertext, "Ciphertext must not be empty"); Assert.notNull(transitContext, "VaultTransitContext must not be null"); - Map request = new LinkedHashMap<>(); - - request.put("ciphertext", ciphertext); + Map request = createRewrapRequest(toCiphertext(ciphertext, transitContext)); applyTransitOptions(transitContext, request); @@ -233,8 +232,23 @@ public class ReactiveVaultTransitTemplate implements ReactiveVaultTransitOperati .map(response -> (String) response.getRequiredData().get("ciphertext")); } + @Override + public Flux rewrap(String keyName, List batchRequest) { + + Assert.hasText(keyName, "Key name must not be empty"); + Assert.notEmpty(batchRequest, "BatchRequest must not be null and must have at least one entry"); + + return Flux.fromIterable(batchRequest) + .map(VaultTransitTemplate::createRewrapRequest) + .collectList() + .flatMap(batch -> this.reactiveVaultOperations.write(String.format("%s/rewrap/%s", this.path, keyName), + Collections.singletonMap("batch_input", batch))) + .flatMapIterable(vaultResponse -> toBatchResults(vaultResponse, batchRequest, Ciphertext::getContext)); + } + @Override public Flux encrypt(String keyName, List batchRequest) { + Assert.hasText(keyName, "Key name must not be empty"); Assert.notEmpty(batchRequest, "BatchRequest must not be null and must have at least one entry"); @@ -247,7 +261,7 @@ public class ReactiveVaultTransitTemplate implements ReactiveVaultTransitOperati .collectList() .flatMap(batch -> this.reactiveVaultOperations.write(String.format("%s/encrypt/%s", this.path, keyName), Collections.singletonMap("batch_input", batch))) - .flatMapIterable(vaultResponse -> toEncryptionResults(vaultResponse, batchRequest)); + .flatMapIterable(vaultResponse -> toBatchResults(vaultResponse, batchRequest, Plaintext::getContext)); } @Override diff --git a/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitOperations.java b/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitOperations.java index e6cc1dda..bbfa76d5 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitOperations.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitOperations.java @@ -42,6 +42,7 @@ import org.springframework.vault.support.VaultTransitKeyCreationRequest; * @author Sven Schürmann * @author Praveendra Singh * @author Luander Ribeiro + * @author Nanne Baars * @see <a href="https://www.vaultproject.io/docs/secrets/transit/index.html">Transit * Secret Backend</a> */ @@ -227,6 +228,7 @@ public interface VaultTransitOperations { * context * @return the rewrapped result in the order of {@code batchRequest} ciphertexts. * @see #rewrap(String, String) + * @since 3.1 */ List<VaultEncryptionResult> rewrap(String keyName, List<Ciphertext> batchRequest); diff --git a/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitTemplate.java b/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitTemplate.java index bba7833a..7d855c88 100644 --- a/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitTemplate.java +++ b/spring-vault-core/src/main/java/org/springframework/vault/core/VaultTransitTemplate.java @@ -22,8 +22,10 @@ import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.function.Function; import com.fasterxml.jackson.annotation.JsonProperty; + import org.springframework.lang.Nullable; import org.springframework.util.Assert; import org.springframework.util.ObjectUtils; @@ -227,7 +229,7 @@ public class VaultTransitTemplate implements VaultTransitOperations { VaultResponse vaultResponse = this.vaultOperations.write(String.format("%s/encrypt/%s", this.path, keyName), Collections.singletonMap("batch_input", batch)); - return toEncryptionResults(vaultResponse, batchRequest); + return toBatchResults(vaultResponse, batchRequest, Plaintext::getContext); } @Override @@ -328,11 +330,7 @@ public class VaultTransitTemplate implements VaultTransitOperations { Assert.hasText(ciphertext, "Ciphertext must not be empty"); Assert.notNull(transitContext, "VaultTransitContext must not be null"); - Map<String, String> request = new LinkedHashMap<>(); - - request.put("ciphertext", ciphertext); - - applyTransitOptions(transitContext, request); + Map<String, String> request = createRewrapRequest(toCiphertext(ciphertext, transitContext)); return (String) this.vaultOperations.write(String.format("%s/rewrap/%s", this.path, keyName), request) .getRequiredData() @@ -341,6 +339,7 @@ public class VaultTransitTemplate implements VaultTransitOperations { @Override public List<VaultEncryptionResult> rewrap(String keyName, List<Ciphertext> batchRequest) { + Assert.hasText(keyName, "Key name must not be empty"); Assert.notEmpty(batchRequest, "BatchRequest must not be null and must have at least one entry"); @@ -348,21 +347,14 @@ public class VaultTransitTemplate implements VaultTransitOperations { for (Ciphertext request : batchRequest) { - Map<String, String> vaultRequest = new LinkedHashMap<>(2); - - vaultRequest.put("ciphertext", request.getCiphertext()); - - if (request.getContext() != null) { - applyTransitOptions(request.getContext(), vaultRequest); - } - + Map<String, String> vaultRequest = createRewrapRequest(request); batch.add(vaultRequest); } VaultResponse vaultResponse = this.vaultOperations.write(String.format("%s/rewrap/%s", this.path, keyName), Collections.singletonMap("batch_input", batch)); - return toRewrappedEncryptionResults(vaultResponse, batchRequest); + return toBatchResults(vaultResponse, batchRequest, Ciphertext::getContext); } @Override @@ -509,45 +501,16 @@ public class VaultTransitTemplate implements VaultTransitOperations { } } - static List<VaultEncryptionResult> toEncryptionResults(VaultResponse vaultResponse, List<Plaintext> batchRequest) { + static <T> List<VaultEncryptionResult> toBatchResults(VaultResponse vaultResponse, List<T> batchRequests, + Function<T, VaultTransitContext> contextExtractor) { - List<VaultEncryptionResult> result = new ArrayList<>(batchRequest.size()); + List<VaultEncryptionResult> result = new ArrayList<>(batchRequests.size()); List<Map<String, String>> batchData = getBatchData(vaultResponse); - for (int i = 0; i < batchRequest.size(); i++) { + for (int i = 0; i < batchRequests.size(); i++) { VaultEncryptionResult encrypted; - Plaintext plaintext = batchRequest.get(i); - if (batchData.size() > i) { - - Map<String, String> data = batchData.get(i); - if (StringUtils.hasText(data.get("error"))) { - encrypted = new VaultEncryptionResult(new VaultException(data.get("error"))); - } - else { - encrypted = new VaultEncryptionResult(toCiphertext(data.get("ciphertext"), plaintext.getContext())); - } - } - else { - encrypted = new VaultEncryptionResult(new VaultException("No result for plaintext #" + i)); - } - - result.add(encrypted); - } - - return result; - } - - static List<VaultEncryptionResult> toRewrappedEncryptionResults(VaultResponse vaultResponse, - List<Ciphertext> batchRequest) { - - List<VaultEncryptionResult> result = new ArrayList<>(batchRequest.size()); - List<Map<String, String>> batchData = getBatchData(vaultResponse); - - for (int i = 0; i < batchRequest.size(); i++) { - - VaultEncryptionResult encrypted; - Ciphertext ciphertext = batchRequest.get(i); + T request = batchRequests.get(i); if (batchData.size() > i) { Map<String, String> data = batchData.get(i); @@ -556,11 +519,11 @@ public class VaultTransitTemplate implements VaultTransitOperations { } else { encrypted = new VaultEncryptionResult( - toCiphertext(data.get("ciphertext"), ciphertext.getContext())); + toCiphertext(data.get("ciphertext"), contextExtractor.apply(request))); } } else { - encrypted = new VaultEncryptionResult(new VaultException("No result for cipher text #" + i)); + encrypted = new VaultEncryptionResult(new VaultException("No result for request #" + i)); } result.add(encrypted); @@ -607,6 +570,15 @@ public class VaultTransitTemplate implements VaultTransitOperations { return new VaultDecryptionResult(Plaintext.empty().with(ciphertext.getContext())); } + static Map<String, String> createRewrapRequest(Ciphertext request) { + + Map<String, String> vaultRequest = new LinkedHashMap<>(2); + vaultRequest.put("ciphertext", request.getCiphertext()); + applyTransitOptions(request.getContext(), vaultRequest); + + return vaultRequest; + } + static Ciphertext toCiphertext(String ciphertext, @Nullable VaultTransitContext context) { return context != null ? Ciphertext.of(ciphertext).with(context) : Ciphertext.of(ciphertext); } diff --git a/spring-vault-core/src/test/java/org/springframework/vault/core/ReactiveVaultTransitIntegrationTests.java b/spring-vault-core/src/test/java/org/springframework/vault/core/ReactiveVaultTransitIntegrationTests.java index e68d77bd..3fe57e6f 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/core/ReactiveVaultTransitIntegrationTests.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/core/ReactiveVaultTransitIntegrationTests.java @@ -15,10 +15,20 @@ */ package org.springframework.vault.core; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import java.util.Objects; +import java.util.stream.Stream; + import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; + import org.springframework.beans.factory.annotation.Autowired; import org.springframework.test.context.ContextConfiguration; import org.springframework.test.context.junit.jupiter.SpringExtension; @@ -38,16 +48,15 @@ import org.springframework.vault.support.VaultTransitKeyCreationRequest; import org.springframework.vault.util.IntegrationTestSupport; import org.springframework.vault.util.RequiresVaultVersion; import org.springframework.vault.util.Version; -import reactor.core.publisher.Flux; -import reactor.core.publisher.Mono; -import reactor.test.StepVerifier; -import java.util.Arrays; -import java.util.Collections; -import java.util.List; -import java.util.Objects; import static org.assertj.core.api.Assertions.assertThat; -import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.*; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.AES256_GCM96_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.BATCH_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.ECDSA521_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.ED25519_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.KEY_EXPORT_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.RSA3072_INTRODUCED_IN_VERSION; +import static org.springframework.vault.core.VaultTransitTemplateIntegrationTests.SIGN_VERIFY_INTRODUCED_IN_VERSION; /** * Integration tests for {@link ReactiveVaultTransitTemplate} using the {@code transit} @@ -441,6 +450,36 @@ public class ReactiveVaultTransitIntegrationTests extends IntegrationTestSupport .verifyComplete(); } + @Test + void encryptAndRewrapInBatchShouldCreateCiphertext() { + + this.reactiveTransitOperations + .createKey("mykey", + VaultTransitKeyCreationRequest.builder().convergentEncryption(true).derived(true).build()) + .as(StepVerifier::create) + .verifyComplete(); + + VaultTransitContext transitRequest = VaultTransitContext.builder() // + .context("blubb".getBytes()) // + .nonce("123456789012".getBytes()) // + .build(); + + String ciphertext1 = this.reactiveTransitOperations.encrypt("mykey", "hello-world".getBytes(), transitRequest) + .block(); + String ciphertext2 = this.reactiveTransitOperations.encrypt("mykey", "hello-vault".getBytes(), transitRequest) + .block(); + this.reactiveTransitOperations.rotate("mykey").as(StepVerifier::create).verifyComplete(); + + List<Ciphertext> batchRequest = Stream.of(ciphertext1, ciphertext2) + .map(ct -> Ciphertext.of(ct).with(transitRequest)) + .toList(); + this.reactiveTransitOperations.rewrap("mykey", batchRequest) + .as(StepVerifier::create) + .assertNext(it -> assertThat(it.get().getCiphertext()).startsWith("vault:v2")) + .expectNextCount(1) + .verifyComplete(); + } + @Test void shouldEncryptBinaryPlaintext() { diff --git a/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTransitTemplateIntegrationTests.java b/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTransitTemplateIntegrationTests.java index d6f8f1ec..780a8cdb 100644 --- a/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTransitTemplateIntegrationTests.java +++ b/spring-vault-core/src/test/java/org/springframework/vault/core/VaultTransitTemplateIntegrationTests.java @@ -55,6 +55,7 @@ import org.springframework.vault.util.IntegrationTestSupport; import org.springframework.vault.util.RequiresVaultVersion; import org.springframework.vault.util.Version; +import static org.assertj.core.api.Assertions.*; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatExceptionOfType; import static org.assertj.core.api.Assertions.fail; @@ -67,6 +68,7 @@ import static org.assertj.core.api.Assertions.fail; * @author Praveendra Singh * @author Luander Ribeiro * @author Mikko Koli + * @author Nanne Baars */ @ExtendWith(SpringExtension.class) @ContextConfiguration(classes = VaultIntegrationTestConfiguration.class) @@ -372,7 +374,7 @@ class VaultTransitTemplateIntegrationTests extends IntegrationTestSupport { assertThat(ciphertext).startsWith("vault:%s:".formatted(expectedKeyPrefix)); } catch (Exception e) { - Assertions.assertThat(expectedKeyPrefix).isNullOrEmpty(); + assertThat(expectedKeyPrefix).isNullOrEmpty(); } } @@ -549,14 +551,11 @@ class VaultTransitTemplateIntegrationTests extends IntegrationTestSupport { String ciphertext2 = this.transitOperations.encrypt("mykey", "hello-vault".getBytes(), transitRequest); this.transitOperations.rotate("mykey"); - List<Ciphertext> batchRequest = List.of(ciphertext1, ciphertext2) - .stream() + List<Ciphertext> batchRequest = Stream.of(ciphertext1, ciphertext2) .map(ct -> Ciphertext.of(ct).with(transitRequest)) .toList(); List<VaultEncryptionResult> rewrappedResult = this.transitOperations.rewrap("mykey", batchRequest); - Assertions.assertThat(rewrappedResult) - .hasSize(2) - .allMatch(result -> result.get().getCiphertext().startsWith("vault:v2")); + assertThat(rewrappedResult).hasSize(2).allMatch(result -> result.get().getCiphertext().startsWith("vault:v2")); } @Test