mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 07:26:36 +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:
@@ -17,6 +17,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||
@@ -105,7 +106,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
||||
|
||||
// 在独立 goroutine 中读取客户端消息(保持连接活跃 + 检测断开)
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
defer close(done)
|
||||
for {
|
||||
_, _, err := conn.ReadMessage()
|
||||
@@ -113,7 +114,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
})
|
||||
|
||||
// 主循环:推送日志
|
||||
for {
|
||||
|
||||
@@ -6,7 +6,7 @@ package message_gateway
|
||||
const (
|
||||
errNameRequired = "name is required"
|
||||
errTypeInvalid = "type must be telegram or qq"
|
||||
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
||||
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
||||
errQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
||||
errChannelNotFound = "channel not found"
|
||||
errChannelProbeFailed = "channel probe failed"
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -76,7 +77,7 @@ var DefaultTrigger = &EventTrigger{}
|
||||
//nolint:contextcheck
|
||||
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
if body == nil {
|
||||
body = make(map[string]any)
|
||||
}
|
||||
@@ -100,7 +101,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
|
||||
flatBody := getFlatBody(body)
|
||||
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
||||
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||
@@ -58,11 +59,11 @@ func ApplyUpdate(c *gin.Context) {
|
||||
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
time.Sleep(time.Second)
|
||||
if err := replaceAndRestart(executable, stagedBinary); err != nil {
|
||||
defaultManager.finishUpgrade()
|
||||
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
@@ -178,10 +178,13 @@ var (
|
||||
// GetDefaultManager yields the global singleton CAPTCHA manager.
|
||||
func GetDefaultManager() *Manager {
|
||||
once.Do(func() {
|
||||
secret := []byte("default-captcha-secret-key-at-least-16-bytes")
|
||||
if config.Config != nil && config.Config.App.SessionSecret != "" {
|
||||
var secret []byte
|
||||
if config.Config != nil && strings.TrimSpace(config.Config.App.SessionSecret) != "" {
|
||||
secret = []byte(config.Config.App.SessionSecret)
|
||||
}
|
||||
if len(secret) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
var store pkgcap.Store
|
||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||
|
||||
@@ -16,6 +16,10 @@ func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
if mgr == nil {
|
||||
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
|
||||
return
|
||||
}
|
||||
|
||||
token := c.GetHeader("X-Cap-Token")
|
||||
if token == "" {
|
||||
|
||||
@@ -44,6 +44,10 @@ func Challenge(c *gin.Context) {
|
||||
}
|
||||
|
||||
mgr := GetDefaultManager()
|
||||
if mgr == nil {
|
||||
response.AbortInternal(c, "captcha is not configured")
|
||||
return
|
||||
}
|
||||
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
|
||||
@@ -77,6 +81,10 @@ func Redeem(c *gin.Context) {
|
||||
}
|
||||
|
||||
mgr := GetDefaultManager()
|
||||
if mgr == nil {
|
||||
response.AbortInternal(c, "captcha is not configured")
|
||||
return
|
||||
}
|
||||
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
|
||||
|
||||
@@ -9,10 +9,12 @@ import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||
@@ -39,6 +41,16 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
|
||||
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "test-captcha-session-secret"
|
||||
once = sync.Once{}
|
||||
defaultManager = nil
|
||||
t.Cleanup(func() {
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
once = sync.Once{}
|
||||
defaultManager = nil
|
||||
})
|
||||
|
||||
r := testhelper.NewTestGinEngine()
|
||||
|
||||
// Mount CAPTCHA API endpoints
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -190,7 +191,7 @@ func startRuntimeSettingsInvalidationListener() {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
@@ -208,5 +209,5 @@ func startRuntimeSettingsInvalidationListener() {
|
||||
InvalidateRuntimeSettings()
|
||||
}
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
+3
-2
@@ -17,6 +17,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
||||
@@ -56,7 +57,7 @@ func startAccessCacheInvalidationListener() {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
pubsub := db.Redis.Subscribe(
|
||||
context.Background(),
|
||||
objectstore.ConfigInvalidationChannel,
|
||||
@@ -69,7 +70,7 @@ func startAccessCacheInvalidationListener() {
|
||||
for range pubsub.Channel() {
|
||||
ResetAccessCaches()
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
||||
|
||||
+14
-5
@@ -12,6 +12,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 (
|
||||
@@ -29,6 +30,7 @@ var (
|
||||
uploadMetaListenerOnce sync.Once
|
||||
uploadMetaListenerCtx context.Context
|
||||
uploadMetaListenerCancel context.CancelFunc
|
||||
uploadMetaListenerDone chan struct{}
|
||||
)
|
||||
|
||||
func uploadMetaRedisKey(id uint64) string {
|
||||
@@ -48,17 +50,20 @@ func ensureUploadMetaCacheListener() {
|
||||
|
||||
func startUploadMetaCacheInvalidationListener() {
|
||||
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
||||
uploadMetaListenerDone = make(chan struct{})
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
util.Go(func() {
|
||||
defer close(uploadMetaListenerDone)
|
||||
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
<-uploadMetaListenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload uploadMetaInvalidationMessage
|
||||
@@ -68,7 +73,7 @@ func startUploadMetaCacheInvalidationListener() {
|
||||
}
|
||||
uploadMetaRAM.Invalidate(payload.ID)
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
||||
@@ -145,7 +150,11 @@ func ResetUploadMetaCacheForTest() {
|
||||
func StopUploadMetaCacheListener() {
|
||||
if uploadMetaListenerCancel != nil {
|
||||
uploadMetaListenerCancel()
|
||||
if uploadMetaListenerDone != nil {
|
||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 db.Redis 不再竞争
|
||||
}
|
||||
uploadMetaListenerCancel = nil
|
||||
uploadMetaListenerDone = nil
|
||||
}
|
||||
uploadMetaListenerOnce = sync.Once{}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
@@ -99,7 +100,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
}()
|
||||
|
||||
//nolint:contextcheck,gosec
|
||||
go func() {
|
||||
util.Go(func() {
|
||||
ticker := time.NewTicker(renewalInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
@@ -114,7 +115,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
active, err := objectstore.LoadConfig(ctx)
|
||||
|
||||
@@ -44,4 +44,5 @@ const (
|
||||
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
||||
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
||||
errSendMailFailed = "发送邮件失败: %w"
|
||||
errLoginRateLimited = "请求过于频繁,请稍后再试"
|
||||
)
|
||||
|
||||
@@ -6,14 +6,17 @@ package user
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
@@ -49,6 +52,48 @@ type updateProfileInput struct {
|
||||
Location string
|
||||
}
|
||||
|
||||
const (
|
||||
loginFailLimitKeyFormat = "login:fail:%s"
|
||||
loginFailLimitMax = 20
|
||||
loginFailLimitWindow = 10 * time.Minute
|
||||
)
|
||||
|
||||
func loginFailLimitKey(ip string) string {
|
||||
return fmt.Sprintf(loginFailLimitKeyFormat, strings.TrimSpace(ip))
|
||||
}
|
||||
|
||||
func loginAttemptsBlocked(ctx context.Context, ip string) bool {
|
||||
if db.Redis == nil {
|
||||
return false
|
||||
}
|
||||
n, err := db.Redis.Get(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Int()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return n >= loginFailLimitMax
|
||||
}
|
||||
|
||||
func recordFailedLogin(ctx context.Context, ip string) {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
key := db.PrefixedKey(loginFailLimitKey(ip))
|
||||
n, err := db.Redis.Incr(ctx, key).Result()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if n == 1 {
|
||||
_ = db.Redis.Expire(ctx, key, loginFailLimitWindow).Err()
|
||||
}
|
||||
}
|
||||
|
||||
func clearFailedLogins(ctx context.Context, ip string) {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Err()
|
||||
}
|
||||
|
||||
func isPasswordLoginEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled)
|
||||
if err != nil {
|
||||
@@ -60,7 +105,7 @@ func isPasswordLoginEnabled(ctx context.Context) bool {
|
||||
func isPasswordRegisterEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
@@ -68,7 +113,7 @@ func isPasswordRegisterEnabled(ctx context.Context) bool {
|
||||
func isRegistrationEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
@@ -161,7 +206,9 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
|
||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||
return false
|
||||
}
|
||||
if storedCode != code {
|
||||
sumGot := sha256.Sum256([]byte(strings.TrimSpace(code)))
|
||||
sumWant := sha256.Sum256([]byte(strings.TrimSpace(storedCode)))
|
||||
if subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) != 1 {
|
||||
return false
|
||||
}
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
|
||||
|
||||
@@ -4,20 +4,17 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/shared"
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -53,41 +50,6 @@ type updateProfileRequest struct {
|
||||
Location string `json:"location"`
|
||||
}
|
||||
|
||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, user.ID)
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
session.Set(oauth.PasswordHashKey, user.Password)
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||
if err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
session.Options(oauth.GetSessionOptions(maxAge))
|
||||
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if isSessionCookie {
|
||||
oauth.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Login 用户密码登录
|
||||
// @Summary 用户密码登录
|
||||
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
||||
@@ -96,7 +58,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
||||
// @Produce json
|
||||
// @Param request body user.loginRequest true "登录请求参数"
|
||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "登录成功,返回用户信息"
|
||||
// @Failure 400 {object} response.Any "用户名或密码错误、帐号已禁用等"
|
||||
// @Failure 400 {object} response.Any "用户名或密码错误"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/user/login [post]
|
||||
func Login(c *gin.Context) {
|
||||
@@ -105,6 +67,10 @@ func Login(c *gin.Context) {
|
||||
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
||||
return
|
||||
}
|
||||
if loginAttemptsBlocked(ctx, c.ClientIP()) {
|
||||
response.AbortBadRequest(c, errLoginRateLimited)
|
||||
return
|
||||
}
|
||||
var req loginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
@@ -118,13 +84,16 @@ func Login(c *gin.Context) {
|
||||
|
||||
user, err := getUserByUsernameOrEmail(ctx, req.Username)
|
||||
if err != nil {
|
||||
pkgu.DummyCheckPassword(req.Password)
|
||||
recordFailedLogin(ctx, c.ClientIP())
|
||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||
return
|
||||
}
|
||||
if !user.IsActive {
|
||||
recordFailedLogin(ctx, c.ClientIP())
|
||||
logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||
response.AbortBadRequest(c, shared.BannedAccount)
|
||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -132,6 +101,7 @@ func Login(c *gin.Context) {
|
||||
isPlaintext := !user.IsPasswordEncrypted()
|
||||
|
||||
if !user.CheckPassword(req.Password) {
|
||||
recordFailedLogin(ctx, c.ClientIP())
|
||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||
return
|
||||
@@ -149,21 +119,19 @@ func Login(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
needChangePassword := isPlaintext
|
||||
|
||||
if isPlaintext {
|
||||
session.Set("need_change_password", true)
|
||||
} else {
|
||||
session.Delete("need_change_password")
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
if err := updateLastLogin(ctx, user); err != nil {
|
||||
response.AbortBadRequest(c, "更新登录时间失败,请稍后再试")
|
||||
return
|
||||
}
|
||||
if err := setLoginSession(ctx, c, user); err != nil {
|
||||
extras := map[string]any{}
|
||||
if isPlaintext {
|
||||
extras["need_change_password"] = true
|
||||
}
|
||||
clearFailedLogins(ctx, c.ClientIP())
|
||||
if err := oauth.SetLoginSession(ctx, c, user, extras); err != nil {
|
||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||
return
|
||||
}
|
||||
@@ -253,7 +221,7 @@ func Register(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
||||
if err := oauth.SetLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||
return
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user