From e1e78c8c4ab60039383d0e71b99e2b92da2a870b Mon Sep 17 00:00:00 2001 From: Spencer Gibb Date: Wed, 4 Oct 2017 18:30:04 -0400 Subject: [PATCH] Add PrincipalNameKeyResolver fixes gh-76 --- spring-cloud-gateway-core/pom.xml | 5 + .../config/GatewayAutoConfiguration.java | 7 + .../RequestRateLimiterWebFilterFactory.java | 27 ++-- .../gateway/filter/ratelimit/KeyResolver.java | 1 - .../ratelimit/PrincipalNameKeyResolver.java | 16 +++ .../cloud/gateway/route/Routes.java | 13 +- ...ncipalNameKeyResolverIntegrationTests.java | 131 ++++++++++++++++++ .../gateway/test/BaseWebClientTests.java | 12 ++ 8 files changed, 198 insertions(+), 14 deletions(-) create mode 100644 spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolver.java create mode 100644 spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolverIntegrationTests.java diff --git a/spring-cloud-gateway-core/pom.xml b/spring-cloud-gateway-core/pom.xml index 1ca84c2a..f18b75dc 100644 --- a/spring-cloud-gateway-core/pom.xml +++ b/spring-cloud-gateway-core/pom.xml @@ -67,6 +67,11 @@ spring-boot-starter-data-redis-reactive true + + org.springframework.boot + spring-boot-starter-security-reactive + true + org.springframework.cloud spring-cloud-starter-eureka diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java index c5cd2c2c..637fb645 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayAutoConfiguration.java @@ -54,6 +54,7 @@ import org.springframework.cloud.gateway.filter.factory.SetResponseHeaderWebFilt import org.springframework.cloud.gateway.filter.factory.SetStatusWebFilterFactory; import org.springframework.cloud.gateway.filter.factory.WebFilterFactory; import org.springframework.cloud.gateway.filter.ratelimit.KeyResolver; +import org.springframework.cloud.gateway.filter.ratelimit.PrincipalNameKeyResolver; import org.springframework.cloud.gateway.filter.ratelimit.RateLimiter; import org.springframework.cloud.gateway.handler.FilteringWebHandler; import org.springframework.cloud.gateway.handler.RoutePredicateHandlerMapping; @@ -322,6 +323,12 @@ public class GatewayAutoConfiguration { return new RemoveResponseHeaderWebFilterFactory(); } + @Bean(name = PrincipalNameKeyResolver.BEAN_NAME) + @ConditionalOnBean(RateLimiter.class) + public PrincipalNameKeyResolver principalNameKeyResolver() { + return new PrincipalNameKeyResolver(); + } + @Bean @ConditionalOnBean({RateLimiter.class, KeyResolver.class}) public RequestRateLimiterWebFilterFactory requestRateLimiterWebFilterFactory(RateLimiter rateLimiter) { diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestRateLimiterWebFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestRateLimiterWebFilterFactory.java index 7ea317c4..34b5706b 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestRateLimiterWebFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RequestRateLimiterWebFilterFactory.java @@ -19,6 +19,7 @@ package org.springframework.cloud.gateway.filter.factory; import org.springframework.beans.BeansException; import org.springframework.cloud.gateway.filter.ratelimit.KeyResolver; +import org.springframework.cloud.gateway.filter.ratelimit.PrincipalNameKeyResolver; import org.springframework.cloud.gateway.filter.ratelimit.RateLimiter; import org.springframework.cloud.gateway.filter.ratelimit.RateLimiter.Response; import org.springframework.context.ApplicationContext; @@ -67,21 +68,25 @@ public class RequestRateLimiterWebFilterFactory implements WebFilterFactory, App // How much bursting do you want to allow? int capacity = args.getInt(BURST_CAPACITY_KEY); - String beanName = args.getString(KEY_RESOLVER_NAME_KEY); + String beanName; + if (args.hasFieldName(KEY_RESOLVER_NAME_KEY)) { + beanName = args.getString(KEY_RESOLVER_NAME_KEY); + } else { + beanName = PrincipalNameKeyResolver.BEAN_NAME; + } KeyResolver keyResolver = this.context.getBean(beanName, KeyResolver.class); return (exchange, chain) -> - keyResolver.resolve(exchange).flatMap(key -> { - Response response = rateLimiter.isAllowed(key, replenishRate, capacity).block(); //FIXME: block() + keyResolver.resolve(exchange).flatMap(key -> + rateLimiter.isAllowed(key, replenishRate, capacity).flatMap(response -> { + //TODO: set some headers for rate, tokens left - //TODO: set some headers for rate, tokens left - - if (response.isAllowed()) { - return chain.filter(exchange); - } - exchange.getResponse().setStatusCode(HttpStatus.TOO_MANY_REQUESTS); - return exchange.getResponse().setComplete(); - }); + if (response.isAllowed()) { + return chain.filter(exchange); + } + exchange.getResponse().setStatusCode(HttpStatus.TOO_MANY_REQUESTS); + return exchange.getResponse().setComplete(); + })); } } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/KeyResolver.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/KeyResolver.java index 8d164bc4..3f3e7e70 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/KeyResolver.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/KeyResolver.java @@ -6,7 +6,6 @@ import reactor.core.publisher.Mono; /** * @author Spencer Gibb */ -//TODO: KeyResolver for exchange.getPrincipal().flatMap(principal -> {}) public interface KeyResolver { Mono resolve(ServerWebExchange exchange); } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolver.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolver.java new file mode 100644 index 00000000..502e9e58 --- /dev/null +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolver.java @@ -0,0 +1,16 @@ +package org.springframework.cloud.gateway.filter.ratelimit; + +import org.springframework.web.server.ServerWebExchange; +import reactor.core.publisher.Mono; + +import java.security.Principal; + +public class PrincipalNameKeyResolver implements KeyResolver { + + public static final String BEAN_NAME = "principalNameKeyResolver"; + + @Override + public Mono resolve(ServerWebExchange exchange) { + return exchange.getPrincipal().map(Principal::getName).switchIfEmpty(Mono.empty()); + } +} diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/route/Routes.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/route/Routes.java index 4f996069..f210fe1e 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/route/Routes.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/route/Routes.java @@ -23,6 +23,7 @@ import java.util.Collection; import java.util.List; import java.util.function.Predicate; +import org.springframework.cloud.gateway.filter.OrderedWebFilter; import org.springframework.cloud.gateway.filter.factory.WebFilterFactories; import org.springframework.web.server.ServerWebExchange; import org.springframework.web.server.WebFilter; @@ -129,12 +130,20 @@ public class Routes { } public WebFilterSpec webFilters(List webFilters) { - this.builder.webFilters(webFilters); + this.addAll(webFilters); return this; } public WebFilterSpec add(WebFilter webFilter) { - this.builder.add(webFilter); + return this.filter(webFilter); + } + + public WebFilterSpec filter(WebFilter webFilter) { + return this.filter(webFilter, 0); + } + + public WebFilterSpec filter(WebFilter webFilter, int order) { + this.builder.add(new OrderedWebFilter(webFilter, order)); return this; } diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolverIntegrationTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolverIntegrationTests.java new file mode 100644 index 00000000..500628ec --- /dev/null +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/PrincipalNameKeyResolverIntegrationTests.java @@ -0,0 +1,131 @@ +package org.springframework.cloud.gateway.filter.ratelimit; + +import java.security.Principal; +import java.util.Collections; +import java.util.Map; + +import org.junit.AfterClass; +import org.junit.Before; +import org.junit.BeforeClass; +import org.junit.Test; +import org.junit.runner.RunWith; +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.boot.web.server.LocalServerPort; +import org.springframework.cloud.gateway.filter.factory.RequestRateLimiterWebFilterFactory; +import org.springframework.cloud.gateway.route.RouteLocator; +import org.springframework.cloud.gateway.route.Routes; +import org.springframework.context.annotation.Bean; +import org.springframework.security.config.web.server.HttpSecurity; +import org.springframework.security.core.userdetails.MapUserDetailsRepository; +import org.springframework.security.core.userdetails.User; +import org.springframework.security.core.userdetails.UserDetails; +import org.springframework.security.web.server.SecurityWebFilterChain; +import org.springframework.test.context.ActiveProfiles; +import org.springframework.test.context.junit4.SpringRunner; +import org.springframework.test.web.reactive.server.WebTestClient; +import org.springframework.util.SocketUtils; +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; + +import static org.springframework.boot.test.context.SpringBootTest.WebEnvironment.DEFINED_PORT; +import static org.springframework.cloud.gateway.filter.factory.RequestRateLimiterWebFilterFactory.BURST_CAPACITY_KEY; +import static org.springframework.cloud.gateway.filter.factory.RequestRateLimiterWebFilterFactory.REPLENISH_RATE_KEY; +import static org.springframework.cloud.gateway.filter.factory.WebFilterFactories.prefixPath; +import static org.springframework.cloud.gateway.handler.predicate.RoutePredicates.path; +import static org.springframework.tuple.TupleBuilder.tuple; +import static org.springframework.web.reactive.function.client.ExchangeFilterFunctions.basicAuthentication; + +import reactor.core.publisher.Mono; + +@RunWith(SpringRunner.class) +@SpringBootTest(webEnvironment = DEFINED_PORT) +@ActiveProfiles("principalname") +public class PrincipalNameKeyResolverIntegrationTests { + @LocalServerPort + protected int port = 0; + + protected WebTestClient client; + protected String baseUri; + + + @BeforeClass + public static void beforeClass() { + System.setProperty("server.port", String.valueOf(SocketUtils.findAvailableTcpPort())); + } + + @AfterClass + public static void afterClass() { + System.clearProperty("server.port"); + } + + @Before + public void setup() { + this.baseUri = "http://localhost:" + port; + this.client = WebTestClient.bindToServer().baseUrl(baseUri).build(); + } + + @Test + public void keyResolverWorks() { + this.client.mutate() + .filter(basicAuthentication("user", "password")) + .build() + .get() + .uri("/myapi/1") + .exchange() + .expectStatus().isOk() + .expectBody().json("{\"user\":\"1\"}"); + } + + + @RestController + @RequestMapping("/downstream") + @EnableAutoConfiguration + @SpringBootConfiguration + protected static class TestConfig { + + @Value("${server.port}") + private int port; + + @RequestMapping("/myapi/{id}") + public Map myapi(@PathVariable String id, Principal principal) { + return Collections.singletonMap(principal.getName(), id); + } + + @Bean + public RouteLocator customRouteLocator(RequestRateLimiterWebFilterFactory rateLimiterFactory) { + return Routes.locator() + .route("protected-throttled") + .uri("http://localhost:"+port) + .predicate(path("/myapi/**")) + .filter(rateLimiterFactory.apply(tuple().of(REPLENISH_RATE_KEY, 1, BURST_CAPACITY_KEY, 1))) + .filter(prefixPath("/downstream")) + .and() + .build(); + } + + @Bean + RateLimiter rateLimiter() { + return (id, replenishRate, burstCapacity) -> Mono.just(new RateLimiter.Response(true, Long.MAX_VALUE)); + } + + @Bean + SecurityWebFilterChain springWebFilterChain(HttpSecurity http) throws Exception { + return http.httpBasic().and() + .authorizeExchange() + .pathMatchers("/myapi/**").authenticated() + .anyExchange().permitAll() + .and() + .build(); + } + + @Bean + public MapUserDetailsRepository userDetailsRepository() { + UserDetails user = User.withUsername("user").password("password").roles("USER").build(); + return new MapUserDetailsRepository(user); + } + } +} diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/test/BaseWebClientTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/test/BaseWebClientTests.java index b97f369a..fdef37e5 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/test/BaseWebClientTests.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/test/BaseWebClientTests.java @@ -38,6 +38,8 @@ import org.springframework.core.annotation.Order; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; import org.springframework.http.codec.multipart.Part; +import org.springframework.security.config.web.server.HttpSecurity; +import org.springframework.security.web.server.SecurityWebFilterChain; import org.springframework.web.bind.annotation.PathVariable; import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; @@ -181,6 +183,16 @@ public class BaseWebClientTests { return chain.filter(exchange); }; } + + + @Bean + SecurityWebFilterChain springWebFilterChain(HttpSecurity http) throws Exception { + return http.authorizeExchange() + //.pathMatchers("/admin/**").hasRole("ADMIN") + .anyExchange().permitAll() + .and() + .build(); + } } protected static class TestRibbonConfig {