diff --git a/spring-ws-core/src/main/java/org/springframework/ws/client/core/WebServiceTemplate.java b/spring-ws-core/src/main/java/org/springframework/ws/client/core/WebServiceTemplate.java index b87a1da8..29c0f996 100644 --- a/spring-ws-core/src/main/java/org/springframework/ws/client/core/WebServiceTemplate.java +++ b/spring-ws-core/src/main/java/org/springframework/ws/client/core/WebServiceTemplate.java @@ -627,8 +627,9 @@ public class WebServiceTemplate extends WebServiceAccessor implements WebService if (!messageContext.hasResponse() && !intercepted) { sendRequest(connection, messageContext.getRequest()); if (hasError(connection, messageContext.getRequest())) { + Object fallback = handleError(connection, messageContext.getRequest()); triggerAfterCompletion(interceptorIndex, messageContext, null); - return (T) handleError(connection, messageContext.getRequest()); + return (T) fallback; } WebServiceMessage response = connection.receive(getMessageFactory()); messageContext.setResponse(response); @@ -637,13 +638,15 @@ public class WebServiceTemplate extends WebServiceAccessor implements WebService if (messageContext.hasResponse()) { if (!hasFault(connection, messageContext.getResponse())) { triggerHandleResponse(interceptorIndex, messageContext); + T result = responseExtractor.extractData(messageContext.getResponse()); triggerAfterCompletion(interceptorIndex, messageContext, null); - return responseExtractor.extractData(messageContext.getResponse()); + return result; } else { triggerHandleFault(interceptorIndex, messageContext); + Object fallback = handleFault(connection, messageContext); triggerAfterCompletion(interceptorIndex, messageContext, null); - return (T) handleFault(connection, messageContext); + return (T) fallback; } } else { @@ -820,7 +823,12 @@ public class WebServiceTemplate extends WebServiceAccessor implements WebService throws WebServiceClientException { if (this.interceptors != null) { for (int i = interceptorIndex; i >= 0; i--) { - this.interceptors[i].afterCompletion(messageContext, ex); + try { + this.interceptors[i].afterCompletion(messageContext, ex); + } + catch (Exception interceptorEx) { + logger.error("ClientInterceptor.afterCompletion threw exception", ex); + } } } } diff --git a/spring-ws-core/src/test/java/org/springframework/ws/client/core/WebServiceTemplateTest.java b/spring-ws-core/src/test/java/org/springframework/ws/client/core/WebServiceTemplateTest.java index 5707ea50..520e791d 100644 --- a/spring-ws-core/src/test/java/org/springframework/ws/client/core/WebServiceTemplateTest.java +++ b/spring-ws-core/src/test/java/org/springframework/ws/client/core/WebServiceTemplateTest.java @@ -16,11 +16,17 @@ package org.springframework.ws.client.core; +import java.io.IOException; import java.net.URI; import javax.xml.transform.Result; import javax.xml.transform.Source; +import javax.xml.transform.TransformerException; +import org.assertj.core.api.AbstractObjectAssert; +import org.assertj.core.api.AbstractThrowableAssert; +import org.assertj.core.api.AssertProvider; +import org.assertj.core.api.Assertions; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -29,6 +35,7 @@ import org.springframework.oxm.Unmarshaller; import org.springframework.ws.MockWebServiceMessage; import org.springframework.ws.MockWebServiceMessageFactory; import org.springframework.ws.WebServiceMessage; +import org.springframework.ws.client.WebServiceClientException; import org.springframework.ws.client.WebServiceTransportException; import org.springframework.ws.client.support.destination.DestinationProvider; import org.springframework.ws.client.support.interceptor.ClientInterceptor; @@ -43,6 +50,8 @@ import org.springframework.xml.transform.StringSource; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatExceptionOfType; import static org.assertj.core.api.Assertions.assertThatIllegalStateException; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.BDDMockito.given; import static org.mockito.Mockito.isA; import static org.mockito.Mockito.isNull; import static org.mockito.Mockito.mock; @@ -418,6 +427,141 @@ public class WebServiceTemplateTest { assertThat(result).isEqualTo(extracted); } + @Test + void afterCompletionInvokedOnlyOnceWithSuccess() throws Exception { + NoOpClientInterceptor clientInterceptor1 = new NoOpClientInterceptor(); + NoOpClientInterceptor clientInterceptor2 = new NoOpClientInterceptor(); + this.template.setInterceptors(new ClientInterceptor[] { clientInterceptor1, clientInterceptor2 }); + + WebServiceMessageCallback requestCallback = mock(WebServiceMessageCallback.class); + requestCallback.doWithMessage(any(WebServiceMessage.class)); + Object extracted = new Object(); + WebServiceMessageExtractor extract = createSimpleExtractor(extracted); + + this.connectionMock.send(isA(WebServiceMessage.class)); + when(this.connectionMock.hasError()).thenReturn(false); + when(this.connectionMock.receive(this.messageFactory)).thenReturn(new MockWebServiceMessage("")); + when(this.connectionMock.hasFault()).thenReturn(false); + this.connectionMock.close(); + + Object result = this.template.sendAndReceive(requestCallback, extract); + + assertThat(result).isEqualTo(extracted); + assertThat(clientInterceptor1).hasHandledExchange(); + assertThat(clientInterceptor2).hasHandledExchange(); + } + + @Test + void afterCompletionInvokedOnlyOnceWitFailureInAfterCompletion() throws Exception { + IllegalStateException testException = new IllegalStateException("test"); + NoOpClientInterceptor clientInterceptor1 = new NoOpClientInterceptor() { + @Override + public void afterCompletion(MessageContext messageContext, Exception ex) throws WebServiceClientException { + super.afterCompletion(messageContext, ex); + throw testException; + } + }; + NoOpClientInterceptor clientInterceptor2 = new NoOpClientInterceptor(); + this.template.setInterceptors(new ClientInterceptor[] { clientInterceptor1, clientInterceptor2 }); + + WebServiceMessageCallback requestCallback = mock(WebServiceMessageCallback.class); + requestCallback.doWithMessage(any(WebServiceMessage.class)); + Object extracted = new Object(); + WebServiceMessageExtractor extract = createSimpleExtractor(extracted); + + this.connectionMock.send(isA(WebServiceMessage.class)); + when(this.connectionMock.hasError()).thenReturn(false); + when(this.connectionMock.receive(this.messageFactory)).thenReturn(new MockWebServiceMessage("")); + when(this.connectionMock.hasFault()).thenReturn(false); + this.connectionMock.close(); + + Object result = this.template.sendAndReceive(requestCallback, extract); + + assertThat(result).isEqualTo(extracted); + assertThat(clientInterceptor1).hasHandledExchange().hasNoCompletionException(); + assertThat(clientInterceptor2).hasHandledExchange().hasNoCompletionException(); + } + + @Test + void afterCompletionInvokedOnlyOnceWithError() throws Exception { + NoOpClientInterceptor clientInterceptor1 = new NoOpClientInterceptor(); + NoOpClientInterceptor clientInterceptor2 = new NoOpClientInterceptor(); + this.template.setInterceptors(new ClientInterceptor[] { clientInterceptor1, clientInterceptor2 }); + + Object extracted = new Object(); + WebServiceMessageExtractor extract = createSimpleExtractor(extracted); + + this.connectionMock.send(isA(WebServiceMessage.class)); + when(this.connectionMock.hasError()).thenReturn(true); + when(this.connectionMock.hasFault()).thenReturn(false); + String errorMessage = "errorMessage"; + when(this.connectionMock.getErrorMessage()).thenReturn(errorMessage); + this.connectionMock.close(); + + assertThatExceptionOfType(WebServiceTransportException.class) + .isThrownBy(() -> this.template.sendAndReceive(null, extract)) + .satisfies(exception -> { + assertThat(clientInterceptor1).hasHandledError().completionException().isSameAs(exception); + assertThat(clientInterceptor2).hasHandledError().completionException().isSameAs(exception); + }); + } + + @Test + void afterCompletionInvokedOnlyOnceWithFault() throws Exception { + NoOpClientInterceptor clientInterceptor1 = new NoOpClientInterceptor(); + NoOpClientInterceptor clientInterceptor2 = new NoOpClientInterceptor(); + this.template.setInterceptors(new ClientInterceptor[] { clientInterceptor1, clientInterceptor2 }); + this.template.setFaultMessageResolver(null); + + WebServiceMessageExtractor extractorMock = createSimpleExtractor(new Object()); + MockWebServiceMessage response = new MockWebServiceMessage(""); + response.setFault(true); + + this.connectionMock.send(isA(WebServiceMessage.class)); + when(this.connectionMock.hasError()).thenReturn(false); + when(this.connectionMock.hasFault()).thenReturn(true); + when(this.connectionMock.receive(this.messageFactory)).thenReturn(response); + this.connectionMock.close(); + + assertThatExceptionOfType(WebServiceTransportException.class) + .isThrownBy(() -> this.template.sendAndReceive(null, extractorMock)) + .satisfies(exception -> { + assertThat(clientInterceptor1).hasHandledFault().completionException().isSameAs(exception); + assertThat(clientInterceptor2).hasHandledFault().completionException().isSameAs(exception); + }); + } + + @Test + void afterCompletionInvokedOnlyOnceWithFaultAndFaultMessageResolver() throws Exception { + NoOpClientInterceptor clientInterceptor1 = new NoOpClientInterceptor(); + NoOpClientInterceptor clientInterceptor2 = new NoOpClientInterceptor(); + this.template.setInterceptors(new ClientInterceptor[] { clientInterceptor1, clientInterceptor2 }); + + WebServiceMessageExtractor extractorMock = createSimpleExtractor(new Object()); + FaultMessageResolver faultMessageResolverMock = mock(FaultMessageResolver.class); + faultMessageResolverMock.resolveFault(isA(WebServiceMessage.class)); + this.template.setFaultMessageResolver(faultMessageResolverMock); + + MockWebServiceMessage response = new MockWebServiceMessage(""); + response.setFault(true); + + this.connectionMock.send(isA(WebServiceMessage.class)); + when(this.connectionMock.hasError()).thenReturn(false); + when(this.connectionMock.hasFault()).thenReturn(true); + when(this.connectionMock.receive(this.messageFactory)).thenReturn(response); + this.connectionMock.close(); + + this.template.sendAndReceive(null, extractorMock); + assertThat(clientInterceptor1).hasHandledFault().hasNoCompletionException(); + assertThat(clientInterceptor2).hasHandledFault().hasNoCompletionException(); + } + + private WebServiceMessageExtractor createSimpleExtractor(T target) throws IOException, TransformerException { + WebServiceMessageExtractor extractor = mock(WebServiceMessageExtractor.class); + given(extractor.extractData(any(WebServiceMessage.class))).willReturn(target); + return extractor; + } + @Test public void testDestinationResolver() throws Exception { @@ -456,4 +600,102 @@ public class WebServiceTemplateTest { assertThat(result).isNull(); } + private static class NoOpClientInterceptor + implements ClientInterceptor, AssertProvider { + + private boolean handledRequest; + + private boolean handledResponse; + + private boolean handledFault; + + private boolean afterCompletion; + + private Exception afterCompletionException; + + @Override + public boolean handleRequest(MessageContext messageContext) throws WebServiceClientException { + if (this.handledRequest) { + throw new IllegalStateException("handleRequest has already been called"); + } + this.handledRequest = true; + return true; + } + + @Override + public boolean handleResponse(MessageContext messageContext) throws WebServiceClientException { + if (this.handledResponse) { + throw new IllegalStateException("handleResponse has already been called"); + } + this.handledResponse = true; + return true; + } + + @Override + public boolean handleFault(MessageContext messageContext) throws WebServiceClientException { + if (this.handledFault) { + throw new IllegalStateException("handleFault has already been called"); + } + this.handledFault = true; + return true; + } + + @Override + public void afterCompletion(MessageContext messageContext, Exception ex) throws WebServiceClientException { + if (this.afterCompletion) { + throw new IllegalStateException("afterCompletion has already been called"); + } + this.afterCompletion = true; + this.afterCompletionException = ex; + } + + @Override + public NoOpClientInterceptorAssert assertThat() { + return new NoOpClientInterceptorAssert(this); + } + + } + + private static class NoOpClientInterceptorAssert + extends AbstractObjectAssert { + + NoOpClientInterceptorAssert(NoOpClientInterceptor actual) { + super(actual, NoOpClientInterceptorAssert.class); + } + + NoOpClientInterceptorAssert hasHandledExchange() { + assertThat(this.actual.handledRequest).isTrue(); + assertThat(this.actual.handledResponse).isTrue(); + assertThat(this.actual.handledFault).isFalse(); + assertThat(this.actual.afterCompletion).isTrue(); + return this.myself; + } + + NoOpClientInterceptorAssert hasHandledError() { + assertThat(this.actual.handledRequest).isTrue(); + assertThat(this.actual.handledResponse).isFalse(); + assertThat(this.actual.handledFault).isFalse(); + assertThat(this.actual.afterCompletion).isTrue(); + return this.myself; + } + + NoOpClientInterceptorAssert hasHandledFault() { + assertThat(this.actual.handledRequest).isTrue(); + assertThat(this.actual.handledResponse).isFalse(); + assertThat(this.actual.handledFault).isTrue(); + assertThat(this.actual.afterCompletion).isTrue(); + return this.myself; + } + + NoOpClientInterceptorAssert hasNoCompletionException() { + assertThat(this.actual.afterCompletionException).isNull(); + return this.myself; + } + + AbstractThrowableAssert completionException() { + return Assertions.assertThat(this.actual.afterCompletionException); + } + + } + }