Add zipWith step to AuthenticationSteps.

AuthenticationSteps now allow for combination of two node results.

Closes gh-275.
This commit is contained in:
Mark Paluch
2018-08-01 10:00:10 +02:00
parent 672962645f
commit c544fc76fd
5 changed files with 256 additions and 52 deletions

View File

@@ -87,7 +87,7 @@ public class AuthenticationSteps {
private static final Node<Object> HEAD = new Node<>();
final List<Node<?>> steps = new ArrayList<>();
final List<Node<?>> steps;
/**
* Create a flow definition using a provided {@link VaultToken}.
@@ -146,6 +146,12 @@ public class AuthenticationSteps {
}
AuthenticationSteps(PathAware pathAware) {
this.steps = getChain(pathAware);
}
static List<Node<?>> getChain(PathAware pathAware) {
List<Node<?>> steps = new ArrayList<>();
PathAware current = pathAware;
do {
@@ -164,6 +170,8 @@ public class AuthenticationSteps {
while (!Objects.equals(current, AuthenticationSteps.HEAD));
Collections.reverse(steps);
return steps;
}
/**
@@ -189,6 +197,20 @@ public class AuthenticationSteps {
return new MapStep<>(mappingFunction, this);
}
/**
* Combine the result from this {@link Node} and another into a {@link Pair}.
*
* @return the next {@link Node}.
* @since 2.1
*/
public <R> Node<Pair<T, R>> zipWith(Node<? extends R> other) {
Assert.notNull(other, "Other node must not be null");
Assert.isInstanceOf(PathAware.class, other, "Other node must be PathAware");
return new ZipStep<>(this, (PathAware) other);
}
/**
* Callback with the current state object.
*
@@ -467,6 +489,32 @@ public class AuthenticationSteps {
}
}
@Value
@EqualsAndHashCode(callSuper = false)
static class ZipStep<L, R> extends Node<Pair<L, R>> implements PathAware {
@NonNull
Node<?> left;
@NonNull
List<Node<?>> right;
public ZipStep(Node<?> left, PathAware right) {
this.left = left;
this.right = getChain(right);
}
@Override
public Node<?> getPrevious() {
return left;
}
@Override
public String toString() {
return "Zip";
}
}
@Value
@EqualsAndHashCode(callSuper = false)
@RequiredArgsConstructor(access = AccessLevel.PACKAGE)
@@ -511,4 +559,52 @@ public class AuthenticationSteps {
interface PathAware {
Node<?> getPrevious();
}
/**
* A tuple of two things.
*
* @param <L>
* @param <R>
* @since 2.1
*/
public static class Pair<L, R> {
private final L left;
private final R right;
private Pair(L left, R right) {
this.left = left;
this.right = right;
}
/**
* Create a new {@link Pair} given {@code left} and {@code right} values.
*
* @param left
* @param right
* @return the {@link Pair}.
*/
public static <L, R> Pair<L, R> of(L left, R right) {
return new Pair<>(left, right);
}
/**
* Type-safe way to get the fist object of this {@link Pair}.
*
* @return The first object
*/
public L getLeft() {
return left;
}
/**
* Type-safe way to get the second object of this {@link Pair}.
*
* @return The second object
*/
public R getRight() {
return right;
}
}
}

View File

@@ -28,7 +28,9 @@ import org.springframework.vault.authentication.AuthenticationSteps.HttpRequestN
import org.springframework.vault.authentication.AuthenticationSteps.MapStep;
import org.springframework.vault.authentication.AuthenticationSteps.Node;
import org.springframework.vault.authentication.AuthenticationSteps.OnNextStep;
import org.springframework.vault.authentication.AuthenticationSteps.Pair;
import org.springframework.vault.authentication.AuthenticationSteps.SupplierStep;
import org.springframework.vault.authentication.AuthenticationSteps.ZipStep;
import org.springframework.vault.client.VaultResponses;
import org.springframework.vault.support.VaultResponse;
import org.springframework.vault.support.VaultToken;
@@ -72,9 +74,32 @@ public class AuthenticationStepsExecutor implements ClientAuthentication {
@SuppressWarnings("unchecked")
public VaultToken login() throws VaultException {
Iterable<Node<?>> steps = chain.steps;
Object state = evaluate(steps);
if (state instanceof VaultToken) {
return (VaultToken) state;
}
if (state instanceof VaultResponse) {
VaultResponse response = (VaultResponse) state;
Assert.state(response.getAuth() != null, "Auth field must not be null");
return LoginTokenUtil.from(response.getAuth());
}
throw new IllegalStateException(String.format(
"Cannot retrieve VaultToken from authentication chain. Got instead %s",
state));
}
private Object evaluate(Iterable<Node<?>> steps) {
Object state = null;
for (Node<?> o : chain.steps) {
for (Node<?> o : steps) {
if (logger.isDebugEnabled()) {
logger.debug(String
@@ -86,15 +111,19 @@ public class AuthenticationStepsExecutor implements ClientAuthentication {
state = doHttpRequest((HttpRequestNode<Object>) o, state);
}
if (o instanceof AuthenticationSteps.MapStep) {
if (o instanceof MapStep) {
state = doMapStep((MapStep<Object, Object>) o, state);
}
if (o instanceof ZipStep) {
state = doZipStep((ZipStep<Object, Object>) o, state);
}
if (o instanceof OnNextStep) {
state = doOnNext((OnNextStep<Object>) o, state);
}
if (o instanceof AuthenticationSteps.SupplierStep<?>) {
if (o instanceof SupplierStep<?>) {
state = doSupplierStep((SupplierStep<Object>) o);
}
@@ -114,21 +143,7 @@ public class AuthenticationStepsExecutor implements ClientAuthentication {
"Authentication execution failed in %s", o), e);
}
}
if (state instanceof VaultToken) {
return (VaultToken) state;
}
if (state instanceof VaultResponse) {
VaultResponse response = (VaultResponse) state;
Assert.state(response.getAuth() != null, "Auth field must not be null");
return LoginTokenUtil.from(response.getAuth());
}
throw new IllegalStateException(String.format(
"Cannot retrieve VaultToken from authentication chain. Got instead %s",
state));
return state;
}
private static Object doSupplierStep(SupplierStep<Object> supplierStep) {
@@ -139,6 +154,12 @@ public class AuthenticationStepsExecutor implements ClientAuthentication {
return o.apply(state);
}
private Object doZipStep(ZipStep<Object, Object> o, Object state) {
Object result = evaluate(o.getRight());
return Pair.of(state, result);
}
private static Object doOnNext(OnNextStep<Object> o, Object state) {
return o.apply(state);
}

View File

@@ -30,7 +30,9 @@ import org.springframework.vault.authentication.AuthenticationSteps.HttpRequestN
import org.springframework.vault.authentication.AuthenticationSteps.MapStep;
import org.springframework.vault.authentication.AuthenticationSteps.Node;
import org.springframework.vault.authentication.AuthenticationSteps.OnNextStep;
import org.springframework.vault.authentication.AuthenticationSteps.Pair;
import org.springframework.vault.authentication.AuthenticationSteps.SupplierStep;
import org.springframework.vault.authentication.AuthenticationSteps.ZipStep;
import org.springframework.vault.support.VaultResponse;
import org.springframework.vault.support.VaultToken;
import org.springframework.web.reactive.function.client.WebClient;
@@ -76,39 +78,7 @@ public class AuthenticationStepsOperator implements VaultTokenSupplier {
@SuppressWarnings("unchecked")
public Mono<VaultToken> getVaultToken() throws VaultException {
Mono<Object> state = Mono.just(Undefinded.INSTANCE);
for (Node<?> o : chain.steps) {
if (logger.isDebugEnabled()) {
logger.debug(String
.format("Executing %s with current state %s", o, state));
}
if (o instanceof HttpRequestNode) {
state = state.flatMap(stateObject -> doHttpRequest(
(HttpRequestNode<Object>) o, stateObject));
}
if (o instanceof AuthenticationSteps.MapStep) {
state = state.map(stateObject -> doMapStep((MapStep<Object, Object>) o,
stateObject));
}
if (o instanceof OnNextStep) {
state = state.doOnNext(stateObject -> doOnNext((OnNextStep<Object>) o,
stateObject));
}
if (o instanceof AuthenticationSteps.SupplierStep<?>) {
state = state
.map(stateObject -> doSupplierStep((SupplierStep<Object>) o));
}
if (logger.isDebugEnabled()) {
logger.debug(String.format("Executed %s with current state %s", o, state));
}
}
Mono<Object> state = createMono(chain.steps);
return state
.map(stateObject -> {
@@ -137,6 +107,49 @@ public class AuthenticationStepsOperator implements VaultTokenSupplier {
"Cannot retrieve VaultToken from authentication chain", t));
}
private Mono<Object> createMono(Iterable<Node<?>> steps) {
Mono<Object> state = Mono.just(Undefinded.INSTANCE);
for (Node<?> o : steps) {
if (logger.isDebugEnabled()) {
logger.debug(String
.format("Executing %s with current state %s", o, state));
}
if (o instanceof HttpRequestNode) {
state = state.flatMap(stateObject -> doHttpRequest(
(HttpRequestNode<Object>) o, stateObject));
}
if (o instanceof MapStep) {
state = state.map(stateObject -> doMapStep((MapStep<Object, Object>) o,
stateObject));
}
if (o instanceof ZipStep) {
state = state.zipWith(doZipStep((ZipStep<Object, Object>) o)).map(
it -> Pair.of(it.getT1(), it.getT2()));
}
if (o instanceof OnNextStep) {
state = state.doOnNext(stateObject -> doOnNext((OnNextStep<Object>) o,
stateObject));
}
if (o instanceof SupplierStep<?>) {
state = state
.map(stateObject -> doSupplierStep((SupplierStep<Object>) o));
}
if (logger.isDebugEnabled()) {
logger.debug(String.format("Executed %s with current state %s", o, state));
}
}
return state;
}
private static Object doSupplierStep(SupplierStep<Object> supplierStep) {
return supplierStep.get();
}
@@ -145,6 +158,10 @@ public class AuthenticationStepsOperator implements VaultTokenSupplier {
return o.apply(state);
}
private Mono<Object> doZipStep(ZipStep<Object, Object> o) {
return createMono(o.getRight());
}
private static Object doOnNext(OnNextStep<Object> o, Object state) {
return o.apply(state);
}

View File

@@ -25,6 +25,7 @@ import org.springframework.http.HttpMethod;
import org.springframework.http.MediaType;
import org.springframework.test.web.client.MockRestServiceServer;
import org.springframework.vault.VaultException;
import org.springframework.vault.authentication.AuthenticationSteps.Node;
import org.springframework.vault.client.VaultClients;
import org.springframework.vault.client.VaultClients.PrefixAwareUriTemplateHandler;
import org.springframework.vault.support.VaultResponse;
@@ -165,6 +166,34 @@ public class AuthenticationStepsExecutorUnitTests {
assertThat(login(steps)).isEqualTo(VaultToken.of("foo-token"));
}
@Test
public void zipWithShouldRequestTwoItems() {
mockRest.expect(requestTo("/auth/login/left"))
.andExpect(method(HttpMethod.POST))
.andRespond(
withSuccess().contentType(MediaType.APPLICATION_JSON).body(
"{" + "\"request_id\": \"left\"}"));
mockRest.expect(requestTo("/auth/login/right"))
.andExpect(method(HttpMethod.POST))
.andRespond(
withSuccess().contentType(MediaType.APPLICATION_JSON).body(
"{" + "\"request_id\": \"right\"}"));
Node<VaultResponse> left = AuthenticationSteps.fromHttpRequest(post(
"/auth/login/left").as(VaultResponse.class));
Node<VaultResponse> right = AuthenticationSteps.fromHttpRequest(post(
"/auth/login/right").as(VaultResponse.class));
AuthenticationSteps steps = left.zipWith(right).login(
it -> VaultToken.of(it.getLeft().getRequestId() + "-"
+ it.getRight().getRequestId()));
assertThat(login(steps)).isEqualTo(VaultToken.of("left-right"));
}
private VaultToken login(AuthenticationSteps steps) {
return new AuthenticationStepsExecutor(steps, restTemplate).login();
}

View File

@@ -27,6 +27,7 @@ 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.vault.authentication.AuthenticationSteps.Node;
import org.springframework.vault.support.VaultResponse;
import org.springframework.vault.support.VaultToken;
import org.springframework.web.reactive.function.client.WebClient;
@@ -103,6 +104,46 @@ public class AuthenticationStepsOperatorUnitTests {
StepVerifier.create(login(steps, webClient)).expectError().verify();
}
@Test
public void zipWithShouldRequestTwoItems() {
ClientHttpRequest leftRequest = new MockClientHttpRequest(HttpMethod.GET,
"/auth/login/left");
MockClientHttpResponse leftResponse = new MockClientHttpResponse(HttpStatus.OK);
leftResponse.getHeaders().setContentType(MediaType.APPLICATION_JSON);
leftResponse.setBody("{" + "\"request_id\": \"left\"}");
ClientHttpRequest rightRequest = new MockClientHttpRequest(HttpMethod.GET,
"/auth/login/right");
MockClientHttpResponse rightResponse = new MockClientHttpResponse(HttpStatus.OK);
rightResponse.getHeaders().setContentType(MediaType.APPLICATION_JSON);
rightResponse.setBody("{" + "\"request_id\": \"right\"}");
ClientHttpConnector connector = (method, uri, fn) -> {
if (uri.toString().contains("left")) {
return fn.apply(leftRequest).then(Mono.just(leftResponse));
}
return fn.apply(rightRequest).then(Mono.just(rightResponse));
};
WebClient webClient = WebClient.builder().clientConnector(connector).build();
Node<VaultResponse> left = AuthenticationSteps.fromHttpRequest(post(
"/auth/login/left").as(VaultResponse.class));
Node<VaultResponse> right = AuthenticationSteps.fromHttpRequest(post(
"/auth/login/right").as(VaultResponse.class));
AuthenticationSteps steps = left.zipWith(right).login(
it -> VaultToken.of(it.getLeft().getRequestId() + "-"
+ it.getRight().getRequestId()));
StepVerifier.create(login(steps, webClient))
.expectNext(VaultToken.of("left-right")).verifyComplete();
}
private Mono<VaultToken> login(AuthenticationSteps steps) {
AuthenticationStepsOperator operator = new AuthenticationStepsOperator(steps,