From 0b60b9bf1d6f34900f6bf879e6f81f1a40d32007 Mon Sep 17 00:00:00 2001 From: Marta Medio Date: Mon, 30 Jan 2023 10:47:30 +0100 Subject: [PATCH 1/2] Fixes CORS config to process all routes. Fix CorsGatewayFilterApplicationListener to process all routes (Bean defined ones too) Fixes gh-2854 --- .../config/GatewayAutoConfiguration.java | 4 +-- .../CorsGatewayFilterApplicationListener.java | 34 +++++++++++-------- .../cloud/gateway/cors/CorsPerRouteTests.java | 33 ++++++++++++++++++ 3 files changed, 54 insertions(+), 17 deletions(-) diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java index f0e560c0..ca2b7a14 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java @@ -272,9 +272,9 @@ public class GatewayAutoConfiguration { @ConditionalOnProperty(name = "spring.cloud.gateway.globalcors.enabled", matchIfMissing = true) public CorsGatewayFilterApplicationListener corsGatewayFilterApplicationListener( GlobalCorsProperties globalCorsProperties, RoutePredicateHandlerMapping routePredicateHandlerMapping, - RouteDefinitionLocator routeDefinitionLocator) { + RouteLocator routeLocator) { return new CorsGatewayFilterApplicationListener(globalCorsProperties, routePredicateHandlerMapping, - routeDefinitionLocator); + routeLocator); } @Bean diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java index ada6c141..afbaff0f 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java @@ -26,8 +26,8 @@ import java.util.Optional; import org.springframework.cloud.gateway.config.GlobalCorsProperties; import org.springframework.cloud.gateway.event.RefreshRoutesEvent; import org.springframework.cloud.gateway.handler.RoutePredicateHandlerMapping; -import org.springframework.cloud.gateway.route.RouteDefinition; -import org.springframework.cloud.gateway.route.RouteDefinitionLocator; +import org.springframework.cloud.gateway.route.Route; +import org.springframework.cloud.gateway.route.RouteLocator; import org.springframework.context.ApplicationListener; import org.springframework.web.cors.CorsConfiguration; @@ -41,7 +41,7 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener private final RoutePredicateHandlerMapping routePredicateHandlerMapping; - private final RouteDefinitionLocator routeDefinitionLocator; + private final RouteLocator routeLocator; private static final String PATH_PREDICATE_NAME = "Path"; @@ -50,22 +50,22 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener private static final String ALL_PATHS = "/**"; public CorsGatewayFilterApplicationListener(GlobalCorsProperties globalCorsProperties, - RoutePredicateHandlerMapping routePredicateHandlerMapping, RouteDefinitionLocator routeDefinitionLocator) { + RoutePredicateHandlerMapping routePredicateHandlerMapping, RouteLocator routeLocator) { this.globalCorsProperties = globalCorsProperties; this.routePredicateHandlerMapping = routePredicateHandlerMapping; - this.routeDefinitionLocator = routeDefinitionLocator; + this.routeLocator = routeLocator; } @Override public void onApplicationEvent(RefreshRoutesEvent event) { - routeDefinitionLocator.getRouteDefinitions().collectList().subscribe(routeDefinitions -> { + routeLocator.getRoutes().collectList().subscribe(routes -> { // pre-populate with pre-existing global cors configurations to combine with. var corsConfigurations = new HashMap<>(globalCorsProperties.getCorsConfigurations()); - routeDefinitions.forEach(routeDefinition -> { - var corsConfiguration = getCorsConfiguration(routeDefinition); + routes.forEach(route -> { + var corsConfiguration = getCorsConfiguration(route); corsConfiguration.ifPresent(configuration -> { - var pathPredicate = getPathPredicate(routeDefinition); + var pathPredicate = getPathPredicate(route); corsConfigurations.put(pathPredicate, configuration); }); }); @@ -74,15 +74,19 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener }); } - private String getPathPredicate(RouteDefinition routeDefinition) { - return routeDefinition.getPredicates().stream() - .filter(predicate -> PATH_PREDICATE_NAME.equals(predicate.getName())).findFirst() - .flatMap(predicate -> predicate.getArgs().values().stream().findFirst()).orElse(ALL_PATHS); + private String getPathPredicate(Route route) { + String predicate = route.getPredicate().toString(); + try { + return predicate.substring(predicate.indexOf("[") + 1, predicate.indexOf("]")); + } + catch (ArrayIndexOutOfBoundsException e) { + return ALL_PATHS; + } } @SuppressWarnings("unchecked") - private Optional getCorsConfiguration(RouteDefinition routeDefinition) { - Map corsMetadata = (Map) routeDefinition.getMetadata().get(METADATA_KEY); + private Optional getCorsConfiguration(Route route) { + Map corsMetadata = (Map) route.getMetadata().get(METADATA_KEY); if (corsMetadata != null) { final CorsConfiguration corsConfiguration = new CorsConfiguration(); diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java index 7584f42e..675a4dc6 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java @@ -20,10 +20,14 @@ import java.util.Map; import org.junit.jupiter.api.Test; +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.route.RouteLocator; +import org.springframework.cloud.gateway.route.builder.RouteLocatorBuilder; import org.springframework.cloud.gateway.test.BaseWebClientTests; +import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Import; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpMethod; @@ -60,6 +64,21 @@ public class CorsPerRouteTests extends BaseWebClientTests { assertThat(responseHeaders.getAccessControlAllowCredentials()) .as(missingHeader(ACCESS_CONTROL_ALLOW_CREDENTIALS)).isEqualTo(true); }); + + testClient.options().uri("/route-test").header("Origin", "another-domain.com") + .header("Access-Control-Request-Method", "GET").exchange().expectBody(Map.class).consumeWith(result -> { + assertThat(result.getResponseBody()).isNull(); + assertThat(result.getStatus()).isEqualTo(HttpStatus.OK); + + HttpHeaders responseHeaders = result.getResponseHeaders(); + assertThat(responseHeaders.getAccessControlAllowOrigin()) + .as(missingHeader(ACCESS_CONTROL_ALLOW_ORIGIN)).isEqualTo("another-domain.com"); + assertThat(responseHeaders.getAccessControlAllowMethods()) + .as(missingHeader(HttpHeaders.ACCESS_CONTROL_ALLOW_METHODS)) + .containsExactlyInAnyOrder(HttpMethod.GET); + assertThat(responseHeaders.getAccessControlMaxAge()).as(missingHeader(ACCESS_CONTROL_MAX_AGE)) + .isEqualTo(50L); + }); } @Test @@ -89,6 +108,20 @@ public class CorsPerRouteTests extends BaseWebClientTests { @Import(DefaultTestConfig.class) public static class TestConfig { + @Value("${test.uri}") + String uri; + + @Bean + public RouteLocator testRouteLocator(RouteLocatorBuilder builder) { + return builder.routes() + .route("cors_route_java_test", + r -> r.path("/route-test/**").filters(f -> f.stripPrefix(1).prefixPath("/httpbin")) + .metadata(Map.of("cors", Map.of("allowedOrigins", "another-domain.com", + "allowedMethods", HttpMethod.GET.name(), "maxAge", 50))) + .uri(uri)) + .build(); + } + } } From a7a4b40e7fe0d6cad7ec485450af018e36309985 Mon Sep 17 00:00:00 2001 From: spencergibb Date: Tue, 14 Feb 2023 15:05:31 -0500 Subject: [PATCH 2/2] Polishes Fix CORS to process all routes. Updates CorsGatewayFilterApplicationListener.getPathPredicate() to use the Predicate.accept(Visitor) pattern rather than String parsing. See gh-2854 --- .../CorsGatewayFilterApplicationListener.java | 30 ++++++++++++++----- .../cloud/gateway/cors/CorsPerRouteTests.java | 9 ++++-- 2 files changed, 29 insertions(+), 10 deletions(-) diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java index afbaff0f..884b72b9 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/cors/CorsGatewayFilterApplicationListener.java @@ -22,16 +22,21 @@ import java.util.HashMap; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; import org.springframework.cloud.gateway.config.GlobalCorsProperties; import org.springframework.cloud.gateway.event.RefreshRoutesEvent; import org.springframework.cloud.gateway.handler.RoutePredicateHandlerMapping; +import org.springframework.cloud.gateway.handler.predicate.PathRoutePredicateFactory; import org.springframework.cloud.gateway.route.Route; import org.springframework.cloud.gateway.route.RouteLocator; import org.springframework.context.ApplicationListener; import org.springframework.web.cors.CorsConfiguration; /** + * This class updates Cors configuration each time a {@link RefreshRoutesEvent} is consumed. + * The {@link Route}'s predicates are inspected for a {@link PathRoutePredicateFactory} and + * the first pattern is used. * @author Fredrich Ombico * @author Abel Salgado Romero */ @@ -43,8 +48,6 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener private final RouteLocator routeLocator; - private static final String PATH_PREDICATE_NAME = "Path"; - private static final String METADATA_KEY = "cors"; private static final String ALL_PATHS = "/**"; @@ -74,14 +77,25 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener }); } + /** + * Finds the first path predicate and first pattern in the config. + * @param route The Route to use. + * @return the first path predicate pattern or /**. + */ private String getPathPredicate(Route route) { - String predicate = route.getPredicate().toString(); - try { - return predicate.substring(predicate.indexOf("[") + 1, predicate.indexOf("]")); - } - catch (ArrayIndexOutOfBoundsException e) { - return ALL_PATHS; + var predicate = route.getPredicate(); + var pathPatterns = new AtomicReference(); + predicate.accept(p -> { + if (p.getConfig() instanceof PathRoutePredicateFactory.Config pathConfig) { + if (!pathConfig.getPatterns().isEmpty()) { + pathPatterns.compareAndSet(null, pathConfig.getPatterns().get(0)); + } + } + }); + if (pathPatterns.get() != null) { + return pathPatterns.get(); } + return ALL_PATHS; } @SuppressWarnings("unchecked") diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java index 675a4dc6..5f186903 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/cors/CorsPerRouteTests.java @@ -64,9 +64,13 @@ public class CorsPerRouteTests extends BaseWebClientTests { assertThat(responseHeaders.getAccessControlAllowCredentials()) .as(missingHeader(ACCESS_CONTROL_ALLOW_CREDENTIALS)).isEqualTo(true); }); + } + @Test + public void testPreFlightCorsRequestJavaConfig() { testClient.options().uri("/route-test").header("Origin", "another-domain.com") - .header("Access-Control-Request-Method", "GET").exchange().expectBody(Map.class).consumeWith(result -> { + .header("Host", "www.javaconfhost.org").header("Access-Control-Request-Method", "GET").exchange() + .expectBody(Map.class).consumeWith(result -> { assertThat(result.getResponseBody()).isNull(); assertThat(result.getStatus()).isEqualTo(HttpStatus.OK); @@ -115,7 +119,8 @@ public class CorsPerRouteTests extends BaseWebClientTests { public RouteLocator testRouteLocator(RouteLocatorBuilder builder) { return builder.routes() .route("cors_route_java_test", - r -> r.path("/route-test/**").filters(f -> f.stripPrefix(1).prefixPath("/httpbin")) + r -> r.host("*.javaconfhost.org").and().path("/route-test/**") + .filters(f -> f.stripPrefix(1).prefixPath("/httpbin")) .metadata(Map.of("cors", Map.of("allowedOrigins", "another-domain.com", "allowedMethods", HttpMethod.GET.name(), "maxAge", 50))) .uri(uri))