diff --git a/README.md b/README.md index 01256d45c..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) @@ -1396,6 +1397,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) +}