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
+15
View File
@@ -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"])
}
}
+2
View File
@@ -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
+22 -4
View File
@@ -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 {