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..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.route.RouteDefinition; -import org.springframework.cloud.gateway.route.RouteDefinitionLocator; +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 */ @@ -41,31 +46,29 @@ public class CorsGatewayFilterApplicationListener implements ApplicationListener private final RoutePredicateHandlerMapping routePredicateHandlerMapping; - private final RouteDefinitionLocator routeDefinitionLocator; - - private static final String PATH_PREDICATE_NAME = "Path"; + private final RouteLocator routeLocator; private static final String METADATA_KEY = "cors"; 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 +77,30 @@ 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); + /** + * 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) { + 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") - 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..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 @@ -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; @@ -62,6 +66,25 @@ public class CorsPerRouteTests extends BaseWebClientTests { }); } + @Test + public void testPreFlightCorsRequestJavaConfig() { + testClient.options().uri("/route-test").header("Origin", "another-domain.com") + .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); + + 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 public void testPreFlightForbiddenCorsRequest() { testClient.get().uri("/cors").header("Origin", "domain.com").header("Access-Control-Request-Method", "GET") @@ -89,6 +112,21 @@ 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.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)) + .build(); + } + } }