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
+3 -1
View File
@@ -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
+47 -6
View File
@@ -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")
}
}
+2
View File
@@ -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
})
+22 -2
View File
@@ -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