mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 23:26:38 +08:00
公开 OAuth 登录/授权入口按会话限制 10 分钟内最多 20 个 state,堵住未授权 Redis 洪水。metric 持平 8。
Result: {"status":"keep","total_issues":8,"golint_canonicalheader":0,"golint_errname":0,"golint_errorlint":1,"golint_forcetypeassert":0,"golint_gosec":0,"golint_intrange":0,"golint_modernize":3,"golint_nilnil":3,"golint_perfsprint":0,"golint_prealloc":0,"golint_recvcheck":1,"golint_usestdlibvars":0,"golint_wastedassign":0,"golint_total":8,"eslint_problems":0,"eslint_errors":0,"eslint_warnings":0,"tsc_errors":0,"measure_s":95,"golint_test_testifylint":0,"golint_test_thelper":0,"golint_test_usetesting":0,"golint_test_total":0,"golint_exhaustive":0,"golint_vetx_total":0,"vitest_failed":0,"vitest_total":126}
This commit is contained in:
@@ -5,6 +5,7 @@ package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -58,6 +59,10 @@ func GetLoginURL(c *gin.Context) {
|
||||
|
||||
userID := GetUserIDFromSession(session)
|
||||
sessionHash := hashSessionToken(token)
|
||||
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
@@ -70,11 +75,10 @@ func GetLoginURL(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
@@ -98,6 +102,24 @@ func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state stri
|
||||
return authConfig.AuthCodeURL(state), nil
|
||||
}
|
||||
|
||||
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||
if db.Redis == nil || sessionHash == "" {
|
||||
return nil
|
||||
}
|
||||
key := db.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
||||
n, err := db.Redis.Incr(ctx, key).Result()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 1 {
|
||||
_ = db.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
||||
}
|
||||
if n > oauthStateLimitMax {
|
||||
return errors.New(errOAuthStateRateLimited)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Authorize 发起指定认证源授权
|
||||
// @Summary 发起指定认证源授权
|
||||
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
|
||||
@@ -147,6 +169,10 @@ func Authorize(c *gin.Context) {
|
||||
}
|
||||
|
||||
sessionHash := hashSessionToken(token)
|
||||
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
@@ -159,7 +185,7 @@ func Authorize(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user