diff --git a/spring-integration-http/src/main/java/org/springframework/integration/http/inbound/HttpRequestHandlingMessagingGateway.java b/spring-integration-http/src/main/java/org/springframework/integration/http/inbound/HttpRequestHandlingMessagingGateway.java index f470c99c58..c6409184f0 100644 --- a/spring-integration-http/src/main/java/org/springframework/integration/http/inbound/HttpRequestHandlingMessagingGateway.java +++ b/spring-integration-http/src/main/java/org/springframework/integration/http/inbound/HttpRequestHandlingMessagingGateway.java @@ -24,8 +24,10 @@ import javax.servlet.ServletException; import javax.servlet.http.HttpServletRequest; import javax.servlet.http.HttpServletResponse; +import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; import org.springframework.http.converter.HttpMessageConverter; import org.springframework.http.server.ServletServerHttpRequest; import org.springframework.http.server.ServletServerHttpResponse; @@ -114,7 +116,27 @@ public class HttpRequestHandlingMessagingGateway extends HttpRequestHandlingEndp response.setStatusCode((HttpStatus) responseContent); } else { - this.writeResponse(responseContent, response, request.getHeaders().getAccept()); + if (responseContent instanceof ResponseEntity) { + ResponseEntity responseEntity = (ResponseEntity) responseContent; + responseContent = responseEntity.getBody(); + response.setStatusCode(responseEntity.getStatusCode()); + + HttpHeaders outputHeaders = response.getHeaders(); + HttpHeaders entityHeaders = responseEntity.getHeaders(); + + if (!entityHeaders.isEmpty()) { + entityHeaders.entrySet().stream() + .filter(entry -> !outputHeaders.containsKey(entry.getKey())) + .forEach(entry -> outputHeaders.put(entry.getKey(), entry.getValue())); + } + } + + if (responseContent != null) { + writeResponse(responseContent, response, request.getHeaders().getAccept()); + } + else { + response.flush(); + } } } else { diff --git a/spring-integration-http/src/test/java/org/springframework/integration/http/HttpProxyScenarioTests.java b/spring-integration-http/src/test/java/org/springframework/integration/http/HttpProxyScenarioTests.java index 1973aee237..8d6afa320f 100644 --- a/spring-integration-http/src/test/java/org/springframework/integration/http/HttpProxyScenarioTests.java +++ b/spring-integration-http/src/test/java/org/springframework/integration/http/HttpProxyScenarioTests.java @@ -19,7 +19,6 @@ package org.springframework.integration.http; import static org.hamcrest.Matchers.instanceOf; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; -import static org.junit.Assert.assertNull; import static org.junit.Assert.assertThat; import static org.mockito.ArgumentMatchers.isNull; @@ -144,8 +143,8 @@ public class HttpProxyScenarioTests { this.handlerAdapter.handle(request, response, handler); - assertNull(response.getHeaderValue("If-Modified-Since")); - assertNull(response.getHeaderValue("If-Unmodified-Since")); + assertEquals(ifModifiedSinceValue, response.getHeaderValue("If-Modified-Since")); + assertEquals(ifUnmodifiedSinceValue, response.getHeaderValue("If-Unmodified-Since")); assertEquals("close", response.getHeaderValue("Connection")); assertEquals(contentDispositionValue, response.getHeader("Content-Disposition")); assertEquals("text/plain", response.getContentType()); diff --git a/spring-integration-http/src/test/java/org/springframework/integration/http/dsl/HttpDslTests.java b/spring-integration-http/src/test/java/org/springframework/integration/http/dsl/HttpDslTests.java index 33a738d34c..48e54cf43c 100644 --- a/spring-integration-http/src/test/java/org/springframework/integration/http/dsl/HttpDslTests.java +++ b/spring-integration-http/src/test/java/org/springframework/integration/http/dsl/HttpDslTests.java @@ -33,16 +33,15 @@ import org.springframework.beans.DirectFieldAccessor; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.http.ResponseEntity; import org.springframework.integration.channel.DirectChannel; import org.springframework.integration.config.EnableIntegration; import org.springframework.integration.dsl.IntegrationFlow; import org.springframework.integration.dsl.IntegrationFlows; -import org.springframework.integration.http.HttpHeaders; import org.springframework.integration.http.outbound.HttpRequestExecutingMessageHandler; import org.springframework.integration.security.channel.ChannelSecurityInterceptor; import org.springframework.integration.security.channel.SecuredChannel; import org.springframework.messaging.MessageChannel; -import org.springframework.messaging.support.MessageBuilder; import org.springframework.security.access.AccessDecisionManager; import org.springframework.security.access.vote.AffirmativeBased; import org.springframework.security.access.vote.RoleVoter; @@ -186,9 +185,7 @@ public class HttpDslTests { return f -> f .transform(Throwable::getCause) .handle((p, h) -> - MessageBuilder.withPayload(p.getResponseBodyAsString()) - .setHeader(HttpHeaders.STATUS_CODE, p.getStatusCode()) - .build()); + new ResponseEntity<>(p.getResponseBodyAsString(), p.getStatusCode())); } @Bean diff --git a/spring-integration-webflux/src/main/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpoint.java b/spring-integration-webflux/src/main/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpoint.java index f6d3221c60..e06967af93 100644 --- a/spring-integration-webflux/src/main/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpoint.java +++ b/spring-integration-webflux/src/main/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpoint.java @@ -16,7 +16,9 @@ package org.springframework.integration.webflux.inbound; +import java.time.Instant; import java.util.ArrayList; +import java.util.Arrays; import java.util.Collections; import java.util.Comparator; import java.util.LinkedHashSet; @@ -35,8 +37,10 @@ import org.springframework.expression.Expression; import org.springframework.expression.spel.support.StandardEvaluationContext; import org.springframework.http.HttpEntity; import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; import org.springframework.http.codec.HttpMessageReader; import org.springframework.http.codec.HttpMessageWriter; import org.springframework.http.codec.ServerCodecConfigurer; @@ -75,6 +79,8 @@ public class WebFluxInboundEndpoint extends BaseHttpInboundEndpoint implements W private static final MediaType MEDIA_TYPE_APPLICATION_ALL = new MediaType("application"); + private static final List SAFE_METHODS = Arrays.asList(HttpMethod.GET, HttpMethod.HEAD); + private ServerCodecConfigurer codecConfigurer = ServerCodecConfigurer.create(); private RequestedContentTypeResolver requestedContentTypeResolver = new HeaderContentTypeResolver(); @@ -318,12 +324,44 @@ public class WebFluxInboundEndpoint extends BaseHttpInboundEndpoint implements W return response.setComplete(); } else { - HttpStatus httpStatus = resolveHttpStatusFromHeaders(replyMessage.getHeaders()); + final HttpStatus httpStatus = resolveHttpStatusFromHeaders(replyMessage.getHeaders()); if (httpStatus != null) { response.setStatusCode(httpStatus); } - return writeResponseBody(exchange, responseContent); + if (responseContent instanceof ResponseEntity) { + return Mono.just((ResponseEntity) responseContent) + .flatMap(e -> { + if (httpStatus == null) { + exchange.getResponse().setStatusCode(e.getStatusCode()); + } + + HttpHeaders entityHeaders = e.getHeaders(); + HttpHeaders responseHeaders = exchange.getResponse().getHeaders(); + + if (!entityHeaders.isEmpty()) { + entityHeaders.entrySet().stream() + .filter(entry -> !responseHeaders.containsKey(entry.getKey())) + .forEach(entry -> responseHeaders.put(entry.getKey(), entry.getValue())); + } + + if (e.getBody() == null) { + return exchange.getResponse().setComplete(); + } + + String etag = entityHeaders.getETag(); + Instant lastModified = Instant.ofEpochMilli(entityHeaders.getLastModified()); + HttpMethod httpMethod = exchange.getRequest().getMethod(); + if (SAFE_METHODS.contains(httpMethod) && exchange.checkNotModified(etag, lastModified)) { + return exchange.getResponse().setComplete(); + } + + return writeResponseBody(exchange, e.getBody()); + }); + } + else { + return writeResponseBody(exchange, responseContent); + } } } diff --git a/spring-integration-webflux/src/test/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpointTests.java b/spring-integration-webflux/src/test/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpointTests.java index c930535d84..124199b6ef 100644 --- a/spring-integration-webflux/src/test/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpointTests.java +++ b/spring-integration-webflux/src/test/java/org/springframework/integration/webflux/inbound/WebFluxInboundEndpointTests.java @@ -27,6 +27,7 @@ import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; import org.springframework.integration.annotation.ServiceActivator; import org.springframework.integration.channel.FluxMessageChannel; import org.springframework.integration.config.EnableIntegration; @@ -83,6 +84,17 @@ public class WebFluxInboundEndpointTests { .jsonPath("$[2].name").isEqualTo("John"); } + @Test + public void testServerInternalErrorRequest() { + this.webTestClient + .get() + .uri("/error") + .accept(MediaType.TEXT_PLAIN) + .exchange() + .expectStatus() + .is5xxServerError(); + } + @Configuration @EnableWebFlux @EnableIntegration @@ -128,6 +140,22 @@ public class WebFluxInboundEndpointTests { return Flux.just(new Person("Jane"), new Person("Jason"), new Person("John")); } + + @Bean + public WebFluxInboundEndpoint errorInboundEndpoint() { + WebFluxInboundEndpoint endpoint = new WebFluxInboundEndpoint(); + RequestMapping requestMapping = new RequestMapping(); + requestMapping.setPathPatterns("/error"); + endpoint.setRequestMapping(requestMapping); + endpoint.setRequestChannelName("errorServiceChannel"); + return endpoint; + } + + @ServiceActivator(inputChannel = "errorServiceChannel") + public ResponseEntity processHttpRequest() { + return new ResponseEntity<>("<500 Internal Server Error,{}>", HttpStatus.INTERNAL_SERVER_ERROR); + } + }