From 5efd73dc94ebad045392066b27041e3277f782fc Mon Sep 17 00:00:00 2001 From: Spencer Gibb Date: Fri, 9 Mar 2018 15:52:42 -0500 Subject: [PATCH] Create alternate RedisRateLimiter constructor --- .../config/GatewayRedisAutoConfiguration.java | 4 +- .../filter/factory/GatewayFilterFactory.java | 1 + .../filter/ratelimit/AbstractRateLimiter.java | 6 ++- .../filter/ratelimit/RedisRateLimiter.java | 46 +++++++++++++++++-- .../gateway/support/ConfigurationUtils.java | 1 + .../cloud/gateway/support/NameUtils.java | 13 +++++- .../RedisRateLimiterConfigTests.java | 46 +++++++++++++------ 7 files changed, 95 insertions(+), 22 deletions(-) diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayRedisAutoConfiguration.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayRedisAutoConfiguration.java index 0f8c3051..e1e31b24 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayRedisAutoConfiguration.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayRedisAutoConfiguration.java @@ -7,6 +7,7 @@ import org.springframework.boot.autoconfigure.AutoConfigureAfter; import org.springframework.boot.autoconfigure.AutoConfigureBefore; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; +import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; import org.springframework.boot.autoconfigure.data.redis.RedisReactiveAutoConfiguration; import org.springframework.cloud.gateway.filter.ratelimit.RedisRateLimiter; import org.springframework.context.annotation.Bean; @@ -57,8 +58,9 @@ class GatewayRedisAutoConfiguration { } @Bean + @ConditionalOnMissingBean public RedisRateLimiter redisRateLimiter(ReactiveRedisTemplate redisTemplate, - @Qualifier("redisRequestRateLimiterScript") RedisScript> redisScript, + @Qualifier(RedisRateLimiter.REDIS_SCRIPT_NAME) RedisScript> redisScript, Validator validator) { return new RedisRateLimiter(redisTemplate, redisScript, validator); } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java index 29de7181..9ef52087 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java @@ -68,6 +68,7 @@ public interface GatewayFilterFactory extends ShortcutConfigurable, Configura GatewayFilter apply(Tuple args); default String name() { + //TODO: deal with proxys return NameUtils.normalizeFilterFactoryName(getClass()); } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/AbstractRateLimiter.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/AbstractRateLimiter.java index 5da0a5de..9ca41404 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/AbstractRateLimiter.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/AbstractRateLimiter.java @@ -28,7 +28,7 @@ import org.springframework.validation.Validator; public abstract class AbstractRateLimiter extends AbstractStatefulConfigurable implements RateLimiter, ApplicationListener { private String configurationPropertyName; - private final Validator validator; + private Validator validator; protected AbstractRateLimiter(Class configClass, String configurationPropertyName, Validator validator) { super(configClass); @@ -44,6 +44,10 @@ public abstract class AbstractRateLimiter extends AbstractStatefulConfigurabl return validator; } + public void setValidator(Validator validator) { + this.validator = validator; + } + @Override public void onApplicationEvent(FilterArgsEvent event) { Map args = event.getArgs(); diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiter.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiter.java index cd1a176b..34aae6bd 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiter.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiter.java @@ -4,11 +4,15 @@ import java.time.Instant; import java.util.ArrayList; import java.util.Arrays; import java.util.List; +import java.util.concurrent.atomic.AtomicBoolean; import javax.validation.constraints.Min; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import org.springframework.beans.BeansException; +import org.springframework.context.ApplicationContext; +import org.springframework.context.ApplicationContextAware; import org.springframework.data.redis.core.ReactiveRedisTemplate; import org.springframework.data.redis.core.script.RedisScript; import org.springframework.validation.Validator; @@ -23,23 +27,51 @@ import reactor.core.publisher.Mono; * * @author Spencer Gibb */ -public class RedisRateLimiter extends AbstractRateLimiter { +public class RedisRateLimiter extends AbstractRateLimiter implements ApplicationContextAware { @Deprecated public static final String REPLENISH_RATE_KEY = "replenishRate"; @Deprecated public static final String BURST_CAPACITY_KEY = "burstCapacity"; + public static final String CONFIGURATION_PROPERTY_NAME = "redis-rate-limiter"; + public static final String REDIS_SCRIPT_NAME = "redisRequestRateLimiterScript"; private Log log = LogFactory.getLog(getClass()); - private final ReactiveRedisTemplate redisTemplate; - private final RedisScript> script; + private ReactiveRedisTemplate redisTemplate; + private RedisScript> script; + private AtomicBoolean initialized = new AtomicBoolean(false); + private Config defaultConfig; public RedisRateLimiter(ReactiveRedisTemplate redisTemplate, RedisScript> script, Validator validator) { super(Config.class, CONFIGURATION_PROPERTY_NAME, validator); this.redisTemplate = redisTemplate; this.script = script; + initialized.compareAndSet(false, true); + } + + public RedisRateLimiter(int defaultReplenishRate, int defaultBurstCapacity) { + super(Config.class, CONFIGURATION_PROPERTY_NAME, null); + this.defaultConfig = new Config() + .setReplenishRate(defaultReplenishRate) + .setBurstCapacity(defaultBurstCapacity); + } + + @Override + @SuppressWarnings("unchecked") + public void setApplicationContext(ApplicationContext context) throws BeansException { + if (initialized.compareAndSet(false, true)) { + this.redisTemplate = context.getBean("stringReactiveRedisTemplate", ReactiveRedisTemplate.class); + this.script = context.getBean(REDIS_SCRIPT_NAME, RedisScript.class); + if (context.getBeanNamesForType(Validator.class).length > 0) { + this.setValidator(context.getBean(Validator.class)); + } + } + } + + /* for testing */ Config getDefaultConfig() { + return defaultConfig; } /** @@ -50,11 +82,17 @@ public class RedisRateLimiter extends AbstractRateLimiter isAllowed(String routeId, String id) { + if (!this.initialized.get()) { + throw new IllegalStateException("RedisRateLimiter is not initialized"); + } Config routeConfig = getConfig().get(routeId); if (routeConfig == null) { - throw new IllegalArgumentException("No Configuration found for route "+ routeId); + if (defaultConfig == null) { + throw new IllegalArgumentException("No Configuration found for route " + routeId); + } + routeConfig = defaultConfig; } // How many requests per second do you want a user to be allowed to do? diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/ConfigurationUtils.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/ConfigurationUtils.java index 8246cfbb..41714f8a 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/ConfigurationUtils.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/ConfigurationUtils.java @@ -46,6 +46,7 @@ public abstract class ConfigurationUtils { } } + @SuppressWarnings("unchecked") public static T getTargetObject(Object candidate) { try { if (AopUtils.isAopProxy(candidate) && (candidate instanceof Advised)) { diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/NameUtils.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/NameUtils.java index 064e61d0..9d34aa60 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/NameUtils.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/NameUtils.java @@ -31,10 +31,19 @@ public class NameUtils { } public static String normalizeRoutePredicateName(Class clazz) { - return clazz.getSimpleName().replace(RoutePredicateFactory.class.getSimpleName(), ""); + return removeGarbage(clazz.getSimpleName().replace(RoutePredicateFactory.class.getSimpleName(), "")); } public static String normalizeFilterFactoryName(Class clazz) { - return clazz.getSimpleName().replace(GatewayFilterFactory.class.getSimpleName(), ""); + return removeGarbage(clazz.getSimpleName().replace(GatewayFilterFactory.class.getSimpleName(), "")); + } + + private static String removeGarbage(String s) { + int garbageIdx = s.indexOf("$Mockito"); + if (garbageIdx > 0) { + return s.substring(0, garbageIdx); + } + + return s; } } diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiterConfigTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiterConfigTests.java index b6e17e09..9d9e8b4c 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiterConfigTests.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ratelimit/RedisRateLimiterConfigTests.java @@ -17,13 +17,13 @@ package org.springframework.cloud.gateway.filter.ratelimit; +import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; import org.springframework.boot.test.context.SpringBootTest; -import org.springframework.cloud.gateway.filter.ratelimit.RedisRateLimiter.Config; import org.springframework.cloud.gateway.route.Route; import org.springframework.cloud.gateway.route.RouteLocator; import org.springframework.cloud.gateway.route.builder.RouteLocatorBuilder; @@ -31,12 +31,9 @@ import org.springframework.context.annotation.Bean; import org.springframework.test.annotation.DirtiesContext; import org.springframework.test.context.ActiveProfiles; import org.springframework.test.context.junit4.SpringRunner; -import org.springframework.web.server.ServerWebExchange; import static org.assertj.core.api.Assertions.assertThat; -import reactor.core.publisher.Mono; - /** * @author Spencer Gibb */ @@ -52,20 +49,40 @@ public class RedisRateLimiterConfigTests { @Autowired private RouteLocator routeLocator; + @Before + public void init() { + System.out.println(); + } + @Test public void redisRateConfiguredFromEnvironment() { - assertFilter("redis_rate_limiter_config_test", 10, 20, PrincipalNameKeyResolver.class); + assertFilter("redis_rate_limiter_config_test", 10, 20, + false); } @Test public void redisRateConfiguredFromJavaAPI() { - assertFilter("custom_redis_rate_limiter", 20, 40, MyKeyResolver.class); + assertFilter("custom_redis_rate_limiter", 20, 40, + false); } - private void assertFilter(String key, int replenishRate, int burstCapacity, Class keyResolverClass) { - assertThat(rateLimiter.getConfig()).containsKey(key); + @Test + public void redisRateConfiguredFromJavaAPIDirectBean() { + assertFilter("alt_custom_redis_rate_limiter", 30, 60, + true); + } - Config config = rateLimiter.getConfig().get(key); + private void assertFilter(String key, int replenishRate, int burstCapacity, + boolean useDefaultConfig) { + RedisRateLimiter.Config config; + + if (useDefaultConfig) { + config = rateLimiter.getDefaultConfig(); + } else { + assertThat(rateLimiter.getConfig()).containsKey(key); + config = rateLimiter.getConfig().get(key); + } + assertThat(config).isNotNull(); assertThat(config.getReplenishRate()).isEqualTo(replenishRate); assertThat(config.getBurstCapacity()).isEqualTo(burstCapacity); @@ -87,15 +104,16 @@ public class RedisRateLimiterConfigTests { rl -> rl.setBurstCapacity(40).setReplenishRate(20)) .and()) .uri("http://localhost")) + .route("alt_custom_redis_rate_limiter", r -> r.path("/custom") + .filters(f -> f.requestRateLimiter(c -> c.setRateLimiter(myRateLimiter()))) + .uri("http://localhost")) .build(); } - } - private static class MyKeyResolver implements KeyResolver { - @Override - public Mono resolve(ServerWebExchange exchange) { - return null; + @Bean + public RedisRateLimiter myRateLimiter() { + return new RedisRateLimiter(30, 60); } } }