From df351cbd33babfe7374a592951d7aba5dce01218 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 28 Aug 2026 20:15:47 +0800 Subject: [PATCH] refactor(core): decouple gin from pkg/util and reduce code duplication --- backend/core/extpoints/migration.go | 27 ++----- backend/core/extpoints/registry.go | 22 ++++++ backend/core/extpoints/schedule.go | 18 +---- backend/core/extpoints/setting.go | 18 +---- backend/core/extpoints/task.go | 18 +---- backend/pkg/cache/disk/cache.go | 5 +- backend/pkg/cache/disk/cache_test.go | 35 +++------ .../gin_context.go => ginutil/context.go} | 3 +- .../plugins/domain/admin/handlers_config.go | 49 +++++++----- backend/plugins/domain/admin/handlers_user.go | 6 +- backend/plugins/domain/admin/middlewares.go | 8 +- backend/plugins/domain/auth/handlers.go | 3 +- backend/plugins/domain/auth/middleware.go | 76 ++++++++++--------- backend/plugins/domain/auth/service.go | 4 +- .../domain/message_gateway/admin_handlers.go | 74 +++++++----------- .../domain/message_gateway/handler_helpers.go | 49 ++++++++++++ .../domain/message_gateway/handlers.go | 4 +- .../domain/message_gateway/push_channels.go | 66 ++++++---------- .../domain/message_gateway/push_handlers.go | 63 +++++++-------- .../domain/message_gateway/repository.go | 44 +++++------ .../plugins/domain/risk_control/middleware.go | 4 +- .../domain/risk_control/middleware_test.go | 4 +- .../domain/upload/filesrv/file_server.go | 28 ++++--- .../domain/upload/handler/file_management.go | 8 +- .../plugins/domain/upload/handler/routers.go | 6 +- .../domain/upload/handler/routers_test.go | 4 +- .../drivers/driver_asynq_cron/plugin.go | 63 ++++++++------- .../drivers/driver_asynq_cron/scheduler.go | 37 +++++---- .../drivers/driver_asynq_worker/plugin.go | 32 +++----- .../drivers/driver_inproc_cron/scheduler.go | 22 +----- .../drivers/driver_inproc_worker/executor.go | 20 +++-- scripts/check_cordis_architecture.sh | 6 +- 32 files changed, 394 insertions(+), 432 deletions(-) create mode 100644 backend/core/extpoints/registry.go rename backend/pkg/{util/gin_context.go => ginutil/context.go} (84%) create mode 100644 backend/plugins/domain/message_gateway/handler_helpers.go diff --git a/backend/core/extpoints/migration.go b/backend/core/extpoints/migration.go index 13d79049..36d9de80 100644 --- a/backend/core/extpoints/migration.go +++ b/backend/core/extpoints/migration.go @@ -19,27 +19,26 @@ type MigrationEntry struct { // MigrationExtension defines the interface for registering and querying plugin migrations. type MigrationExtension interface { Register(pluginID string, fsys fs.FS, dir ...string) + Unregister(pluginID string) bool Entries() []MigrationEntry Get(pluginID string) (MigrationEntry, bool) - Unregister(pluginID string) bool } -// MigrationRegistry collects and stores migration entries from plugins. +// MigrationRegistry implements MigrationExtension. type MigrationRegistry struct { mu sync.RWMutex entries []MigrationEntry lookup map[string]MigrationEntry } -// NewMigrationRegistry creates a new migration registry. +// NewMigrationRegistry creates a new MigrationRegistry. func NewMigrationRegistry() *MigrationRegistry { return &MigrationRegistry{ lookup: make(map[string]MigrationEntry), } } -// Register adds a migration entry for a plugin. -// If dir is not specified, it defaults to "migrations". +// Register registers an embedded migration filesystem for a plugin. func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string) { m.mu.Lock() defer m.mu.Unlock() @@ -72,21 +71,9 @@ func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string) // Unregister removes a registered migration entry by plugin ID. func (m *MigrationRegistry) Unregister(pluginID string) bool { - m.mu.Lock() - defer m.mu.Unlock() - - if _, exists := m.lookup[pluginID]; !exists { - return false - } - - delete(m.lookup, pluginID) - for i, e := range m.entries { - if e.PluginID == pluginID { - m.entries = append(m.entries[:i], m.entries[i+1:]...) - break - } - } - return true + return unregisterEntry(&m.mu, m.lookup, &m.entries, pluginID, func(e MigrationEntry) bool { + return e.PluginID == pluginID + }) } // Entries returns a copy of all registered migration entries in registration order. diff --git a/backend/core/extpoints/registry.go b/backend/core/extpoints/registry.go new file mode 100644 index 00000000..9ecf03b9 --- /dev/null +++ b/backend/core/extpoints/registry.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package extpoints + +import ( + "slices" + "sync" +) + +func unregisterEntry[T any](mu *sync.RWMutex, lookup map[string]T, list *[]T, key string, matches func(T) bool) bool { + mu.Lock() + defer mu.Unlock() + + if _, exists := lookup[key]; !exists { + return false + } + + delete(lookup, key) + *list = slices.DeleteFunc(*list, matches) + return true +} diff --git a/backend/core/extpoints/schedule.go b/backend/core/extpoints/schedule.go index 0c819041..663f69a6 100644 --- a/backend/core/extpoints/schedule.go +++ b/backend/core/extpoints/schedule.go @@ -88,21 +88,9 @@ func (s *ScheduleRegistry) RegisterCron(spec, taskType string, payload any, opts // Unregister removes a registered schedule definition by its task type. func (s *ScheduleRegistry) Unregister(taskType string) bool { - s.mu.Lock() - defer s.mu.Unlock() - - if _, exists := s.lookup[taskType]; !exists { - return false - } - - delete(s.lookup, taskType) - for i, item := range s.schedules { - if item.TaskType == taskType { - s.schedules = append(s.schedules[:i], s.schedules[i+1:]...) - break - } - } - return true + return unregisterEntry(&s.mu, s.lookup, &s.schedules, taskType, func(item ScheduleDefinition) bool { + return item.TaskType == taskType + }) } // Schedules returns a copy of all registered ScheduleDefinitions. diff --git a/backend/core/extpoints/setting.go b/backend/core/extpoints/setting.go index 713caf0e..2232ebc2 100644 --- a/backend/core/extpoints/setting.go +++ b/backend/core/extpoints/setting.go @@ -65,21 +65,9 @@ func (s *SettingRegistry) Register(schema SettingSchema) { // Unregister removes a registered SettingSchema by its key. func (s *SettingRegistry) Unregister(key string) bool { - s.mu.Lock() - defer s.mu.Unlock() - - if _, exists := s.lookup[key]; !exists { - return false - } - - delete(s.lookup, key) - for i, item := range s.schemas { - if item.Key == key { - s.schemas = append(s.schemas[:i], s.schemas[i+1:]...) - break - } - } - return true + return unregisterEntry(&s.mu, s.lookup, &s.schemas, key, func(item SettingSchema) bool { + return item.Key == key + }) } // Schemas returns a copy of all registered SettingSchemas. diff --git a/backend/core/extpoints/task.go b/backend/core/extpoints/task.go index cc8650f1..7dafe3c9 100644 --- a/backend/core/extpoints/task.go +++ b/backend/core/extpoints/task.go @@ -107,21 +107,9 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) // Unregister removes a registered task definition by its pattern. func (t *TaskRegistry) Unregister(pattern string) bool { - t.mu.Lock() - defer t.mu.Unlock() - - if _, exists := t.lookup[pattern]; !exists { - return false - } - - delete(t.lookup, pattern) - for i, item := range t.tasks { - if item.Pattern == pattern { - t.tasks = append(t.tasks[:i], t.tasks[i+1:]...) - break - } - } - return true + return unregisterEntry(&t.mu, t.lookup, &t.tasks, pattern, func(item TaskDefinition) bool { + return item.Pattern == pattern + }) } // Tasks returns a copy of all registered TaskDefinitions. diff --git a/backend/pkg/cache/disk/cache.go b/backend/pkg/cache/disk/cache.go index a90edb52..a4514957 100644 --- a/backend/pkg/cache/disk/cache.go +++ b/backend/pkg/cache/disk/cache.go @@ -5,6 +5,7 @@ package disk import ( + "Wavelet/pkg/util" "container/list" "encoding/binary" "errors" @@ -100,7 +101,9 @@ var ( func Default() *Cache { defaultCacheOnce.Do(func() { defaultCache = New("uploads/diskcache") - go defaultCache.StartCleanupWorker(defaultCleanupInterval) + util.Go(func() { + defaultCache.StartCleanupWorker(defaultCleanupInterval) + }) }) return defaultCache } diff --git a/backend/pkg/cache/disk/cache_test.go b/backend/pkg/cache/disk/cache_test.go index 107d3c44..03b363d2 100644 --- a/backend/pkg/cache/disk/cache_test.go +++ b/backend/pkg/cache/disk/cache_test.go @@ -5,15 +5,12 @@ package disk import ( "bytes" - "os" "testing" "time" ) func TestDiskCacheBasic(t *testing.T) { - testDir := "uploads/test_diskcache_basic" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) defer func() { _ = c.Clear() }() @@ -55,9 +52,7 @@ func TestDiskCacheBasic(t *testing.T) { } func TestDiskCacheTTL(t *testing.T) { - testDir := "uploads/test_diskcache_ttl" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) defer func() { _ = c.Clear() }() @@ -91,9 +86,7 @@ func TestDiskCacheTTL(t *testing.T) { } func TestDiskCacheExpirationPolicies(t *testing.T) { - testDir := "uploads/test_diskcache_expiration_policies" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) defer func() { _ = c.Clear() }() @@ -102,14 +95,14 @@ func TestDiskCacheExpirationPolicies(t *testing.T) { if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil { t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err) } - if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil { - t.Fatalf("Set(custom, 100ms) returned error: %v", err) + if err := c.Set("custom", []byte("custom"), 150*time.Millisecond); err != nil { + t.Fatalf("Set(custom, 150ms) returned error: %v", err) } if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil { t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err) } - time.Sleep(75 * time.Millisecond) + time.Sleep(80 * time.Millisecond) if _, err := c.Get("default"); err != ErrCacheMiss { t.Errorf("Get(default) error = %v, want ErrCacheMiss", err) @@ -121,7 +114,7 @@ func TestDiskCacheExpirationPolicies(t *testing.T) { t.Errorf("Get(permanent) returned error: %v", err) } - time.Sleep(50 * time.Millisecond) + time.Sleep(100 * time.Millisecond) if _, err := c.Get("custom"); err != ErrCacheMiss { t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err) @@ -132,9 +125,7 @@ func TestDiskCacheExpirationPolicies(t *testing.T) { } func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) { - testDir := "uploads/test_diskcache_no_expiration_reload" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil { @@ -154,9 +145,7 @@ func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) { } func TestDiskCacheLRUEviction(t *testing.T) { - testDir := "uploads/test_diskcache_lru" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) defer func() { _ = c.Clear() }() @@ -186,10 +175,8 @@ func TestDiskCacheLRUEviction(t *testing.T) { t.Errorf("k2 should exist: %v", err) } - // Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20. - // This should evict the oldest item. Since k1 was accessed, but then k2 was accessed, - // wait, let's access k1 again to make it the most recently used, so k2 becomes oldest! - _, _ = c.Get("k1") // k1 is now MRU, k2 is LRU + // Access k1 again to make it MRU, k2 becomes LRU + _, _ = c.Get("k1") err = c.Set("k3", []byte("v3"), DefaultExpiration) if err != nil { diff --git a/backend/pkg/util/gin_context.go b/backend/pkg/ginutil/context.go similarity index 84% rename from backend/pkg/util/gin_context.go rename to backend/pkg/ginutil/context.go index 0b31aafc..3ea9ec0d 100644 --- a/backend/pkg/util/gin_context.go +++ b/backend/pkg/ginutil/context.go @@ -1,7 +1,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package util +// Package ginutil provides helper utilities for Gin web framework contexts. +package ginutil import "github.com/gin-gonic/gin" diff --git a/backend/plugins/domain/admin/handlers_config.go b/backend/plugins/domain/admin/handlers_config.go index 6ce71d43..9aed1061 100644 --- a/backend/plugins/domain/admin/handlers_config.go +++ b/backend/plugins/domain/admin/handlers_config.go @@ -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) { diff --git a/backend/plugins/domain/admin/handlers_user.go b/backend/plugins/domain/admin/handlers_user.go index db954242..bbe95710 100644 --- a/backend/plugins/domain/admin/handlers_user.go +++ b/backend/plugins/domain/admin/handlers_user.go @@ -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 diff --git a/backend/plugins/domain/admin/middlewares.go b/backend/plugins/domain/admin/middlewares.go index 9632c7aa..f9b66307 100644 --- a/backend/plugins/domain/admin/middlewares.go +++ b/backend/plugins/domain/admin/middlewares.go @@ -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 diff --git a/backend/plugins/domain/auth/handlers.go b/backend/plugins/domain/auth/handlers.go index 3fe020d1..d8b149d6 100644 --- a/backend/plugins/domain/auth/handlers.go +++ b/backend/plugins/domain/auth/handlers.go @@ -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 diff --git a/backend/plugins/domain/auth/middleware.go b/backend/plugins/domain/auth/middleware.go index 2ef4c130..b89ae183 100644 --- a/backend/plugins/domain/auth/middleware.go +++ b/backend/plugins/domain/auth/middleware.go @@ -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 } diff --git a/backend/plugins/domain/auth/service.go b/backend/plugins/domain/auth/service.go index dc936f36..1ef3c2e7 100644 --- a/backend/plugins/domain/auth/service.go +++ b/backend/plugins/domain/auth/service.go @@ -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 } } diff --git a/backend/plugins/domain/message_gateway/admin_handlers.go b/backend/plugins/domain/message_gateway/admin_handlers.go index 885cd10d..a44871f9 100644 --- a/backend/plugins/domain/message_gateway/admin_handlers.go +++ b/backend/plugins/domain/message_gateway/admin_handlers.go @@ -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()) diff --git a/backend/plugins/domain/message_gateway/handler_helpers.go b/backend/plugins/domain/message_gateway/handler_helpers.go new file mode 100644 index 00000000..1908b8ae --- /dev/null +++ b/backend/plugins/domain/message_gateway/handler_helpers.go @@ -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)) +} diff --git a/backend/plugins/domain/message_gateway/handlers.go b/backend/plugins/domain/message_gateway/handlers.go index e01a5bb7..22949af0 100644 --- a/backend/plugins/domain/message_gateway/handlers.go +++ b/backend/plugins/domain/message_gateway/handlers.go @@ -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. diff --git a/backend/plugins/domain/message_gateway/push_channels.go b/backend/plugins/domain/message_gateway/push_channels.go index 37726b63..175d1ad7 100644 --- a/backend/plugins/domain/message_gateway/push_channels.go +++ b/backend/plugins/domain/message_gateway/push_channels.go @@ -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()) diff --git a/backend/plugins/domain/message_gateway/push_handlers.go b/backend/plugins/domain/message_gateway/push_handlers.go index 41ed9af8..086065d6 100644 --- a/backend/plugins/domain/message_gateway/push_handlers.go +++ b/backend/plugins/domain/message_gateway/push_handlers.go @@ -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)) diff --git a/backend/plugins/domain/message_gateway/repository.go b/backend/plugins/domain/message_gateway/repository.go index 3221cf7b..cd8ed007 100644 --- a/backend/plugins/domain/message_gateway/repository.go +++ b/backend/plugins/domain/message_gateway/repository.go @@ -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 清理启用通知事件的缓存。 diff --git a/backend/plugins/domain/risk_control/middleware.go b/backend/plugins/domain/risk_control/middleware.go index 44cfd82d..631467b2 100644 --- a/backend/plugins/domain/risk_control/middleware.go +++ b/backend/plugins/domain/risk_control/middleware.go @@ -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 } diff --git a/backend/plugins/domain/risk_control/middleware_test.go b/backend/plugins/domain/risk_control/middleware_test.go index c8af3113..f1bd7bb6 100644 --- a/backend/plugins/domain/risk_control/middleware_test.go +++ b/backend/plugins/domain/risk_control/middleware_test.go @@ -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()) diff --git a/backend/plugins/domain/upload/filesrv/file_server.go b/backend/plugins/domain/upload/filesrv/file_server.go index 276d8d30..34d0922f 100644 --- a/backend/plugins/domain/upload/filesrv/file_server.go +++ b/backend/plugins/domain/upload/filesrv/file_server.go @@ -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 } diff --git a/backend/plugins/domain/upload/handler/file_management.go b/backend/plugins/domain/upload/handler/file_management.go index 1e25371e..eb495270 100644 --- a/backend/plugins/domain/upload/handler/file_management.go +++ b/backend/plugins/domain/upload/handler/file_management.go @@ -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) diff --git a/backend/plugins/domain/upload/handler/routers.go b/backend/plugins/domain/upload/handler/routers.go index 6f0b44dc..9305fcb5 100644 --- a/backend/plugins/domain/upload/handler/routers.go +++ b/backend/plugins/domain/upload/handler/routers.go @@ -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") diff --git a/backend/plugins/domain/upload/handler/routers_test.go b/backend/plugins/domain/upload/handler/routers_test.go index 22d391e4..36ab068c 100644 --- a/backend/plugins/domain/upload/handler/routers_test.go +++ b/backend/plugins/domain/upload/handler/routers_test.go @@ -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() } diff --git a/backend/plugins/drivers/driver_asynq_cron/plugin.go b/backend/plugins/drivers/driver_asynq_cron/plugin.go index 8823e46a..a57e2105 100644 --- a/backend/plugins/drivers/driver_asynq_cron/plugin.go +++ b/backend/plugins/drivers/driver_asynq_cron/plugin.go @@ -139,34 +139,7 @@ func (p *Plugin) Start(_ context.Context) error { } if p.scheduler == nil { - opts := p.schedulerOpts - if opts == nil { - opts = &asynq.SchedulerOpts{ - Location: p.location, - } - } else if opts.Location == nil && p.location != nil { - opts.Location = p.location - } - - opt := p.redisOpt - if opt == nil { - if RedisOpt != nil { - opt = RedisOpt - } else { - redisCfg := config.Config.Redis - addr := "127.0.0.1:6379" - if len(redisCfg.Addrs) > 0 && redisCfg.Addrs[0] != "" { - addr = redisCfg.Addrs[0] - } - opt = asynq.RedisClientOpt{ - Addr: addr, - Username: redisCfg.Username, - Password: redisCfg.Password, - DB: redisCfg.DB, - } - } - } - p.scheduler = asynq.NewScheduler(opt, opts) + p.scheduler = p.initScheduler() } if p.coreCtx != nil && p.coreCtx.Schedules() != nil { @@ -268,3 +241,37 @@ func buildAsynqOptions(opts map[string]any) []asynq.Option { return res } + +func (p *Plugin) initScheduler() *asynq.Scheduler { + opts := p.schedulerOpts + if opts == nil { + opts = &asynq.SchedulerOpts{ + Location: p.location, + } + } else if opts.Location == nil && p.location != nil { + opts.Location = p.location + } + + opt := p.resolveRedisOpt() + return asynq.NewScheduler(opt, opts) +} + +func (p *Plugin) resolveRedisOpt() asynq.RedisConnOpt { + if p.redisOpt != nil { + return p.redisOpt + } + if RedisOpt != nil { + return RedisOpt + } + redisCfg := config.Config.Redis + addr := "127.0.0.1:6379" + if len(redisCfg.Addrs) > 0 && redisCfg.Addrs[0] != "" { + addr = redisCfg.Addrs[0] + } + return asynq.RedisClientOpt{ + Addr: addr, + Username: redisCfg.Username, + Password: redisCfg.Password, + DB: redisCfg.DB, + } +} diff --git a/backend/plugins/drivers/driver_asynq_cron/scheduler.go b/backend/plugins/drivers/driver_asynq_cron/scheduler.go index 1841b5cc..5915186c 100644 --- a/backend/plugins/drivers/driver_asynq_cron/scheduler.go +++ b/backend/plugins/drivers/driver_asynq_cron/scheduler.go @@ -111,21 +111,7 @@ func ReloadScheduler() error { // 4. 遍历并注册任务 taskSvc := getTaskService() for _, s := range schedules { - taskName := s.TaskType - maxRetry := 3 - queue := "default" - - if taskSvc != nil { - if meta, ok := taskSvc.GetTaskMeta(s.TaskType); ok { - taskName = meta.Name - if meta.MaxRetry > 0 { - maxRetry = meta.MaxRetry - } - if meta.Queue != "" { - queue = meta.Queue - } - } - } + taskName, maxRetry, queue := resolveTaskScheduleMeta(taskSvc, s.TaskType) // 构造 Asynq 载荷 t := asynq.NewTask(taskName, []byte(s.Payload)) @@ -160,3 +146,24 @@ func waitForStop(done, signals <-chan struct{}) bool { return true } } + +func resolveTaskScheduleMeta(taskSvc contracts.TaskService, taskType string) (name string, maxRetry int, queue string) { + name = taskType + maxRetry = 3 + queue = "default" + if taskSvc == nil { + return + } + meta, ok := taskSvc.GetTaskMeta(taskType) + if !ok { + return + } + name = meta.Name + if meta.MaxRetry > 0 { + maxRetry = meta.MaxRetry + } + if meta.Queue != "" { + queue = meta.Queue + } + return +} diff --git a/backend/plugins/drivers/driver_asynq_worker/plugin.go b/backend/plugins/drivers/driver_asynq_worker/plugin.go index b702aed9..5fd31bc0 100644 --- a/backend/plugins/drivers/driver_asynq_worker/plugin.go +++ b/backend/plugins/drivers/driver_asynq_worker/plugin.go @@ -366,27 +366,8 @@ func (s *taskServiceImpl) ListExecutions(ctx context.Context, taskType, status s return nil, 0, err } res := make([]contracts.TaskExecutionDTO, 0, len(rows)) - for _, r := range rows { - res = append(res, contracts.TaskExecutionDTO{ - ID: r.ID, - TaskID: r.TaskID, - TaskType: r.TaskType, - TaskName: r.TaskName, - Status: string(r.Status), - Retryable: r.Retryable, - MaxRetry: r.MaxRetry, - RetryCount: r.RetryCount, - Log: r.Log, - ErrorMessage: r.ErrorMessage, - Result: r.Result, - StartedAt: r.StartedAt, - FinishedAt: r.FinishedAt, - Duration: r.Duration, - Payload: r.Payload, - TriggeredBy: r.TriggeredBy, - CreatedAt: r.CreatedAt, - UpdatedAt: r.UpdatedAt, - }) + for i := range rows { + res = append(res, toTaskExecutionDTO(&rows[i])) } return res, total, nil } @@ -416,7 +397,12 @@ func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contrac if err != nil { return nil, err } - return &contracts.TaskExecutionDTO{ + dto := toTaskExecutionDTO(exec) + return &dto, nil +} + +func toTaskExecutionDTO(exec *TaskExecution) contracts.TaskExecutionDTO { + return contracts.TaskExecutionDTO{ ID: exec.ID, TaskID: exec.TaskID, TaskType: exec.TaskType, @@ -435,5 +421,5 @@ func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contrac TriggeredBy: exec.TriggeredBy, CreatedAt: exec.CreatedAt, UpdatedAt: exec.UpdatedAt, - }, nil + } } diff --git a/backend/plugins/drivers/driver_inproc_cron/scheduler.go b/backend/plugins/drivers/driver_inproc_cron/scheduler.go index 05b75a34..8b3b33f0 100644 --- a/backend/plugins/drivers/driver_inproc_cron/scheduler.go +++ b/backend/plugins/drivers/driver_inproc_cron/scheduler.go @@ -12,6 +12,7 @@ import ( "encoding/json" "errors" "fmt" + "strings" "sync" "time" @@ -78,7 +79,7 @@ func (s *inprocScheduler) registerJob(ctx context.Context, def extpoints.Schedul spec := def.Spec taskType := def.TaskType - fields := len(cronFields(spec)) + fields := len(strings.Fields(spec)) cronSpec := spec if fields == standardCronFields { cronSpec = "0 " + spec @@ -141,22 +142,3 @@ func invokeHandler(ctx context.Context, handler any, payload []byte) error { return fmt.Errorf("unsupported handler type: %T", handler) } } - -func cronFields(s string) []string { - var fields []string - var current []rune - for _, r := range s { - if r == ' ' || r == '\t' { - if len(current) > 0 { - fields = append(fields, string(current)) - current = nil - } - } else { - current = append(current, r) - } - } - if len(current) > 0 { - fields = append(fields, string(current)) - } - return fields -} diff --git a/backend/plugins/drivers/driver_inproc_worker/executor.go b/backend/plugins/drivers/driver_inproc_worker/executor.go index 6c47acca..7fcc9f2f 100644 --- a/backend/plugins/drivers/driver_inproc_worker/executor.go +++ b/backend/plugins/drivers/driver_inproc_worker/executor.go @@ -101,7 +101,7 @@ func (q *InprocQueue) Start(ctx context.Context) { q.wg.Add(1) util.Go(func() { defer q.wg.Done() - q.workerLoop() + q.workerLoop(ctx) }) } } @@ -128,21 +128,23 @@ func (q *InprocQueue) Stop(ctx context.Context) error { } } -func (q *InprocQueue) workerLoop() { +func (q *InprocQueue) workerLoop(ctx context.Context) { for { select { case <-q.stopCh: return + case <-ctx.Done(): + return case msg, ok := <-q.queue: if !ok { return } - q.executeTask(msg) + q.executeTask(ctx, msg) } } } -func (q *InprocQueue) executeTask(msg TaskMessage) { +func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) { if q.taskReg == nil { return } @@ -157,7 +159,7 @@ func (q *InprocQueue) executeTask(msg TaskMessage) { timeout = 5 * time.Minute } - taskCtx, cancel := context.WithTimeout(q.baseCtx, timeout) + taskCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() err := invokeHandler(taskCtx, td.Handler, msg.Payload) @@ -165,7 +167,13 @@ func (q *InprocQueue) executeTask(msg TaskMessage) { msg.RetryLeft-- // Retry with backoff util.Go(func() { - time.Sleep(defaultRetryBackoff) + select { + case <-time.After(defaultRetryBackoff): + case <-q.stopCh: + return + case <-ctx.Done(): + return + } if q.running.Load() { select { case q.queue <- msg: diff --git a/scripts/check_cordis_architecture.sh b/scripts/check_cordis_architecture.sh index 40921d57..77df1624 100755 --- a/scripts/check_cordis_architecture.sh +++ b/scripts/check_cordis_architecture.sh @@ -106,12 +106,12 @@ else log_pass "backend/pkg/ 零插件依赖" fi -# 3.2 pkg/util/ 严禁导入 ORM / Session 框架 -UTIL_FRAMEWORK_IMPORTS=$(rg -n '"gorm.io/gorm"|"github.com/gorilla/sessions"' \ +# 3.2 pkg/util/ 严禁导入 Gin / ORM / Session 框架 +UTIL_FRAMEWORK_IMPORTS=$(rg -n '"gorm.io/gorm"|"github.com/gorilla/sessions"|"github.com/gin-gonic/gin"' \ "${BACKEND_DIR}/pkg/util/" --glob '*.go' -g '!*_test.go' || true) if [ -n "${UTIL_FRAMEWORK_IMPORTS}" ]; then - log_fail "backend/pkg/util/ 必须保持纯粹,禁止导入 gorm、sessions 等数据库/会话框架包:" + log_fail "backend/pkg/util/ 必须保持纯粹,禁止导入 gin、gorm、sessions 等 Web/数据库/会话框架包:" echo "${UTIL_FRAMEWORK_IMPORTS}" >&2 else log_pass "backend/pkg/util/ 保持纯净无状态"