refactor(core): decouple gin from pkg/util and reduce code duplication

This commit is contained in:
ryan
2026-08-28 20:15:47 +08:00
parent c52b7c4abf
commit df351cbd33
32 changed files with 394 additions and 432 deletions
+28 -21
View File
@@ -438,31 +438,38 @@ func maskSensitiveConfig(key, value string) string {
case ConfigKeySMTPPassword:
return maskedConfigValue
case ConfigKeyStorageConfig:
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
if cfg.S3.SecretAccessKey != "" {
cfg.S3.SecretAccessKey = maskedConfigValue
}
if cfg.R2.SecretAccessKey != "" {
cfg.R2.SecretAccessKey = maskedConfigValue
}
if cfg.MinIO.SecretAccessKey != "" {
cfg.MinIO.SecretAccessKey = maskedConfigValue
}
if cfg.OSS.SecretAccessKey != "" {
cfg.OSS.SecretAccessKey = maskedConfigValue
}
if cfg.WebDAV.Password != "" {
cfg.WebDAV.Password = maskedConfigValue
}
if val, err := json.Marshal(cfg); err == nil {
return string(val)
}
}
return maskStorageConfig(value)
}
return value
}
func maskStorageConfig(value string) string {
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err != nil {
return value
}
if cfg.S3.SecretAccessKey != "" {
cfg.S3.SecretAccessKey = maskedConfigValue
}
if cfg.R2.SecretAccessKey != "" {
cfg.R2.SecretAccessKey = maskedConfigValue
}
if cfg.MinIO.SecretAccessKey != "" {
cfg.MinIO.SecretAccessKey = maskedConfigValue
}
if cfg.OSS.SecretAccessKey != "" {
cfg.OSS.SecretAccessKey = maskedConfigValue
}
if cfg.WebDAV.Password != "" {
cfg.WebDAV.Password = maskedConfigValue
}
val, err := json.Marshal(cfg)
if err != nil {
return value
}
return string(val)
}
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value, currentConfig string) (string, error) {
@@ -5,9 +5,9 @@ package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"errors"
"net/http"
"strconv"
@@ -262,7 +262,7 @@ func DeleteUser(c *gin.Context) {
return
}
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
@@ -373,7 +373,7 @@ func UpdateUser(c *gin.Context) {
return
}
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
+4 -4
View File
@@ -5,10 +5,10 @@ package admin
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/pkg/util"
"github.com/gin-gonic/gin"
)
@@ -19,15 +19,15 @@ func LoginAdminRequired() gin.HandlerFunc {
ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
return
}
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired)
return
+2 -1
View File
@@ -5,6 +5,7 @@ package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
@@ -438,7 +439,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
// UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) {
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true
+39 -37
View File
@@ -5,9 +5,9 @@ package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/pkg/util"
"context"
"crypto/sha256"
"encoding/hex"
@@ -25,43 +25,45 @@ func hashToken(token string) string {
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err == nil {
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err == nil && user != nil && user.IsActive {
return user, tokenRecord, nil
if err != nil || tokenRecord == nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
var userRow contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&userRow).Error; err != nil {
return nil, nil, err
}
user = &userRow
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
var userRow contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
return nil, nil, err
}
SetCachedUser(ctx, userRow.ID, &userRow)
return &userRow, tokenRecord, nil
return user, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
var tokenStr string
// Check token in headers
tokenStr := c.GetHeader("X-Access-Token")
if tokenStr == "" {
tokenFromQuery := c.Query("token")
if tokenFromQuery != "" {
tokenStr = tokenFromQuery
} else {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
@@ -74,8 +76,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
if user.Username == SystemUsername {
return nil, errors.New("system user is not allowed to login")
}
util.SetToContext(c, contracts.AuthTokenAuthKey, true)
util.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
@@ -96,8 +98,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
SetCachedUser(ctx, userID, user)
}
util.SetToContext(c, contracts.AuthTokenAuthKey, false)
util.SetToContext(c, contracts.AuthTokenAdminKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New("system user is not allowed to login")
@@ -119,7 +121,7 @@ func LoginRequired() gin.HandlerFunc {
}
LogForAudit(c.Request.Context(), user, c)
util.SetToContext(c, contracts.AuthUserObjKey, user)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
@@ -136,8 +138,8 @@ func AdminRequired() gin.HandlerFunc {
return
}
isTokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
@@ -152,7 +154,7 @@ func AdminRequired() gin.HandlerFunc {
}
LogForAudit(c.Request.Context(), user, c)
util.SetToContext(c, contracts.AuthUserObjKey, user)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
@@ -165,7 +167,7 @@ func LoginAdminRequired() gin.HandlerFunc {
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
+2 -2
View File
@@ -5,7 +5,7 @@ package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"Wavelet/pkg/ginutil"
"context"
"errors"
"sync"
@@ -29,7 +29,7 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := util.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
return u, nil
}
}
@@ -40,6 +40,23 @@ func ListAdminChannels(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(rows))
}
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return 0, false
}
return id, true
}
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
fallback(c, err.Error())
}
// CreateAdminChannel creates a messaging channel.
// @Summary Create message gateway channel
// @Description Creates a Telegram or QQ channel with encrypted credentials
@@ -52,17 +69,7 @@ func ListAdminChannels(c *gin.Context) {
// @Failure 400 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels [post]
func CreateAdminChannel(c *gin.Context) {
var req CreateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := createChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
handleJSONRequest(c, createChannel)
}
// UpdateAdminChannel patches a messaging channel. Empty secrets keep the previous values.
@@ -79,26 +86,9 @@ func CreateAdminChannel(c *gin.Context) {
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id} [patch]
func UpdateAdminChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updateChannel(c.Request.Context(), id, req)
if err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
handleEntityUpdate(c, parseAdminChannelID, updateChannel, func(c *gin.Context, err error) {
handleAdminChannelError(c, err, response.AbortBadRequest)
})
}
// DeleteAdminChannel removes a channel and its bindings/pairing codes.
@@ -112,17 +102,12 @@ func UpdateAdminChannel(c *gin.Context) {
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id} [delete]
func DeleteAdminChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := deleteChannel(c.Request.Context(), id); err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
handleAdminChannelError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -140,17 +125,12 @@ func DeleteAdminChannel(c *gin.Context) {
// @Failure 404 {object} response.Any
// @Router /api/v1/admin/message-gateway/channels/{id}/test [post]
func TestAdminChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
id, ok := parseAdminChannelID(c)
if !ok {
return
}
if err := probeChannel(c.Request.Context(), id); err != nil {
if err.Error() == errChannelNotFound {
response.AbortNotFound(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
handleAdminChannelError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -0,0 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"Wavelet/pkg/response"
"context"
"net/http"
"github.com/gin-gonic/gin"
)
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
res, err := handler(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(res))
}
func handleEntityUpdate[Req any, Res any](
c *gin.Context,
parseID func(*gin.Context) (uint64, bool),
updater func(ctx context.Context, id uint64, req Req) (Res, error),
onErr func(*gin.Context, error),
) {
id, ok := parseID(c)
if !ok {
return
}
var req Req
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := updater(c.Request.Context(), id, req)
if err != nil {
onErr(c, err)
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
@@ -5,8 +5,8 @@ package message_gateway
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"errors"
"net/http"
"strconv"
@@ -15,7 +15,7 @@ import (
)
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
return util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
}
// ListChannels lists enabled channels a user can bind.
@@ -214,20 +214,26 @@ type CreatePushChannelRequest struct {
Enabled bool `json:"enabled"`
}
func parsePushChannelID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return 0, false
}
return id, true
}
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
fallback(c, err.Error())
}
// CreatePushChannel creates a push channel.
func CreatePushChannel(c *gin.Context) {
var req CreatePushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := createPushChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
handleJSONRequest(c, createPushChannel)
}
// UpdatePushChannelRequest is the update channel request payload.
@@ -242,44 +248,20 @@ type UpdatePushChannelRequest struct {
// UpdatePushChannel updates a push channel.
func UpdatePushChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdatePushChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := updatePushChannel(c.Request.Context(), id, req)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
handleEntityUpdate(c, parsePushChannelID, updatePushChannel, func(c *gin.Context, err error) {
handlePushChannelNotFoundError(c, err, response.AbortInternal)
})
}
// DeletePushChannel deletes a push channel.
func DeletePushChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
id, ok := parsePushChannelID(c)
if !ok {
return
}
if err := deletePushChannel(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
handlePushChannelNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -56,36 +56,37 @@ func ListBuiltInPushEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(GetBuiltInEvents()))
}
func parsePushEventID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return 0, false
}
return id, true
}
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
fallback(c, err.Error())
}
// CreatePushEvent creates a new push event configuration.
func CreatePushEvent(c *gin.Context) {
var req CreatePushEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
event, err := createPushEvent(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(event))
handleJSONRequest(c, createPushEvent)
}
// DeletePushEvent deletes a push event configuration by ID.
func DeletePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
id, ok := parsePushEventID(c)
if !ok {
return
}
if err := deletePushEvent(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortInternal(c, err.Error())
handlePushEventNotFoundError(c, err, response.AbortInternal)
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -93,9 +94,8 @@ func DeletePushEvent(c *gin.Context) {
// UpdatePushEvent updates an existing push event.
func UpdatePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
id, ok := parsePushEventID(c)
if !ok {
return
}
@@ -106,11 +106,7 @@ func UpdatePushEvent(c *gin.Context) {
}
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -118,19 +114,14 @@ func UpdatePushEvent(c *gin.Context) {
// TogglePushEvent toggles the enabled state of a push event.
func TogglePushEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
id, ok := parsePushEventID(c)
if !ok {
return
}
enabled, err := togglePushEvent(c.Request.Context(), id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
return
}
c.JSON(http.StatusOK, response.OK(enabled))
@@ -217,25 +217,31 @@ func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
return nil
}
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
cacheKey := "push:channel:active:" + name
var channel PushChannel
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
var val T
if cache := getCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &channel); err == nil {
return &channel, nil
if err := cache.Get(ctx, cacheKey, &val); err == nil {
return &val, nil
}
}
if err := getDB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
db := getDB(ctx)
if err := query(db, &val); err != nil {
return nil, err
}
if cache := getCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, channel, activePushChannelCacheTTL)
_ = cache.Set(ctx, cacheKey, val, ttl)
}
return &channel, nil
return &val, nil
}
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
return getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *PushChannel) error {
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
})
}
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
@@ -329,23 +335,9 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
cacheKey := "push:event:active:" + key
var event PushEvent
if cache := getCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &event); err == nil {
return &event, nil
}
}
if err := getDB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
return nil, err
}
if cache := getCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, event, activePushEventCacheTTL)
}
return &event, nil
return getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *PushEvent) error {
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
})
}
// DeleteActivePushEventCache 清理启用通知事件的缓存。
@@ -7,9 +7,9 @@ package risk_control
import (
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/idgen"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control/logstore"
"encoding/json"
"net/http"
@@ -42,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
c.Next()
// 3. 后置身份检查:仅记录通过认证的请求
userObj, exists := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
userObj, exists := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if !exists || userObj == nil {
return
}
@@ -7,8 +7,8 @@ import (
"Wavelet/core/contracts"
"Wavelet/pkg/batchwriter"
"Wavelet/pkg/config"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/testhelper"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control"
"Wavelet/plugins/domain/risk_control/logstore"
"context"
@@ -99,7 +99,7 @@ func TestRiskControlMiddleware(t *testing.T) {
r := gin.New()
r.Use(func(c *gin.Context) {
user := &contracts.UserDTO{ID: 12345}
util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user)
ginutil.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user)
c.Next()
})
r.Use(risk_control.RiskControlMiddleware())
@@ -23,8 +23,7 @@ import (
"sync"
pkgcache "Wavelet/pkg/cache/disk"
pkgutil "Wavelet/pkg/util"
"Wavelet/pkg/ginutil"
uploadstorage "Wavelet/plugins/domain/upload/storage"
@@ -301,7 +300,7 @@ func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, e
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUserID uint64
var isAdmin bool
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
currUserID = u.ID
isAdmin = u.IsAdmin
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
@@ -329,14 +328,19 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
return checkPrivateFileOwner(c, upload.UserID)
}
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
if authSvc := shared.GetAuthService(c); authSvc != nil {
if _, err := authSvc.GetCurrentUser(c); err != nil {
return err
}
}
}
if cache.IsFilePublic(c.Request.Context(), upload.Type) {
return nil
}
return nil
if _, ok := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok {
return nil
}
authSvc := shared.GetAuthService(c)
if authSvc == nil {
return nil
}
_, err := authSvc.GetCurrentUser(c)
return err
}
@@ -5,8 +5,8 @@ package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/ingest"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/repository"
@@ -168,7 +168,7 @@ type listMyFilesResponse struct {
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
@@ -215,7 +215,7 @@ func ListMyFiles(c *gin.Context) {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
@@ -262,7 +262,7 @@ type updateMyFileRequest struct {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
@@ -29,11 +29,11 @@ import (
"strconv"
"strings"
"Wavelet/pkg/ginutil"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
pkgutil "Wavelet/pkg/util"
uploadstorage "Wavelet/plugins/domain/upload/storage"
)
@@ -64,7 +64,7 @@ func UploadFile(c *gin.Context) {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
currUser, _ := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
@@ -5,8 +5,8 @@ package handler
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
"archive/zip"
@@ -42,7 +42,7 @@ func setupTestRouter(authUser *contracts.UserDTO) *gin.Engine {
authMiddleware := func(c *gin.Context) {
if authUser != nil {
util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser)
ginutil.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser)
}
c.Next()
}