Merge branch 'mmedio/fix-cors-per-route'

This commit is contained in:
spencergibb
2023-02-14 15:07:43 -05:00
3 changed files with 75 additions and 19 deletions

View File

@@ -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

View File

@@ -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();

View File

@@ -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();
}
}
}