From cf9295b88757c29817da2327ab3fe74dff32bf2c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?St=C3=A9phane=20Nicoll?= Date: Tue, 8 Apr 2025 10:52:59 +0200 Subject: [PATCH] Invoke ClientInterceptor#afterCompletion once per request This commit reviews how afterCompletion is invoked. First of all, any exception thrown by an interceptor does not bubble up the stack to give a chance for other interceptors to clean things up. Then, rather than calling afterCompletion before handling a fault, an error, or the response body, it is invoked after. This gives a chance to any exception to break the flow and handle the completion in the catch block. Tests have been added to ensure that the callback is only invoked once. Closes gh-1054 --- .../ws/client/core/WebServiceTemplate.java | 16 +- .../client/core/WebServiceTemplateTest.java | 242 ++++++++++++++++++ 2 files changed, 254 insertions(+), 4 deletions(-) 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); + } + + } + }