diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactory.java index e5b8966c..33923e44 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactory.java @@ -19,12 +19,18 @@ package org.springframework.cloud.gateway.filter.factory; import java.util.List; import java.util.Map; +import reactor.core.publisher.Mono; + import org.springframework.cloud.gateway.filter.GatewayFilter; +import org.springframework.cloud.gateway.filter.GatewayFilterChain; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; import org.springframework.http.server.reactive.ServerHttpRequest; import org.springframework.util.unit.DataSize; import org.springframework.util.unit.DataUnit; +import org.springframework.web.server.ServerWebExchange; + +import static org.springframework.cloud.gateway.support.GatewayToStringStyler.filterToStringCreator; /** * This filter validates the size of each Request Header in the request. If size of any of @@ -46,28 +52,37 @@ public class RequestHeaderSizeGatewayFilterFactory extends @Override public GatewayFilter apply(RequestHeaderSizeGatewayFilterFactory.Config config) { - return (exchange, chain) -> { - ServerHttpRequest request = exchange.getRequest(); - HttpHeaders headers = request.getHeaders(); - Long headerSizeInBytes = 0L; + return new GatewayFilter() { + @Override + public Mono filter(ServerWebExchange exchange, GatewayFilterChain chain) { + ServerHttpRequest request = exchange.getRequest(); + HttpHeaders headers = request.getHeaders(); + Long headerSizeInBytes = 0L; - for (Map.Entry> headerEntry : headers.entrySet()) { - List values = headerEntry.getValue(); - for (String value : values) { - headerSizeInBytes += Long.valueOf(value.getBytes().length); + for (Map.Entry> headerEntry : headers.entrySet()) { + List values = headerEntry.getValue(); + for (String value : values) { + headerSizeInBytes += Long.valueOf(value.getBytes().length); + } } + + if (headerSizeInBytes > config.getMaxSize().toBytes()) { + exchange.getResponse() + .setStatusCode(HttpStatus.REQUEST_HEADER_FIELDS_TOO_LARGE); + exchange.getResponse().getHeaders().add("errorMessage", + getErrorMessage(headerSizeInBytes, config.getMaxSize())); + return exchange.getResponse().setComplete(); + + } + + return chain.filter(exchange); } - if (headerSizeInBytes > config.getMaxSize().toBytes()) { - exchange.getResponse() - .setStatusCode(HttpStatus.REQUEST_HEADER_FIELDS_TOO_LARGE); - exchange.getResponse().getHeaders().add("errorMessage", - getErrorMessage(headerSizeInBytes, config.getMaxSize())); - return exchange.getResponse().setComplete(); - + @Override + public String toString() { + return filterToStringCreator(RequestHeaderSizeGatewayFilterFactory.this) + .append("max", config.getMaxSize()).toString(); } - - return chain.filter(exchange); }; } diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactoryTest.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactoryTest.java index de1338cd..f8a142b8 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactoryTest.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RequestHeaderSizeGatewayFilterFactoryTest.java @@ -23,6 +23,8 @@ import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.cloud.gateway.filter.GatewayFilter; +import org.springframework.cloud.gateway.filter.factory.RequestHeaderSizeGatewayFilterFactory.Config; import org.springframework.cloud.gateway.route.RouteLocator; import org.springframework.cloud.gateway.route.builder.RouteLocatorBuilder; import org.springframework.cloud.gateway.test.BaseWebClientTests; @@ -34,6 +36,7 @@ import org.springframework.test.context.junit4.SpringRunner; import org.springframework.util.unit.DataSize; import org.springframework.util.unit.DataUnit; +import static org.assertj.core.api.Assertions.assertThat; import static org.springframework.boot.test.context.SpringBootTest.WebEnvironment.RANDOM_PORT; /** @@ -55,6 +58,14 @@ public class RequestHeaderSizeGatewayFilterFactoryTest extends BaseWebClientTest .expectHeader().valueMatches("errorMessage", responseMesssage); } + @Test + public void toStringFormat() { + Config config = new Config(); + config.setMaxSize(DataSize.ofBytes(1000L)); + GatewayFilter filter = new RequestHeaderSizeGatewayFilterFactory().apply(config); + assertThat(filter.toString()).contains("max").contains("1000B"); + } + @EnableAutoConfiguration @SpringBootConfiguration @Import(DefaultTestConfig.class)