mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
feat(core): sync framework security hardening and accessibility improvements
- add util.Go with panic recovery for background goroutines - add util.EscapeLike and explicit ESCAPE clause for SQL LIKE queries - add DummyCheckPassword and subtle.ConstantTimeCompare against timing attacks - enforce session ID rotation upon login/oauth callback to prevent session fixation - add sliding window login failure rate limiting and oauth state rate limiting - fix redis client capture race in pubsub listeners and wait on stop channel - adjust global --primary to oklch(51.1% 0.262 276.966) for WCAG AA contrast - fix semantic heading levels and missing aria-labels across UI components - document security, concurrency, and a11y standards in AGENTS.md
This commit is contained in:
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -31,10 +32,12 @@ var (
|
||||
tokenListenerOnce sync.Once
|
||||
tokenListenerCtx context.Context
|
||||
tokenListenerCancel context.CancelFunc
|
||||
tokenListenerDone chan struct{}
|
||||
|
||||
userListenerOnce sync.Once
|
||||
userListenerCtx context.Context
|
||||
userListenerCancel context.CancelFunc
|
||||
userListenerDone chan struct{}
|
||||
)
|
||||
|
||||
func tokenCacheKey(tokenHash string) string {
|
||||
@@ -54,17 +57,22 @@ func ensureTokenCacheListener() {
|
||||
|
||||
func startTokenCacheInvalidationListener() {
|
||||
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
|
||||
tokenListenerDone = make(chan struct{})
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(tokenListenerCtx, oauthTokenInvalidationChannel)
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := tokenListenerCtx
|
||||
defer close(tokenListenerDone)
|
||||
|
||||
pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
<-tokenListenerCtx.Done()
|
||||
util.Go(func() {
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
tokenHash := msg.Payload
|
||||
@@ -74,7 +82,7 @@ func startTokenCacheInvalidationListener() {
|
||||
tokenRAM.Invalidate(tokenHash)
|
||||
}
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
|
||||
@@ -93,17 +101,22 @@ func ensureUserCacheListener() {
|
||||
|
||||
func startUserCacheInvalidationListener() {
|
||||
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
|
||||
userListenerDone = make(chan struct{})
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(userListenerCtx, oauthUserInvalidationChannel)
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := userListenerCtx
|
||||
defer close(userListenerDone)
|
||||
|
||||
pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
<-userListenerCtx.Done()
|
||||
util.Go(func() {
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
userIDStr := msg.Payload
|
||||
@@ -113,7 +126,7 @@ func startUserCacheInvalidationListener() {
|
||||
userRAM.Invalidate(userID)
|
||||
}
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
|
||||
@@ -213,13 +226,21 @@ func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
||||
func StopOauthCacheListener() {
|
||||
if tokenListenerCancel != nil {
|
||||
tokenListenerCancel()
|
||||
if tokenListenerDone != nil {
|
||||
<-tokenListenerDone
|
||||
}
|
||||
tokenListenerCancel = nil
|
||||
tokenListenerDone = nil
|
||||
}
|
||||
tokenListenerOnce = sync.Once{}
|
||||
|
||||
if userListenerCancel != nil {
|
||||
userListenerCancel()
|
||||
if userListenerDone != nil {
|
||||
<-userListenerDone
|
||||
}
|
||||
userListenerCancel = nil
|
||||
userListenerDone = nil
|
||||
}
|
||||
userListenerOnce = sync.Once{}
|
||||
}
|
||||
|
||||
@@ -24,6 +24,8 @@ const (
|
||||
const (
|
||||
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
||||
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
||||
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
|
||||
oauthStateLimitMax = 10
|
||||
)
|
||||
|
||||
// OAuth 授权用途常量
|
||||
|
||||
@@ -19,4 +19,5 @@ const (
|
||||
errAuthSourceDisabled = "认证源未启用"
|
||||
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
||||
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
|
||||
)
|
||||
|
||||
@@ -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,7 +75,7 @@ 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
|
||||
}
|
||||
@@ -98,6 +103,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 +170,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 +186,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
|
||||
}
|
||||
|
||||
@@ -176,7 +176,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
|
||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
||||
if err := SetLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -195,7 +195,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
||||
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||
if regErr != nil {
|
||||
registrationEnabled = true
|
||||
registrationEnabled = false
|
||||
}
|
||||
|
||||
if !registrationEnabled {
|
||||
|
||||
@@ -95,6 +95,25 @@ func (m *mockRedisClient) Del(ctx context.Context, keys ...string) *redis.IntCmd
|
||||
return cmd
|
||||
}
|
||||
|
||||
func (m *mockRedisClient) Incr(ctx context.Context, key string) *redis.IntCmd {
|
||||
cmd := redis.NewIntCmd(ctx)
|
||||
n := int64(1)
|
||||
if raw, ok := m.store[key]; ok {
|
||||
fmt.Sscan(raw, &n)
|
||||
n++
|
||||
}
|
||||
m.store[key] = fmt.Sprintf("%d", n)
|
||||
cmd.SetVal(n)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func (m *mockRedisClient) Expire(ctx context.Context, key string, expiration time.Duration) *redis.BoolCmd {
|
||||
cmd := redis.NewBoolCmd(ctx)
|
||||
_, ok := m.store[key]
|
||||
cmd.SetVal(ok)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func (m *mockRedisClient) Scan(ctx context.Context, cursor uint64, match string, count int64) *redis.ScanCmd {
|
||||
cmd := redis.NewScanCmd(ctx, nil, cursor, match, count)
|
||||
var keys []string
|
||||
@@ -689,6 +708,11 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
dbConn := setupTestDB(t)
|
||||
mockRedis := newMockRedisClient()
|
||||
seedTestAuthSource(t, dbConn)
|
||||
dbConn.Create(&model.SystemConfig{
|
||||
Key: model.ConfigKeyRegistrationEnabled,
|
||||
Value: "true",
|
||||
Type: "system",
|
||||
})
|
||||
|
||||
var state string
|
||||
|
||||
@@ -815,14 +839,14 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
}
|
||||
|
||||
t.Run("OIDC login when registration disabled - need bind", func(t *testing.T) {
|
||||
// Disable registration in database
|
||||
dbConn.Create(&model.SystemConfig{
|
||||
Key: model.ConfigKeyRegistrationEnabled,
|
||||
Value: "false",
|
||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "false")
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled)
|
||||
t.Cleanup(func() {
|
||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "true")
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled)
|
||||
})
|
||||
defer func() {
|
||||
dbConn.Where("key = ?", model.ConfigKeyRegistrationEnabled).Delete(&model.SystemConfig{})
|
||||
}()
|
||||
|
||||
var state4 string
|
||||
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, &state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
gsessions "github.com/gorilla/sessions"
|
||||
)
|
||||
|
||||
// GetUserIDFromSession 从 Session 中提取用户 ID
|
||||
@@ -47,11 +48,28 @@ func hashSessionToken(token string) string {
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
||||
func rotateSessionID(s sessions.Session) {
|
||||
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
|
||||
if sess := inner.Session(); sess != nil {
|
||||
sess.ID = ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SetLoginSession writes the authenticated user into a freshly rotated session.
|
||||
func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extras ...map[string]any) error {
|
||||
session := sessions.Default(c)
|
||||
session.Clear()
|
||||
rotateSessionID(session)
|
||||
|
||||
session.Set(UserIDKey, user.ID)
|
||||
session.Set(UserNameKey, user.Username)
|
||||
session.Set(PasswordHashKey, user.Password)
|
||||
if len(extras) > 0 {
|
||||
for key, value := range extras[0] {
|
||||
session.Set(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := config.Config.App.SessionAge
|
||||
|
||||
Reference in New Issue
Block a user