feat(auth): implement decoupled sliding-window rate limiting for login and oauth

This commit is contained in:
ryan
2026-09-02 22:15:21 +08:00
parent 39f02b5d7a
commit 9632604958
23 changed files with 927 additions and 50 deletions
+119
View File
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package limiter provides in-memory rate limiting utilities.
package limiter
import (
"Wavelet/core/contracts"
"context"
"sync"
"time"
)
type memoryEntry struct {
timestamps []time.Time
lastSeen time.Time
}
func (e *memoryEntry) prune(cutoff time.Time) {
validIdx := len(e.timestamps)
for i, ts := range e.timestamps {
if ts.After(cutoff) {
validIdx = i
break
}
}
if validIdx > 0 && validIdx <= len(e.timestamps) {
e.timestamps = e.timestamps[validIdx:]
}
}
func (e *memoryEntry) calcBlockedResult(limit int, period time.Duration, now time.Time) *contracts.RateLimitResult {
currentCount := len(e.timestamps)
if currentCount == 0 {
return &contracts.RateLimitResult{
Allowed: false,
Remaining: limit,
ResetAfter: period,
RetryAfter: 0,
}
}
oldest := e.timestamps[0]
retryAfter := max(0, oldest.Add(period).Sub(now))
newest := e.timestamps[currentCount-1]
resetAfter := max(0, newest.Add(period).Sub(now))
return &contracts.RateLimitResult{
Allowed: false,
Remaining: limit - currentCount,
ResetAfter: resetAfter,
RetryAfter: retryAfter,
}
}
// MemoryLimiter implements contracts.LimiterService using an in-memory sliding window algorithm.
type MemoryLimiter struct {
mu sync.Mutex
entries map[string]*memoryEntry
}
// NewMemoryLimiter creates a new in-memory rate limiter.
func NewMemoryLimiter() *MemoryLimiter {
return &MemoryLimiter{
entries: make(map[string]*memoryEntry),
}
}
// Allow checks whether 1 event for key is permitted under rate.
func (m *MemoryLimiter) Allow(ctx context.Context, key string, rate contracts.Rate) (*contracts.RateLimitResult, error) {
return m.AllowN(ctx, key, rate, 1)
}
// AllowN checks whether n events for key are permitted under rate.
func (m *MemoryLimiter) AllowN(_ context.Context, key string, rate contracts.Rate, n int) (*contracts.RateLimitResult, error) {
if rate.Limit <= 0 || rate.Period <= 0 || n <= 0 {
return &contracts.RateLimitResult{Allowed: true}, nil
}
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
cutoff := now.Add(-rate.Period)
entry, ok := m.entries[key]
if !ok {
entry = &memoryEntry{}
m.entries[key] = entry
}
entry.lastSeen = now
entry.prune(cutoff)
if len(entry.timestamps)+n > rate.Limit {
return entry.calcBlockedResult(rate.Limit, rate.Period, now), nil
}
for i := 0; i < n; i++ {
entry.timestamps = append(entry.timestamps, now)
}
remaining := max(0, rate.Limit-len(entry.timestamps))
return &contracts.RateLimitResult{
Allowed: true,
Remaining: remaining,
ResetAfter: rate.Period,
RetryAfter: 0,
}, nil
}
// Reset clears rate limit state for key.
func (m *MemoryLimiter) Reset(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.entries, key)
return nil
}
+119
View File
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package limiter
import (
"Wavelet/core/contracts"
"context"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMemoryLimiter_Basic(t *testing.T) {
ctx := context.Background()
lim := NewMemoryLimiter()
rate := contracts.Rate{
Limit: 3,
Period: 100 * time.Millisecond,
}
// 1st request
res, err := lim.Allow(ctx, "test_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
assert.Equal(t, 2, res.Remaining)
// 2nd request
res, err = lim.Allow(ctx, "test_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
assert.Equal(t, 1, res.Remaining)
// 3rd request
res, err = lim.Allow(ctx, "test_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
assert.Equal(t, 0, res.Remaining)
// 4th request - should be blocked
res, err = lim.Allow(ctx, "test_key", rate)
require.NoError(t, err)
assert.False(t, res.Allowed)
assert.Equal(t, 0, res.Remaining)
assert.Greater(t, res.RetryAfter, time.Duration(0))
// Reset
err = lim.Reset(ctx, "test_key")
require.NoError(t, err)
// Immediately allowed after reset
res, err = lim.Allow(ctx, "test_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
assert.Equal(t, 2, res.Remaining)
}
func TestMemoryLimiter_WindowSlide(t *testing.T) {
ctx := context.Background()
lim := NewMemoryLimiter()
rate := contracts.Rate{
Limit: 2,
Period: 50 * time.Millisecond,
}
res, err := lim.Allow(ctx, "slide_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
res, err = lim.Allow(ctx, "slide_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
res, err = lim.Allow(ctx, "slide_key", rate)
require.NoError(t, err)
assert.False(t, res.Allowed)
// Wait for window to slide
time.Sleep(60 * time.Millisecond)
res, err = lim.Allow(ctx, "slide_key", rate)
require.NoError(t, err)
assert.True(t, res.Allowed)
}
func TestMemoryLimiter_Concurrency(t *testing.T) {
ctx := context.Background()
lim := NewMemoryLimiter()
rate := contracts.Rate{
Limit: 100,
Period: time.Second,
}
var wg sync.WaitGroup
allowedCount := int32(0)
var mu sync.Mutex
for i := 0; i < 200; i++ {
wg.Add(1)
go func() {
defer wg.Done()
res, err := lim.Allow(ctx, "concurrent_key", rate)
if err == nil && res.Allowed {
mu.Lock()
allowedCount++
mu.Unlock()
}
}()
}
wg.Wait()
assert.Equal(t, int32(100), allowedCount)
}