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:
ryan
2026-08-27 23:01:28 +08:00
parent b66cf3ae9c
commit ae3b792e16
63 changed files with 570 additions and 176 deletions
+33 -12
View File
@@ -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{}
}
+2
View File
@@ -24,6 +24,8 @@ const (
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
oauthStateLimitMax = 10
)
// OAuth 授权用途常量
+1
View File
@@ -19,4 +19,5 @@ const (
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
)
+29 -2
View File
@@ -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
}
+2 -2
View File
@@ -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 {
+31 -7
View File
@@ -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")
+19 -1
View File
@@ -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