Create alternate RedisRateLimiter constructor

This commit is contained in:
Spencer Gibb
2018-03-09 15:52:42 -05:00
parent 34fd45dc4a
commit 5efd73dc94
7 changed files with 95 additions and 22 deletions

View File

@@ -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<String, String> redisTemplate,
@Qualifier("redisRequestRateLimiterScript") RedisScript<List<Long>> redisScript,
@Qualifier(RedisRateLimiter.REDIS_SCRIPT_NAME) RedisScript<List<Long>> redisScript,
Validator validator) {
return new RedisRateLimiter(redisTemplate, redisScript, validator);
}

View File

@@ -68,6 +68,7 @@ public interface GatewayFilterFactory<C> extends ShortcutConfigurable, Configura
GatewayFilter apply(Tuple args);
default String name() {
//TODO: deal with proxys
return NameUtils.normalizeFilterFactoryName(getClass());
}

View File

@@ -28,7 +28,7 @@ import org.springframework.validation.Validator;
public abstract class AbstractRateLimiter<C> extends AbstractStatefulConfigurable<C> implements RateLimiter<C>, ApplicationListener<FilterArgsEvent> {
private String configurationPropertyName;
private final Validator validator;
private Validator validator;
protected AbstractRateLimiter(Class<C> configClass, String configurationPropertyName, Validator validator) {
super(configClass);
@@ -44,6 +44,10 @@ public abstract class AbstractRateLimiter<C> extends AbstractStatefulConfigurabl
return validator;
}
public void setValidator(Validator validator) {
this.validator = validator;
}
@Override
public void onApplicationEvent(FilterArgsEvent event) {
Map<String, Object> args = event.getArgs();

View File

@@ -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<RedisRateLimiter.Config> {
public class RedisRateLimiter extends AbstractRateLimiter<RedisRateLimiter.Config> 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<String, String> redisTemplate;
private final RedisScript<List<Long>> script;
private ReactiveRedisTemplate<String, String> redisTemplate;
private RedisScript<List<Long>> script;
private AtomicBoolean initialized = new AtomicBoolean(false);
private Config defaultConfig;
public RedisRateLimiter(ReactiveRedisTemplate<String, String> redisTemplate,
RedisScript<List<Long>> 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<RedisRateLimiter.Confi
@Override
@SuppressWarnings("unchecked")
public Mono<Response> 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?

View File

@@ -46,6 +46,7 @@ public abstract class ConfigurationUtils {
}
}
@SuppressWarnings("unchecked")
public static <T> T getTargetObject(Object candidate) {
try {
if (AopUtils.isAopProxy(candidate) && (candidate instanceof Advised)) {

View File

@@ -31,10 +31,19 @@ public class NameUtils {
}
public static String normalizeRoutePredicateName(Class<? extends RoutePredicateFactory> clazz) {
return clazz.getSimpleName().replace(RoutePredicateFactory.class.getSimpleName(), "");
return removeGarbage(clazz.getSimpleName().replace(RoutePredicateFactory.class.getSimpleName(), ""));
}
public static String normalizeFilterFactoryName(Class<? extends GatewayFilterFactory> 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;
}
}

View File

@@ -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<? extends KeyResolver> 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<String> resolve(ServerWebExchange exchange) {
return null;
@Bean
public RedisRateLimiter myRateLimiter() {
return new RedisRateLimiter(30, 60);
}
}
}