mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 15:06:36 +08:00
feat(auth): implement decoupled sliding-window rate limiting for login and oauth
This commit is contained in:
@@ -125,6 +125,21 @@ func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||
if sessionHash == "" {
|
||||
return nil
|
||||
}
|
||||
if limiter := getLimiter(ctx); limiter != nil {
|
||||
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
|
||||
res, err := limiter.Allow(ctx, key, contracts.Rate{
|
||||
Limit: oauthStateLimitMax,
|
||||
Period: OAuthStateCacheKeyExpiration,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !res.Allowed {
|
||||
return errors.New(errOAuthStateRateLimited)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
cache := getCache(ctx)
|
||||
if cache == nil {
|
||||
return nil
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package auth_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/limiter"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestOAuthRateLimiting(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
ctx := core.NewContext(context.Background())
|
||||
ctx.Config().SetSource(core.NewMapSource(nil))
|
||||
require.NoError(t, ctx.Config().Resolve())
|
||||
|
||||
testDB := setupTestDB(t)
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
|
||||
// Provide in-memory limiter service
|
||||
memLimiter := limiter.NewMemoryLimiter()
|
||||
core.Provide[contracts.LimiterService](ctx, memLimiter)
|
||||
|
||||
require.NoError(t, auth.New().Apply(ctx))
|
||||
|
||||
// Create an active OIDC source
|
||||
authSrc := auth.AuthSource{
|
||||
ID: 1,
|
||||
Name: "google",
|
||||
Type: "oidc",
|
||||
DisplayName: "Google",
|
||||
ClientID: "client-id-123",
|
||||
ClientSecret: "client-secret-456",
|
||||
OpenIDDiscoveryURL: "https://accounts.google.com",
|
||||
IsActive: true,
|
||||
}
|
||||
require.NoError(t, testDB.Create(&authSrc).Error)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(response.ErrorHandlerMiddleware())
|
||||
store := cookie.NewStore([]byte("test-session-secret-123"))
|
||||
router.Use(sessions.Sessions("wavelet_session_id", store))
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root()))
|
||||
c.Next()
|
||||
})
|
||||
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))
|
||||
for _, m := range rd.Middlewares {
|
||||
if h, ok := m.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := m.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
for _, raw := range rd.Handlers {
|
||||
if h, ok := raw.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := raw.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
router.Handle(rd.Method, rd.Path, handlers...)
|
||||
}
|
||||
|
||||
// 10 state slots are allowed per session (oauthStateLimitMax = 10)
|
||||
// We'll simulate 10 requests with the same cookie
|
||||
var cookies []*http.Cookie
|
||||
for i := 1; i <= 10; i++ {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil)
|
||||
for _, ck := range cookies {
|
||||
req.AddCookie(ck)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
if len(w.Result().Cookies()) > 0 {
|
||||
cookies = w.Result().Cookies()
|
||||
}
|
||||
}
|
||||
|
||||
// 11th request for the same session should be rate limited
|
||||
{
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil)
|
||||
for _, ck := range cookies {
|
||||
req.AddCookie(ck)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
var resp map[string]any
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "请求授权过于频繁,请稍后重试", resp["error_msg"])
|
||||
}
|
||||
}
|
||||
@@ -89,9 +89,11 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
|
||||
core.Bind[contracts.DBService](ctx, setDBService)
|
||||
core.Bind[contracts.CacheService](ctx, setCacheService)
|
||||
core.Bind[contracts.LimiterService](ctx, setLimiterService)
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
setCacheService(nil)
|
||||
setLimiterService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
|
||||
@@ -62,6 +62,13 @@ func hashToken(token string) string {
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
type testSystemConfig struct {
|
||||
Key string `gorm:"primaryKey"`
|
||||
Value string
|
||||
}
|
||||
|
||||
func (testSystemConfig) TableName() string { return "w_system_configs" }
|
||||
|
||||
func setupTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
|
||||
@@ -73,6 +80,7 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
||||
&testAccessToken{},
|
||||
&auth.AuthSource{},
|
||||
&auth.ExternalAccount{},
|
||||
&testSystemConfig{},
|
||||
))
|
||||
|
||||
return testDB
|
||||
|
||||
@@ -15,10 +15,12 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
limiterMu sync.RWMutex
|
||||
limiterSvc contracts.LimiterService
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
@@ -33,6 +35,12 @@ func setCacheService(s contracts.CacheService) {
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
func setLimiterService(s contracts.LimiterService) {
|
||||
limiterMu.Lock()
|
||||
defer limiterMu.Unlock()
|
||||
limiterSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
@@ -56,6 +64,16 @@ func getCache(ctx context.Context) contracts.CacheService {
|
||||
return s
|
||||
}
|
||||
|
||||
func getLimiter(ctx context.Context) contracts.LimiterService {
|
||||
if s, err := core.InjectFrom[contracts.LimiterService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
limiterMu.RLock()
|
||||
s := limiterSvc
|
||||
limiterMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
|
||||
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, error) {
|
||||
var row struct {
|
||||
|
||||
@@ -8,7 +8,9 @@ const (
|
||||
errInvalidParams = "无效的请求参数"
|
||||
errUserNotFound = "用户不存在"
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
errPasswordMismatch = "用户名或密码错误"
|
||||
errPasswordMismatch = "用户名或密码错误"
|
||||
errTooManyLoginAttempts = "登录尝试过于频繁,请稍后重试"
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
errOldPasswordIncorrect = "原密码不正确"
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -79,15 +81,16 @@ func invalidateTokenCache(ctx context.Context, tokenHash string) {
|
||||
}
|
||||
}
|
||||
|
||||
// Login 用户密码登录
|
||||
// @Summary 用户密码登录
|
||||
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
||||
// Login 用户登录
|
||||
// @Summary 用户登录
|
||||
// @Description 使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.loginRequest true "登录请求参数"
|
||||
// @Success 200 {object} response.Any "登录成功,返回用户信息"
|
||||
// @Failure 400 {object} response.Any "用户名或密码错误"
|
||||
// @Failure 429 {object} response.Any "登录尝试过于频繁"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/user/login [post]
|
||||
func Login(c *gin.Context) {
|
||||
@@ -97,8 +100,28 @@ func Login(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
user, err := GetUserByUsername(c.Request.Context(), req.Username)
|
||||
ctx := c.Request.Context()
|
||||
clientIP := c.ClientIP()
|
||||
rateKey := "auth:login:ip:" + clientIP
|
||||
if clientIP == "" {
|
||||
rateKey = "auth:login:user:" + req.Username
|
||||
}
|
||||
|
||||
limiter := getLimiter(ctx)
|
||||
if limiter != nil {
|
||||
res, err := limiter.Allow(ctx, rateKey, contracts.Rate{
|
||||
Limit: 5,
|
||||
Period: 1 * time.Minute,
|
||||
})
|
||||
if err == nil && !res.Allowed {
|
||||
response.AbortTooManyRequests(c, errTooManyLoginAttempts)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
user, err := GetUserByUsername(ctx, req.Username)
|
||||
if err != nil {
|
||||
util.DummyCheckPassword(req.Password)
|
||||
response.AbortUnauthorized(c, errPasswordMismatch)
|
||||
return
|
||||
}
|
||||
@@ -108,9 +131,13 @@ func Login(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if limiter != nil {
|
||||
_ = limiter.Reset(ctx, rateKey)
|
||||
}
|
||||
|
||||
if user.ID == 0 {
|
||||
newID := idgen.NextUint64ID()
|
||||
if err := getDB(c.Request.Context()).Model(&User{}).Where("username = ?", user.Username).Update("id", newID).Error; err == nil {
|
||||
if err := getDB(ctx).Model(&User{}).Where("username = ?", user.Username).Update("id", newID).Error; err == nil {
|
||||
user.ID = newID
|
||||
}
|
||||
}
|
||||
@@ -122,7 +149,7 @@ func Login(c *gin.Context) {
|
||||
user.NeedChangePassword = needChange
|
||||
sess.Set("need_change_password", needChange)
|
||||
if err := sess.Save(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "save session failed on login: %v", err)
|
||||
logger.ErrorF(ctx, "save session failed on login: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(user))
|
||||
@@ -137,6 +164,7 @@ func Login(c *gin.Context) {
|
||||
// @Param request body user.registerRequest true "注册请求参数"
|
||||
// @Success 200 {object} response.Any "注册并登录成功,返回用户信息"
|
||||
// @Failure 400 {object} response.Any "参数错误、用户名已存在或注册已关闭"
|
||||
// @Failure 429 {object} response.Any "注册尝试过于频繁"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/user/register [post]
|
||||
func Register(c *gin.Context) {
|
||||
@@ -146,6 +174,19 @@ func Register(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
clientIP := c.ClientIP()
|
||||
if limiter := getLimiter(ctx); limiter != nil && clientIP != "" {
|
||||
res, err := limiter.Allow(ctx, "auth:register:ip:"+clientIP, contracts.Rate{
|
||||
Limit: 10,
|
||||
Period: 1 * time.Minute,
|
||||
})
|
||||
if err == nil && !res.Allowed {
|
||||
response.AbortTooManyRequests(c, errTooManyLoginAttempts)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
newUser := &User{
|
||||
Username: req.Username,
|
||||
Email: req.Email,
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/limiter"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"Wavelet/plugins/domain/user"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestUserLoginRateLimiting(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
ctx := core.NewContext(context.Background())
|
||||
ctx.Config().SetSource(core.NewMapSource(nil))
|
||||
require.NoError(t, ctx.Config().Resolve())
|
||||
|
||||
testDB := setupTestDB(t)
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
|
||||
// Provide in-memory limiter service
|
||||
memLimiter := limiter.NewMemoryLimiter()
|
||||
core.Provide[contracts.LimiterService](ctx, memLimiter)
|
||||
|
||||
require.NoError(t, auth.New().Apply(ctx))
|
||||
require.NoError(t, user.New().Apply(ctx))
|
||||
|
||||
// Create a test user
|
||||
userSvc, err := core.Inject[contracts.UserService](ctx)
|
||||
require.NoError(t, err)
|
||||
createdUser, err := userSvc.CreateUser(context.Background(), contracts.CreateUserRequest{
|
||||
Username: "ratelimit_user",
|
||||
Password: "CorrectPassword123!",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, createdUser)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(response.ErrorHandlerMiddleware())
|
||||
store := cookie.NewStore([]byte("test-secret-key-session"))
|
||||
router.Use(sessions.Sessions("wavelet_session", store))
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root()))
|
||||
c.Next()
|
||||
})
|
||||
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))
|
||||
for _, m := range rd.Middlewares {
|
||||
if h, ok := m.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := m.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
for _, raw := range rd.Handlers {
|
||||
if h, ok := raw.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := raw.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
router.Handle(rd.Method, rd.Path, handlers...)
|
||||
}
|
||||
|
||||
loginBody, _ := json.Marshal(map[string]string{
|
||||
"username": "ratelimit_user",
|
||||
"password": "WrongPassword!",
|
||||
})
|
||||
|
||||
// Make 5 failed login attempts (Limit is 5)
|
||||
for i := 1; i <= 5; i++ {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = "192.168.1.100:12345"
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code, "attempt %d should be 401 Unauthorized", i)
|
||||
}
|
||||
|
||||
// 6th attempt from the same IP should be blocked with 429 Too Many Requests
|
||||
{
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = "192.168.1.100:12345"
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusTooManyRequests, w.Code, "6th attempt should be 429 Too Many Requests")
|
||||
|
||||
var resp map[string]any
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "登录尝试过于频繁,请稍后重试", resp["error_msg"])
|
||||
}
|
||||
|
||||
// Another IP is not blocked
|
||||
{
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = "192.168.1.101:12345"
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code, "different IP should receive 401, not 429")
|
||||
}
|
||||
}
|
||||
@@ -79,10 +79,12 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
core.Bind[contracts.DBService](ctx, SetDBService)
|
||||
core.Bind[contracts.CacheService](ctx, SetCacheService)
|
||||
core.Bind[contracts.TaskService](ctx, SetTaskService)
|
||||
core.Bind[contracts.LimiterService](ctx, SetLimiterService)
|
||||
ctx.OnDispose(func() error {
|
||||
SetDBService(nil)
|
||||
SetCacheService(nil)
|
||||
SetTaskService(nil)
|
||||
SetLimiterService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
|
||||
@@ -18,8 +18,10 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
limiterMu sync.RWMutex
|
||||
limiterSvc contracts.LimiterService
|
||||
)
|
||||
|
||||
// SetDBService sets the active DBService contract for the user domain plugin.
|
||||
@@ -29,6 +31,13 @@ func SetDBService(s contracts.DBService) {
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetLimiterService sets the active LimiterService contract for the user domain plugin.
|
||||
func SetLimiterService(s contracts.LimiterService) {
|
||||
limiterMu.Lock()
|
||||
defer limiterMu.Unlock()
|
||||
limiterSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
@@ -44,6 +53,17 @@ func getDB(ctx context.Context) *gorm.DB {
|
||||
return nil
|
||||
}
|
||||
|
||||
func getLimiter(ctx context.Context) contracts.LimiterService {
|
||||
if s, err := core.InjectFrom[contracts.LimiterService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
|
||||
limiterMu.RLock()
|
||||
s := limiterSvc
|
||||
limiterMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetUserByID 通过 ID 获取用户
|
||||
func GetUserByID(ctx context.Context, id uint64) (*User, error) {
|
||||
var u User
|
||||
|
||||
+59
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cache
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"context"
|
||||
|
||||
"github.com/go-redis/redis_rate/v10"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
type redisLimiterImpl struct {
|
||||
limiter *redis_rate.Limiter
|
||||
keyPrefix string
|
||||
}
|
||||
|
||||
func newRedisLimiter(client redis.UniversalClient, keyPrefix string) contracts.LimiterService {
|
||||
return &redisLimiterImpl{
|
||||
limiter: redis_rate.NewLimiter(client),
|
||||
keyPrefix: keyPrefix,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *redisLimiterImpl) prefixedKey(key string) string {
|
||||
if r.keyPrefix != "" {
|
||||
return r.keyPrefix + "limiter:" + key
|
||||
}
|
||||
return PrefixedKey("limiter:" + key)
|
||||
}
|
||||
|
||||
func (r *redisLimiterImpl) Allow(ctx context.Context, key string, rate contracts.Rate) (*contracts.RateLimitResult, error) {
|
||||
return r.AllowN(ctx, key, rate, 1)
|
||||
}
|
||||
|
||||
func (r *redisLimiterImpl) AllowN(ctx context.Context, key string, rate contracts.Rate, n int) (*contracts.RateLimitResult, error) {
|
||||
limit := redis_rate.Limit{
|
||||
Rate: rate.Limit,
|
||||
Period: rate.Period,
|
||||
Burst: rate.Limit,
|
||||
}
|
||||
|
||||
res, err := r.limiter.AllowN(ctx, r.prefixedKey(key), limit, n)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &contracts.RateLimitResult{
|
||||
Allowed: res.Allowed > 0,
|
||||
Remaining: res.Remaining,
|
||||
ResetAfter: res.ResetAfter,
|
||||
RetryAfter: res.RetryAfter,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *redisLimiterImpl) Reset(ctx context.Context, key string) error {
|
||||
return r.limiter.Reset(ctx, r.prefixedKey(key))
|
||||
}
|
||||
+84
@@ -0,0 +1,84 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cache_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/infra/cache"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestRedisLimiterService(t *testing.T) {
|
||||
mr, err := miniredis.Run()
|
||||
require.NoError(t, err)
|
||||
defer mr.Close()
|
||||
|
||||
rdb := redis.NewClient(&redis.Options{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
defer func() { _ = rdb.Close() }()
|
||||
|
||||
p := cache.New(
|
||||
cache.WithRedis(rdb),
|
||||
cache.WithKeyPrefix("test:"),
|
||||
)
|
||||
ctx := core.NewContext(context.Background())
|
||||
ctx.Config().SetSource(core.NewMapSource(map[string]any{
|
||||
"redis.enabled": true,
|
||||
}))
|
||||
require.NoError(t, ctx.Config().Resolve())
|
||||
require.NoError(t, p.Apply(ctx))
|
||||
|
||||
limiter, err := core.Inject[contracts.LimiterService](ctx)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, limiter)
|
||||
|
||||
testCtx := context.Background()
|
||||
rate := contracts.Rate{
|
||||
Limit: 3,
|
||||
Period: time.Minute,
|
||||
}
|
||||
|
||||
// 1st request
|
||||
res, err := limiter.Allow(testCtx, "user:123", rate)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allowed)
|
||||
assert.Equal(t, 2, res.Remaining)
|
||||
|
||||
// 2nd request
|
||||
res, err = limiter.Allow(testCtx, "user:123", rate)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allowed)
|
||||
assert.Equal(t, 1, res.Remaining)
|
||||
|
||||
// 3rd request
|
||||
res, err = limiter.Allow(testCtx, "user:123", rate)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allowed)
|
||||
assert.Equal(t, 0, res.Remaining)
|
||||
|
||||
// 4th request - blocked
|
||||
res, err = limiter.Allow(testCtx, "user:123", 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 = limiter.Reset(testCtx, "user:123")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Allowed again after reset
|
||||
res, err = limiter.Allow(testCtx, "user:123", rate)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allowed)
|
||||
}
|
||||
+10
@@ -8,6 +8,7 @@ import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/limiter"
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"encoding/json"
|
||||
@@ -139,6 +140,15 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
}
|
||||
|
||||
core.Provide[contracts.CacheService](ctx, svc)
|
||||
|
||||
var limiterSvc contracts.LimiterService
|
||||
if redisClient != nil {
|
||||
limiterSvc = newRedisLimiter(redisClient, p.keyPrefix)
|
||||
} else {
|
||||
limiterSvc = limiter.NewMemoryLimiter()
|
||||
}
|
||||
core.Provide[contracts.LimiterService](ctx, limiterSvc)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ package cache_memory
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/limiter"
|
||||
)
|
||||
|
||||
const defaultRAMCapacity = 10000
|
||||
@@ -79,5 +80,6 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
}
|
||||
|
||||
core.Provide[contracts.CacheService](ctx, svc)
|
||||
core.Provide[contracts.LimiterService](ctx, limiter.NewMemoryLimiter())
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -78,4 +78,14 @@ func TestCacheMemoryPlugin(t *testing.T) {
|
||||
var tempVal string
|
||||
err = cacheSvc.Get(reqCtx, "temp_key", &tempVal)
|
||||
assert.ErrorIs(t, err, contracts.ErrCacheMiss)
|
||||
|
||||
// 6. LimiterService
|
||||
limiterSvc, err := core.Inject[contracts.LimiterService](ctx)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, limiterSvc)
|
||||
|
||||
rateRes, err := limiterSvc.Allow(reqCtx, "key_a", contracts.Rate{Limit: 2, Period: time.Minute})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, rateRes.Allowed)
|
||||
assert.Equal(t, 1, rateRes.Remaining)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user