Merge branch 'mmedio/fix-cors-per-route'
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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<String>();
|
||||
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<CorsConfiguration> getCorsConfiguration(RouteDefinition routeDefinition) {
|
||||
Map<String, Object> corsMetadata = (Map<String, Object>) routeDefinition.getMetadata().get(METADATA_KEY);
|
||||
private Optional<CorsConfiguration> getCorsConfiguration(Route route) {
|
||||
Map<String, Object> corsMetadata = (Map<String, Object>) route.getMetadata().get(METADATA_KEY);
|
||||
if (corsMetadata != null) {
|
||||
final CorsConfiguration corsConfiguration = new CorsConfiguration();
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user