From 397e4ff4233445218c5cfde598b221dedc484671 Mon Sep 17 00:00:00 2001 From: kiran malsetty Date: Thu, 6 Aug 2026 11:21:32 -0400 Subject: [PATCH 1/3] Add calendar-aligned MONTH rate limit, gated by USE_CALENDAR_MONTH_RATE_LIMIT A unit: month rate limit is computed as a fixed 60*60*24*30 second window counted from the Unix epoch, so it neither aligns with real calendar months nor accounts for months of different lengths (RDGRS-1999). Add a USE_CALENDAR_MONTH_RATE_LIMIT setting (default false) that, when enabled, buckets MONTH cache keys by UTC calendar month, sets their TTL/expiration to the actual time remaining until month end, and reports that same value as the reset duration. Defaults to false so existing MONTH limits keep their current reset behavior unless explicitly opted in. Signed-off-by: kiran malsetty --- README.md | 13 +++ src/limiter/base_limiter.go | 17 +++- src/limiter/cache_key.go | 20 +++- src/memcached/cache_impl.go | 6 +- src/redis/cache_impl.go | 1 + src/redis/fixed_cache_impl.go | 6 +- src/service/ratelimit.go | 4 +- src/settings/settings.go | 7 ++ src/utils/time.go | 31 ++++++ src/utils/utilities.go | 21 +++- test/config/basic_config.yaml | 6 ++ test/limiter/base_limiter_test.go | 62 +++++++++--- test/memcached/cache_impl_test.go | 139 +++++++++++++++----------- test/redis/bench_test.go | 2 +- test/redis/fixed_cache_impl_test.go | 150 ++++++++++++++++------------ test/utils/utilities_test.go | 99 ++++++++++++++++++ 16 files changed, 436 insertions(+), 148 deletions(-) diff --git a/README.md b/README.md index 01256d45c..4e9459eb8 100644 --- a/README.md +++ b/README.md @@ -1396,6 +1396,19 @@ configured address and reaches the current master. The failing command still ret to the caller; only the connection handling changes. Applies to both the main and the per-second Redis clients. +## Calendar-aligned MONTH rate limits + +1. `USE_CALENDAR_MONTH_RATE_LIMIT` : (default is "false") + +By default, a `unit: month` rate limit uses a fixed 30-day window counted from the Unix epoch, +which does not line up with real calendar months (it drifts, and treats every month as 30 days +regardless of its actual length). + +Setting `USE_CALENDAR_MONTH_RATE_LIMIT` to `"true"` switches `MONTH` limits to a true calendar +month window instead: the cache key bucket, TTL/expiration, and reported reset time all cover +the 1st through the last day of the month (UTC). This is opt-in because it changes when +existing `MONTH` limits reset and is therefore not enabled by default. + # Memcache Experimental Memcache support has been added as an alternative to Redis in v1.5. diff --git a/src/limiter/base_limiter.go b/src/limiter/base_limiter.go index 3ee5ce8e5..8c013080f 100644 --- a/src/limiter/base_limiter.go +++ b/src/limiter/base_limiter.go @@ -22,6 +22,9 @@ type BaseRateLimiter struct { localCache *freecache.Cache nearLimitRatio float32 StatsManager stats.Manager + // useCalendarMonth gates the MONTH-unit fix (calendar-aligned window + // instead of a fixed 30-day divider) for expiration/TTL computations. + useCalendarMonth bool } type LimitInfo struct { @@ -61,6 +64,12 @@ func (this *BaseRateLimiter) GenerateCacheKeys(request *pb.RateLimitRequest, return cacheKeys } +// ExpirationSeconds returns the number of seconds, evaluated from the current +// time, until the given rate limit unit's window ends. +func (this *BaseRateLimiter) ExpirationSeconds(unit pb.RateLimitResponse_RateLimit_Unit) int64 { + return utils.ExpirationSeconds(unit, this.timeSource, this.useCalendarMonth) +} + // Returns `true` in case local cache is enabled and contains value for provided cache key, `false` otherwise. func (this *BaseRateLimiter) IsOverLimitWithLocalCache(key string) bool { if this.localCache != nil { @@ -116,7 +125,7 @@ func (this *BaseRateLimiter) GetResponseDescriptorStatus(key string, limitInfo * // similar to mongo_1h, mongo_2h, etc. In the hour 1 (0h0m - 0h59m), the cache key is mongo_1h, we start // to get ratelimited in the 50th minute, the ttl of local_cache will be set as 1 hour(0h50m-1h49m). // In the time of 1h1m, since the cache key becomes different (mongo_2h), it won't get ratelimited. - err := this.localCache.Set([]byte(key), []byte{}, int(utils.UnitToDivider(limitInfo.limit.Limit.Unit))) + err := this.localCache.Set([]byte(key), []byte{}, int(this.ExpirationSeconds(limitInfo.limit.Limit.Unit))) if err != nil { logger.Errorf("Failing to set local cache key: %s", key) } @@ -144,15 +153,17 @@ func (this *BaseRateLimiter) GetResponseDescriptorStatus(key string, limitInfo * func NewBaseRateLimit(timeSource utils.TimeSource, jitterRand *rand.Rand, expirationJitterMaxSeconds int64, localCache *freecache.Cache, nearLimitRatio float32, cacheKeyPrefix string, statsManager stats.Manager, + useCalendarMonth bool, ) *BaseRateLimiter { return &BaseRateLimiter{ timeSource: timeSource, JitterRand: jitterRand, ExpirationJitterMaxSeconds: expirationJitterMaxSeconds, - cacheKeyGenerator: NewCacheKeyGenerator(cacheKeyPrefix), + cacheKeyGenerator: NewCacheKeyGenerator(cacheKeyPrefix, useCalendarMonth), localCache: localCache, nearLimitRatio: nearLimitRatio, StatsManager: statsManager, + useCalendarMonth: useCalendarMonth, } } @@ -205,7 +216,7 @@ func (this *BaseRateLimiter) generateResponseDescriptorStatus(responseCode pb.Ra Code: responseCode, CurrentLimit: limit, LimitRemaining: limitRemaining, - DurationUntilReset: utils.CalculateReset(&limit.Unit, this.timeSource), + DurationUntilReset: utils.CalculateReset(&limit.Unit, this.timeSource, this.useCalendarMonth), } } else { return &pb.RateLimitResponse_DescriptorStatus{ diff --git a/src/limiter/cache_key.go b/src/limiter/cache_key.go index 8b2056fc5..3ece5a2c0 100644 --- a/src/limiter/cache_key.go +++ b/src/limiter/cache_key.go @@ -14,13 +14,17 @@ import ( type CacheKeyGenerator struct { prefix string + // useCalendarMonth gates bucketing MONTH-unit limits by real calendar + // month (UTC) instead of the legacy fixed 30-day divider. + useCalendarMonth bool // bytes.Buffer pool used to efficiently generate cache keys. bufferPool sync.Pool } -func NewCacheKeyGenerator(prefix string) CacheKeyGenerator { +func NewCacheKeyGenerator(prefix string, useCalendarMonth bool) CacheKeyGenerator { return CacheKeyGenerator{ - prefix: prefix, + prefix: prefix, + useCalendarMonth: useCalendarMonth, bufferPool: sync.Pool{ New: func() interface{} { return new(bytes.Buffer) @@ -78,8 +82,16 @@ func (this *CacheKeyGenerator) GenerateCacheKey( b.WriteByte('_') } - divider := utils.UnitToDivider(limit.Limit.Unit) - b.WriteString(strconv.FormatInt((now/divider)*divider, 10)) + var bucketStart int64 + if this.useCalendarMonth && limit.Limit.Unit == pb.RateLimitResponse_RateLimit_MONTH { + // Calendar months vary in length, so bucket by the start of the + // current UTC calendar month rather than a fixed-size divider. + bucketStart = utils.MonthStartUnix(now) + } else { + divider := utils.UnitToDivider(limit.Limit.Unit) + bucketStart = (now / divider) * divider + } + b.WriteString(strconv.FormatInt(bucketStart, 10)) return CacheKey{ Key: b.String(), diff --git a/src/memcached/cache_impl.go b/src/memcached/cache_impl.go index eb65b63d1..21296cef4 100644 --- a/src/memcached/cache_impl.go +++ b/src/memcached/cache_impl.go @@ -161,7 +161,7 @@ func (this *rateLimitMemcacheImpl) increaseAsync(cacheKeys []limiter.CacheKey, i _, err := this.client.Increment(cacheKey.Key, hitsAddends[i]) if err == memcache.ErrCacheMiss { - expirationSeconds := utils.UnitToDivider(limits[i].Limit.Unit) + expirationSeconds := this.baseRateLimiter.ExpirationSeconds(limits[i].Limit.Unit) if this.expirationJitterMaxSeconds > 0 { expirationSeconds += this.jitterRand.Int63n(this.expirationJitterMaxSeconds) } @@ -304,6 +304,7 @@ func runAsync(task func()) { func NewRateLimitCacheImpl(client Client, timeSource utils.TimeSource, jitterRand *rand.Rand, expirationJitterMaxSeconds int64, localCache *freecache.Cache, statsManager stats.Manager, nearLimitRatio float32, cacheKeyPrefix string, + useCalendarMonth bool, ) limiter.RateLimitCache { return &rateLimitMemcacheImpl{ client: client, @@ -312,7 +313,7 @@ func NewRateLimitCacheImpl(client Client, timeSource utils.TimeSource, jitterRan expirationJitterMaxSeconds: expirationJitterMaxSeconds, localCache: localCache, nearLimitRatio: nearLimitRatio, - baseRateLimiter: limiter.NewBaseRateLimit(timeSource, jitterRand, expirationJitterMaxSeconds, localCache, nearLimitRatio, cacheKeyPrefix, statsManager), + baseRateLimiter: limiter.NewBaseRateLimit(timeSource, jitterRand, expirationJitterMaxSeconds, localCache, nearLimitRatio, cacheKeyPrefix, statsManager, useCalendarMonth), } } @@ -328,5 +329,6 @@ func NewRateLimitCacheImplFromSettings(s settings.Settings, timeSource utils.Tim statsManager, s.NearLimitRatio, s.CacheKeyPrefix, + s.UseCalendarMonthRateLimit, ) } diff --git a/src/redis/cache_impl.go b/src/redis/cache_impl.go index 9b9827809..f56971b4a 100644 --- a/src/redis/cache_impl.go +++ b/src/redis/cache_impl.go @@ -46,5 +46,6 @@ func NewRateLimiterCacheImplFromSettings(ctx context.Context, s settings.Setting s.CacheKeyPrefix, statsManager, s.StopCacheKeyIncrementWhenOverlimit, + s.UseCalendarMonthRateLimit, ), closer } diff --git a/src/redis/fixed_cache_impl.go b/src/redis/fixed_cache_impl.go index 9e8918593..5f3a7fa77 100644 --- a/src/redis/fixed_cache_impl.go +++ b/src/redis/fixed_cache_impl.go @@ -165,7 +165,7 @@ func (this *fixedRateLimitCacheImpl) DoLimit( logger.Debugf("looking up cache key: %s", cacheKey.Key) - expirationSeconds := utils.UnitToDivider(limits[i].Limit.Unit) + expirationSeconds := this.baseRateLimiter.ExpirationSeconds(limits[i].Limit.Unit) if this.baseRateLimiter.ExpirationJitterMaxSeconds > 0 { expirationSeconds += this.baseRateLimiter.JitterRand.Int63n(this.baseRateLimiter.ExpirationJitterMaxSeconds) } @@ -225,12 +225,12 @@ func (this *fixedRateLimitCacheImpl) Flush() {} func NewFixedRateLimitCacheImpl(client Client, perSecondClient Client, timeSource utils.TimeSource, jitterRand *rand.Rand, expirationJitterMaxSeconds int64, localCache *freecache.Cache, nearLimitRatio float32, cacheKeyPrefix string, statsManager stats.Manager, - stopCacheKeyIncrementWhenOverlimit bool, + stopCacheKeyIncrementWhenOverlimit bool, useCalendarMonth bool, ) limiter.RateLimitCache { return &fixedRateLimitCacheImpl{ client: client, perSecondClient: perSecondClient, stopCacheKeyIncrementWhenOverlimit: stopCacheKeyIncrementWhenOverlimit, - baseRateLimiter: limiter.NewBaseRateLimit(timeSource, jitterRand, expirationJitterMaxSeconds, localCache, nearLimitRatio, cacheKeyPrefix, statsManager), + baseRateLimiter: limiter.NewBaseRateLimit(timeSource, jitterRand, expirationJitterMaxSeconds, localCache, nearLimitRatio, cacheKeyPrefix, statsManager, useCalendarMonth), } } diff --git a/src/service/ratelimit.go b/src/service/ratelimit.go index 9bd9241dc..e9fbfdc69 100644 --- a/src/service/ratelimit.go +++ b/src/service/ratelimit.go @@ -54,6 +54,7 @@ type service struct { globalShadowMode bool globalQuotaMode bool responseDynamicMetadataEnabled bool + useCalendarMonthRateLimit bool } func (this *service) SetConfig(updateEvent provider.ConfigUpdateEvent, healthyWithAtLeastOneConfigLoad bool) { @@ -90,6 +91,7 @@ func (this *service) SetConfig(updateEvent provider.ConfigUpdateEvent, healthyWi this.globalShadowMode = rlSettings.GlobalShadowMode this.globalQuotaMode = rlSettings.GlobalQuotaMode this.responseDynamicMetadataEnabled = rlSettings.ResponseDynamicMetadata + this.useCalendarMonthRateLimit = rlSettings.UseCalendarMonthRateLimit if rlSettings.RateLimitResponseHeadersEnabled { this.customHeadersEnabled = true @@ -393,7 +395,7 @@ func (this *service) rateLimitResetHeader( ) *core.HeaderValue { return &core.HeaderValue{ Key: this.customHeaderResetHeader, - Value: strconv.FormatInt(utils.CalculateReset(&descriptor.CurrentLimit.Unit, this.customHeaderClock).GetSeconds(), 10), + Value: strconv.FormatInt(utils.CalculateReset(&descriptor.CurrentLimit.Unit, this.customHeaderClock, this.useCalendarMonthRateLimit).GetSeconds(), 10), } } diff --git a/src/settings/settings.go b/src/settings/settings.go index e57317e45..ffb547083 100644 --- a/src/settings/settings.go +++ b/src/settings/settings.go @@ -110,6 +110,13 @@ type Settings struct { CacheKeyPrefix string `envconfig:"CACHE_KEY_PREFIX" default:""` BackendType string `envconfig:"BACKEND_TYPE" default:"redis"` StopCacheKeyIncrementWhenOverlimit bool `envconfig:"STOP_CACHE_KEY_INCREMENT_WHEN_OVERLIMIT" default:"false"` + // UseCalendarMonthRateLimit switches MONTH-unit rate limits to a true calendar + // month window (the 1st through the last day of the month, UTC) for cache key + // bucketing, TTL/expiration, and the reported reset time. Defaults to false, + // which preserves the legacy behavior of a fixed 30-day rolling window counted + // from the Unix epoch, so enabling this for existing MONTH limits changes when + // they reset and is opt-in. + UseCalendarMonthRateLimit bool `envconfig:"USE_CALENDAR_MONTH_RATE_LIMIT" default:"false"` // Settings for optional returning of custom headers RateLimitResponseHeadersEnabled bool `envconfig:"LIMIT_RESPONSE_HEADERS_ENABLED" default:"false"` diff --git a/src/utils/time.go b/src/utils/time.go index e7978cc6c..b67039b93 100644 --- a/src/utils/time.go +++ b/src/utils/time.go @@ -24,6 +24,37 @@ func (this *timeSourceImpl) UnixNow() int64 { return time.Now().Unix() } +// expiryUntilMonthEnd returns the duration remaining until the start of the +// next calendar month, evaluated in UTC so the result does not depend on the +// server's local timezone or DST transitions. +func expiryUntilMonthEnd(now time.Time) time.Duration { + // Always operate in UTC to avoid timezone/DST drift + nowUTC := now.UTC() + // Calculate the start of the next month in UTC + nextMonth := nowUTC.AddDate(0, 1, -nowUTC.Day()+1) + nextMonthStart := time.Date( + nextMonth.Year(), nextMonth.Month(), 1, + 0, 0, 0, 0, time.UTC, + ) + // Return the duration between now and the next month boundary + return nextMonthStart.Sub(nowUTC) +} + +// MonthExpirationSeconds returns the number of seconds remaining until the +// end of the calendar month (UTC) containing the instant represented by +// nowUnix. Used as the TTL/expiration for a MONTH-unit rate limit entry. +func MonthExpirationSeconds(nowUnix int64) int64 { + return int64(expiryUntilMonthEnd(time.Unix(nowUnix, 0)).Seconds()) +} + +// MonthStartUnix returns the Unix timestamp (UTC) of the first moment of the +// calendar month containing the instant represented by nowUnix. Used to +// bucket a MONTH-unit rate limit's cache key by calendar month. +func MonthStartUnix(nowUnix int64) int64 { + t := time.Unix(nowUnix, 0).UTC() + return time.Date(t.Year(), t.Month(), 1, 0, 0, 0, 0, time.UTC).Unix() +} + // rand for jitter. type lockedSource struct { lk sync.Mutex diff --git a/src/utils/utilities.go b/src/utils/utilities.go index fb398b764..4abc54f72 100644 --- a/src/utils/utilities.go +++ b/src/utils/utilities.go @@ -38,10 +38,25 @@ func UnitToDivider(unit pb.RateLimitResponse_RateLimit_Unit) int64 { panic("should not get here") } -func CalculateReset(unit *pb.RateLimitResponse_RateLimit_Unit, timeSource TimeSource) *durationpb.Duration { +// ExpirationSeconds returns the number of seconds, evaluated from the +// current time, until the given rate limit unit's window ends. When +// useCalendarMonth is true, MONTH reflects the actual calendar-aligned +// window (the 1st through the last day of the month, UTC) instead of the +// fixed-length UnitToDivider approximation. +func ExpirationSeconds(unit pb.RateLimitResponse_RateLimit_Unit, timeSource TimeSource, useCalendarMonth bool) int64 { + if useCalendarMonth && unit == pb.RateLimitResponse_RateLimit_MONTH { + return MonthExpirationSeconds(timeSource.UnixNow()) + } + return UnitToDivider(unit) +} + +func CalculateReset(unit *pb.RateLimitResponse_RateLimit_Unit, timeSource TimeSource, useCalendarMonth bool) *durationpb.Duration { + nowUnix := timeSource.UnixNow() + if useCalendarMonth && *unit == pb.RateLimitResponse_RateLimit_MONTH { + return &durationpb.Duration{Seconds: MonthExpirationSeconds(nowUnix)} + } sec := UnitToDivider(*unit) - now := timeSource.UnixNow() - return &durationpb.Duration{Seconds: sec - now%sec} + return &durationpb.Duration{Seconds: sec - nowUnix%sec} } // Mask credentials from a redis connection string like diff --git a/test/config/basic_config.yaml b/test/config/basic_config.yaml index 8992966fc..0b63be906 100644 --- a/test/config/basic_config.yaml +++ b/test/config/basic_config.yaml @@ -69,3 +69,9 @@ descriptors: rate_limit: unit: minute requests_per_unit: 70 + + - key: key8 + rate_limit: + name: key8_rate_limit + unit: month + requests_per_unit: 200 diff --git a/test/limiter/base_limiter_test.go b/test/limiter/base_limiter_test.go index 7bc404079..dadf47ce8 100644 --- a/test/limiter/base_limiter_test.go +++ b/test/limiter/base_limiter_test.go @@ -2,7 +2,9 @@ package limiter import ( "math/rand" + "strconv" "testing" + "time" mockstats "github.com/envoyproxy/ratelimit/test/mocks/stats" @@ -27,7 +29,7 @@ func TestGenerateCacheKeys(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) timeSource.EXPECT().UnixNow().Return(int64(1234)) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm, false) request := common.NewRateLimitRequest("domain", [][][2]string{{{"key", "value"}}}, 1) limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal(uint64(0), limits[0].Stats.TotalHits.Value()) @@ -46,7 +48,7 @@ func TestGenerateCacheKeysPrefix(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) timeSource.EXPECT().UnixNow().Return(int64(1234)) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "prefix:", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "prefix:", sm, false) request := common.NewRateLimitRequest("domain", [][][2]string{{{"key", "value"}}}, 1) limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal(uint64(0), limits[0].Stats.TotalHits.Value()) @@ -56,6 +58,40 @@ func TestGenerateCacheKeysPrefix(t *testing.T) { assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) } +func TestGenerateCacheKeysMonth(t *testing.T) { + assert := assert.New(t) + controller := gomock.NewController(t) + defer controller.Finish() + timeSource := mock_utils.NewMockTimeSource(controller) + jitterSource := mock_utils.NewMockJitterRandSource(controller) + statsStore := stats.NewStore(stats.NewNullSink(), false) + sm := mockstats.NewMockStatManager(statsStore) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm, true) + request := common.NewRateLimitRequest("domain", [][][2]string{{{"key", "value"}}}, 1) + limits := []*config.RateLimit{config.NewRateLimit(200, pb.RateLimitResponse_RateLimit_MONTH, sm.NewStats("key_value"), false, false, false, "", nil, false)} + + // 2024-01-01T00:00:00Z is the start of the January bucket. + monthStart := time.Date(2024, time.January, 1, 0, 0, 0, 0, time.UTC).Unix() + // 2024-01-31T23:59:59Z is still within January, so it must map to the same bucket. + endOfJanuary := time.Date(2024, time.January, 31, 23, 59, 59, 0, time.UTC).Unix() + // 2024-02-01T00:00:00Z is the start of the February bucket, and must differ. + startOfFebruary := time.Date(2024, time.February, 1, 0, 0, 0, 0, time.UTC).Unix() + + timeSource.EXPECT().UnixNow().Return(monthStart) + cacheKeysStart := baseRateLimit.GenerateCacheKeys(request, limits, []uint64{1}) + + timeSource.EXPECT().UnixNow().Return(endOfJanuary) + cacheKeysEndOfMonth := baseRateLimit.GenerateCacheKeys(request, limits, []uint64{1}) + + timeSource.EXPECT().UnixNow().Return(startOfFebruary) + cacheKeysNextMonth := baseRateLimit.GenerateCacheKeys(request, limits, []uint64{1}) + + assert.Equal(cacheKeysStart[0].Key, cacheKeysEndOfMonth[0].Key) + assert.NotEqual(cacheKeysStart[0].Key, cacheKeysNextMonth[0].Key) + assert.Equal("domain_key_value_"+strconv.FormatInt(monthStart, 10), cacheKeysStart[0].Key) + assert.Equal("domain_key_value_"+strconv.FormatInt(startOfFebruary, 10), cacheKeysNextMonth[0].Key) +} + func TestGenerateCacheKeysWithShareThreshold(t *testing.T) { assert := assert.New(t) controller := gomock.NewController(t) @@ -65,7 +101,7 @@ func TestGenerateCacheKeysWithShareThreshold(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) timeSource.EXPECT().UnixNow().Return(int64(1234)).AnyTimes() - baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm, false) // Test 1: Simple case - different values with same wildcard prefix generate same cache key limit := config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("files_files/*"), false, false, false, "", nil, false) @@ -144,7 +180,7 @@ func TestOverLimitWithLocalCache(t *testing.T) { localCache := freecache.NewCache(100) localCache.Set([]byte("key"), []byte("value"), 100) sm := mockstats.NewMockStatManager(stats.NewStore(stats.NewNullSink(), false)) - baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, localCache, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, localCache, 0.8, "", sm, false) // Returns true, as local cache contains over limit value for the key. assert.Equal(true, baseRateLimit.IsOverLimitWithLocalCache("key")) } @@ -154,11 +190,11 @@ func TestNoOverLimitWithLocalCache(t *testing.T) { controller := gomock.NewController(t) defer controller.Finish() sm := mockstats.NewMockStatManager(stats.NewStore(stats.NewNullSink(), false)) - baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, nil, 0.8, "", sm, false) // Returns false, as local cache is nil. assert.Equal(false, baseRateLimit.IsOverLimitWithLocalCache("domain_key_value_1234")) localCache := freecache.NewCache(100) - baseRateLimitWithLocalCache := limiter.NewBaseRateLimit(nil, nil, 3600, localCache, 0.8, "", sm) + baseRateLimitWithLocalCache := limiter.NewBaseRateLimit(nil, nil, 3600, localCache, 0.8, "", sm, false) // Returns false, as local cache does not contain value for cache key. assert.Equal(false, baseRateLimitWithLocalCache.IsOverLimitWithLocalCache("domain_key_value_1234")) } @@ -168,7 +204,7 @@ func TestGetResponseStatusEmptyKey(t *testing.T) { controller := gomock.NewController(t) defer controller.Finish() sm := mockstats.NewMockStatManager(stats.NewStore(stats.NewNullSink(), false)) - baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(nil, nil, 3600, nil, 0.8, "", sm, false) responseStatus := baseRateLimit.GetResponseDescriptorStatus("", nil, false, 1) assert.Equal(pb.RateLimitResponse_OK, responseStatus.GetCode()) assert.Equal(uint32(0), responseStatus.GetLimitRemaining()) @@ -182,7 +218,7 @@ func TestGetResponseStatusOverLimitWithLocalCache(t *testing.T) { timeSource.EXPECT().UnixNow().Return(int64(1234)) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm, false) limits := []*config.RateLimit{config.NewRateLimit(5, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 6, 4, 5) // As `isOverLimitWithLocalCache` is passed as `true`, immediate response is returned with no checks of the limits. @@ -204,7 +240,7 @@ func TestGetResponseStatusOverLimitWithLocalCacheShadowMode(t *testing.T) { timeSource.EXPECT().UnixNow().Return(int64(1234)) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm, false) // This limit is in ShadowMode limits := []*config.RateLimit{config.NewRateLimit(5, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, true, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 6, 4, 5) @@ -229,7 +265,7 @@ func TestGetResponseStatusOverLimit(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) localCache := freecache.NewCache(100) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, localCache, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, localCache, 0.8, "", sm, false) limits := []*config.RateLimit{config.NewRateLimit(5, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 7, 4, 5) responseStatus := baseRateLimit.GetResponseDescriptorStatus("key", limitInfo, false, 1) @@ -254,7 +290,7 @@ func TestGetResponseStatusOverLimitShadowMode(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) localCache := freecache.NewCache(100) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, localCache, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, localCache, 0.8, "", sm, false) // Key is in shadow_mode: true limits := []*config.RateLimit{config.NewRateLimit(5, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, true, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 7, 4, 5) @@ -277,7 +313,7 @@ func TestGetResponseStatusBelowLimit(t *testing.T) { timeSource.EXPECT().UnixNow().Return(int64(1234)) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm, false) limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 6, 9, 10) responseStatus := baseRateLimit.GetResponseDescriptorStatus("key", limitInfo, false, 1) @@ -298,7 +334,7 @@ func TestGetResponseStatusBelowLimitShadowMode(t *testing.T) { timeSource.EXPECT().UnixNow().Return(int64(1234)) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm) + baseRateLimit := limiter.NewBaseRateLimit(timeSource, nil, 3600, nil, 0.8, "", sm, false) limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, true, false, "", nil, false)} limitInfo := limiter.NewRateLimitInfo(limits[0], 2, 6, 9, 10) responseStatus := baseRateLimit.GetResponseDescriptorStatus("key", limitInfo, false, 1) diff --git a/test/memcached/cache_impl_test.go b/test/memcached/cache_impl_test.go index 606a9d844..d784a9a6e 100644 --- a/test/memcached/cache_impl_test.go +++ b/test/memcached/cache_impl_test.go @@ -44,7 +44,7 @@ func TestMemcached(t *testing.T) { client := mock_memcached.NewMockClient(controller) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "", false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) client.EXPECT().GetMulti([]string{"domain_key_value_1234"}).Return( @@ -56,8 +56,9 @@ func TestMemcached(t *testing.T) { limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -74,7 +75,8 @@ func TestMemcached(t *testing.T) { [][][2]string{ {{"key2", "value2"}}, {{"key2", "value2"}, {"subkey2", "subvalue2"}}, - }, 1) + }, 1, + ) limits = []*config.RateLimit{ nil, config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_MINUTE, sm.NewStats("key2_value2_subkey2_subvalue2"), false, false, false, "", nil, false), @@ -82,9 +84,10 @@ func TestMemcached(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ {Code: pb.RateLimitResponse_OK, CurrentLimit: nil, LimitRemaining: 0}, - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[1].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[1].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[1].Stats.NearLimit.Value()) @@ -109,17 +112,19 @@ func TestMemcached(t *testing.T) { [][][2]string{ {{"key3", "value3"}}, {{"key3", "value3"}, {"subkey3", "subvalue3"}}, - }, []uint64{1, 2}) + }, []uint64{1, 2}, + ) limits = []*config.RateLimit{ config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_HOUR, sm.NewStats("key3_value3"), false, false, false, "", nil, false), config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_DAY, sm.NewStats("key3_value3_subkey3_subvalue3"), false, false, false, "", nil, false), } assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -141,7 +146,7 @@ func TestMemcachedGetError(t *testing.T) { client := mock_memcached.NewMockClient(controller) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "", false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) client.EXPECT().GetMulti([]string{"domain_key_value_1234"}).Return( @@ -153,8 +158,9 @@ func TestMemcachedGetError(t *testing.T) { limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -171,8 +177,9 @@ func TestMemcachedGetError(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value1"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -229,7 +236,7 @@ func TestOverLimitWithLocalCache(t *testing.T) { sink := &common.TestStatSink{} statsStore := stats.NewStore(sink, true) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, localCache, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, localCache, sm, 0.8, "", false) localCacheStats := limiter.NewLocalCacheStats(localCache, statsStore.Scope("localcache")) // Test Near Limit Stats. Under Near Limit Ratio @@ -247,9 +254,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -268,9 +276,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -289,9 +298,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -307,9 +317,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { client.EXPECT().Increment("domain_key4_value4_997200", uint64(1)).Times(0) assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(4), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -331,7 +342,7 @@ func TestNearLimit(t *testing.T) { client := mock_memcached.NewMockClient(controller) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "", false) // Test Near Limit Stats. Under Near Limit Ratio timeSource.EXPECT().UnixNow().Return(int64(1000000)).MaxTimes(3) @@ -348,9 +359,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -365,9 +377,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -383,9 +396,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -403,8 +417,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key5_value5"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 15, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 15, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -421,8 +436,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(8, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key6_value6"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -439,8 +455,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key7_value7"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(3), limits[0].Stats.NearLimit.Value()) @@ -457,8 +474,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key8_value8"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -475,8 +493,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key9_value9"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(7), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(4), limits[0].Stats.NearLimit.Value()) @@ -493,8 +512,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key10_value10"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(3), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -513,7 +533,7 @@ func TestMemcacheWithJitter(t *testing.T) { jitterSource := mock_utils.NewMockJitterRandSource(controller) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, rand.New(jitterSource), 3600, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, rand.New(jitterSource), 3600, nil, sm, 0.8, "", false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) jitterSource.EXPECT().Int63().Return(int64(100)) @@ -522,7 +542,8 @@ func TestMemcacheWithJitter(t *testing.T) { client.EXPECT().GetMulti([]string{"domain_key_value_1234"}).Return(nil, nil) // First increment attempt will fail client.EXPECT().Increment("domain_key_value_1234", uint64(1)).Return( - uint64(0), memcache.ErrCacheMiss) + uint64(0), memcache.ErrCacheMiss, + ) // Add succeeds client.EXPECT().Add( &memcache.Item{ @@ -537,8 +558,9 @@ func TestMemcacheWithJitter(t *testing.T) { limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -556,14 +578,15 @@ func TestMemcacheAdd(t *testing.T) { client := mock_memcached.NewMockClient(controller) statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "", false) // Test a race condition with the initial add timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) client.EXPECT().GetMulti([]string{"domain_key_value_1234"}).Return(nil, nil) client.EXPECT().Increment("domain_key_value_1234", uint64(1)).Return( - uint64(0), memcache.ErrCacheMiss) + uint64(0), memcache.ErrCacheMiss, + ) // Add fails, must have been a race condition client.EXPECT().Add( &memcache.Item{ @@ -574,14 +597,16 @@ func TestMemcacheAdd(t *testing.T) { ).Return(memcache.ErrNotStored) // Should work the second time, since some other client added the key. client.EXPECT().Increment("domain_key_value_1234", uint64(1)).Return( - uint64(2), nil) + uint64(2), nil, + ) request := common.NewRateLimitRequest("domain", [][][2]string{{{"key", "value"}}}, 1) limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -591,7 +616,8 @@ func TestMemcacheAdd(t *testing.T) { timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) client.EXPECT().GetMulti([]string{"domain_key2_value2_1200"}).Return(nil, nil) client.EXPECT().Increment("domain_key2_value2_1200", uint64(1)).Return( - uint64(0), memcache.ErrCacheMiss) + uint64(0), memcache.ErrCacheMiss, + ) client.EXPECT().Add( &memcache.Item{ Key: "domain_key2_value2_1200", @@ -604,8 +630,9 @@ func TestMemcacheAdd(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_MINUTE, sm.NewStats("key2_value2"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 9, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -665,7 +692,7 @@ func TestMemcachedTracer(t *testing.T) { statsStore := stats.NewStore(stats.NewNullSink(), false) sm := mockstats.NewMockStatManager(statsStore) - cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "") + cache := memcached.NewRateLimitCacheImpl(client, timeSource, nil, 0, nil, sm, 0.8, "", false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) client.EXPECT().GetMulti([]string{"domain_key_value_1234"}).Return( diff --git a/test/redis/bench_test.go b/test/redis/bench_test.go index 5b7b968d7..c1535239a 100644 --- a/test/redis/bench_test.go +++ b/test/redis/bench_test.go @@ -47,7 +47,7 @@ func BenchmarkParallelDoLimit(b *testing.B) { client := redis.NewClientImpl(context.Background(), statsStore, false, "", "tcp", "single", "127.0.0.1:6379", poolSize, pipelineWindow, pipelineLimit, nil, false, nil, 10*time.Second, "", "", time.Second, 30*time.Second, 0, false) defer client.Close() - cache := redis.NewFixedRateLimitCacheImpl(client, nil, utils.NewTimeSourceImpl(), rand.New(utils.NewLockedSource(time.Now().Unix())), 10, nil, 0.8, "", sm, true) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, utils.NewTimeSourceImpl(), rand.New(utils.NewLockedSource(time.Now().Unix())), 10, nil, 0.8, "", sm, true, false) request := common.NewRateLimitRequest("domain", [][][2]string{{{"key", "value"}}}, 1) limits := []*config.RateLimit{config.NewRateLimit(1000000000, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} diff --git a/test/redis/fixed_cache_impl_test.go b/test/redis/fixed_cache_impl_test.go index 1de91d012..8aab18f1f 100644 --- a/test/redis/fixed_cache_impl_test.go +++ b/test/redis/fixed_cache_impl_test.go @@ -54,9 +54,9 @@ func testRedis(usePerSecondRedis bool) func(*testing.T) { timeSource := mock_utils.NewMockTimeSource(controller) var cache limiter.RateLimitCache if usePerSecondRedis { - cache = redis.NewFixedRateLimitCacheImpl(client, perSecondClient, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false) + cache = redis.NewFixedRateLimitCacheImpl(client, perSecondClient, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false, false) } else { - cache = redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false) + cache = redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false, false) } timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) @@ -75,8 +75,9 @@ func testRedis(usePerSecondRedis bool) func(*testing.T) { limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -94,7 +95,8 @@ func testRedis(usePerSecondRedis bool) func(*testing.T) { [][][2]string{ {{"key2", "value2"}}, {{"key2", "value2"}, {"subkey2", "subvalue2"}}, - }, []uint64{0, 1}) + }, []uint64{0, 1}, + ) limits = []*config.RateLimit{ nil, config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_MINUTE, sm.NewStats("key2_value2_subkey2_subvalue2"), false, false, false, "", nil, false), @@ -102,9 +104,10 @@ func testRedis(usePerSecondRedis bool) func(*testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ {Code: pb.RateLimitResponse_OK, CurrentLimit: nil, LimitRemaining: 0}, - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[1].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[1].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[1].Stats.NearLimit.Value()) @@ -125,17 +128,19 @@ func testRedis(usePerSecondRedis bool) func(*testing.T) { [][][2]string{ {{"key3", "value3"}}, {{"key3", "value3"}, {"subkey3", "subvalue3"}}, - }, []uint64{0, 1}) + }, []uint64{0, 1}, + ) limits = []*config.RateLimit{ config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_HOUR, sm.NewStats("key3_value3"), false, false, false, "", nil, false), config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_DAY, sm.NewStats("key3_value3_subkey3_subvalue3"), false, false, false, "", nil, false), } assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(0), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -196,7 +201,7 @@ func TestOverLimitWithLocalCache(t *testing.T) { sink := common.NewTestStatSink() statsStore := gostats.NewStore(sink, false) sm := stats.NewMockStatManager(statsStore) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, false) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, false, false) localCacheScopeName := "localcache" localCacheStats := limiter.NewLocalCacheStats(localCache, statsStore.Scope(localCacheScopeName)) @@ -216,9 +221,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -237,9 +243,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -258,9 +265,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -277,9 +285,10 @@ func TestOverLimitWithLocalCache(t *testing.T) { "EXPIRE", "domain_key4_value4_997200", int64(3600)).Times(0) assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(4), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -299,7 +308,7 @@ func TestNearLimit(t *testing.T) { timeSource := mock_utils.NewMockTimeSource(controller) statsStore := gostats.NewStore(gostats.NewNullSink(), false) sm := stats.NewMockStatManager(statsStore) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false, false) // Test Near Limit Stats. Under Near Limit Ratio timeSource.EXPECT().UnixNow().Return(int64(1000000)).MaxTimes(3) @@ -316,9 +325,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -333,9 +343,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -351,9 +362,10 @@ func TestNearLimit(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -370,8 +382,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key5_value5"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 15, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 15, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -387,8 +400,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(8, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key6_value6"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -404,8 +418,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key7_value7"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(3), limits[0].Stats.NearLimit.Value()) @@ -421,8 +436,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key8_value8"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(1), limits[0].Stats.NearLimit.Value()) @@ -438,8 +454,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(20, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key9_value9"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(7), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(2), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(4), limits[0].Stats.NearLimit.Value()) @@ -455,8 +472,9 @@ func TestNearLimit(t *testing.T) { limits = []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key10_value10"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(3), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -473,7 +491,7 @@ func TestRedisWithJitter(t *testing.T) { jitterSource := mock_utils.NewMockJitterRandSource(controller) statsStore := gostats.NewStore(gostats.NewNullSink(), false) sm := stats.NewMockStatManager(statsStore) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm, false) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(jitterSource), 3600, nil, 0.8, "", sm, false, false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) jitterSource.EXPECT().Int63().Return(int64(100)) @@ -485,8 +503,9 @@ func TestRedisWithJitter(t *testing.T) { limits := []*config.RateLimit{config.NewRateLimit(10, pb.RateLimitResponse_RateLimit_SECOND, sm.NewStats("key_value"), false, false, false, "", nil, false)} assert.Equal( - []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}}, - cache.DoLimit(context.Background(), request, limits)) + []*pb.RateLimitResponse_DescriptorStatus{{Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 5, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}}, + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.NearLimit.Value()) @@ -504,7 +523,7 @@ func TestOverLimitWithLocalCacheShadowRule(t *testing.T) { sink := common.NewTestStatSink() statsStore := gostats.NewStore(sink, false) sm := stats.NewMockStatManager(statsStore) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, false) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, false, false) localCacheScopeName := "localcache" localCacheStats := limiter.NewLocalCacheStats(localCache, statsStore.Scope(localCacheScopeName)) @@ -524,9 +543,10 @@ func TestOverLimitWithLocalCacheShadowRule(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -545,9 +565,10 @@ func TestOverLimitWithLocalCacheShadowRule(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(2), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -567,9 +588,10 @@ func TestOverLimitWithLocalCacheShadowRule(t *testing.T) { // The result should be OK since limit is in ShadowMode assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(1), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -589,9 +611,10 @@ func TestOverLimitWithLocalCacheShadowRule(t *testing.T) { // The result should be OK since limit is in ShadowMode assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) // Even if you hit the local cache, other metrics should increase normally. assert.Equal(uint64(4), limits[0].Stats.TotalHits.Value()) @@ -618,7 +641,7 @@ func TestRedisTracer(t *testing.T) { client := mock_redis.NewMockClient(controller) timeSource := mock_utils.NewMockTimeSource(controller) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, nil, 0.8, "", sm, false, false) timeSource.EXPECT().UnixNow().Return(int64(1234)).MaxTimes(3) @@ -647,7 +670,7 @@ func TestOverLimitWithStopCacheKeyIncrementWhenOverlimitConfig(t *testing.T) { sink := common.NewTestStatSink() statsStore := gostats.NewStore(sink, false) sm := stats.NewMockStatManager(statsStore) - cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, true) + cache := redis.NewFixedRateLimitCacheImpl(client, nil, timeSource, rand.New(rand.NewSource(1)), 0, localCache, 0.8, "", sm, true, false) localCacheScopeName := "localcache" localCacheStats := limiter.NewLocalCacheStats(localCache, statsStore.Scope(localCacheScopeName)) @@ -675,10 +698,11 @@ func TestOverLimitWithStopCacheKeyIncrementWhenOverlimitConfig(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[1].Limit, LimitRemaining: 3, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 4, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[1].Limit, LimitRemaining: 3, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(1), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -709,10 +733,11 @@ func TestOverLimitWithStopCacheKeyIncrementWhenOverlimitConfig(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[1].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[1].Limit, LimitRemaining: 1, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(3), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) @@ -741,10 +766,11 @@ func TestOverLimitWithStopCacheKeyIncrementWhenOverlimitConfig(t *testing.T) { assert.Equal( []*pb.RateLimitResponse_DescriptorStatus{ - {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource)}, - {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource)}, + {Code: pb.RateLimitResponse_OK, CurrentLimit: limits[0].Limit, LimitRemaining: 2, DurationUntilReset: utils.CalculateReset(&limits[0].Limit.Unit, timeSource, false)}, + {Code: pb.RateLimitResponse_OVER_LIMIT, CurrentLimit: limits[1].Limit, LimitRemaining: 0, DurationUntilReset: utils.CalculateReset(&limits[1].Limit.Unit, timeSource, false)}, }, - cache.DoLimit(context.Background(), request, limits)) + cache.DoLimit(context.Background(), request, limits), + ) assert.Equal(uint64(5), limits[0].Stats.TotalHits.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimit.Value()) assert.Equal(uint64(0), limits[0].Stats.OverLimitWithLocalCache.Value()) diff --git a/test/utils/utilities_test.go b/test/utils/utilities_test.go index aa3768d48..cadd2983f 100644 --- a/test/utils/utilities_test.go +++ b/test/utils/utilities_test.go @@ -2,10 +2,14 @@ package utils_test import ( "testing" + "time" + pb "github.com/envoyproxy/go-control-plane/envoy/service/ratelimit/v3" + gomock "github.com/golang/mock/gomock" "github.com/stretchr/testify/assert" "github.com/envoyproxy/ratelimit/src/utils" + mock_utils "github.com/envoyproxy/ratelimit/test/mocks/utils" ) func TestMaskCredentialsInUrl(t *testing.T) { @@ -39,3 +43,98 @@ func TestMaskCredentialsInUrlSentinel(t *testing.T) { expected = "foob@r,redis://*****@redis1:6379,redis://*****@redis2:6379" assert.Equal(t, expected, utils.MaskCredentialsInUrl(url)) } + +func TestCalculateResetMonthEndOfMonth(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + timeSource := mock_utils.NewMockTimeSource(controller) + // 2024-01-31T23:00:00Z, one hour before the February rollover. + now := time.Date(2024, time.January, 31, 23, 0, 0, 0, time.UTC).Unix() + timeSource.EXPECT().UnixNow().Return(now) + + unit := pb.RateLimitResponse_RateLimit_MONTH + reset := utils.CalculateReset(&unit, timeSource, true) + + assert.EqualValues(t, (1 * time.Hour).Seconds(), reset.Seconds) +} + +func TestCalculateResetMonthDisabledUsesLegacyDivider(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + timeSource := mock_utils.NewMockTimeSource(controller) + // With the feature flag off, MONTH must keep behaving like the legacy + // fixed 30-day divider, regardless of the actual calendar date. + now := time.Date(2024, time.January, 31, 23, 0, 0, 0, time.UTC).Unix() + timeSource.EXPECT().UnixNow().Return(now) + + unit := pb.RateLimitResponse_RateLimit_MONTH + reset := utils.CalculateReset(&unit, timeSource, false) + + sec := utils.UnitToDivider(unit) + assert.EqualValues(t, sec-now%sec, reset.Seconds) +} + +func TestCalculateResetMonthLeapYear(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + timeSource := mock_utils.NewMockTimeSource(controller) + // 2024 is a leap year, so February has 29 days: Feb 28 -> Mar 1 is 2 days away. + now := time.Date(2024, time.February, 28, 0, 0, 0, 0, time.UTC).Unix() + timeSource.EXPECT().UnixNow().Return(now) + + unit := pb.RateLimitResponse_RateLimit_MONTH + reset := utils.CalculateReset(&unit, timeSource, true) + + assert.EqualValues(t, (48 * time.Hour).Seconds(), reset.Seconds) +} + +func TestMonthStartUnix(t *testing.T) { + midMonth := time.Date(2024, time.January, 15, 12, 30, 0, 0, time.UTC).Unix() + expected := time.Date(2024, time.January, 1, 0, 0, 0, 0, time.UTC).Unix() + assert.Equal(t, expected, utils.MonthStartUnix(midMonth)) + + // A non-UTC instant must still bucket by its UTC calendar month. + inTokyo := time.Date(2024, time.February, 1, 5, 0, 0, 0, time.FixedZone("JST", 9*60*60)).Unix() + expected = time.Date(2024, time.January, 1, 0, 0, 0, 0, time.UTC).Unix() + assert.Equal(t, expected, utils.MonthStartUnix(inTokyo)) +} + +func TestExpirationSecondsMonth(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + timeSource := mock_utils.NewMockTimeSource(controller) + now := time.Date(2024, time.February, 28, 0, 0, 0, 0, time.UTC).Unix() + timeSource.EXPECT().UnixNow().Return(now) + + seconds := utils.ExpirationSeconds(pb.RateLimitResponse_RateLimit_MONTH, timeSource, true) + assert.EqualValues(t, (48 * time.Hour).Seconds(), seconds) +} + +func TestExpirationSecondsMonthDisabledUsesLegacyDivider(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + // No UnixNow() expectation is set: with the feature flag off, MONTH must + // not consult the time source at all, matching the legacy UnitToDivider + // behavior exactly. + timeSource := mock_utils.NewMockTimeSource(controller) + + seconds := utils.ExpirationSeconds(pb.RateLimitResponse_RateLimit_MONTH, timeSource, false) + assert.EqualValues(t, 60*60*24*30, seconds) +} + +func TestExpirationSecondsNonMonthDoesNotUseTimeSource(t *testing.T) { + controller := gomock.NewController(t) + defer controller.Finish() + + // No UnixNow() expectation is set: a fixed-length unit must not consult + // the time source at all, matching the pre-existing UnitToDivider behavior. + timeSource := mock_utils.NewMockTimeSource(controller) + + seconds := utils.ExpirationSeconds(pb.RateLimitResponse_RateLimit_DAY, timeSource, true) + assert.EqualValues(t, 60*60*24, seconds) +} From bd0e138f401469e98931905e71f5a04e4e251564 Mon Sep 17 00:00:00 2001 From: kiran malsetty Date: Thu, 6 Aug 2026 12:40:07 -0400 Subject: [PATCH 2/3] Trigger CI Signed-off-by: kiran malsetty From 751de87e674b062d00ffd166230e2636a22cb805 Mon Sep 17 00:00:00 2001 From: kiran malsetty Date: Fri, 7 Aug 2026 08:40:00 -0400 Subject: [PATCH 3/3] docs: regenerate README TOC for new calendar-month section Signed-off-by: kiran malsetty --- README.md | 1 + 1 file changed, 1 insertion(+) diff --git a/README.md b/README.md index 4e9459eb8..95738b98e 100644 --- a/README.md +++ b/README.md @@ -70,6 +70,7 @@ - [Two Redis Instances](#two-redis-instances) - [Health Checking for Redis Active Connection](#health-checking-for-redis-active-connection) - [Recovering from a failover (READONLY errors)](#recovering-from-a-failover-readonly-errors) + - [Calendar-aligned MONTH rate limits](#calendar-aligned-month-rate-limits) - [Memcache](#memcache) - [Custom headers](#custom-headers) - [Tracing](#tracing)