Create alternate RedisRateLimiter constructor
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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?
|
||||
|
||||
@@ -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)) {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user