From 1b2e083aec1279094db4fc310ff917553fff3b9f Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 18 Jun 2026 12:12:49 +0800 Subject: [PATCH] refactor(api): extract repository layer and thin HTTP handlers Introduce internal/repository for data access and cache-backed system config reads. Move business logic into logics.go across admin push, user, template, cache, system_config, and upload/handler packages. Remove Gin from internal/util by relocating request-scoped helpers to oauth/gin_context.go. Propagate request context for config lookups in user flows. Slim model entities and delete model-level DB/cache helpers. Wire handlers to logics/repository so targeted packages no longer call db.DB directly. Update admin router tests to use ErrorHandlerMiddleware. --- .../apps/admin/auth_source/routers_test.go | 3 +- internal/apps/admin/cache/logics.go | 14 + internal/apps/admin/cache/routers.go | 41 +- internal/apps/admin/logs/utils.go | 4 +- internal/apps/admin/middlewares.go | 7 +- internal/apps/admin/push/channels.go | 147 +------ internal/apps/admin/push/events.go | 148 +------ internal/apps/admin/push/logics.go | 413 ++++++++++++++++++ internal/apps/admin/push/push_test.go | 3 +- internal/apps/admin/push/routers.go | 280 ++---------- internal/apps/admin/push/task_listener.go | 41 +- internal/apps/admin/push/tasks.go | 44 +- internal/apps/admin/system_config/logics.go | 128 ++++++ internal/apps/admin/system_config/routers.go | 152 ++----- .../apps/admin/system_config/routers_test.go | 37 +- internal/apps/admin/task/routers_test.go | 3 +- internal/apps/admin/template/logics.go | 78 ++++ internal/apps/admin/template/routers.go | 120 ++--- internal/apps/admin/template/routers_test.go | 6 +- internal/apps/admin/updater/logics.go | 11 +- internal/apps/admin/user/logics.go | 106 +++++ internal/apps/admin/user/routers.go | 201 ++------- internal/apps/admin/user/routers_test.go | 6 +- internal/apps/cap/routers_test.go | 3 +- internal/apps/cap/runtime_settings.go | 7 +- internal/apps/cap/runtime_settings_test.go | 9 +- .../apps/config/public_config_cache_test.go | 15 +- internal/apps/config/routers.go | 11 +- .../apps/config/system_config_cache_test.go | 54 +-- internal/apps/oauth/auth_source_resolver.go | 9 +- .../context.go => apps/oauth/gin_context.go} | 10 +- internal/apps/oauth/handler_callback.go | 5 +- internal/apps/oauth/middlewares.go | 14 +- internal/apps/oauth/oauth_test.go | 13 +- internal/apps/oauth/routers.go | 4 +- internal/apps/oauth/session_context.go | 3 +- internal/apps/risk_control/middleware.go | 3 +- internal/apps/risk_control/middleware_test.go | 3 +- internal/apps/upload/cache/access_cache.go | 5 +- .../apps/upload/cache/access_cache_test.go | 9 +- internal/apps/upload/filesrv/file_server.go | 6 +- .../apps/upload/handler/file_management.go | 149 ++----- internal/apps/upload/handler/logics.go | 199 +++++++++ internal/apps/upload/handler/routers.go | 106 +---- internal/apps/upload/handler/routers_test.go | 11 +- internal/apps/upload/handler/stats.go | 5 +- internal/apps/user/access_tokens.go | 18 +- internal/apps/user/logics.go | 67 +-- internal/apps/user/logics_test.go | 7 +- internal/apps/user/routers.go | 46 +- internal/apps/user/routers_test.go | 11 +- internal/apps/user/tasks.go | 10 +- internal/db/migrator/migrator.go | 4 +- internal/db/migrator/migrator_test.go | 7 +- internal/diskcache/cache.go | 10 +- internal/diskcache/cache_test.go | 3 +- internal/model/push_channel.go | 62 +-- internal/model/push_event.go | 50 +-- internal/model/system_configs.go | 200 +-------- internal/model/templates.go | 18 - internal/model/users.go | 14 +- internal/repository/push_channel.go | 105 +++++ internal/repository/push_event.go | 124 ++++++ internal/repository/push_history.go | 54 +++ internal/repository/system_config.go | 198 +++++++++ internal/repository/system_config_admin.go | 85 ++++ .../system_config_cache.go | 20 +- internal/repository/template.go | 59 +++ internal/repository/upload.go | 118 +++++ internal/repository/upload_stat.go | 20 + internal/repository/user.go | 172 ++++++++ internal/router/middlewares.go | 5 +- internal/router/middlewares_test.go | 3 +- internal/storage/config.go | 7 +- internal/storage/storage.go | 4 +- internal/testhelper/test_helper.go | 5 +- internal/util/custom_types.go | 1 + 77 files changed, 2370 insertions(+), 1783 deletions(-) create mode 100644 internal/apps/admin/cache/logics.go create mode 100644 internal/apps/admin/push/logics.go create mode 100644 internal/apps/admin/system_config/logics.go create mode 100644 internal/apps/admin/template/logics.go create mode 100644 internal/apps/admin/user/logics.go rename internal/{util/context.go => apps/oauth/gin_context.go} (68%) create mode 100644 internal/apps/upload/handler/logics.go create mode 100644 internal/repository/push_channel.go create mode 100644 internal/repository/push_event.go create mode 100644 internal/repository/push_history.go create mode 100644 internal/repository/system_config.go create mode 100644 internal/repository/system_config_admin.go rename internal/{model => repository}/system_config_cache.go (79%) create mode 100644 internal/repository/template.go create mode 100644 internal/repository/upload.go create mode 100644 internal/repository/upload_stat.go create mode 100644 internal/repository/user.go diff --git a/internal/apps/admin/auth_source/routers_test.go b/internal/apps/admin/auth_source/routers_test.go index 8d9d3478..15112ec4 100644 --- a/internal/apps/admin/auth_source/routers_test.go +++ b/internal/apps/admin/auth_source/routers_test.go @@ -13,7 +13,6 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "github.com/Rain-kl/Wavelet/internal/common/response") @@ -26,7 +25,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine { // Mock authentication middleware adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() }) diff --git a/internal/apps/admin/cache/logics.go b/internal/apps/admin/cache/logics.go new file mode 100644 index 00000000..1f841ca6 --- /dev/null +++ b/internal/apps/admin/cache/logics.go @@ -0,0 +1,14 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache + +import ( + "context" + + "github.com/Rain-kl/Wavelet/internal/repository" +) + +func saveOrUpdateConfig(ctx context.Context, key, value string) error { + return repository.SaveOrUpdateSystemConfig(ctx, key, value) +} \ No newline at end of file diff --git a/internal/apps/admin/cache/routers.go b/internal/apps/admin/cache/routers.go index 3b98a37b..ef99d2b8 100644 --- a/internal/apps/admin/cache/routers.go +++ b/internal/apps/admin/cache/routers.go @@ -4,18 +4,16 @@ // Package cache provides HTTP handlers for managing disk cache. package cache -import ("context" - "errors" +import ( "net/http" "strconv" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" "github.com/gin-gonic/gin" - "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) type updateCacheConfigRequest struct { MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"` @@ -62,25 +60,21 @@ func UpdateCacheConfig(c *gin.Context) { ctx := c.Request.Context() - // Update Max Size if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil { response.AbortInternal(c, err.Error()) return } - // Update Default TTL if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil { response.AbortInternal(c, err.Error()) return } - // Update LRU Enabled if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil { response.AbortInternal(c, err.Error()) return } - // Trigger hot reloading in global cache diskcache.GetGlobalCache().ReloadConfig(ctx) c.JSON(http.StatusOK, response.OKNil()) @@ -103,31 +97,4 @@ func ClearCache(c *gin.Context) { return } c.JSON(http.StatusOK, response.OKNil()) -} - -func saveOrUpdateConfig(ctx context.Context, key string, value string) error { - var sc model.SystemConfig - err := db.DB(ctx).Where("key = ?", key).First(&sc).Error - if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return err - } - - if errors.Is(err, gorm.ErrRecordNotFound) { - sc = model.SystemConfig{ - Key: key, - Value: value, - Type: "system", - Visibility: 0, - } - if err := db.DB(ctx).Create(&sc).Error; err != nil { - return err - } - } else { - sc.Value = value - if err := db.DB(ctx).Save(&sc).Error; err != nil { - return err - } - } - - return model.InvalidateSystemConfigCache(ctx, key) -} +} \ No newline at end of file diff --git a/internal/apps/admin/logs/utils.go b/internal/apps/admin/logs/utils.go index 92232c64..3c93edcc 100644 --- a/internal/apps/admin/logs/utils.go +++ b/internal/apps/admin/logs/utils.go @@ -13,6 +13,7 @@ import ( "github.com/gorilla/websocket" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" ) // getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击 @@ -32,8 +33,7 @@ func getUpgrader() *websocket.Upgrader { // 2. 检查配置的允许跨域 Origin (Check allowed origins in system config) ctx := r.Context() - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" { + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" { originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/") allowedOrigins := strings.Split(sc.Value, ",") for _, allowed := range allowedOrigins { diff --git a/internal/apps/admin/middlewares.go b/internal/apps/admin/middlewares.go index 3bd1c25f..61fc86bf 100644 --- a/internal/apps/admin/middlewares.go +++ b/internal/apps/admin/middlewares.go @@ -7,7 +7,6 @@ package admin import ( "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/Rain-kl/Wavelet/pkg/logger" otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" @@ -22,11 +21,11 @@ func LoginAdminRequired() gin.HandlerFunc { ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired") defer span.End() - user, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) // 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限 - if tokenAuth, _ := util.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth { - tokenAdmin, _ := util.GetFromContext[bool](c, oauth.TokenAdminKey) + if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth { + tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey) if !tokenAdmin { response.AbortNotFound(c, TokenAdminRequired) return diff --git a/internal/apps/admin/push/channels.go b/internal/apps/admin/push/channels.go index 36ab1489..2da12daf 100644 --- a/internal/apps/admin/push/channels.go +++ b/internal/apps/admin/push/channels.go @@ -3,19 +3,20 @@ package push -import ("encoding/json" +import ( + "encoding/json" "errors" "net/http" "strconv" "strings" - "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/model" pkgpush "github.com/Rain-kl/Wavelet/pkg/push" "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" + "github.com/Rain-kl/Wavelet/internal/model" +) // ListChannelDefinitions 获取各种消息通道的表单配置定义列表 // @Summary 获取所有消息通道配置字段定义 @@ -38,9 +39,8 @@ func ListChannelDefinitions(c *gin.Context) { // @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表" // @Router /api/v1/admin/push/channels [get] func ListChannels(c *gin.Context) { - ctx := c.Request.Context() - var channels []model.PushChannel - if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil { + channels, err := listPushChannels(c.Request.Context()) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -75,41 +75,11 @@ func CreateChannel(c *gin.Context) { return } - ctx := c.Request.Context() - - var count int64 - if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", req.Name).Count(&count).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - if count > 0 { - response.AbortBadRequest(c, "channel name already exists") - return - } - - channel := model.PushChannel{ - Name: req.Name, - Description: req.Description, - Type: req.Type, - Token: req.Token, - URL: req.URL, - Other: req.Other, - Enabled: req.Enabled, - } - - if err := channel.Validate(); err != nil { + channel, err := createPushChannel(c.Request.Context(), req) + if err != nil { response.AbortBadRequest(c, err.Error()) return } - - if err := db.DB(ctx).Create(&channel).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除渠道缓存 - model.DeleteActivePushChannelCache(ctx, channel.Name) - c.JSON(http.StatusOK, response.OK(channel)) } @@ -135,8 +105,7 @@ type UpdateChannelRequest struct { // @Success 200 {object} response.Any{data=model.PushChannel} "更新成功" // @Router /api/v1/admin/push/channels/{id} [put] func UpdateChannel(c *gin.Context) { - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) + id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { response.AbortBadRequest(c, "invalid channel id") return @@ -148,10 +117,8 @@ func UpdateChannel(c *gin.Context) { return } - ctx := c.Request.Context() - - var channel model.PushChannel - if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { + channel, err := updatePushChannel(c.Request.Context(), id, req) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, "channel not found") return @@ -159,27 +126,6 @@ func UpdateChannel(c *gin.Context) { response.AbortInternal(c, err.Error()) return } - - channel.Description = req.Description - channel.Type = req.Type - channel.Token = req.Token - channel.URL = req.URL - channel.Other = req.Other - channel.Enabled = req.Enabled - - if err := channel.Validate(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := db.DB(ctx).Save(&channel).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除渠道缓存 - model.DeleteActivePushChannelCache(ctx, channel.Name) - c.JSON(http.StatusOK, response.OK(channel)) } @@ -193,16 +139,13 @@ func UpdateChannel(c *gin.Context) { // @Success 200 {object} response.Any "删除成功" // @Router /api/v1/admin/push/channels/{id} [delete] func DeleteChannel(c *gin.Context) { - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) + id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { response.AbortBadRequest(c, "invalid channel id") return } - ctx := c.Request.Context() - var channel model.PushChannel - if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { + if err := deletePushChannel(c.Request.Context(), id); err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, "channel not found") return @@ -210,15 +153,6 @@ func DeleteChannel(c *gin.Context) { response.AbortInternal(c, err.Error()) return } - - if err := db.DB(ctx).Delete(&channel).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除渠道缓存 - model.DeleteActivePushChannelCache(ctx, channel.Name) - c.JSON(http.StatusOK, response.OKNil()) } @@ -250,26 +184,12 @@ func TestChannel(c *gin.Context) { } ctx := c.Request.Context() - var url, token, other, channelType string - - if req.Name != "" { - var channel model.PushChannel - if err := db.DB(ctx).Where("name = ?", req.Name).First(&channel).Error; err != nil { - response.AbortBadRequest(c, "channel not found") - return - } - url = channel.URL - token = channel.Token - other = channel.Other - channelType = channel.Type - } else { - url = req.URL - token = req.Token - other = req.Other - channelType = req.Type + url, token, other, channelType, err := loadChannelForTest(ctx, req) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return } - // 对邮件类型应用全局配置作为回退 if channelType == channelEmail { url, token, other = resolveSMTPConfig(ctx, url, token, other) } @@ -282,7 +202,6 @@ func TestChannel(c *gin.Context) { Type: channelType, Enabled: true, } - if err := tempChannel.Validate(); err != nil { response.AbortBadRequest(c, err.Error()) return @@ -291,34 +210,16 @@ func TestChannel(c *gin.Context) { var config pkgpush.Config var renderedJSON string - switch channelType { case channelLark: - config = pkgpush.Config{ - Channel: channelLark, - URL: url, - Secret: token, - } + config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token} renderedJSON = other case channelEmail: - config = pkgpush.Config{ - Channel: channelEmail, - URL: url, - Key: token, - Secret: other, - } + config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other} case channelTelegram: - config = pkgpush.Config{ - Channel: channelTelegram, - URL: url, - Secret: token, - Key: other, - } + config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other} default: - config = pkgpush.Config{ - Channel: channelCustom, - URL: url, - } + config = pkgpush.Config{Channel: channelCustom, URL: url} customPushReq := CustomPushRequest{ Title: "通道测试通知", Content: "这是一条来自系统的消息通道连通性测试消息。", @@ -340,12 +241,10 @@ func TestChannel(c *gin.Context) { }, Template: renderedJSON, } - if err := enqueuePushTask(ctx, payload); err != nil { response.AbortInternal(c, err.Error()) return } - c.JSON(http.StatusOK, response.OKNil()) } @@ -376,4 +275,4 @@ func renderCustomPayload(template string, req CustomPushRequest) string { result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL)) result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To)) return result -} +} \ No newline at end of file diff --git a/internal/apps/admin/push/events.go b/internal/apps/admin/push/events.go index 4e699097..18772a94 100644 --- a/internal/apps/admin/push/events.go +++ b/internal/apps/admin/push/events.go @@ -9,12 +9,9 @@ import ( "encoding/json" "errors" "fmt" - "strconv" "strings" - "sync" - - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/pkg/logger" pkgpush "github.com/Rain-kl/Wavelet/pkg/push" @@ -73,30 +70,7 @@ type EventTrigger struct{} // DefaultTrigger is the singleton instance of EventTrigger. var DefaultTrigger = &EventTrigger{} -var ( - systemUser *model.User - systemOnce sync.Once -) - -func getSystemUser(ctx context.Context) *model.User { - systemOnce.Do(func() { - var u model.User - if err := db.DB(ctx).Where("username = ?", "system").First(&u).Error; err == nil { - systemUser = &u - } else { - systemUser = &model.User{ - ID: 999, - Username: "system", - Nickname: "系统", - Email: "", - } - } - }) - return systemUser -} - // Trigger receives event metadata and processes the event notification dispatch asynchronously. -// It automatically enqueues tasks using a background goroutine and avoids blocking the calling thread. // //nolint:contextcheck func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) { @@ -109,8 +83,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map body["user"] = getSystemUser(asyncCtx) } - // 1. Check if the event is enabled (try Redis cache first) - eventPtr, err := model.GetActivePushEventByKey(asyncCtx, meta.Key) + eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return @@ -119,22 +92,16 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map return } event := *eventPtr - if len(event.Channels) == 0 { return } - // 2. Build and render notification message flatBody := getFlatBody(body) msg, _ := t.buildMessage(&event, meta, flatBody, body) - - // 3. Enqueue tasks for each matching channel 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) { var msg NotificationMessage renderedTemplate := "" @@ -222,13 +189,11 @@ func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[str func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) { for _, channelName := range event.Channels { - // 检查是不是自定义数据库渠道 (使用 Redis 缓存优先) - customChannel, err := model.GetActivePushChannelByName(ctx, channelName) + customChannel, err := repository.GetActivePushChannelByName(ctx, channelName) if err == nil { t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) continue } - logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) } } @@ -251,32 +216,15 @@ func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, m switch channel.Type { case channelLark: - config = pkgpush.Config{ - Channel: channelLark, - URL: channel.URL, - Secret: channel.Token, // Feishu Bot Sign Secret - } - renderedTemplate = channel.Other // Optional custom template/card for lark + config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token} + renderedTemplate = channel.Other case channelEmail: url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other) - config = pkgpush.Config{ - Channel: channelEmail, - URL: url, // SMTP host:port - Key: token, // SMTP Username - Secret: other, // SMTP Password - } + config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other} case channelTelegram: - config = pkgpush.Config{ - Channel: channelTelegram, - URL: channel.URL, - Secret: channel.Token, // Telegram Bot Token - Key: channel.Other, // Default Chat ID - } - default: // custom - config = pkgpush.Config{ - Channel: channelCustom, - URL: channel.URL, - } + config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} + default: + config = pkgpush.Config{Channel: channelCustom, URL: channel.URL} customPushReq := CustomPushRequest{ Title: msg.Title, Content: msg.Content, @@ -301,8 +249,6 @@ func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, m } } - - func enqueuePushTask(ctx context.Context, payload SendPayload) error { payloadBytes, err := json.Marshal(payload) if err != nil { @@ -348,48 +294,23 @@ func resolveTarget(ctx context.Context, target string, flatBody map[string]any, } resolved := resolveDynamicKeyword(target, flatBody) - - // 2. 如果包含 @,说明已经是个邮箱,直接返回 if strings.Contains(resolved, "@") { return resolved } - - // 2.5 如果为特殊的系统虚拟用户,自动映射为首位管理员 if val, matched := resolveSystemTarget(ctx, resolved, channel); matched { return val } - // 3. 不包含 @,说明可能是用户 ID 或用户名。我们需要从数据库中查询对应用户 - var user model.User - found := false - - // 尝试作为用户 ID 查询(纯数字) - if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { - if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err == nil { - found = true - } - } - - // 如果没有按 ID 查到,尝试作为用户名查询 - if !found { - if err := db.DB(ctx).Where("username = ?", resolved).First(&user).Error; err == nil { - found = true - } - } - - // 4. 根据查询结果 and 推送渠道进行转换 + user, found := resolveTargetUser(ctx, resolved, channel) if !found { return resolved } - if channel == channelEmail && user.Email != "" { return user.Email } - if channel != channelEmail && user.Username != "" { return user.Username } - return resolved } @@ -418,51 +339,4 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string { } } return target -} - -func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) { - if resolved != "系统" && resolved != "system" && resolved != "0" { - return "", false - } - var adminUser model.User - if err := db.DB(ctx).Where("is_admin = ?", true).Order("id asc").First(&adminUser).Error; err != nil { - return resolved, true - } - if channel == channelEmail && adminUser.Email != "" { - return adminUser.Email, true - } - if channel != channelEmail && adminUser.Username != "" { - return adminUser.Username, true - } - return resolved, true -} - -// resolveSMTPConfig resolves SMTP configuration by falling back to system-wide global configuration if any inputs are blank. -func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) { - if url != "" && token != "" { - return url, token, other - } - var smtpHost, smtpPort, smtpUser, smtpPass model.SystemConfig - _ = smtpHost.GetByKey(ctx, model.ConfigKeySMTPHost) - _ = smtpPort.GetByKey(ctx, model.ConfigKeySMTPPort) - _ = smtpUser.GetByKey(ctx, model.ConfigKeySMTPUsername) - _ = smtpPass.GetByKey(ctx, model.ConfigKeySMTPPassword) - - if smtpHost.Value == "" || smtpUser.Value == "" { - return url, token, other - } - port := smtpPort.Value - if port == "" { - port = "587" - } - if url == "" { - url = smtpHost.Value + ":" + port - } - if token == "" { - token = smtpUser.Value - } - if other == "" { - other = smtpPass.Value - } - return url, token, other -} +} \ No newline at end of file diff --git a/internal/apps/admin/push/logics.go b/internal/apps/admin/push/logics.go new file mode 100644 index 00000000..df9f99df --- /dev/null +++ b/internal/apps/admin/push/logics.go @@ -0,0 +1,413 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "encoding/json" + "errors" + "strconv" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/task" + pkgpush "github.com/Rain-kl/Wavelet/pkg/push" + "gorm.io/gorm" +) + +type smtpConfig struct { + Host string + Port string + Username string + Password string +} + +func loadSMTPConfig(ctx context.Context) smtpConfig { + host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost) + port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort) + user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername) + pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword) + return smtpConfig{ + Host: host.Value, + Port: port.Value, + Username: user.Value, + Password: pass.Value, + } +} + +func syncBuiltInEvents(ctx context.Context) error { + for _, meta := range BuiltInEvents { + _, err := repository.GetPushEventByKey(ctx, meta.Key) + if errors.Is(err, gorm.ErrRecordNotFound) { + var defaultTemplateStr string + if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { + defaultTemplateStr = string(defaultTemplateBytes) + } + event := model.PushEvent{ + EventKey: meta.Key, + Name: meta.Name, + Channels: []string{}, + Targets: []string{}, + Template: defaultTemplateStr, + Enabled: false, + } + if err := repository.CreatePushEvent(ctx, &event); err != nil { + return err + } + } else if err != nil { + return err + } + } + return nil +} + +func listPushEvents(ctx context.Context) ([]model.PushEvent, error) { + return repository.ListPushEvents(ctx) +} + +func createPushEvent(ctx context.Context, req CreateEventRequest) (model.PushEvent, error) { + eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req) + if err != nil { + return model.PushEvent{}, err + } + + count, err := repository.CountPushEventsByKey(ctx, eventKey) + if err != nil { + return model.PushEvent{}, err + } + if count > 0 { + return model.PushEvent{}, errors.New("this notification event is already configured") + } + + templateStr := strings.TrimSpace(req.Template) + if templateStr == "" { + templateStr = string(defaultTemplateBytes) + } else { + var tempMap map[string]any + if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { + return model.PushEvent{}, errors.New("custom template is not a valid JSON format") + } + } + + channels := req.Channels + if channels == nil { + channels = []string{} + } + targets := req.Targets + if targets == nil { + targets = []string{} + } + + event := model.PushEvent{ + EventKey: eventKey, + Name: eventName, + TaskType: req.TaskType, + Channels: channels, + Targets: targets, + Template: templateStr, + Enabled: req.Enabled, + } + if err := event.Validate(); err != nil { + return model.PushEvent{}, err + } + if err := repository.CreatePushEvent(ctx, &event); err != nil { + return model.PushEvent{}, err + } + return event, nil +} + +func deletePushEvent(ctx context.Context, id uint64) error { + event, err := repository.GetPushEventByID(ctx, id) + if err != nil { + return err + } + return repository.DeletePushEvent(ctx, &event) +} + +func updatePushEvent(ctx context.Context, id uint64, req UpdateEventRequest) error { + event, err := repository.GetPushEventByID(ctx, id) + if err != nil { + return err + } + + event.Channels = req.Channels + event.Targets = req.Targets + event.Template = req.Template + event.Enabled = req.Enabled + if err := event.Validate(); err != nil { + return err + } + return repository.SavePushEvent(ctx, &event) +} + +func togglePushEvent(ctx context.Context, id uint64) (bool, error) { + event, err := repository.GetPushEventByID(ctx, id) + if err != nil { + return false, err + } + + enabled := !event.Enabled + if enabled && len(event.Channels) == 0 { + return false, errors.New("cannot enable event without any push channels configured") + } + if err := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil { + return false, err + } + return enabled, nil +} + +func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) { + return repository.ListPushHistories(ctx, filter) +} + +func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) { + if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") { + return + } + smtp := loadSMTPConfig(ctx) + if smtp.Host == "" || smtp.Username == "" { + return + } + port := smtp.Port + if port == "" { + port = "587" + } + cfg.URL = smtp.Host + ":" + port + cfg.Key = smtp.Username + cfg.Secret = smtp.Password +} + +func listPushChannels(ctx context.Context) ([]model.PushChannel, error) { + return repository.ListPushChannels(ctx) +} + +func createPushChannel(ctx context.Context, req CreateChannelRequest) (model.PushChannel, error) { + count, err := repository.CountPushChannelsByName(ctx, req.Name) + if err != nil { + return model.PushChannel{}, err + } + if count > 0 { + return model.PushChannel{}, errors.New("channel name already exists") + } + + channel := model.PushChannel{ + Name: req.Name, + Description: req.Description, + Type: req.Type, + Token: req.Token, + URL: req.URL, + Other: req.Other, + Enabled: req.Enabled, + } + if err := channel.Validate(); err != nil { + return model.PushChannel{}, err + } + if err := repository.CreatePushChannel(ctx, &channel); err != nil { + return model.PushChannel{}, err + } + return channel, nil +} + +func updatePushChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (model.PushChannel, error) { + channel, err := repository.GetPushChannelByID(ctx, id) + if err != nil { + return model.PushChannel{}, err + } + + channel.Description = req.Description + channel.Type = req.Type + channel.Token = req.Token + channel.URL = req.URL + channel.Other = req.Other + channel.Enabled = req.Enabled + if err := channel.Validate(); err != nil { + return model.PushChannel{}, err + } + if err := repository.SavePushChannel(ctx, &channel); err != nil { + return model.PushChannel{}, err + } + return channel, nil +} + +func deletePushChannel(ctx context.Context, id uint64) error { + channel, err := repository.GetPushChannelByID(ctx, id) + if err != nil { + return err + } + return repository.DeletePushChannel(ctx, &channel) +} + +func loadChannelForTest(ctx context.Context, req TestChannelRequest) (string, string, string, string, error) { + if req.Name != "" { + channel, err := repository.GetPushChannelByName(ctx, req.Name) + if err != nil { + return "", "", "", "", errors.New("channel not found") + } + return channel.URL, channel.Token, channel.Other, channel.Type, nil + } + return req.URL, req.Token, req.Other, req.Type, nil +} + +func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) { + return repository.ListActivePushEventsByTaskType(ctx, taskType) +} + +func loadUserFromPayload(ctx context.Context, data map[string]any) any { + if u, exists := data["user"]; exists && u != nil { + return u + } + + if userID, ok := extractUserID(data); ok && userID > 0 { + if user, err := repository.GetUserByID(ctx, userID); err == nil { + return &user + } + } + + if username := extractUsername(data); username != "" { + if user, err := repository.GetUserByUsername(ctx, username); err == nil { + return &user + } + } + return nil +} + +func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error { + title := req.Body.Title + content := req.Body.Content + level := req.Body.Level + if title == "" { + title = "系统通知" + } + if level == "" { + level = defaultLevelInfo + } + + target := req.Target + if target == "" { + if req.Config.URL != "" { + target = req.Config.URL + const maxTargetLen = 50 + const truncatedLen = 47 + if len(target) > maxTargetLen { + target = target[:truncatedLen] + "..." + } + } else { + target = "default" + } + } + + history := model.PushHistory{ + EventKey: req.EventKey, + Channel: req.Config.Channel, + Target: target, + Title: title, + Content: content, + Level: level, + Status: status, + ErrorMsg: errMsg, + } + return repository.CreatePushHistory(ctx, &history) +} + +func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) { + found := false + var user model.User + + if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { + if u, err := repository.GetUserByID(ctx, id); err == nil { + user = u + found = true + } + } + if !found { + if u, err := repository.GetUserByUsername(ctx, resolved); err == nil { + user = u + found = true + } + } + return user, found +} + +func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) { + if resolved != "系统" && resolved != "system" && resolved != "0" { + return "", false + } + adminUser, err := repository.GetFirstAdminUser(ctx) + if err != nil { + return resolved, true + } + if channel == channelEmail && adminUser.Email != "" { + return adminUser.Email, true + } + if channel != channelEmail && adminUser.Username != "" { + return adminUser.Username, true + } + return resolved, true +} + +func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) { + if url != "" && token != "" { + return url, token, other + } + smtp := loadSMTPConfig(ctx) + if smtp.Host == "" || smtp.Username == "" { + return url, token, other + } + port := smtp.Port + if port == "" { + port = "587" + } + if url == "" { + url = smtp.Host + ":" + port + } + if token == "" { + token = smtp.Username + } + if other == "" { + other = smtp.Password + } + return url, token, other +} + +func getSystemUser(ctx context.Context) *model.User { + user := repository.GetSystemUser(ctx) + return &user +} + +func getEventInfo(req CreateEventRequest) (string, string, []byte, error) { + if req.TaskType != "" { + meta := task.GetTaskMetaByAsynqTask(req.TaskType) + if meta == nil { + return "", "", nil, errors.New("unsupported task type") + } + eventKey := "task_completed:" + req.TaskType + eventName := "任务完成: " + meta.Name + defaultTemplate := NotificationMessage{ + Title: "任务完成: " + meta.Name, + Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", + Level: defaultLevelInfo, + } + defaultTemplateBytes, err := json.Marshal(defaultTemplate) + if err != nil { + return "", "", nil, err + } + return eventKey, eventName, defaultTemplateBytes, nil + } + + if req.EventKey == "" { + return "", "", nil, errors.New("either event_key or task_type must be provided") + } + + meta, found := findBuiltInEvent(req.EventKey) + if !found { + return "", "", nil, errors.New("unsupported built-in event key") + } + + defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) + if err != nil { + return "", "", nil, err + } + return req.EventKey, meta.Name, defaultTemplateBytes, nil +} \ No newline at end of file diff --git a/internal/apps/admin/push/push_test.go b/internal/apps/admin/push/push_test.go index c3782e2f..5761255b 100644 --- a/internal/apps/admin/push/push_test.go +++ b/internal/apps/admin/push/push_test.go @@ -16,7 +16,6 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" pkgpush "github.com/Rain-kl/Wavelet/pkg/push" "github.com/alicebob/miniredis/v2" "github.com/gin-gonic/gin" @@ -104,7 +103,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine { adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, "user_obj", authUser) + oauth.SetToContext(c, "user_obj", authUser) } c.Next() }) diff --git a/internal/apps/admin/push/routers.go b/internal/apps/admin/push/routers.go index 1cc50d0c..674338cc 100644 --- a/internal/apps/admin/push/routers.go +++ b/internal/apps/admin/push/routers.go @@ -4,22 +4,21 @@ // Package push defines push notification HTTP routes. package push -import ("context" - "encoding/json" +import ( + "context" "errors" "fmt" "net/http" "strconv" - "strings" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/task" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/push" "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) // UpdateEventRequest 更新事件请求参数 type UpdateEventRequest struct { @@ -37,30 +36,7 @@ type TestPushRequest struct { // SyncEvents automatically registers/updates built-in events in the database. func SyncEvents(ctx context.Context) error { - for _, meta := range BuiltInEvents { - var event model.PushEvent - err := db.DB(ctx).Where("event_key = ?", meta.Key).First(&event).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - var defaultTemplateStr string - if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { - defaultTemplateStr = string(defaultTemplateBytes) - } - event = model.PushEvent{ - EventKey: meta.Key, - Name: meta.Name, - Channels: []string{}, - Targets: []string{}, - Template: defaultTemplateStr, - Enabled: false, - } - if err := db.DB(ctx).Create(&event).Error; err != nil { - return err - } - } else if err != nil { - return err - } - } - return nil + return syncBuiltInEvents(ctx) } // ListEvents 获取通知事件列表 @@ -73,9 +49,8 @@ func SyncEvents(ctx context.Context) error { // @Router /api/v1/admin/push/events [get] func ListEvents(c *gin.Context) { ctx := c.Request.Context() - - var events []model.PushEvent - if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil { + events, err := listPushEvents(ctx) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -113,46 +88,6 @@ func ListBuiltInEvents(c *gin.Context) { c.JSON(http.StatusOK, response.OK(BuiltInEvents)) } -func getEventInfo(req CreateEventRequest) (string, string, []byte, error) { - if req.TaskType != "" { - // 1. 检查关联任务是否存在 - meta := task.GetTaskMetaByAsynqTask(req.TaskType) - if meta == nil { - return "", "", nil, errors.New("unsupported task type") - } - eventKey := "task_completed:" + req.TaskType - eventName := "任务完成: " + meta.Name - - defaultTemplate := NotificationMessage{ - Title: "任务完成: " + meta.Name, - Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", - Level: defaultLevelInfo, - } - defaultTemplateBytes, err := json.Marshal(defaultTemplate) - if err != nil { - return "", "", nil, err - } - return eventKey, eventName, defaultTemplateBytes, nil - } - - if req.EventKey == "" { - return "", "", nil, errors.New("either event_key or task_type must be provided") - } - - // 1. 检查内置事件是否存在 - meta, found := findBuiltInEvent(req.EventKey) - if !found { - return "", "", nil, errors.New("unsupported built-in event key") - } - - defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) - if err != nil { - return "", "", nil, err - } - - return req.EventKey, meta.Name, defaultTemplateBytes, nil -} - // CreateEvent 创建通知事件 // @Summary 创建通知事件 // @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限 @@ -170,70 +105,11 @@ func CreateEvent(c *gin.Context) { return } - ctx := c.Request.Context() - - eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req) + event, err := createPushEvent(c.Request.Context(), req) if err != nil { response.AbortBadRequest(c, err.Error()) return } - - // 2. 检查是否已经创建过该事件的配置 - var count int64 - if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", eventKey).Count(&count).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - if count > 0 { - response.AbortBadRequest(c, "this notification event is already configured") - return - } - - // 3. 模板处理 - templateStr := strings.TrimSpace(req.Template) - if templateStr == "" { - templateStr = string(defaultTemplateBytes) - } else { - var tempMap map[string]any - if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { - response.AbortBadRequest(c, "custom template is not a valid JSON format") - return - } - } - - // 4. 创建事件记录 - channels := req.Channels - if channels == nil { - channels = []string{} - } - targets := req.Targets - if targets == nil { - targets = []string{} - } - - event := model.PushEvent{ - EventKey: eventKey, - Name: eventName, - TaskType: req.TaskType, - Channels: channels, - Targets: targets, - Template: templateStr, - Enabled: req.Enabled, - } - - if err := event.Validate(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := db.DB(ctx).Create(&event).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除旧事件缓存 - model.DeleteActivePushEventCache(ctx, event.EventKey) - c.JSON(http.StatusOK, response.OK(event)) } @@ -247,32 +123,20 @@ func CreateEvent(c *gin.Context) { // @Success 200 {object} response.Any{data=string} "删除成功" // @Router /api/v1/admin/push/events/{id} [delete] func DeleteEvent(c *gin.Context) { - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) + id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { response.AbortBadRequest(c, "invalid event id") return } - ctx := c.Request.Context() - var event model.PushEvent - if err := db.DB(ctx).First(&event, id).Error; err != nil { + if err := deletePushEvent(c.Request.Context(), id); err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, "notification event not found") - } else { - response.AbortInternal(c, err.Error()) + return } - return - } - - if err := db.DB(ctx).Delete(&event).Error; err != nil { response.AbortInternal(c, err.Error()) return } - - // 缓存一致性:清除事件缓存 - model.DeleteActivePushEventCache(ctx, event.EventKey) - c.JSON(http.StatusOK, response.OKNil()) } @@ -288,8 +152,7 @@ func DeleteEvent(c *gin.Context) { // @Success 200 {object} response.Any{data=string} "修改成功" // @Router /api/v1/admin/push/events/{id} [put] func UpdateEvent(c *gin.Context) { - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) + id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { response.AbortBadRequest(c, "invalid event id") return @@ -301,34 +164,14 @@ func UpdateEvent(c *gin.Context) { return } - var event model.PushEvent - if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil { + if err := updatePushEvent(c.Request.Context(), id, req); err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, "notification event not found") - } else { - response.AbortInternal(c, err.Error()) + return } - return - } - - event.Channels = req.Channels - event.Targets = req.Targets - event.Template = req.Template - event.Enabled = req.Enabled - - if err := event.Validate(); err != nil { response.AbortBadRequest(c, err.Error()) return } - - if err := db.DB(c.Request.Context()).Save(&event).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除事件缓存 - model.DeleteActivePushEventCache(c.Request.Context(), event.EventKey) - c.JSON(http.StatusOK, response.OKNil()) } @@ -342,37 +185,22 @@ func UpdateEvent(c *gin.Context) { // @Success 200 {object} response.Any{data=string} "切换成功" // @Router /api/v1/admin/push/events/{id}/toggle [post] func ToggleEvent(c *gin.Context) { - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) + id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { response.AbortBadRequest(c, "invalid event id") return } - var event model.PushEvent - if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil { + enabled, err := togglePushEvent(c.Request.Context(), id) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, "notification event not found") - } else { - response.AbortInternal(c, err.Error()) + return } + response.AbortBadRequest(c, err.Error()) return } - - event.Enabled = !event.Enabled - if event.Enabled && len(event.Channels) == 0 { - response.AbortBadRequest(c, "cannot enable event without any push channels configured") - return - } - if err := db.DB(c.Request.Context()).Model(&event).Update("enabled", event.Enabled).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - // 缓存一致性:清除事件缓存 - model.DeleteActivePushEventCache(c.Request.Context(), event.EventKey) - - c.JSON(http.StatusOK, response.OK(event.Enabled)) + c.JSON(http.StatusOK, response.OK(enabled)) } // pushHistoriesResponse 推送历史分页响应 @@ -396,37 +224,22 @@ type pushHistoriesResponse struct { // @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表" // @Router /api/v1/admin/push/histories [get] func ListHistories(c *gin.Context) { - pageStr := c.DefaultQuery("page", "1") - pageSizeStr := c.DefaultQuery("page_size", "20") - eventKey := c.Query("event_key") - status := c.Query("status") - - page, err := strconv.Atoi(pageStr) - if err != nil || page < 1 { + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) + if page < 1 { page = 1 } - pageSize, err := strconv.Atoi(pageSizeStr) - if err != nil || pageSize < 1 { + if pageSize < 1 { pageSize = 20 } - query := db.DB(c.Request.Context()).Model(&model.PushHistory{}).Order("created_at DESC") - if eventKey != "" { - query = query.Where("event_key = ?", eventKey) - } - if status != "" { - query = query.Where("status = ?", status) - } - - var total int64 - if err := query.Count(&total).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - var results []model.PushHistory - offset := (page - 1) * pageSize - if err := query.Offset(offset).Limit(pageSize).Find(&results).Error; err != nil { + total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{ + EventKey: c.Query("event_key"), + Status: c.Query("status"), + Page: page, + PageSize: pageSize, + }) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -459,44 +272,21 @@ func TestPush(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } - - // 校验配置 if err := pusher.ValidateConfig(req.Config); err != nil { response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err)) return } - // 邮件渠道需要从系统设置中拉取发件人 SMTP 信息做测试 (除非配了独立的) - if req.Config.Channel == channelEmail && (req.Config.URL == "" || req.Config.Key == "") { - var smtpHost, smtpPort, smtpUser, smtpPass model.SystemConfig - ctx := c.Request.Context() - _ = smtpHost.GetByKey(ctx, model.ConfigKeySMTPHost) - _ = smtpPort.GetByKey(ctx, model.ConfigKeySMTPPort) - _ = smtpUser.GetByKey(ctx, model.ConfigKeySMTPUsername) - _ = smtpPass.GetByKey(ctx, model.ConfigKeySMTPPassword) - - if smtpHost.Value != "" && smtpUser.Value != "" { - port := smtpPort.Value - if port == "" { - port = "587" - } - req.Config.URL = smtpHost.Value + ":" + port - req.Config.Key = smtpUser.Value - req.Config.Secret = smtpPass.Value - } - } + applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config) testBody := map[string]any{ keyTitle: "测试通道推送", keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", keyLevel: defaultLevelInfo, } - - err = pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil) - if err != nil { + if err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil { response.AbortBadRequest(c, err.Error()) return } - c.JSON(http.StatusOK, response.OKNil()) -} +} \ No newline at end of file diff --git a/internal/apps/admin/push/task_listener.go b/internal/apps/admin/push/task_listener.go index 5dafd187..657d28e1 100644 --- a/internal/apps/admin/push/task_listener.go +++ b/internal/apps/admin/push/task_listener.go @@ -9,7 +9,6 @@ import ( "strconv" "time" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -20,21 +19,16 @@ func RegisterTaskListeners() { task.OnTaskCompleted(handleTaskCompleted) } -// handleTaskCompleted handles task completions and triggers appropriate push events. func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) { - // Query all active push events configured for this task type - var events []model.PushEvent - err := db.DB(ctx).Where("task_type = ? AND enabled = ?", execution.TaskType, true).Find(&events).Error + events, err := listActivePushEventsByTaskType(ctx, execution.TaskType) if err != nil { logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err) return } - if len(events) == 0 { return } - // Build the notification body context body := map[string]any{ "task_id": execution.TaskID, "task_name": execution.TaskName, @@ -43,20 +37,17 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re "task_duration": execution.Duration, "time": time.Now().Format("2006-01-02 15:04:05"), } - if execErr != nil { body["task_error"] = execErr.Error() } else { body["task_error"] = "" } - if result != nil { body["task_result"] = result.Message } else { body["task_result"] = "" } - // Parse payload parameters if it is valid JSON var payloadMap map[string]any if execution.Payload != "" { if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil { @@ -64,8 +55,6 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re extractUserFromMap(ctx, payloadMap, body) } } - - // Parse result detail parameters if it is valid JSON if result != nil && result.Detail != "" { var detailMap map[string]any if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil { @@ -74,7 +63,6 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re } } - // Trigger notifications for all configured events for _, event := range events { meta := EventMetadata{ Key: event.EventKey, @@ -85,35 +73,15 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re } } -// extractUserFromMap tries to find user information from a map and load the full User model. func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) { if u, exists := body["user"]; exists && u != nil { return } - - if uVal, ok := data["user"]; ok && uVal != nil { - body["user"] = uVal - return - } - - if userID, ok := extractUserID(data); ok && userID > 0 { - var u model.User - if err := db.DB(ctx).Where("id = ?", userID).First(&u).Error; err == nil { - body["user"] = &u - return - } - } - - if username := extractUsername(data); username != "" { - var u model.User - if err := db.DB(ctx).Where("username = ?", username).First(&u).Error; err == nil { - body["user"] = &u - return - } + if user := loadUserFromPayload(ctx, data); user != nil { + body["user"] = user } } -// extractUserID extracts and validates a userID from map fields. func extractUserID(data map[string]any) (uint64, bool) { for _, k := range []string{"user_id", "userId", "uid"} { val, ok := data[k] @@ -144,7 +112,6 @@ func extractUserID(data map[string]any) (uint64, bool) { return 0, false } -// extractUsername extracts a username string from map fields. func extractUsername(data map[string]any) string { for _, k := range []string{"username", "user_name"} { if val, ok := data[k]; ok && val != nil { @@ -154,4 +121,4 @@ func extractUsername(data map[string]any) string { } } return "" -} +} \ No newline at end of file diff --git a/internal/apps/admin/push/tasks.go b/internal/apps/admin/push/tasks.go index ec13a038..235bf673 100644 --- a/internal/apps/admin/push/tasks.go +++ b/internal/apps/admin/push/tasks.go @@ -9,10 +9,7 @@ import ( "encoding/json" "errors" "fmt" - "time" - "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/pkg/push" ) @@ -118,46 +115,7 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskRe } func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) { - title := req.Body.Title - content := req.Body.Content - level := req.Body.Level - - if title == "" { - title = "系统通知" - } - if level == "" { - level = defaultLevelInfo - } - - target := req.Target - if target == "" { - // 如果目标人为空 (例如 webhook bot),用其地址填充前缀或默认词作为归档 - if req.Config.URL != "" { - target = req.Config.URL - // 隐藏敏感 URL 细节 - //nolint:mnd - if len(target) > 50 { - target = target[:47] + "..." - } - } else { - target = "default" - } - } - - history := model.PushHistory{ - EventKey: req.EventKey, - Channel: req.Config.Channel, - Target: target, - Title: title, - Content: content, - Level: level, - Status: status, - ErrorMsg: errMsg, - CreatedAt: time.Now(), - } - - // 记录到数据库 - if dbErr := db.DB(ctx).Create(&history).Error; dbErr != nil { + if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil { task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr) } } diff --git a/internal/apps/admin/system_config/logics.go b/internal/apps/admin/system_config/logics.go new file mode 100644 index 00000000..04630e82 --- /dev/null +++ b/internal/apps/admin/system_config/logics.go @@ -0,0 +1,128 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package system_config + +import ( + "context" + "encoding/json" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/storage" + "github.com/Rain-kl/Wavelet/pkg/logger" + "gorm.io/gorm" +) + +func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error { + exists, err := repository.SystemConfigExists(ctx, req.Key) + if err != nil { + return err + } + if exists { + return errors.New(ConfigKeyExists) + } + + config := model.SystemConfig{ + Key: req.Key, + Value: req.Value, + Type: req.Type, + Visibility: req.Visibility, + Description: req.Description, + } + if err := repository.CreateSystemConfig(ctx, &config); err != nil { + return err + } + + invalidateSystemConfigCaches(ctx, req.Key) + if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { + logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err) + } + return nil +} + +func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) { + return repository.ListAdminSystemConfigs(ctx, configType) +} + +func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) { + return repository.GetAdminSystemConfigByKey(ctx, key) +} + +func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error { + config, err := repository.GetAdminSystemConfigByKey(ctx, key) + if err != nil { + return err + } + + var originalDriver storage.Driver + if key == model.ConfigKeyStorageConfig { + var currentCfg storage.Config + if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil { + originalDriver = currentCfg.Driver + } + + validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value) + if err != nil { + return err + } + req.Value = validatedVal + } + + if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + updates := map[string]any{ + "description": req.Description, + } + if req.Visibility != nil { + updates["visibility"] = *req.Visibility + config.Visibility = *req.Visibility + } + if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue { + updates["value"] = req.Value + config.Value = req.Value + } + if err := tx.Model(&config).Updates(updates).Error; err != nil { + return err + } + resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value) + return nil + }); err != nil { + return err + } + + invalidateCachesAfterConfigUpdate(ctx, key) + return nil +} + +func resolveStorageMigrationTasksOnDirectDriverUpdate( + ctx context.Context, + tx *gorm.DB, + key string, + originalDriver storage.Driver, + newValue string, +) { + if key != model.ConfigKeyStorageConfig || originalDriver == "" { + return + } + + var newCfg storage.Config + if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil { + return + } + if newCfg.Driver != originalDriver { + return + } + + if err := tx.Model(&model.TaskExecution{}). + Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed). + Updates(map[string]any{ + "status": model.TaskExecutionStatusSucceeded, + "result": "存储配置直接更新,故障迁移任务自动标记为已解决", + "finished_at": time.Now(), + }).Error; err != nil { + logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err) + } +} \ No newline at end of file diff --git a/internal/apps/admin/system_config/routers.go b/internal/apps/admin/system_config/routers.go index 89f15096..ed92c7a9 100644 --- a/internal/apps/admin/system_config/routers.go +++ b/internal/apps/admin/system_config/routers.go @@ -10,7 +10,7 @@ import ( "errors" "fmt" "net/http" - "time" + "strings" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -18,8 +18,8 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/cap" "github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/common/response" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/pkg/logger" mail "github.com/Rain-kl/Wavelet/pkg/mail" @@ -64,42 +64,15 @@ func CreateSystemConfig(c *gin.Context) { return } - // 检查配置键是否已存在 - var existing model.SystemConfig - if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil { - response.AbortBadRequest(c, ConfigKeyExists) - return - } else if !errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortInternal(c, err.Error()) - return - } - - config := model.SystemConfig{ - Key: req.Key, - Value: req.Value, - Type: req.Type, - Visibility: req.Visibility, - Description: req.Description, - } - - if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error { - // 创建配置 - if err := tx.Create(&config).Error; err != nil { - return err + if err := createSystemConfig(c.Request.Context(), req); err != nil { + if err.Error() == ConfigKeyExists { + response.AbortBadRequest(c, ConfigKeyExists) + return } - - return nil - }); err != nil { response.AbortInternal(c, err.Error()) return } - invalidateSystemConfigCaches(c.Request.Context(), req.Key) - - if err := model.InvalidateVisibleSystemConfigsCache(c.Request.Context()); err != nil { - logger.WarnF(c.Request.Context(), "清理公共配置列表缓存失败: %v", err) - } - c.JSON(http.StatusOK, response.OKNil()) } @@ -116,14 +89,8 @@ func CreateSystemConfig(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/system-configs [get] func ListSystemConfigs(c *gin.Context) { - configType := c.Query("type") - query := db.DB(c.Request.Context()).Order("created_at DESC") - if configType != "" { - query = query.Where("type = ?", configType) - } - - var configs []model.SystemConfig - if err := query.Find(&configs).Error; err != nil { + configs, err := listSystemConfigs(c.Request.Context(), c.Query("type")) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -149,8 +116,8 @@ func ListSystemConfigs(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/system-configs/{key} [get] func GetSystemConfig(c *gin.Context) { - var config model.SystemConfig - if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&config).Error; err != nil { + config, err := getSystemConfig(c.Request.Context(), c.Param("key")) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, SystemConfigNotFound) } else { @@ -188,101 +155,24 @@ func UpdateSystemConfig(c *gin.Context) { } key := c.Param("key") - - // 检查配置是否存在 - var config model.SystemConfig - if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil { + if err := updateSystemConfig(c.Request.Context(), key, req); err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { response.AbortNotFound(c, SystemConfigNotFound) - } else { - response.AbortInternal(c, err.Error()) + return } - return - } - - var originalDriver storage.Driver - if key == model.ConfigKeyStorageConfig { - var currentCfg storage.Config - if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil { - originalDriver = currentCfg.Driver - } - - validatedVal, err := validateAndMergeStorageConfig(c.Request.Context(), req.Value, config.Value) - if err != nil { + if isStorageConfigValidationError(err) { response.AbortBadRequest(c, err.Error()) return } - req.Value = validatedVal - } - - if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error { - // 更新配置 - updates := map[string]interface{}{ - "description": req.Description, - } - if req.Visibility != nil { - updates["visibility"] = *req.Visibility - config.Visibility = *req.Visibility - } - if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue { - updates["value"] = req.Value - config.Value = req.Value - } - if err := tx.Model(&config).Updates(updates).Error; err != nil { - return err - } - - resolveStorageMigrationTasksOnDirectDriverUpdate( - c.Request.Context(), - tx, - key, - originalDriver, - req.Value, - ) - - return nil - }); err != nil { response.AbortInternal(c, err.Error()) return } - invalidateCachesAfterConfigUpdate(c.Request.Context(), key) - c.JSON(http.StatusOK, response.OKNil()) } -func resolveStorageMigrationTasksOnDirectDriverUpdate( - ctx context.Context, - tx *gorm.DB, - key string, - originalDriver storage.Driver, - newValue string, -) { - if key != model.ConfigKeyStorageConfig || originalDriver == "" { - return - } - - var newCfg storage.Config - if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil { - return - } - if newCfg.Driver != originalDriver { - return - } - - if err := tx.Model(&model.TaskExecution{}). - Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed). - Updates(map[string]any{ - "status": model.TaskExecutionStatusSucceeded, - "result": "存储配置直接更新,故障迁移任务自动标记为已解决", - "finished_at": time.Now(), - }).Error; err != nil { - logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err) - } -} - func invalidateSystemConfigCaches(ctx context.Context, key string) { - if err := model.InvalidateSystemConfigCache(ctx, key); err != nil { + if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil { logger.WarnF(ctx, "清理系统配置缓存失败: %v", err) } if cap.IsRuntimeConfigKey(key) { @@ -304,7 +194,7 @@ func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) { upload.PublishAccessCacheInvalidation(ctx) } - if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil { + if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err) } } @@ -345,8 +235,7 @@ func TestSMTP(c *gin.Context) { password := req.SMTPPassword if password == maskedConfigValue { - var sc model.SystemConfig - if err := sc.GetByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil { + if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil { password = sc.Value } } @@ -375,6 +264,15 @@ func TestSMTP(c *gin.Context) { c.JSON(http.StatusOK, response.OK(resp)) } +func isStorageConfigValidationError(err error) bool { + msg := err.Error() + return strings.HasPrefix(msg, "解析") || + strings.HasPrefix(msg, "验证") || + strings.HasPrefix(msg, "初始化测试") || + strings.HasPrefix(msg, "存储连通性") || + strings.HasPrefix(msg, "序列化") +} + func maskSensitiveConfig(key, value string) string { if value == "" { return value diff --git a/internal/apps/admin/system_config/routers_test.go b/internal/apps/admin/system_config/routers_test.go index c47bc7de..fd423736 100644 --- a/internal/apps/admin/system_config/routers_test.go +++ b/internal/apps/admin/system_config/routers_test.go @@ -4,7 +4,8 @@ package system_config -import ("bufio" +import ( + "bufio" "bytes" "context" "encoding/json" @@ -17,24 +18,24 @@ import ("bufio" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) const expectedDefaultConfigsCount = 30 func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() }) @@ -86,22 +87,22 @@ func TestCreateSystemConfig(t *testing.T) { // Verify caches are invalidated after create and repopulate on read _, err = db.Redis.HGet( context.Background(), - db.PrefixedKey(model.SystemConfigRedisHashKey), + db.PrefixedKey(repository.SystemConfigRedisHashKey), "custom_key", ).Result() if err == nil { t.Fatal("expected redis cache miss immediately after create") } - var loaded model.SystemConfig - if err := loaded.GetByKey(context.Background(), "custom_key"); err != nil { - t.Fatalf("GetByKey(custom_key) error = %v", err) + loaded, err := repository.GetSystemConfigByKey(context.Background(), "custom_key") + if err != nil { + t.Fatalf("GetSystemConfigByKey(custom_key) error = %v", err) } if loaded.Value != "custom_value" { - t.Errorf("GetByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value") + t.Errorf("GetSystemConfigByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value") } if loaded.Visibility != model.ConfigVisibilityVisible { - t.Errorf("GetByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible) + t.Errorf("GetSystemConfigByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible) } }) @@ -245,22 +246,22 @@ func TestUpdateSystemConfig(t *testing.T) { // Verify caches are invalidated after update and repopulate on read _, err := db.Redis.HGet( context.Background(), - db.PrefixedKey(model.SystemConfigRedisHashKey), + db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName, ).Result() if err == nil { t.Fatal("expected redis cache miss immediately after update") } - var loaded model.SystemConfig - if err := loaded.GetByKey(context.Background(), model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) error = %v", err) + loaded, err := repository.GetSystemConfigByKey(context.Background(), model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err) } if loaded.Value != "Super Site Name" { - t.Errorf("GetByKey(site_name).Value = %q, want %q", loaded.Value, "Super Site Name") + t.Errorf("GetSystemConfigByKey(site_name).Value = %q, want %q", loaded.Value, "Super Site Name") } if loaded.Visibility != model.ConfigVisibilityHidden { - t.Errorf("GetByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden) + t.Errorf("GetSystemConfigByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden) } }) diff --git a/internal/apps/admin/task/routers_test.go b/internal/apps/admin/task/routers_test.go index 5c118ff4..a0465777 100644 --- a/internal/apps/admin/task/routers_test.go +++ b/internal/apps/admin/task/routers_test.go @@ -20,7 +20,6 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "github.com/hibiken/asynq" "github.com/stretchr/testify/assert" @@ -51,7 +50,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine { // Mock authentication middleware adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() }) diff --git a/internal/apps/admin/template/logics.go b/internal/apps/admin/template/logics.go new file mode 100644 index 00000000..b94b4670 --- /dev/null +++ b/internal/apps/admin/template/logics.go @@ -0,0 +1,78 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package template + +import ( + "context" + "errors" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" +) + +func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) { + exists, err := repository.TemplateExistsByKey(ctx, req.Key) + if err != nil { + return model.Template{}, err + } + if exists { + return model.Template{}, errors.New(TemplateKeyExists) + } + + tmpl := model.Template{ + Key: req.Key, + Name: req.Name, + Type: req.Type, + Subject: req.Subject, + Content: req.Content, + Description: req.Description, + IsSystem: false, + } + if err := tmpl.Validate(); err != nil { + return model.Template{}, err + } + if err := repository.CreateTemplate(ctx, &tmpl); err != nil { + return model.Template{}, err + } + return tmpl, nil +} + +func listTemplates(ctx context.Context) ([]model.Template, error) { + return repository.ListTemplates(ctx) +} + +func getTemplate(ctx context.Context, key string) (model.Template, error) { + return repository.GetTemplateByKey(ctx, key) +} + +func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) { + tmpl, err := repository.GetTemplateByKey(ctx, key) + if err != nil { + return model.Template{}, err + } + + tmpl.Name = req.Name + tmpl.Type = req.Type + tmpl.Subject = req.Subject + tmpl.Content = req.Content + tmpl.Description = req.Description + if err := tmpl.Validate(); err != nil { + return model.Template{}, err + } + if err := repository.SaveTemplate(ctx, &tmpl); err != nil { + return model.Template{}, err + } + return tmpl, nil +} + +func deleteTemplate(ctx context.Context, key string) error { + tmpl, err := repository.GetTemplateByKey(ctx, key) + if err != nil { + return err + } + if tmpl.IsSystem { + return errors.New(SystemTemplateCannotDelete) + } + return repository.DeleteTemplate(ctx, &tmpl) +} \ No newline at end of file diff --git a/internal/apps/admin/template/routers.go b/internal/apps/admin/template/routers.go index 09d0dc0d..a07f7675 100644 --- a/internal/apps/admin/template/routers.go +++ b/internal/apps/admin/template/routers.go @@ -3,15 +3,15 @@ package template -import ("errors" +import ( + "errors" "net/http" - "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/model" "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) // CreateTemplateRequest 创建模板请求 type CreateTemplateRequest struct { @@ -32,6 +32,24 @@ type UpdateTemplateRequest struct { Description string `json:"description" binding:"max=255"` } +func abortTemplateLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + response.AbortNotFound(c, TemplateNotFound) + return true + } + msg := err.Error() + switch msg { + case TemplateKeyExists, SystemTemplateCannotDelete: + response.AbortBadRequest(c, msg) + return true + } + response.AbortInternal(c, msg) + return true +} + // CreateTemplate 创建模板 // @Summary 创建模板 // @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限 @@ -53,33 +71,8 @@ func CreateTemplate(c *gin.Context) { return } - // 检查模板 Key 是否已存在 - var existing model.Template - if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil { - response.AbortBadRequest(c, TemplateKeyExists) - return - } else if !errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortInternal(c, err.Error()) - return - } - - tmpl := model.Template{ - Key: req.Key, - Name: req.Name, - Type: req.Type, - Subject: req.Subject, - Content: req.Content, - Description: req.Description, - IsSystem: false, - } - - if err := tmpl.Validate(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := db.DB(c.Request.Context()).Create(&tmpl).Error; err != nil { - response.AbortInternal(c, err.Error()) + tmpl, err := createTemplate(c.Request.Context(), req) + if abortTemplateLogicError(c, err) { return } @@ -98,8 +91,8 @@ func CreateTemplate(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/templates [get] func ListTemplates(c *gin.Context) { - var templates []model.Template - if err := db.DB(c.Request.Context()).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil { + templates, err := listTemplates(c.Request.Context()) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -121,13 +114,8 @@ func ListTemplates(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/templates/{key} [get] func GetTemplate(c *gin.Context) { - var tmpl model.Template - if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&tmpl).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, TemplateNotFound) - } else { - response.AbortInternal(c, err.Error()) - } + tmpl, err := getTemplate(c.Request.Context(), c.Param("key")) + if abortTemplateLogicError(c, err) { return } @@ -157,32 +145,8 @@ func UpdateTemplate(c *gin.Context) { return } - key := c.Param("key") - - // 检查模板是否存在 - var tmpl model.Template - if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, TemplateNotFound) - } else { - response.AbortInternal(c, err.Error()) - } - return - } - - tmpl.Name = req.Name - tmpl.Type = req.Type - tmpl.Subject = req.Subject - tmpl.Content = req.Content - tmpl.Description = req.Description - - if err := tmpl.Validate(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := db.DB(c.Request.Context()).Save(&tmpl).Error; err != nil { - response.AbortInternal(c, err.Error()) + tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req) + if abortTemplateLogicError(c, err) { return } @@ -204,29 +168,9 @@ func UpdateTemplate(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/templates/{key} [delete] func DeleteTemplate(c *gin.Context) { - key := c.Param("key") - - // 检查模板是否存在 - var tmpl model.Template - if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, TemplateNotFound) - } else { - response.AbortInternal(c, err.Error()) - } - return - } - - // 限制系统模板删除 - if tmpl.IsSystem { - response.AbortBadRequest(c, SystemTemplateCannotDelete) - return - } - - if err := db.DB(c.Request.Context()).Delete(&tmpl).Error; err != nil { - response.AbortInternal(c, err.Error()) + if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) { return } c.JSON(http.StatusOK, response.OKNil()) -} +} \ No newline at end of file diff --git a/internal/apps/admin/template/routers_test.go b/internal/apps/admin/template/routers_test.go index b78f72c9..bf669ac8 100644 --- a/internal/apps/admin/template/routers_test.go +++ b/internal/apps/admin/template/routers_test.go @@ -12,20 +12,18 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "github.com/Rain-kl/Wavelet/internal/common/response") func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() }) diff --git a/internal/apps/admin/updater/logics.go b/internal/apps/admin/updater/logics.go index 9a48238e..8a5ac9d4 100644 --- a/internal/apps/admin/updater/logics.go +++ b/internal/apps/admin/updater/logics.go @@ -23,6 +23,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/buildinfo" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "golang.org/x/mod/semver" ) @@ -225,19 +226,19 @@ func (m *manager) fetchRelease(ctx context.Context, repository string) (githubRe } func loadRepository(ctx context.Context) (string, error) { - var config model.SystemConfig - if err := config.GetByKey(ctx, model.ConfigKeyUpdateUpstreamRepository); err != nil { + config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository) + if err != nil { return "", fmt.Errorf("%s: %w", errInvalidRepository, err) } return parseRepository(config.Value) } func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) { - repository, err := loadRepository(ctx) + upstreamRepo, err := loadRepository(ctx) if err != nil { return Status{}, releaseAsset{}, err } - release, asset, err := m.fetchRelease(ctx, repository) + release, asset, err := m.fetchRelease(ctx, upstreamRepo) if err != nil { return Status{}, releaseAsset{}, err } @@ -259,7 +260,7 @@ func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) { ReleaseNotes: release.Body, ReleaseURL: release.HTMLURL, PublishedAt: release.Published.Format(time.RFC3339), - UpstreamRepository: repository, + UpstreamRepository: upstreamRepo, AssetName: asset.Name, Platform: runtime.GOOS + "/" + runtime.GOARCH, }, asset, nil diff --git a/internal/apps/admin/user/logics.go b/internal/apps/admin/user/logics.go new file mode 100644 index 00000000..fb263923 --- /dev/null +++ b/internal/apps/admin/user/logics.go @@ -0,0 +1,106 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db/idgen" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" +) + +func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) { + return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{ + UserID: req.UserID, + Username: strings.TrimSpace(req.Username), + Page: req.Page, + PageSize: req.PageSize, + }) +} + +func getUserDetail(ctx context.Context, id uint64) (model.User, error) { + return repository.GetAdminUserDetail(ctx, id) +} + +func updateUserStatus(ctx context.Context, id uint64, active bool) error { + flags, err := repository.GetUserAdminFlags(ctx, id) + if err != nil { + return err + } + if !active && flags.IsAdmin { + return errors.New(cannotDisable) + } + return repository.UpdateUserActive(ctx, id, active) +} + +func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { + if currentUserID == targetID { + return errors.New(cannotDeleteSelf) + } + flags, err := repository.GetUserAdminFlags(ctx, targetID) + if err != nil { + return err + } + if flags.IsAdmin { + return errors.New(cannotDelete) + } + return repository.DeleteUserWithRelations(ctx, targetID) +} + +func createUser(ctx context.Context, req createUserRequest) (model.User, error) { + req.Username = strings.TrimSpace(req.Username) + req.Nickname = strings.TrimSpace(req.Nickname) + req.Password = strings.TrimSpace(req.Password) + req.Email = strings.TrimSpace(req.Email) + + if req.Username == "" { + return model.User{}, errors.New(usernameRequired) + } + if req.Email == "" { + return model.User{}, errors.New(emailRequired) + } + if len(req.Password) < minPasswordLength { + return model.User{}, errors.New(passwordTooShort) + } + + count, err := repository.CountUsersByUsername(ctx, req.Username) + if err != nil { + return model.User{}, err + } + if count > 0 { + return model.User{}, errors.New(usernameExists) + } + + emailCount, err := repository.CountUsersByEmail(ctx, req.Email) + if err != nil { + return model.User{}, err + } + if emailCount > 0 { + return model.User{}, errors.New(emailExists) + } + + newUser := model.User{ + ID: idgen.NextUint64ID(), + Username: req.Username, + Nickname: req.Nickname, + Email: req.Email, + IsActive: req.IsActive, + IsAdmin: req.IsAdmin, + LastLoginAt: time.Time{}, + } + if newUser.Nickname == "" { + newUser.Nickname = req.Username + } + if err := newUser.SetEncryptedPassword(req.Password); err != nil { + return model.User{}, err + } + if err := repository.CreateUser(ctx, &newUser); err != nil { + return model.User{}, err + } + return newUser, nil +} \ No newline at end of file diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index 509882d9..36acff97 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -5,16 +5,13 @@ package user import ( + "errors" "net/http" "strconv" - "strings" "time" "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -85,6 +82,31 @@ func toUser(u model.User) user { } } +func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + response.AbortNotFound(c, notFoundMsg) + return true + } + msg := err.Error() + for _, m := range badRequestMsgs { + if msg == m { + response.AbortBadRequest(c, msg) + return true + } + } + for _, m := range forbiddenMsgs { + if msg == m { + response.AbortForbidden(c, msg) + return true + } + } + response.AbortInternal(c, msg) + return true +} + // ListUsers 获取用户列表 // @Summary 获取用户列表 // @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限 @@ -98,7 +120,6 @@ func toUser(u model.User) user { // @Failure 403 {object} response.Any "无管理员权限" // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/users [get] -// ListUsers 获取用户列表 func ListUsers(c *gin.Context) { var req listUsersRequest if err := c.ShouldBindQuery(&req); err != nil { @@ -106,34 +127,8 @@ func ListUsers(c *gin.Context) { return } - var modelUsers []model.User - var total int64 - - query := db.DB(c.Request.Context()).Model(&model.User{}) - - username := strings.TrimSpace(req.Username) - - if req.UserID != nil { - query = query.Where("id = ?", *req.UserID) - } - - if username != "" { - query = query.Where("username LIKE ?", username+"%") - } - - if err := query.Count(&total).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - - offset := (req.Page - 1) * req.PageSize - if err := query. - Select("id, username, nickname, avatar_url, is_active, is_admin, " + - "last_login_at, created_at, updated_at"). - Order("id DESC"). - Offset(offset). - Limit(req.PageSize). - Find(&modelUsers).Error; err != nil { + total, modelUsers, err := listUsers(c.Request.Context(), req) + if err != nil { response.AbortInternal(c, err.Error()) return } @@ -169,17 +164,8 @@ func GetUser(c *gin.Context) { return } - var targetUser model.User - if err := db.DB(c.Request.Context()). - Select("id, username, nickname, email, avatar_url, is_active, is_admin, "+ - "bio, phone, gender, website, location, last_login_at, created_at, updated_at"). - Where("id = ?", id). - First(&targetUser).Error; err != nil { - if err == gorm.ErrRecordNotFound { - response.AbortNotFound(c, userNotFound) - return - } - response.AbortInternal(c, err.Error()) + targetUser, err := getUserDetail(c.Request.Context(), id) + if abortUserLogicError(c, err, userNotFound, nil, nil) { return } @@ -219,32 +205,10 @@ func UpdateUserStatus(c *gin.Context) { return } - var targetUser struct { - ID uint64 `gorm:"column:id"` - IsAdmin bool `gorm:"column:is_admin"` - } - if err := db.DB(c.Request.Context()). - Model(&model.User{}). - Select("id, is_admin"). - Where("id = ?", id). - First(&targetUser).Error; err != nil { - if err == gorm.ErrRecordNotFound { - response.AbortNotFound(c, userNotFound) + if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil { + if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) { return } - response.AbortInternal(c, err.Error()) - return - } - - if !req.IsActive && targetUser.IsAdmin { - response.AbortForbidden(c, cannotDisable) - return - } - - if err := db.DB(c.Request.Context()). - Model(&model.User{}). - Where("id = ?", id). - Update("is_active", req.IsActive).Error; err != nil { response.AbortInternal(c, updateUserFailed) return } @@ -272,43 +236,11 @@ func DeleteUser(c *gin.Context) { return } - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) - if currUser != nil && currUser.ID == id { - response.AbortForbidden(c, cannotDeleteSelf) - return - } - - var targetUser struct { - ID uint64 `gorm:"column:id"` - IsAdmin bool `gorm:"column:is_admin"` - } - if err := db.DB(c.Request.Context()). - Model(&model.User{}). - Select("id, is_admin"). - Where("id = ?", id). - First(&targetUser).Error; err != nil { - if err == gorm.ErrRecordNotFound { - response.AbortNotFound(c, userNotFound) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil { + if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) { return } - response.AbortInternal(c, err.Error()) - return - } - - if targetUser.IsAdmin { - response.AbortForbidden(c, cannotDelete) - return - } - - if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("user_id = ?", id).Delete(&model.AccessToken{}).Error; err != nil { - return err - } - if err := tx.Where("user_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil { - return err - } - return tx.Where("id = ?", id).Delete(&model.User{}).Error - }); err != nil { response.AbortInternal(c, deleteUserFailed) return } @@ -347,67 +279,10 @@ func CreateUser(c *gin.Context) { return } - req.Username = strings.TrimSpace(req.Username) - req.Nickname = strings.TrimSpace(req.Nickname) - req.Password = strings.TrimSpace(req.Password) - req.Email = strings.TrimSpace(req.Email) - - if req.Username == "" { - response.AbortBadRequest(c, usernameRequired) - return - } - if req.Email == "" { - response.AbortBadRequest(c, emailRequired) - return - } - if len(req.Password) < minPasswordLength { - response.AbortBadRequest(c, passwordTooShort) - return - } - - ctx := c.Request.Context() - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - if count > 0 { - response.AbortBadRequest(c, usernameExists) - return - } - - var emailCount int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&emailCount).Error; err != nil { - response.AbortInternal(c, err.Error()) - return - } - if emailCount > 0 { - response.AbortBadRequest(c, emailExists) - return - } - - newUser := model.User{ - ID: idgen.NextUint64ID(), - Username: req.Username, - Nickname: req.Nickname, - Email: req.Email, - IsActive: req.IsActive, - IsAdmin: req.IsAdmin, - LastLoginAt: time.Time{}, - } - if newUser.Nickname == "" { - newUser.Nickname = req.Username - } - - if err := newUser.SetEncryptedPassword(req.Password); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - if err := db.DB(ctx).Create(&newUser).Error; err != nil { - response.AbortInternal(c, err.Error()) + newUser, err := createUser(c.Request.Context(), req) + if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) { return } c.JSON(http.StatusOK, response.OK(toUser(newUser))) -} +} \ No newline at end of file diff --git a/internal/apps/admin/user/routers_test.go b/internal/apps/admin/user/routers_test.go index 697959c8..c2e2cc2c 100644 --- a/internal/apps/admin/user/routers_test.go +++ b/internal/apps/admin/user/routers_test.go @@ -14,20 +14,18 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "github.com/Rain-kl/Wavelet/internal/common/response") func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware adminGroup.Use(func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() }) diff --git a/internal/apps/cap/routers_test.go b/internal/apps/cap/routers_test.go index 53798674..b1ee629c 100644 --- a/internal/apps/cap/routers_test.go +++ b/internal/apps/cap/routers_test.go @@ -15,6 +15,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" pkgcap "github.com/Rain-kl/Wavelet/pkg/cap" ) @@ -68,7 +69,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) { if err != nil { t.Fatalf("failed to enable cap_login_enabled in DB: %v", err) } - if err := model.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil { + if err := repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil { t.Fatalf("InvalidateSystemConfigCache() error = %v", err) } InvalidateRuntimeSettings() diff --git a/internal/apps/cap/runtime_settings.go b/internal/apps/cap/runtime_settings.go index 4b920f7d..55966dac 100644 --- a/internal/apps/cap/runtime_settings.go +++ b/internal/apps/cap/runtime_settings.go @@ -16,6 +16,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" ) const ( @@ -130,7 +131,7 @@ func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, er } func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) { - configs, err := model.ListSystemConfigsByKeys(ctx, runtimeConfigKeys) + configs, err := repository.ListSystemConfigsByKeys(ctx, runtimeConfigKeys) if err != nil { return RuntimeSettings{}, err } @@ -190,7 +191,7 @@ func startRuntimeSettingsInvalidationListener() { } go func() { - pubsub := db.Redis.Subscribe(context.Background(), model.SystemConfigInvalidationChannel) + pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel) defer func() { _ = pubsub.Close() }() @@ -208,4 +209,4 @@ func startRuntimeSettingsInvalidationListener() { } } }() -} \ No newline at end of file +} diff --git a/internal/apps/cap/runtime_settings_test.go b/internal/apps/cap/runtime_settings_test.go index fd7d8d3b..f04ca5ee 100644 --- a/internal/apps/cap/runtime_settings_test.go +++ b/internal/apps/cap/runtime_settings_test.go @@ -10,6 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -19,7 +20,7 @@ func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) { ctx := context.Background() ResetRuntimeSettingsForTest() - model.ResetSystemConfigRAMCacheForTest() + repository.ResetSystemConfigRAMCacheForTest() first, err := CurrentSettings(ctx) if err != nil { @@ -34,7 +35,7 @@ func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) { Update("value", "4").Error; err != nil { t.Fatalf("Update(cap_challenge_count) error = %v", err) } - if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil { + if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil { t.Fatalf("InvalidateSystemConfigCache() error = %v", err) } InvalidateRuntimeSettings() @@ -64,7 +65,7 @@ func TestProtectionEnabledReflectsLoginSwitch(t *testing.T) { Update("value", "true").Error; err != nil { t.Fatalf("Update(cap_login_enabled) error = %v", err) } - if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil { + if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil { t.Fatalf("InvalidateSystemConfigCache() error = %v", err) } InvalidateRuntimeSettings() @@ -115,4 +116,4 @@ func TestInstallTestRuntimeSettings(t *testing.T) { if settings.ChallengeCount != 2 { t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, 2) } -} \ No newline at end of file +} diff --git a/internal/apps/config/public_config_cache_test.go b/internal/apps/config/public_config_cache_test.go index bf0a1850..635af846 100644 --- a/internal/apps/config/public_config_cache_test.go +++ b/internal/apps/config/public_config_cache_test.go @@ -9,6 +9,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -17,11 +18,11 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) { defer cleanup() ctx := context.Background() - if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil { + if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err) } - if _, err := model.ListVisibleSystemConfigs(ctx); err != nil { + if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil { t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err) } @@ -35,7 +36,7 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) { t.Fatalf("Create(cache_probe_public_key) error = %v", err) } - cached, err := model.ListVisibleSystemConfigs(ctx) + cached, err := repository.ListVisibleSystemConfigs(ctx) if err != nil { t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err) } @@ -45,7 +46,7 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) { } } - exists, err := db.Redis.Exists(ctx, db.PrefixedKey(model.SystemConfigVisibleListRedisKey)).Result() + exists, err := db.Redis.Exists(ctx, db.PrefixedKey(repository.SystemConfigVisibleListRedisKey)).Result() if err != nil { t.Fatalf("Redis.Exists() error = %v", err) } @@ -53,11 +54,11 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) { t.Fatal("expected visible config list cache key to exist") } - if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil { + if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err) } - refreshed, err := model.ListVisibleSystemConfigs(ctx) + refreshed, err := repository.ListVisibleSystemConfigs(ctx) if err != nil { t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err) } @@ -72,4 +73,4 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) { if !found { t.Fatal("refreshed visible config list should include newly created public config") } -} \ No newline at end of file +} diff --git a/internal/apps/config/routers.go b/internal/apps/config/routers.go index 7226e7dc..8664a549 100644 --- a/internal/apps/config/routers.go +++ b/internal/apps/config/routers.go @@ -5,12 +5,15 @@ // Package config 提供公开配置查询接口 package config -import ("net/http" +import ( + "net/http" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) // GetPublicConfig 获取公共配置 // @Summary 获取公共配置 @@ -22,7 +25,7 @@ import ("net/http" // @Router /api/v1/config/public [get] func GetPublicConfig(c *gin.Context) { ctx := c.Request.Context() - configs, err := model.ListVisibleSystemConfigs(ctx) + configs, err := repository.ListVisibleSystemConfigs(ctx) if err != nil { response.AbortInternal(c, err.Error()) return @@ -45,7 +48,7 @@ func GetPublicConfig(c *gin.Context) { // @Router /robots.txt [get] func GetRobotsTXT(c *gin.Context) { ctx := c.Request.Context() - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled) + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled) content := "User-Agent: *\nDisallow: /\n" if err == nil && enabled { content = "User-Agent: *\nAllow: /\n" diff --git a/internal/apps/config/system_config_cache_test.go b/internal/apps/config/system_config_cache_test.go index 237e4248..f01897f2 100644 --- a/internal/apps/config/system_config_cache_test.go +++ b/internal/apps/config/system_config_cache_test.go @@ -12,6 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -20,17 +21,17 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) { defer cleanup() ctx := context.Background() - model.ResetSystemConfigRAMCacheForTest() - if err := model.InvalidateAllSystemConfigCaches(ctx); err != nil { + repository.ResetSystemConfigRAMCacheForTest() + if err := repository.InvalidateAllSystemConfigCaches(ctx); err != nil { t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err) } - var warm model.SystemConfig - if err := warm.GetByKey(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) warm error = %v", err) + warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err) } if warm.Value != "Wavelet" { - t.Fatalf("GetByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet") + t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet") } if err := dbConn.Model(&model.SystemConfig{}). @@ -38,31 +39,31 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) { Update("value", "ram_probe_value").Error; err != nil { t.Fatalf("Update(site_name) error = %v", err) } - if err := db.HDel(ctx, model.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil { + if err := db.HDel(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil { t.Fatalf("HDel(site_name) error = %v", err) } - var cached model.SystemConfig - if err := cached.GetByKey(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) cached error = %v", err) + cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err) } if cached.Value != "Wavelet" { - t.Fatalf("GetByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet") + t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet") } - if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { + if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err) } - var refreshed model.SystemConfig - if err := refreshed.GetByKey(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) refreshed error = %v", err) + refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err) } if refreshed.Value != "ram_probe_value" { - t.Fatalf("GetByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value") + t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value") } - exists, err := db.Redis.HExists(ctx, db.PrefixedKey(model.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result() + exists, err := db.Redis.HExists(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result() if err != nil { t.Fatalf("HExists(site_name) error = %v", err) } @@ -76,16 +77,17 @@ func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) { defer cleanup() ctx := context.Background() - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) error = %v", err) + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err) } + _ = sc - if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { + if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err) } - _, err := db.Redis.HGet(ctx, db.PrefixedKey(model.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result() + _, err = db.Redis.HGet(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result() if !errors.Is(err, redis.Nil) { t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err) } @@ -96,11 +98,11 @@ func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) { t.Fatalf("Update(site_name) error = %v", err) } - var refreshed model.SystemConfig - if err := refreshed.GetByKey(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("GetByKey(site_name) refreshed error = %v", err) + refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err) } if refreshed.Value != "after_invalidate" { - t.Fatalf("GetByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate") + t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate") } } \ No newline at end of file diff --git a/internal/apps/oauth/auth_source_resolver.go b/internal/apps/oauth/auth_source_resolver.go index 58dc7218..2f157c60 100644 --- a/internal/apps/oauth/auth_source_resolver.go +++ b/internal/apps/oauth/auth_source_resolver.go @@ -9,12 +9,13 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/coreos/go-oidc/v3/oidc" "golang.org/x/oauth2" ) func isOIDCLoginEnabled(ctx context.Context) bool { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) if err != nil { return true } @@ -37,7 +38,7 @@ func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSourc } func activeLoginSources(ctx context.Context) []AuthSourceView { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) if err == nil && !enabled { return nil } @@ -62,8 +63,8 @@ func activeLoginSources(ctx context.Context) []AuthSourceView { } func getFrontendLoginRedirectURL(ctx context.Context) (string, error) { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" { + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) + if err != nil || strings.TrimSpace(sc.Value) == "" { return "", errors.New(errServerAddressMissing) } return strings.TrimRight(sc.Value, "/") + "/login", nil diff --git a/internal/util/context.go b/internal/apps/oauth/gin_context.go similarity index 68% rename from internal/util/context.go rename to internal/apps/oauth/gin_context.go index 5650f4d7..7143a363 100644 --- a/internal/util/context.go +++ b/internal/apps/oauth/gin_context.go @@ -1,13 +1,11 @@ -// Copyright 2025 linux.do // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package util 提供通用工具函数 -package util +package oauth import "github.com/gin-gonic/gin" -// GetFromContext 从上下文获取指定类型的值 +// GetFromContext 从 Gin 请求上下文获取指定类型的值。 func GetFromContext[T any](c *gin.Context, key string) (T, bool) { value, exists := c.Get(key) if !exists { @@ -18,7 +16,7 @@ func GetFromContext[T any](c *gin.Context, key string) (T, bool) { return typed, ok } -// SetToContext 设置值到上下文 +// SetToContext 设置值到 Gin 请求上下文。 func SetToContext[T any](c *gin.Context, key string, value T) { c.Set(key, value) -} +} \ No newline at end of file diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go index fe8584c6..d80ff250 100644 --- a/internal/apps/oauth/handler_callback.go +++ b/internal/apps/oauth/handler_callback.go @@ -15,6 +15,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "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/pkg/logger" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" @@ -188,7 +189,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth // handleCallbackRegister 处理 OAuth 回调中的自动注册流程 // 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) { - registrationEnabled, regErr := model.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) + registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) if regErr != nil { registrationEnabled = true } @@ -223,4 +224,4 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) return user, true -} \ No newline at end of file +} diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 5eacbbda..6ad871fa 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -12,7 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" + otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" @@ -62,8 +62,8 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { if user.Username == "system" { return nil, errors.New("system user is not allowed to login") } - util.SetToContext(c, TokenAuthKey, true) - util.SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) + SetToContext(c, TokenAuthKey, true) + SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) return user, nil } } @@ -91,8 +91,8 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { } // set keys in context for session auth - util.SetToContext(c, TokenAuthKey, false) - util.SetToContext(c, TokenAdminKey, false) + SetToContext(c, TokenAuthKey, false) + SetToContext(c, TokenAdminKey, false) // 强行阻止 system 用户任何会话/Token 鉴权通过 if user.Username == "system" { @@ -119,7 +119,7 @@ func LoginRequired() gin.HandlerFunc { LogForAudit(ctx, user, c) // set user info - util.SetToContext(c, UserObjKey, user) + SetToContext(c, UserObjKey, user) // next c.Next() @@ -129,7 +129,7 @@ func LoginRequired() gin.HandlerFunc { // DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点 func DisallowTokenAuth() gin.HandlerFunc { return func(c *gin.Context) { - if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth { + if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth { response.AbortForbidden(c, ErrTokenAuthNotAllowed) return } diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 6d58fbf1..8d2ea2d2 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -35,6 +35,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/Rain-kl/Wavelet/internal/util" ) @@ -1094,7 +1095,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { Key: model.ConfigKeyOIDCLoginEnabled, Value: "false", }) - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) if wLoginDisabled.Code != http.StatusBadRequest { t.Errorf("expected 400 when OIDC globally disabled, got %d", wLoginDisabled.Code) @@ -1102,7 +1103,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Re-enable globally, but deactivate source dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) @@ -1113,7 +1114,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // --- 2. Test Authorize enforcement --- // Deactivate globally again dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil) @@ -1124,7 +1125,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // --- 3. Test Callback enforcement --- // Set up a valid state beforehand (when enabled) dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) @@ -1149,7 +1150,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Now disable OIDC globally and attempt callback dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ "Content-Type": "application/json", @@ -1160,7 +1161,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Enable globally but deactivate source and attempt callback dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) // Since callback deletes state, we need to generate state again diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go index 6eeba2f3..41fc6ccc 100644 --- a/internal/apps/oauth/routers.go +++ b/internal/apps/oauth/routers.go @@ -10,7 +10,7 @@ import ("net/http" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/internal/common/response") @@ -60,7 +60,7 @@ func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo { // @Router /api/v1/user-info [get] // @Router /api/v1/user/self [get] func UserInfo(c *gin.Context) { - user, _ := util.GetFromContext[*model.User](c, UserObjKey) + user, _ := GetFromContext[*model.User](c, UserObjKey) session := sessions.Default(c) needChange := session.Get("need_change_password") == true diff --git a/internal/apps/oauth/session_context.go b/internal/apps/oauth/session_context.go index 42be0c55..672d514d 100644 --- a/internal/apps/oauth/session_context.go +++ b/internal/apps/oauth/session_context.go @@ -10,6 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" "github.com/google/uuid" @@ -56,7 +57,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro maxAge := config.Config.App.SessionAge isSessionCookie := false - ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) + ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) if err == nil { switch { case ttlHours == -1: diff --git a/internal/apps/risk_control/middleware.go b/internal/apps/risk_control/middleware.go index ee5daf06..2e3e3fe3 100644 --- a/internal/apps/risk_control/middleware.go +++ b/internal/apps/risk_control/middleware.go @@ -14,7 +14,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" ) @@ -39,7 +38,7 @@ func RiskControlMiddleware() gin.HandlerFunc { c.Next() // 3. 后置身份检查:仅记录通过认证的请求 - userObj, exists := util.GetFromContext[*model.User](c, oauth.UserObjKey) + userObj, exists := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) if !exists || userObj == nil { return } diff --git a/internal/apps/risk_control/middleware_test.go b/internal/apps/risk_control/middleware_test.go index 58d1227f..ac82a40e 100644 --- a/internal/apps/risk_control/middleware_test.go +++ b/internal/apps/risk_control/middleware_test.go @@ -14,7 +14,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" ) @@ -51,7 +50,7 @@ func TestRiskControlMiddleware(t *testing.T) { r.Use(func(c *gin.Context) { // Mock authentication middleware placing user in context user := &model.User{ID: 12345} - util.SetToContext(c, oauth.UserObjKey, user) + oauth.SetToContext(c, oauth.UserObjKey, user) c.Next() }) r.Use(RiskControlMiddleware()) diff --git a/internal/apps/upload/cache/access_cache.go b/internal/apps/upload/cache/access_cache.go index eff0681d..7d6e18ea 100644 --- a/internal/apps/upload/cache/access_cache.go +++ b/internal/apps/upload/cache/access_cache.go @@ -15,6 +15,7 @@ import ( uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/storage" ) @@ -112,8 +113,8 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} { } func parseFileAccessWhitelist(ctx context.Context) []string { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyFileAccessWhitelist); err != nil || sc.Value == "" { + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyFileAccessWhitelist) + if err != nil || sc.Value == "" { return []string{shared.DefaultPublicUploadType} } diff --git a/internal/apps/upload/cache/access_cache_test.go b/internal/apps/upload/cache/access_cache_test.go index 1e7fb3d5..80b9180b 100644 --- a/internal/apps/upload/cache/access_cache_test.go +++ b/internal/apps/upload/cache/access_cache_test.go @@ -8,10 +8,11 @@ import ( "testing" "time" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -67,10 +68,10 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) { if err := dbConn.Save(&sc).Error; err != nil { t.Fatalf("save whitelist config: %v", err) } - if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil { + if err := db.HSetJSON(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil { t.Fatalf("refresh whitelist redis cache: %v", err) } - model.ResetSystemConfigRAMCacheForTest() + repository.ResetSystemConfigRAMCacheForTest() ResetAccessCaches() if !IsFilePublic(ctx, "attachment") { @@ -97,4 +98,4 @@ func TestAccessCacheTTLExpires(t *testing.T) { if !IsFilePublic(ctx, "avatar") { t.Fatal("expected whitelist reload after TTL expiration") } -} \ No newline at end of file +} diff --git a/internal/apps/upload/filesrv/file_server.go b/internal/apps/upload/filesrv/file_server.go index d4a3156a..330f85a3 100644 --- a/internal/apps/upload/filesrv/file_server.go +++ b/internal/apps/upload/filesrv/file_server.go @@ -25,7 +25,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" - apputil "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "golang.org/x/sync/singleflight" @@ -282,7 +282,7 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error { var currUser *model.User var err error - if u, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil { + if u, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil { currUser = u } else { currUser, err = oauth.GetUserFromRequest(c) @@ -306,7 +306,7 @@ func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error { } if !cache.IsFilePublic(c.Request.Context(), upload.Type) { - if _, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); !ok { + if _, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); !ok { if _, err := oauth.GetUserFromRequest(c); err != nil { return err } diff --git a/internal/apps/upload/handler/file_management.go b/internal/apps/upload/handler/file_management.go index 2f005c2d..b733db5c 100644 --- a/internal/apps/upload/handler/file_management.go +++ b/internal/apps/upload/handler/file_management.go @@ -4,22 +4,17 @@ package handler import ( - "errors" "net/http" - "sort" "strconv" - "strings" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/common/response" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" - apputil "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/gin-gonic/gin" - "gorm.io/gorm" ) type listFilesRequest struct { @@ -69,31 +64,15 @@ func ListFiles(c *gin.Context) { req.PageSize = 20 } - query := db.DB(ctx).Model(&model.Upload{}). - Where("status != ?", model.UploadStatusDeleted) - - if req.UserID != 0 { - query = query.Where("user_id = ?", req.UserID) - } - if req.Keyword != "" { - query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(req.Keyword)+"%") - } - if req.Type != "" { - query = query.Where("type = ?", req.Type) - } - if req.Extension != "" { - query = query.Where("extension = ?", strings.ToLower(req.Extension)) - } - - var total int64 - if err := query.Count(&total).Error; err != nil { - response.AbortBadRequest(c, shared.ErrQueryFileCountFailed) - return - } - - var items []model.Upload - offset := (req.Page - 1) * req.PageSize - if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil { + total, items, err := listUploadFiles(ctx, repository.UploadListFilter{ + UserID: req.UserID, + Keyword: req.Keyword, + Type: req.Type, + Extension: req.Extension, + Page: req.Page, + PageSize: req.PageSize, + }) + if err != nil { response.AbortBadRequest(c, shared.ErrQueryFileListFailed) return } @@ -130,20 +109,14 @@ func DeleteFile(c *gin.Context) { return } - var upload model.Upload - if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { + if _, err := softDeleteUpload(ctx, uploadID); err != nil { + if isRecordNotFound(err) { c.AbortWithStatus(http.StatusNotFound) return } - response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) - return - } - if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil { response.AbortBadRequest(c, shared.ErrDeleteFileFailed) return } - uploadstats.RecordUploadStatsRemove(ctx, &upload) c.JSON(http.StatusOK, response.OKNil()) } @@ -159,16 +132,12 @@ func DeleteFile(c *gin.Context) { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/uploads/types [get] func GetDistinctUploadTypes(c *gin.Context) { - var dbTypes []string - if err := db.DB(c.Request.Context()).Model(&model.Upload{}). - Where("type IS NOT NULL AND type != ''"). - Distinct(). - Pluck("type", &dbTypes).Error; err != nil { + types, err := listDistinctUploadTypes(c.Request.Context()) + if err != nil { response.AbortInternal(c, err.Error()) return } - sort.Strings(dbTypes) - c.JSON(http.StatusOK, response.OK(dbTypes)) + c.JSON(http.StatusOK, response.OK(types)) } type listMyFilesRequest struct { @@ -201,7 +170,7 @@ type listMyFilesResponse struct { // @Failure 401 {object} response.Any "未登录" // @Router /api/v1/upload/my [get] func ListMyFiles(c *gin.Context) { - currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() var req listMyFilesRequest @@ -216,28 +185,14 @@ func ListMyFiles(c *gin.Context) { req.PageSize = 20 } - query := db.DB(ctx).Model(&model.Upload{}). - Where("user_id = ? AND status != ?", currUser.ID, model.UploadStatusDeleted) - - if req.Keyword != "" { - query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(req.Keyword)+"%") - } - if req.Type != "" { - query = query.Where("type = ?", req.Type) - } - if req.Extension != "" { - query = query.Where("extension = ?", strings.ToLower(req.Extension)) - } - - var total int64 - if err := query.Count(&total).Error; err != nil { - response.AbortBadRequest(c, shared.ErrQueryFileCountFailed) - return - } - - var items []model.Upload - offset := (req.Page - 1) * req.PageSize - if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil { + total, items, err := listMyUploadFiles(ctx, currUser.ID, repository.UploadListFilter{ + Keyword: req.Keyword, + Type: req.Type, + Extension: req.Extension, + Page: req.Page, + PageSize: req.PageSize, + }) + if err != nil { response.AbortBadRequest(c, shared.ErrQueryFileListFailed) return } @@ -262,7 +217,7 @@ func ListMyFiles(c *gin.Context) { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [delete] func DeleteMyFile(c *gin.Context) { - currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) @@ -275,26 +230,18 @@ func DeleteMyFile(c *gin.Context) { return } - var upload model.Upload - if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { + if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil { + if isRecordNotFound(err) { c.AbortWithStatus(http.StatusNotFound) return } - response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) - return - } - - if upload.UserID != currUser.ID { - c.AbortWithStatus(http.StatusForbidden) - return - } - - if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil { + if err == errUploadForbidden { + c.AbortWithStatus(http.StatusForbidden) + return + } response.AbortBadRequest(c, shared.ErrDeleteFileFailed) return } - uploadstats.RecordUploadStatsRemove(ctx, &upload) c.JSON(http.StatusOK, response.OKNil()) } @@ -317,7 +264,7 @@ type updateMyFileRequest struct { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [put] func UpdateMyFile(c *gin.Context) { - currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) @@ -336,34 +283,18 @@ func UpdateMyFile(c *gin.Context) { return } - var upload model.Upload - if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { + upload, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req)) + if err != nil { + if isRecordNotFound(err) { c.AbortWithStatus(http.StatusNotFound) return } - response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) - return - } - - if upload.UserID != currUser.ID { - c.AbortWithStatus(http.StatusForbidden) - return - } - - updates := make(map[string]any) - if req.FileName != "" { - updates["file_name"] = req.FileName - } - if req.AccessMode != nil { - updates["access_mode"] = *req.AccessMode - } - - if len(updates) > 0 { - if err := db.DB(ctx).Model(&upload).Updates(updates).Error; err != nil { - response.AbortBadRequest(c, "更新文件记录失败") + if err == errUploadForbidden { + c.AbortWithStatus(http.StatusForbidden) return } + response.AbortBadRequest(c, "更新文件记录失败") + return } c.JSON(http.StatusOK, response.OK(upload)) diff --git a/internal/apps/upload/handler/logics.go b/internal/apps/upload/handler/logics.go new file mode 100644 index 00000000..7556c6cf --- /dev/null +++ b/internal/apps/upload/handler/logics.go @@ -0,0 +1,199 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package handler + +import ( + "bytes" + "context" + "errors" + "sort" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" + uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" + "github.com/Rain-kl/Wavelet/internal/db/idgen" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/storage" + "github.com/Rain-kl/Wavelet/pkg/logger" + "gorm.io/gorm" +) + +func listUploadFiles(ctx context.Context, filter repository.UploadListFilter) (int64, []model.Upload, error) { + return repository.ListUploads(ctx, filter) +} + +func listMyUploadFiles(ctx context.Context, userID uint64, filter repository.UploadListFilter) (int64, []model.Upload, error) { + filter.UserID = userID + return repository.ListUploads(ctx, filter) +} + +func softDeleteUpload(ctx context.Context, uploadID uint64) (model.Upload, error) { + upload, err := repository.GetActiveUploadByID(ctx, uploadID) + if err != nil { + return model.Upload{}, err + } + if err := repository.SoftDeleteUpload(ctx, &upload); err != nil { + return model.Upload{}, err + } + uploadstats.RecordUploadStatsRemove(ctx, &upload) + return upload, nil +} + +func softDeleteOwnedUpload(ctx context.Context, userID, uploadID uint64) (model.Upload, error) { + upload, err := repository.GetActiveUploadByID(ctx, uploadID) + if err != nil { + return model.Upload{}, err + } + if upload.UserID != userID { + return model.Upload{}, errUploadForbidden + } + if err := repository.SoftDeleteUpload(ctx, &upload); err != nil { + return model.Upload{}, err + } + uploadstats.RecordUploadStatsRemove(ctx, &upload) + return upload, nil +} + +func listDistinctUploadTypes(ctx context.Context) ([]string, error) { + types, err := repository.ListDistinctUploadTypes(ctx) + if err != nil { + return nil, err + } + sort.Strings(types) + return types, nil +} + +type updateMyUploadInput struct { + FileName string + AccessMode *int +} + +func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updateMyUploadInput) (model.Upload, error) { + upload, err := repository.GetActiveUploadByID(ctx, uploadID) + if err != nil { + return model.Upload{}, err + } + if upload.UserID != userID { + return model.Upload{}, errUploadForbidden + } + + updates := make(map[string]any) + if input.FileName != "" { + updates["file_name"] = input.FileName + } + if input.AccessMode != nil { + updates["access_mode"] = *input.AccessMode + } + if err := repository.UpdateUpload(ctx, &upload, updates); err != nil { + return model.Upload{}, err + } + if name, ok := updates["file_name"].(string); ok { + upload.FileName = name + } + if mode, ok := updates["access_mode"].(int); ok { + upload.AccessMode = mode + } + return upload, nil +} + +func listUploadsForBatchDownload(ctx context.Context, ids []uint64) ([]model.Upload, error) { + return repository.ListUploadsByIDs(ctx, ids) +} + +type instantUploadInput struct { + UserID uint64 + FileHash string + Size int64 + MimeType string + Extension string + OrigName string + UploadType string + AccessMode int +} + +func createInstantUpload(ctx context.Context, existing model.Upload, input instantUploadInput) (model.Upload, error) { + newUpload := model.Upload{ + ID: idgen.NextUint64ID(), + UserID: input.UserID, + FileName: input.OrigName, + FilePath: existing.FilePath, + FileSize: input.Size, + MimeType: input.MimeType, + Extension: input.Extension, + Hash: input.FileHash, + StorageDriver: existing.StorageDriver, + Type: input.UploadType, + Status: model.UploadStatusUsed, + AccessMode: input.AccessMode, + Metadata: existing.Metadata, + } + if err := repository.CreateUpload(ctx, &newUpload); err != nil { + return model.Upload{}, err + } + uploadstats.RecordUploadStatsAdd(ctx, &newUpload) + logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", newUpload.ID, existing.FilePath) + return newUpload, nil +} + +func findReusableUpload(ctx context.Context, hash string, size int64) (model.Upload, error) { + return repository.FindReusableUploadByHash(ctx, hash, size) +} + +func saveNewUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) error { + if err := repository.CreateUpload(ctx, upload); err != nil { + backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver)) + if backendErr == nil { + if deleteErr := backend.Delete(ctx, filePath); deleteErr != nil { + logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr) + } + } + return err + } + uploadstats.RecordUploadStatsAdd(ctx, upload) + return nil +} + +func loadUploadStats(ctx context.Context) ([]model.UploadStat, error) { + return repository.ListUploadStats(ctx) +} + +var errUploadForbidden = errors.New("upload forbidden") + +func storeUploadObject(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, error) { + if uploadstorage.ReadOnly(ctx) { + return "", "", errors.New(shared.ErrStorageReadOnly) + } + driver, backend, err := storage.Active(ctx) + if err != nil { + logger.ErrorF(ctx, "初始化活动存储失败: %v", err) + return "", "", errors.New(shared.ErrSaveFileFailed) + } + result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType) + if err != nil { + logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err) + return "", "", errors.New(shared.ErrSaveFileFailed) + } + meta.Bucket = result.Bucket + return string(driver), result.Key, nil +} + +func validateUploadAllowedExtension(ctx context.Context, ext string) string { + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUploadAllowedExtensions) + if err != nil || sc.Value == "" { + return "" + } + allowedExts := strings.Split(strings.ToLower(sc.Value), ",") + for _, allowedExt := range allowedExts { + if strings.TrimSpace(allowedExt) == ext { + return "" + } + } + return shared.ErrUnsupportedFormat +} + +func isRecordNotFound(err error) bool { + return errors.Is(err, gorm.ErrRecordNotFound) +} \ No newline at end of file diff --git a/internal/apps/upload/handler/routers.go b/internal/apps/upload/handler/routers.go index 7f0953e7..e1442bcf 100644 --- a/internal/apps/upload/handler/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -26,16 +26,12 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/common/response" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/storage" - apputil "github.com/Rain-kl/Wavelet/internal/util" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -68,7 +64,7 @@ func UploadFile(c *gin.Context) { c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize) - currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() header, err := c.FormFile("file") @@ -95,7 +91,7 @@ func UploadFile(c *gin.Context) { ext = "bin" } - if errMsg := validateUploadExtension(ctx, ext); errMsg != "" { + if errMsg := validateUploadAllowedExtension(ctx, ext); errMsg != "" { response.AbortBadRequest(c, errMsg) return } @@ -142,9 +138,9 @@ func UploadFile(c *gin.Context) { id := idgen.NextUint64ID() subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext) - storageDriver, subPath, errMsg := storeUploadFile(ctx, subPath, size, mimeType, &buf, &meta) - if errMsg != "" { - response.AbortBadRequest(c, errMsg) + storageDriver, subPath, err := storeUploadObject(ctx, subPath, size, mimeType, &buf, &meta) + if err != nil { + response.AbortBadRequest(c, err.Error()) return } @@ -164,8 +160,8 @@ func UploadFile(c *gin.Context) { Metadata: meta, } - if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" { - response.AbortBadRequest(c, err) + if err := saveNewUploadRecord(ctx, &newUpload, storageDriver, subPath); err != nil { + response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed) return } @@ -253,8 +249,8 @@ func BatchDownloadFiles(c *gin.Context) { ids = append(ids, id) } - var uploads []model.Upload - if err := db.DB(ctx).Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).Find(&uploads).Error; err != nil { + uploads, err := listUploadsForBatchDownload(ctx, ids) + if err != nil { response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed) return } @@ -325,27 +321,8 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) { return accessMode, "" } -func validateUploadExtension(ctx context.Context, ext string) string { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" { - allowedExts := strings.Split(strings.ToLower(sc.Value), ",") - allowed := false - for _, allowedExt := range allowedExts { - if strings.TrimSpace(allowedExt) == ext { - allowed = true - break - } - } - if !allowed { - return shared.ErrUnsupportedFormat - } - } - return "" -} - func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string, accessMode int) (bool, error) { - var existing model.Upload - err := db.DB(ctx).Where("hash = ? AND file_size = ? AND status IN (?, ?)", fileHash, size, model.UploadStatusPending, model.UploadStatusUsed).First(&existing).Error + existing, err := findReusableUpload(ctx, fileHash, size) if err != nil { return false, err } @@ -354,52 +331,24 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, return true, nil } - id := idgen.NextUint64ID() - newUpload := model.Upload{ - ID: id, - UserID: currUser.ID, - FileName: origName, - FilePath: existing.FilePath, - FileSize: size, - MimeType: mimeType, - Extension: ext, - Hash: fileHash, - StorageDriver: existing.StorageDriver, - Type: c.DefaultPostForm("type", "generic"), - Status: model.UploadStatusUsed, - AccessMode: accessMode, - Metadata: existing.Metadata, - } - - if err := db.DB(ctx).Create(&newUpload).Error; err != nil { + newUpload, err := createInstantUpload(ctx, existing, instantUploadInput{ + UserID: currUser.ID, + FileHash: fileHash, + Size: size, + MimeType: mimeType, + Extension: ext, + OrigName: origName, + UploadType: c.DefaultPostForm("type", "generic"), + AccessMode: accessMode, + }) + if err != nil { response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed) return true, err } - uploadstats.RecordUploadStatsAdd(ctx, &newUpload) - - logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath) c.JSON(http.StatusOK, response.OK(newUpload)) return true, nil } -func storeUploadFile(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) { - if uploadstorage.ReadOnly(ctx) { - return "", "", shared.ErrStorageReadOnly - } - driver, backend, err := storage.Active(ctx) - if err != nil { - logger.ErrorF(ctx, "初始化活动存储失败: %v", err) - return "", "", shared.ErrSaveFileFailed - } - result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType) - if err != nil { - logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err) - return "", "", shared.ErrSaveFileFailed - } - meta.Bucket = result.Bucket - return string(driver), result.Key, "" -} - func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) { var meta model.UploadMetadata metadataStr := c.DefaultPostForm("metadata", "") @@ -422,16 +371,3 @@ func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) return mimeType } -func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) string { - if err := db.DB(ctx).Create(upload).Error; err != nil { - backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver)) - if backendErr == nil { - if deleteErr := backend.Delete(ctx, filePath); deleteErr != nil { - logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr) - } - } - return shared.ErrSaveUploadRecordFailed - } - uploadstats.RecordUploadStatsAdd(ctx, upload) - return "" -} \ No newline at end of file diff --git a/internal/apps/upload/handler/routers_test.go b/internal/apps/upload/handler/routers_test.go index 1dbf9d4e..d944d1ea 100644 --- a/internal/apps/upload/handler/routers_test.go +++ b/internal/apps/upload/handler/routers_test.go @@ -21,13 +21,13 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - "github.com/Rain-kl/Wavelet/internal/common/response" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" + "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" ) @@ -43,7 +43,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine { authMiddleware := func(c *gin.Context) { if authUser != nil { - util.SetToContext(c, oauth.UserObjKey, authUser) + oauth.SetToContext(c, oauth.UserObjKey, authUser) } c.Next() } @@ -297,8 +297,8 @@ func TestUploadFile(t *testing.T) { dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc) sc.Value = "jpg,png,webp,txt" dbConn.Save(&sc) - _ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc) - model.ResetSystemConfigRAMCacheForTest() + _ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc) + repository.ResetSystemConfigRAMCacheForTest() contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{ "type": "document", @@ -998,4 +998,3 @@ func TestUserUploadManagement(t *testing.T) { } }) } - diff --git a/internal/apps/upload/handler/stats.go b/internal/apps/upload/handler/stats.go index 88a0a753..01474849 100644 --- a/internal/apps/upload/handler/stats.go +++ b/internal/apps/upload/handler/stats.go @@ -9,7 +9,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/common/response" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/gin-gonic/gin" ) @@ -48,8 +47,8 @@ type fileStatsResponse struct { func GetFileStats(c *gin.Context) { ctx := c.Request.Context() - var stats []model.UploadStat - if err := db.DB(ctx).Find(&stats).Error; err != nil { + stats, err := loadUploadStats(ctx) + if err != nil { response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go index 28e98644..235aabae 100644 --- a/internal/apps/user/access_tokens.go +++ b/internal/apps/user/access_tokens.go @@ -5,17 +5,19 @@ // Package user 提供用户认证与帐户管理功能 package user -import ("net/http" +import ( + "net/http" "strconv" "strings" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) type createTokenRequest struct { Name string `json:"name"` @@ -38,7 +40,7 @@ type tokenResponse struct { // @Router /api/v1/user/access-tokens [get] // ListAccessTokens 获取当前用户的 AccessToken 列表 func ListAccessTokens(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() var tokens []model.AccessToken @@ -62,7 +64,7 @@ func ListAccessTokens(c *gin.Context) { // @Failure 400 {object} response.Any "参数错误或超限" // @Router /api/v1/user/access-tokens [post] func CreateAccessToken(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() var req createTokenRequest @@ -85,7 +87,7 @@ func CreateAccessToken(c *gin.Context) { // 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5) maxLimit := 5 - if val, err := model.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil { + if val, err := repository.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil { maxLimit = val } @@ -140,7 +142,7 @@ func CreateAccessToken(c *gin.Context) { // @Failure 400 {object} response.Any "参数错误" // @Router /api/v1/user/access-tokens/{id} [delete] func DeleteAccessToken(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() idStr := c.Param("id") @@ -175,7 +177,7 @@ func DeleteAccessToken(c *gin.Context) { // @Failure 400 {object} response.Any "参数错误" // @Router /api/v1/user/access-tokens/{id}/rotate [post] func RotateAccessToken(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() idStr := c.Param("id") diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 37fc259e..1fc9db2f 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -14,6 +14,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/task" pkgu "github.com/Rain-kl/Wavelet/pkg/util" ) @@ -47,8 +48,32 @@ type updateProfileInput struct { Location string } +func isPasswordLoginEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled) + if err != nil { + return true + } + return enabled +} + +func isPasswordRegisterEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled) + if err != nil { + return true + } + return enabled +} + +func isRegistrationEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) + if err != nil { + return true + } + return enabled +} + func isEmailLoginVerificationEnabled(ctx context.Context) bool { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled) + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled) if err != nil { return false } @@ -56,7 +81,7 @@ func isEmailLoginVerificationEnabled(ctx context.Context) bool { } func isEmailRegisterVerificationEnabled(ctx context.Context) bool { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled) + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled) if err != nil { return false } @@ -64,26 +89,14 @@ func isEmailRegisterVerificationEnabled(ctx context.Context) bool { } func isSMTPConfigured(ctx context.Context) bool { - var host, port, username, password string - - var scHost model.SystemConfig - if err := scHost.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil { - host = scHost.Value + scHost, errHost := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost) + scPort, errPort := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort) + scUser, errUser := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername) + scPass, errPass := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword) + if errHost != nil || errPort != nil || errUser != nil || errPass != nil { + return false } - var scPort model.SystemConfig - if err := scPort.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil { - port = scPort.Value - } - var scUser model.SystemConfig - if err := scUser.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil { - username = scUser.Value - } - var scPass model.SystemConfig - if err := scPass.GetByKey(ctx, model.ConfigKeySMTPPassword); err == nil { - password = scPass.Value - } - - return host != "" && port != "" && username != "" && password != "" + return scHost.Value != "" && scPort.Value != "" && scUser.Value != "" && scPass.Value != "" } func generateVerificationCode() (string, error) { @@ -114,11 +127,11 @@ func sendEmailVerificationCode(ctx context.Context, email, scene, templateName s codeKey := getEmailCodeKey(scene, email) cooldownKey := getEmailCooldownKey(scene, email) - emailSubject, emailBody, err := model.RenderTemplate( - ctx, - templateName, - map[string]any{"Code": code}, - ) + tmpl, err := repository.GetTemplateByKey(ctx, templateName) + if err != nil { + return fmt.Errorf("模板 %s 不存在或不可用: %w", templateName, err) + } + emailSubject, emailBody, err := tmpl.Render(map[string]any{"Code": code}) if err != nil { return fmt.Errorf(errRenderEmailTemplateFailed, err) } @@ -271,4 +284,4 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn return nil, err } return &dbUser, nil -} \ No newline at end of file +} diff --git a/internal/apps/user/logics_test.go b/internal/apps/user/logics_test.go index 127cbc6f..fd8de420 100644 --- a/internal/apps/user/logics_test.go +++ b/internal/apps/user/logics_test.go @@ -10,6 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -39,7 +40,7 @@ func TestProcessLoginEmailVerificationSMTPFallback(t *testing.T) { Update("value", "").Error; err != nil { t.Fatalf("clear SMTP host failed: %v", err) } - if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { t.Fatalf("invalidate system config cache failed: %v", err) } @@ -104,7 +105,7 @@ func TestProcessLoginEmailVerificationEmptyEmailFallback(t *testing.T) { t.Fatalf("set %s failed: %v", cfg.key, err) } } - if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { t.Fatalf("invalidate system config cache failed: %v", err) } @@ -148,4 +149,4 @@ func TestProcessLoginEmailVerificationInvalidCode(t *testing.T) { if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired { t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result) } -} \ No newline at end of file +} diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index 43092dbf..b404281a 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -3,23 +3,24 @@ package user -import ("context" +import ( + "context" "net/http" "strings" "time" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/common" + "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/listener" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/common/response" ) type loginRequest struct { @@ -53,30 +54,6 @@ type updateProfileRequest struct { Location string `json:"location"` } -func isPasswordLoginEnabled() bool { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled) - if err != nil { - return true - } - return enabled -} - -func isPasswordRegisterEnabled() bool { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordRegisterEnabled) - if err != nil { - return true - } - return enabled -} - -func isRegistrationEnabled() bool { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyRegistrationEnabled) - if err != nil { - return true - } - return enabled -} - func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error { session := sessions.Default(c) session.Set(oauth.UserIDKey, user.ID) @@ -87,7 +64,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro maxAge := config.Config.App.SessionAge isSessionCookie := false - ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) + ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) if err == nil { switch { case ttlHours == -1: @@ -124,7 +101,8 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro // @Failure 500 {object} response.Any "服务内部错误" // @Router /api/v1/user/login [post] func Login(c *gin.Context) { - if !isPasswordLoginEnabled() { + ctx := c.Request.Context() + if !isPasswordLoginEnabled(ctx) { response.AbortBadRequest(c, errPasswordLoginDisabled) return } @@ -140,7 +118,6 @@ func Login(c *gin.Context) { } var user model.User - ctx := c.Request.Context() if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil { logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP()) response.AbortBadRequest(c, errUsernameOrPasswordWrong) @@ -211,7 +188,8 @@ func Login(c *gin.Context) { // @Failure 500 {object} response.Any "服务内部错误" // @Router /api/v1/user/register [post] func Register(c *gin.Context) { - if !isRegistrationEnabled() || !isPasswordRegisterEnabled() { + ctx := c.Request.Context() + if !isRegistrationEnabled(ctx) || !isPasswordRegisterEnabled(ctx) { response.AbortBadRequest(c, errRegistrationDisabled) return } @@ -242,8 +220,6 @@ func Register(c *gin.Context) { return } - ctx := c.Request.Context() - // 邮箱注册验证校验 if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil { response.AbortBadRequest(c, err.Error()) @@ -344,7 +320,7 @@ func ChangePassword(c *gin.Context) { return } - userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) if userObj == nil { response.AbortUnauthorized(c, errLoginRequired) return @@ -443,7 +419,7 @@ func UpdateProfile(c *gin.Context) { return } - userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) if userObj == nil { response.AbortUnauthorized(c, errLoginRequired) return diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index b0fce8e4..23f959df 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -3,7 +3,8 @@ package user -import ("bytes" +import ( + "bytes" "context" "encoding/json" "net/http" @@ -16,12 +17,14 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) func setupUserTestRouter(t *testing.T) *gin.Engine { t.Helper() @@ -281,7 +284,7 @@ func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) { } // 2.5 Invalidate the system config cache in Redis - if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { t.Fatalf("invalidate system config cache failed: %v", err) } @@ -387,7 +390,7 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) { } // Invalidate the system config cache in Redis - if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil { + if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { t.Fatalf("invalidate system config cache failed: %v", err) } diff --git a/internal/apps/user/tasks.go b/internal/apps/user/tasks.go index 99aafdac..01d9796c 100644 --- a/internal/apps/user/tasks.go +++ b/internal/apps/user/tasks.go @@ -12,6 +12,7 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/pkg/mail" ) @@ -111,17 +112,16 @@ func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.T var smtpUsername string var smtpPassword string - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil { + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost); err == nil { smtpHost = sc.Value } - if err := sc.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil { + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort); err == nil { smtpPortVal = sc.Value } - if err := sc.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil { + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername); err == nil { smtpUsername = sc.Value } - if err := sc.GetByKey(ctx, model.ConfigKeySMTPPassword); err == nil { + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword); err == nil { smtpPassword = sc.Value } diff --git a/internal/db/migrator/migrator.go b/internal/db/migrator/migrator.go index 74e37184..a28d3b4f 100644 --- a/internal/db/migrator/migrator.go +++ b/internal/db/migrator/migrator.go @@ -12,7 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/pressly/goose/v3" ) @@ -69,7 +69,7 @@ func Migrate() { } func clearSystemConfigCache() { - if err := model.InvalidateAllSystemConfigCaches(context.Background()); err != nil { + if err := repository.InvalidateAllSystemConfigCaches(context.Background()); err != nil { log.Printf("[%s] clear system config cache failed: %v\n", dbType(), err) } } diff --git a/internal/db/migrator/migrator_test.go b/internal/db/migrator/migrator_test.go index a355e2ec..56a1241f 100644 --- a/internal/db/migrator/migrator_test.go +++ b/internal/db/migrator/migrator_test.go @@ -10,6 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/alicebob/miniredis/v2" "github.com/glebarez/sqlite" "github.com/redis/go-redis/v9" @@ -107,13 +108,13 @@ func TestMigrateClearsStaleSystemConfigCache(t *testing.T) { Value: "true", Type: "system", } - if err := db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeyCapLoginEnabled, &staleConfig); err != nil { + if err := db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, model.ConfigKeyCapLoginEnabled, &staleConfig); err != nil { t.Fatalf("HSetJSON() error = %v", err) } Migrate() - exists, err := db.Redis.Exists(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Result() + exists, err := db.Redis.Exists(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Result() if err != nil { t.Fatalf("Redis.Exists() error = %v", err) } @@ -121,7 +122,7 @@ func TestMigrateClearsStaleSystemConfigCache(t *testing.T) { t.Fatalf("system config cache exists = %d, want 0", exists) } - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled) + enabled, err := repository.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled) if err != nil { t.Fatalf("GetBoolByKey(%s) error = %v", model.ConfigKeyCapLoginEnabled, err) } diff --git a/internal/diskcache/cache.go b/internal/diskcache/cache.go index b8d96835..c52bfaee 100644 --- a/internal/diskcache/cache.go +++ b/internal/diskcache/cache.go @@ -12,6 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" pkgcache "github.com/Rain-kl/Wavelet/pkg/cache/disk" ) @@ -69,27 +70,24 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) { } // 1. Max Size - var scMaxSize model.SystemConfig maxSizeMB := int64(defaultMaxSizeMB) - if err := scMaxSize.GetByKey(ctx, model.ConfigKeyDiskCacheMaxSizeMB); err == nil && scMaxSize.Value != "" { + if scMaxSize, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheMaxSizeMB); err == nil && scMaxSize.Value != "" { if val, err := strconv.ParseInt(scMaxSize.Value, 10, 64); err == nil && val > 0 { maxSizeMB = val } } // 2. Default TTL - var scTTL model.SystemConfig ttlMinutes := int64(defaultTTLMinutes) - if err := scTTL.GetByKey(ctx, model.ConfigKeyDiskCacheTTLMinutes); err == nil && scTTL.Value != "" { + if scTTL, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheTTLMinutes); err == nil && scTTL.Value != "" { if val, err := strconv.ParseInt(scTTL.Value, 10, 64); err == nil && val >= 0 { ttlMinutes = val } } // 3. LRU Enabled - var scLRU model.SystemConfig lruEnabled := true - if err := scLRU.GetByKey(ctx, model.ConfigKeyDiskCacheLRUEnabled); err == nil && scLRU.Value != "" { + if scLRU, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheLRUEnabled); err == nil && scLRU.Value != "" { if val, err := strconv.ParseBool(scLRU.Value); err == nil { lruEnabled = val } diff --git a/internal/diskcache/cache_test.go b/internal/diskcache/cache_test.go index 4f5cce32..c995f989 100644 --- a/internal/diskcache/cache_test.go +++ b/internal/diskcache/cache_test.go @@ -10,6 +10,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" ) @@ -31,7 +32,7 @@ func TestDiskCacheReloadConfig(t *testing.T) { // Invalidate Redis config cache to force DB reload if db.Redis != nil { - db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)) + db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)) } // Reload config diff --git a/internal/model/push_channel.go b/internal/model/push_channel.go index 515a0dd2..8227cd5b 100644 --- a/internal/model/push_channel.go +++ b/internal/model/push_channel.go @@ -4,15 +4,11 @@ package model import ( - "context" "encoding/json" "errors" "regexp" "strings" "time" - - "github.com/Rain-kl/Wavelet/internal/db" - "gorm.io/gorm" ) const ( @@ -70,8 +66,6 @@ func (pc *PushChannel) Validate() error { return errors.New("request URL/address is required") } - // For custom and lark, we must enforce https:// URL prefix for security. - // For email, it is an SMTP host:port, so no need for https:// prefix. if pc.Type != TypeEmail && !strings.HasPrefix(pc.URL, "https://") { return errors.New("request URL must use HTTPS protocol for security reasons") } @@ -102,58 +96,4 @@ func validateJSON(s string) error { return nil } return errors.New("payload schema must be a valid JSON format") -} - -// GetPushChannelByName 根据名称获取消息通道 -func GetPushChannelByName(ctx context.Context, name string) (*PushChannel, error) { - var channel PushChannel - err := db.DB(ctx).Where("name = ?", name).First(&channel).Error - if err != nil { - return nil, err - } - return &channel, nil -} - -const activePushChannelCacheTTL = 24 * time.Hour - -// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取) -func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) { - cacheKey := "push:channel:active:" + name - var channel PushChannel - if db.Redis != nil { - if err := db.GetJSON(ctx, cacheKey, &channel); err == nil { - return &channel, nil - } - } - - err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error - if err != nil { - return nil, err - } - - if db.Redis != nil { - // 缓存有效时间设置为 24 小时 - _ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL) - } - - return &channel, nil -} - -// DeleteActivePushChannelCache 清理启用消息通道的缓存 -func DeleteActivePushChannelCache(ctx context.Context, name string) { - if db.Redis != nil { - _ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err() - } -} - -// AfterSave GORM 保存后钩子,用于自动清理缓存 -func (pc *PushChannel) AfterSave(tx *gorm.DB) error { - DeleteActivePushChannelCache(tx.Statement.Context, pc.Name) - return nil -} - -// AfterDelete GORM 删除后钩子,用于自动清理缓存 -func (pc *PushChannel) AfterDelete(tx *gorm.DB) error { - DeleteActivePushChannelCache(tx.Statement.Context, pc.Name) - return nil -} +} \ No newline at end of file diff --git a/internal/model/push_event.go b/internal/model/push_event.go index 7e4c5c89..d655af35 100644 --- a/internal/model/push_event.go +++ b/internal/model/push_event.go @@ -4,13 +4,9 @@ package model import ( - "context" "errors" "strings" "time" - - "github.com/Rain-kl/Wavelet/internal/db" - "gorm.io/gorm" ) // PushEvent 系统通知事件模型 @@ -51,48 +47,4 @@ func (pe *PushEvent) Validate() error { return errors.New("cannot enable event without any push channels configured") } return nil -} - -const activePushEventCacheTTL = 24 * time.Hour - -// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取) -func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) { - cacheKey := "push:event:active:" + key - var event PushEvent - if db.Redis != nil { - if err := db.GetJSON(ctx, cacheKey, &event); err == nil { - return &event, nil - } - } - - err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error - if err != nil { - return nil, err - } - - if db.Redis != nil { - // 缓存有效时间设置为 24 小时 - _ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL) - } - - return &event, nil -} - -// DeleteActivePushEventCache 清理启用通知事件的缓存 -func DeleteActivePushEventCache(ctx context.Context, key string) { - if db.Redis != nil { - _ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err() - } -} - -// AfterSave GORM 保存后钩子,用于自动清理缓存 -func (pe *PushEvent) AfterSave(tx *gorm.DB) error { - DeleteActivePushEventCache(tx.Statement.Context, pe.EventKey) - return nil -} - -// AfterDelete GORM 删除后钩子,用于自动清理缓存 -func (pe *PushEvent) AfterDelete(tx *gorm.DB) error { - DeleteActivePushEventCache(tx.Statement.Context, pe.EventKey) - return nil -} +} \ No newline at end of file diff --git a/internal/model/system_configs.go b/internal/model/system_configs.go index 61775469..09d99867 100644 --- a/internal/model/system_configs.go +++ b/internal/model/system_configs.go @@ -3,19 +3,7 @@ package model -import ( - "context" - "encoding/json" - "errors" - "fmt" - "strconv" - "time" - - "github.com/redis/go-redis/v9" - "github.com/shopspring/decimal" - - "github.com/Rain-kl/Wavelet/internal/db" -) +import "time" // 配置键常量 - 所有系统配置的 key 定义 const ( @@ -51,13 +39,6 @@ const ( ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON) ) -const ( - // SystemConfigRedisHashKey Redis Hash key,存储所有系统配置 - SystemConfigRedisHashKey = "system:system_configs" - // SystemConfigVisibleListRedisKey Redis key,缓存所有 visibility=1 的公共配置列表 - SystemConfigVisibleListRedisKey = "system:visible_configs" -) - const ( // ConfigVisibilityHidden 表示配置不通过公共配置接口暴露 ConfigVisibilityHidden = 0 @@ -79,181 +60,4 @@ type SystemConfig struct { // TableName 表名 func (SystemConfig) TableName() string { return "w_system_configs" -} - -// GetByKey 通过 key 查询配置(带 RAM + Redis 缓存) -func (sc *SystemConfig) GetByKey(ctx context.Context, key string) error { - ensureSystemConfigCacheListener() - - if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok { - *sc = cloneSystemConfig(cached) - return nil - } - - if db.Redis != nil { - if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, sc); err == nil { - systemConfigRAMCache.Set(key, cloneSystemConfig(*sc)) - return nil - } else if !errors.Is(err, redis.Nil) { - // Redis 服务错误,返回错误 - return err - } - } - - // 查数据库 - database := db.DB(ctx) - if database == nil { - return errors.New(errDatabaseNotInitialized) - } - - if err := database.Where("key = ?", key).First(sc).Error; err != nil { - return err - } - - populateSystemConfigCache(ctx, *sc) - - return nil -} - -// ListSystemConfigsByKeys loads multiple config keys in one database round trip. -// Keys already present in the process-local RAM cache are returned without querying PostgreSQL. -func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]SystemConfig, error) { - if len(keys) == 0 { - return map[string]SystemConfig{}, nil - } - - ensureSystemConfigCacheListener() - - result := make(map[string]SystemConfig, len(keys)) - missing := make([]string, 0, len(keys)) - for _, key := range keys { - if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok { - result[key] = cloneSystemConfig(cached) - continue - } - missing = append(missing, key) - } - - if len(missing) == 0 { - return result, nil - } - - database := db.DB(ctx) - if database == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - - var configs []SystemConfig - if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil { - return nil, err - } - - for i := range configs { - populateSystemConfigCache(ctx, configs[i]) - result[configs[i].Key] = cloneSystemConfig(configs[i]) - } - - return result, nil -} - -// InvalidateVisibleSystemConfigsCache clears the cached public config list. -func InvalidateVisibleSystemConfigsCache(ctx context.Context) error { - if db.Redis == nil { - return nil - } - return db.Redis.Del(ctx, db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err() -} - -// ListVisibleSystemConfigs 查询所有可通过公共配置接口暴露的配置(带 Redis 列表缓存) -func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) { - if db.Redis != nil { - var cached []SystemConfig - if err := db.GetJSON(ctx, SystemConfigVisibleListRedisKey, &cached); err == nil { - return cached, nil - } else if !errors.Is(err, redis.Nil) { - return nil, err - } - } - - database := db.DB(ctx) - if database == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - - var configs []SystemConfig - if err := database.Where("visibility = ?", ConfigVisibilityVisible).Find(&configs).Error; err != nil { - return nil, err - } - - if db.Redis != nil { - _ = db.SetJSON(ctx, SystemConfigVisibleListRedisKey, configs, 0) - } - - return configs, nil -} - -// GetIntByKey 通过 key 查询配置并转换为 int 类型 -func GetIntByKey(ctx context.Context, key string) (int, error) { - var sc SystemConfig - if err := sc.GetByKey(ctx, key); err != nil { - return 0, err - } - - value, err := strconv.Atoi(sc.Value) - if err != nil { - return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err) - } - - return value, nil -} - -// GetDecimalByKey 通过 key 查询配置并转换为 decimal.Decimal 类型 -// precision 指定保留的小数位数,多余的小数会被裁剪 -func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) { - var sc SystemConfig - if err := sc.GetByKey(ctx, key); err != nil { - return decimal.Zero, err - } - - value, err := decimal.NewFromString(sc.Value) - if err != nil { - return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err) - } - - // 裁剪到指定小数位数 - return value.Truncate(precision), nil -} - -// GetBoolByKey 通过 key 查询配置并转换为 bool 类型 -func GetBoolByKey(ctx context.Context, key string) (bool, error) { - var sc SystemConfig - if err := sc.GetByKey(ctx, key); err != nil { - return false, err - } - - value, err := strconv.ParseBool(sc.Value) - if err != nil { - return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err) - } - - return value, nil -} - -// GetMenuDisplayConfig 获取目录显示配置,解析为 map[string]bool -func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) { - var sc SystemConfig - if err := sc.GetByKey(ctx, ConfigKeyMenuDisplayConfig); err != nil { - return nil, err - } - - config := make(map[string]bool) - if sc.Value == "" || sc.Value == "{}" { - return config, nil - } - - if err := json.Unmarshal([]byte(sc.Value), &config); err != nil { - return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err) - } - - return config, nil -} +} \ No newline at end of file diff --git a/internal/model/templates.go b/internal/model/templates.go index a62159e9..28f78ad1 100644 --- a/internal/model/templates.go +++ b/internal/model/templates.go @@ -5,14 +5,10 @@ package model import ( "bytes" - "context" "errors" - "fmt" "strings" "text/template" "time" - - "github.com/Rain-kl/Wavelet/internal/db" ) // Template 邮件/消息模板实体 @@ -90,17 +86,3 @@ func (t *Template) Render(data any) (string, string, error) { return subject, bodyBuf.String(), nil } - -// RenderTemplate 渲染指定模板。模板不存在或渲染失败时返回错误,由调用方决定如何处理。 -func RenderTemplate(ctx context.Context, key string, data any) (string, string, error) { - var t Template - if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err != nil { - return "", "", fmt.Errorf(errTemplateUnavailable, key, err) - } - - subject, body, err := t.Render(data) - if err != nil { - return "", "", fmt.Errorf(errTemplateRenderFailed, key, err) - } - return subject, body, nil -} diff --git a/internal/model/users.go b/internal/model/users.go index b786d714..179c1ea5 100644 --- a/internal/model/users.go +++ b/internal/model/users.go @@ -139,12 +139,7 @@ func (u *User) assignIDIfMissing() error { } // CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验) -func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error { - enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled) - if err == nil && !enabled { - return errors.New(errRegistrationDisabled) - } - +func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error { now := time.Now() userID := oauthInfo.GetID() newUser := User{ @@ -169,12 +164,7 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser } // RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验) -func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error { - enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled) - if err == nil && !enabled { - return errors.New(errRegistrationDisabled) - } - +func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error { // 检查用户名冲突 var count int64 if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil { diff --git a/internal/repository/push_channel.go b/internal/repository/push_channel.go new file mode 100644 index 00000000..db4dd4bd --- /dev/null +++ b/internal/repository/push_channel.go @@ -0,0 +1,105 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const activePushChannelCacheTTL = 24 * time.Hour + +// ListPushChannels returns all push channels ordered by creation time descending. +func ListPushChannels(ctx context.Context) ([]model.PushChannel, error) { + var channels []model.PushChannel + if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil { + return nil, err + } + return channels, nil +} + +// GetPushChannelByID loads a push channel by primary key. +func GetPushChannelByID(ctx context.Context, id uint64) (model.PushChannel, error) { + var channel model.PushChannel + if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { + return model.PushChannel{}, err + } + return channel, nil +} + +// GetPushChannelByName 根据名称获取消息通道。 +func GetPushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) { + var channel model.PushChannel + if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil { + return nil, err + } + return &channel, nil +} + +// CountPushChannelsByName returns how many channels share the given name. +func CountPushChannelsByName(ctx context.Context, name string) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CreatePushChannel persists a new channel and invalidates cache. +func CreatePushChannel(ctx context.Context, channel *model.PushChannel) error { + if err := db.DB(ctx).Create(channel).Error; err != nil { + return err + } + DeleteActivePushChannelCache(ctx, channel.Name) + return nil +} + +// SavePushChannel updates a channel and invalidates cache. +func SavePushChannel(ctx context.Context, channel *model.PushChannel) error { + if err := db.DB(ctx).Save(channel).Error; err != nil { + return err + } + DeleteActivePushChannelCache(ctx, channel.Name) + return nil +} + +// DeletePushChannel removes a channel and invalidates cache. +func DeletePushChannel(ctx context.Context, channel *model.PushChannel) error { + if err := db.DB(ctx).Delete(channel).Error; err != nil { + return err + } + DeleteActivePushChannelCache(ctx, channel.Name) + return nil +} + +// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。 +func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) { + cacheKey := "push:channel:active:" + name + var channel model.PushChannel + if db.Redis != nil { + if err := db.GetJSON(ctx, cacheKey, &channel); err == nil { + return &channel, nil + } + } + + if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil { + return nil, err + } + + if db.Redis != nil { + _ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL) + } + + return &channel, nil +} + +// DeleteActivePushChannelCache 清理启用消息通道的缓存。 +func DeleteActivePushChannelCache(ctx context.Context, name string) { + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err() + } +} \ No newline at end of file diff --git a/internal/repository/push_event.go b/internal/repository/push_event.go new file mode 100644 index 00000000..959f51ba --- /dev/null +++ b/internal/repository/push_event.go @@ -0,0 +1,124 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const activePushEventCacheTTL = 24 * time.Hour + +// ListPushEvents returns all push events ordered by creation time descending. +func ListPushEvents(ctx context.Context) ([]model.PushEvent, error) { + var events []model.PushEvent + if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil { + return nil, err + } + return events, nil +} + +// GetPushEventByID loads a push event by primary key. +func GetPushEventByID(ctx context.Context, id uint64) (model.PushEvent, error) { + var event model.PushEvent + if err := db.DB(ctx).First(&event, id).Error; err != nil { + return model.PushEvent{}, err + } + return event, nil +} + +// GetPushEventByKey loads a push event by event key. +func GetPushEventByKey(ctx context.Context, key string) (model.PushEvent, error) { + var event model.PushEvent + if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil { + return model.PushEvent{}, err + } + return event, nil +} + +// CountPushEventsByKey returns how many events use the given event key. +func CountPushEventsByKey(ctx context.Context, key string) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CreatePushEvent persists a new push event and invalidates cache. +func CreatePushEvent(ctx context.Context, event *model.PushEvent) error { + if err := db.DB(ctx).Create(event).Error; err != nil { + return err + } + DeleteActivePushEventCache(ctx, event.EventKey) + return nil +} + +// SavePushEvent updates a push event and invalidates cache. +func SavePushEvent(ctx context.Context, event *model.PushEvent) error { + if err := db.DB(ctx).Save(event).Error; err != nil { + return err + } + DeleteActivePushEventCache(ctx, event.EventKey) + return nil +} + +// UpdatePushEventEnabled toggles the enabled flag for a push event. +func UpdatePushEventEnabled(ctx context.Context, event *model.PushEvent, enabled bool) error { + event.Enabled = enabled + if err := db.DB(ctx).Model(event).Update("enabled", enabled).Error; err != nil { + return err + } + DeleteActivePushEventCache(ctx, event.EventKey) + return nil +} + +// DeletePushEvent removes a push event and invalidates cache. +func DeletePushEvent(ctx context.Context, event *model.PushEvent) error { + if err := db.DB(ctx).Delete(event).Error; err != nil { + return err + } + DeleteActivePushEventCache(ctx, event.EventKey) + return nil +} + +// ListActivePushEventsByTaskType returns enabled events bound to a task type. +func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) { + var events []model.PushEvent + if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil { + return nil, err + } + return events, nil +} + +// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。 +func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) { + cacheKey := "push:event:active:" + key + var event model.PushEvent + if db.Redis != nil { + if err := db.GetJSON(ctx, cacheKey, &event); err == nil { + return &event, nil + } + } + + if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil { + return nil, err + } + + if db.Redis != nil { + _ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL) + } + + return &event, nil +} + +// DeleteActivePushEventCache 清理启用通知事件的缓存。 +func DeleteActivePushEventCache(ctx context.Context, key string) { + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err() + } +} \ No newline at end of file diff --git a/internal/repository/push_history.go b/internal/repository/push_history.go new file mode 100644 index 00000000..fb1679bb --- /dev/null +++ b/internal/repository/push_history.go @@ -0,0 +1,54 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// PushHistoryListFilter filters push history pagination queries. +type PushHistoryListFilter struct { + EventKey string + Status string + Page int + PageSize int +} + +// ListPushHistories returns paginated push history records. +func ListPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []model.PushHistory, error) { + query := db.DB(ctx).Model(&model.PushHistory{}).Order("created_at DESC") + if filter.EventKey != "" { + query = query.Where("event_key = ?", filter.EventKey) + } + if filter.Status != "" { + query = query.Where("status = ?", filter.Status) + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return 0, nil, err + } + + var results []model.PushHistory + offset := (filter.Page - 1) * filter.PageSize + if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil { + return 0, nil, err + } + + return total, results, nil +} + +// CreatePushHistory persists a push history audit record. +func CreatePushHistory(ctx context.Context, history *model.PushHistory) error { + return db.DB(ctx).Create(history).Error +} + +// PushHistoryQuery returns a scoped query builder for push histories. +func PushHistoryQuery(ctx context.Context) *gorm.DB { + return db.DB(ctx).Model(&model.PushHistory{}) +} \ No newline at end of file diff --git a/internal/repository/system_config.go b/internal/repository/system_config.go new file mode 100644 index 00000000..f0ea0a2a --- /dev/null +++ b/internal/repository/system_config.go @@ -0,0 +1,198 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package repository provides data access with caching and persistence boundaries. +package repository + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strconv" + + "github.com/redis/go-redis/v9" + "github.com/shopspring/decimal" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + errDatabaseNotInitialized = "database not initialized" + errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w" + errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w" + errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w" + errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w" +) + +// GetSystemConfigByKey 通过 key 查询配置(带 RAM + Redis 缓存)。 +func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) { + ensureSystemConfigCacheListener() + + if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok { + return cloneSystemConfig(cached), nil + } + + var sc model.SystemConfig + if db.Redis != nil { + if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, &sc); err == nil { + systemConfigRAMCache.Set(key, cloneSystemConfig(sc)) + return sc, nil + } else if !errors.Is(err, redis.Nil) { + return model.SystemConfig{}, err + } + } + + database := db.DB(ctx) + if database == nil { + return model.SystemConfig{}, errors.New(errDatabaseNotInitialized) + } + + if err := database.Where("key = ?", key).First(&sc).Error; err != nil { + return model.SystemConfig{}, err + } + + populateSystemConfigCache(ctx, sc) + return sc, nil +} + +// ListSystemConfigsByKeys loads multiple config keys in one database round trip. +func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) { + if len(keys) == 0 { + return map[string]model.SystemConfig{}, nil + } + + ensureSystemConfigCacheListener() + + result := make(map[string]model.SystemConfig, len(keys)) + missing := make([]string, 0, len(keys)) + for _, key := range keys { + if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok { + result[key] = cloneSystemConfig(cached) + continue + } + missing = append(missing, key) + } + + if len(missing) == 0 { + return result, nil + } + + database := db.DB(ctx) + if database == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + var configs []model.SystemConfig + if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil { + return nil, err + } + + for i := range configs { + populateSystemConfigCache(ctx, configs[i]) + result[configs[i].Key] = cloneSystemConfig(configs[i]) + } + + return result, nil +} + +// InvalidateVisibleSystemConfigsCache clears the cached public config list. +func InvalidateVisibleSystemConfigsCache(ctx context.Context) error { + if db.Redis == nil { + return nil + } + return db.Redis.Del(ctx, db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err() +} + +// ListVisibleSystemConfigs 查询所有可通过公共配置接口暴露的配置(带 Redis 列表缓存)。 +func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) { + if db.Redis != nil { + var cached []model.SystemConfig + if err := db.GetJSON(ctx, SystemConfigVisibleListRedisKey, &cached); err == nil { + return cached, nil + } else if !errors.Is(err, redis.Nil) { + return nil, err + } + } + + database := db.DB(ctx) + if database == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + var configs []model.SystemConfig + if err := database.Where("visibility = ?", model.ConfigVisibilityVisible).Find(&configs).Error; err != nil { + return nil, err + } + + if db.Redis != nil { + _ = db.SetJSON(ctx, SystemConfigVisibleListRedisKey, configs, 0) + } + + return configs, nil +} + +// GetIntByKey 通过 key 查询配置并转换为 int 类型。 +func GetIntByKey(ctx context.Context, key string) (int, error) { + sc, err := GetSystemConfigByKey(ctx, key) + if err != nil { + return 0, err + } + + value, err := strconv.Atoi(sc.Value) + if err != nil { + return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err) + } + + return value, nil +} + +// GetDecimalByKey 通过 key 查询配置并转换为 decimal.Decimal 类型。 +func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) { + sc, err := GetSystemConfigByKey(ctx, key) + if err != nil { + return decimal.Zero, err + } + + value, err := decimal.NewFromString(sc.Value) + if err != nil { + return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err) + } + + return value.Truncate(precision), nil +} + +// GetBoolByKey 通过 key 查询配置并转换为 bool 类型。 +func GetBoolByKey(ctx context.Context, key string) (bool, error) { + sc, err := GetSystemConfigByKey(ctx, key) + if err != nil { + return false, err + } + + value, err := strconv.ParseBool(sc.Value) + if err != nil { + return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err) + } + + return value, nil +} + +// GetMenuDisplayConfig 获取目录显示配置,解析为 map[string]bool。 +func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) { + sc, err := GetSystemConfigByKey(ctx, model.ConfigKeyMenuDisplayConfig) + if err != nil { + return nil, err + } + + config := make(map[string]bool) + if sc.Value == "" || sc.Value == "{}" { + return config, nil + } + + if err := json.Unmarshal([]byte(sc.Value), &config); err != nil { + return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err) + } + + return config, nil +} \ No newline at end of file diff --git a/internal/repository/system_config_admin.go b/internal/repository/system_config_admin.go new file mode 100644 index 00000000..a79e9d39 --- /dev/null +++ b/internal/repository/system_config_admin.go @@ -0,0 +1,85 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// ListAdminSystemConfigs returns all configs, optionally filtered by type. +func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) { + query := db.DB(ctx).Order("created_at DESC") + if configType != "" { + query = query.Where("type = ?", configType) + } + var configs []model.SystemConfig + if err := query.Find(&configs).Error; err != nil { + return nil, err + } + return configs, nil +} + +// GetAdminSystemConfigByKey loads a config directly from PostgreSQL. +func GetAdminSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) { + var config model.SystemConfig + if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil { + return model.SystemConfig{}, err + } + return config, nil +} + +// SystemConfigExists reports whether a config key already exists. +func SystemConfigExists(ctx context.Context, key string) (bool, error) { + var existing model.SystemConfig + err := db.DB(ctx).Where("key = ?", key).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + if err != nil { + return false, err + } + return true, nil +} + +// CreateSystemConfig persists a new system config row. +func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error { + return db.DB(ctx).Create(config).Error +} + +// UpdateSystemConfigFields applies partial updates to a system config row. +func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error { + return db.DB(ctx).Model(config).Updates(updates).Error +} + +// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache. +func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error { + var sc model.SystemConfig + err := db.DB(ctx).Where("key = ?", key).First(&sc).Error + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + if errors.Is(err, gorm.ErrRecordNotFound) { + sc = model.SystemConfig{ + Key: key, + Value: value, + Type: "system", + Visibility: model.ConfigVisibilityHidden, + } + if err := db.DB(ctx).Create(&sc).Error; err != nil { + return err + } + } else { + sc.Value = value + if err := db.DB(ctx).Save(&sc).Error; err != nil { + return err + } + } + return InvalidateSystemConfigCache(ctx, key) +} \ No newline at end of file diff --git a/internal/model/system_config_cache.go b/internal/repository/system_config_cache.go similarity index 79% rename from internal/model/system_config_cache.go rename to internal/repository/system_config_cache.go index e354d7b9..f983f928 100644 --- a/internal/model/system_config_cache.go +++ b/internal/repository/system_config_cache.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -9,14 +9,20 @@ import ( "sync" "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/cache/ram" ) const ( // SystemConfigInvalidationChannel broadcasts RAM cache eviction across nodes. SystemConfigInvalidationChannel = "system:config_invalidation" - systemConfigInvalidateAllToken = "*" - systemConfigRAMMaximumSize = 512 + // SystemConfigRedisHashKey Redis Hash key,存储所有系统配置。 + SystemConfigRedisHashKey = "system:system_configs" + // SystemConfigVisibleListRedisKey Redis key,缓存所有 visibility=1 的公共配置列表。 + SystemConfigVisibleListRedisKey = "system:visible_configs" + + systemConfigInvalidateAllToken = "*" + systemConfigRAMMaximumSize = 512 ) type systemConfigInvalidationMessage struct { @@ -24,7 +30,7 @@ type systemConfigInvalidationMessage struct { } var ( - systemConfigRAMCache = ram.MustNew[string, SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) + systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) systemConfigListenerOnce sync.Once ) @@ -58,11 +64,11 @@ func startSystemConfigCacheInvalidationListener() { }() } -func cloneSystemConfig(sc SystemConfig) SystemConfig { +func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig { return sc } -func populateSystemConfigCache(ctx context.Context, sc SystemConfig) { +func populateSystemConfigCache(ctx context.Context, sc model.SystemConfig) { systemConfigRAMCache.Set(sc.Key, cloneSystemConfig(sc)) if db.Redis != nil { _ = db.HSetJSON(ctx, SystemConfigRedisHashKey, sc.Key, &sc) @@ -81,7 +87,6 @@ func publishSystemConfigRAMInvalidation(ctx context.Context, key string) { } // InvalidateSystemConfigCache evicts one config key from local RAM and Redis. -// It also publishes cluster-wide RAM invalidation when Redis is available. func InvalidateSystemConfigCache(ctx context.Context, key string) error { ensureSystemConfigCacheListener() @@ -96,7 +101,6 @@ func InvalidateSystemConfigCache(ctx context.Context, key string) error { } // InvalidateAllSystemConfigCaches evicts all config entries from local RAM and Redis. -// It also publishes cluster-wide RAM invalidation when Redis is available. func InvalidateAllSystemConfigCaches(ctx context.Context) error { ensureSystemConfigCacheListener() diff --git a/internal/repository/template.go b/internal/repository/template.go new file mode 100644 index 00000000..2c826c5e --- /dev/null +++ b/internal/repository/template.go @@ -0,0 +1,59 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// ListTemplates returns all templates ordered by system flag and creation time. +func ListTemplates(ctx context.Context) ([]model.Template, error) { + var templates []model.Template + if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil { + return nil, err + } + return templates, nil +} + +// GetTemplateByKey loads a template by its key. +func GetTemplateByKey(ctx context.Context, key string) (model.Template, error) { + var tmpl model.Template + if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil { + return model.Template{}, err + } + return tmpl, nil +} + +// TemplateExistsByKey reports whether a template key is already taken. +func TemplateExistsByKey(ctx context.Context, key string) (bool, error) { + var existing model.Template + err := db.DB(ctx).Where("key = ?", key).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + if err != nil { + return false, err + } + return true, nil +} + +// CreateTemplate persists a new template. +func CreateTemplate(ctx context.Context, tmpl *model.Template) error { + return db.DB(ctx).Create(tmpl).Error +} + +// SaveTemplate updates an existing template. +func SaveTemplate(ctx context.Context, tmpl *model.Template) error { + return db.DB(ctx).Save(tmpl).Error +} + +// DeleteTemplate removes a template record. +func DeleteTemplate(ctx context.Context, tmpl *model.Template) error { + return db.DB(ctx).Delete(tmpl).Error +} \ No newline at end of file diff --git a/internal/repository/upload.go b/internal/repository/upload.go new file mode 100644 index 00000000..62fed841 --- /dev/null +++ b/internal/repository/upload.go @@ -0,0 +1,118 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "strings" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// UploadListFilter filters paginated upload queries. +type UploadListFilter struct { + UserID uint64 + Keyword string + Type string + Extension string + Page int + PageSize int +} + +// ListUploads returns paginated upload records matching the filter. +func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.Upload, error) { + query := db.DB(ctx).Model(&model.Upload{}). + Where("status != ?", model.UploadStatusDeleted) + + if filter.UserID != 0 { + query = query.Where("user_id = ?", filter.UserID) + } + if filter.Keyword != "" { + query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(filter.Keyword)+"%") + } + if filter.Type != "" { + query = query.Where("type = ?", filter.Type) + } + if filter.Extension != "" { + query = query.Where("extension = ?", strings.ToLower(filter.Extension)) + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return 0, nil, err + } + + var items []model.Upload + offset := (filter.Page - 1) * filter.PageSize + if err := query.Order("created_at DESC").Offset(offset).Limit(filter.PageSize).Find(&items).Error; err != nil { + return 0, nil, err + } + return total, items, nil +} + +// GetActiveUploadByID loads a non-deleted upload by ID. +func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) { + var upload model.Upload + if err := db.DB(ctx).Where("id = ? AND status != ?", id, model.UploadStatusDeleted).First(&upload).Error; err != nil { + return model.Upload{}, err + } + return upload, nil +} + +// SoftDeleteUpload marks an upload as deleted. +func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error { + return db.DB(ctx).Model(upload).Update("status", model.UploadStatusDeleted).Error +} + +// UpdateUpload applies partial field updates to an upload record. +func UpdateUpload(ctx context.Context, upload *model.Upload, updates map[string]any) error { + if len(updates) == 0 { + return nil + } + return db.DB(ctx).Model(upload).Updates(updates).Error +} + +// ListDistinctUploadTypes returns all distinct non-empty upload business types. +func ListDistinctUploadTypes(ctx context.Context) ([]string, error) { + var types []string + if err := db.DB(ctx).Model(&model.Upload{}). + Where("type IS NOT NULL AND type != ''"). + Distinct(). + Pluck("type", &types).Error; err != nil { + return nil, err + } + return types, nil +} + +// FindReusableUploadByHash finds an existing upload with the same hash and size. +func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (model.Upload, error) { + var existing model.Upload + err := db.DB(ctx). + Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, model.UploadStatusPending, model.UploadStatusUsed). + First(&existing).Error + return existing, err +} + +// CreateUpload persists a new upload record. +func CreateUpload(ctx context.Context, upload *model.Upload) error { + return db.DB(ctx).Create(upload).Error +} + +// ListUploadsByIDs returns active uploads matching the given IDs. +func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]model.Upload, error) { + var uploads []model.Upload + if err := db.DB(ctx). + Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed). + Find(&uploads).Error; err != nil { + return nil, err + } + return uploads, nil +} + +// UploadQuery returns a scoped GORM query for uploads. +func UploadQuery(ctx context.Context) *gorm.DB { + return db.DB(ctx).Model(&model.Upload{}) +} \ No newline at end of file diff --git a/internal/repository/upload_stat.go b/internal/repository/upload_stat.go new file mode 100644 index 00000000..7977e7ba --- /dev/null +++ b/internal/repository/upload_stat.go @@ -0,0 +1,20 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListUploadStats returns all upload statistics rows. +func ListUploadStats(ctx context.Context) ([]model.UploadStat, error) { + var stats []model.UploadStat + if err := db.DB(ctx).Find(&stats).Error; err != nil { + return nil, err + } + return stats, nil +} \ No newline at end of file diff --git a/internal/repository/user.go b/internal/repository/user.go new file mode 100644 index 00000000..277d3ca3 --- /dev/null +++ b/internal/repository/user.go @@ -0,0 +1,172 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// GetUserByID loads an active user by ID. +func GetUserByID(ctx context.Context, id uint64) (model.User, error) { + var user model.User + if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// GetUserByUsername loads a user by username. +func GetUserByUsername(ctx context.Context, username string) (model.User, error) { + var user model.User + if err := db.DB(ctx).Where("username = ?", username).First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// GetSystemUser loads the built-in system user, or returns a synthetic fallback. +func GetSystemUser(ctx context.Context) model.User { + var user model.User + if err := db.DB(ctx).Where("username = ?", "system").First(&user).Error; err == nil { + return user + } + return model.User{ + ID: 999, + Username: "system", + Nickname: "系统", + } +} + +// GetFirstAdminUser loads the earliest admin user. +func GetFirstAdminUser(ctx context.Context) (model.User, error) { + var user model.User + if err := db.DB(ctx).Where("is_admin = ?", true).Order("id asc").First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// AdminUserListFilter filters admin user list queries. +type AdminUserListFilter struct { + UserID *uint64 + Username string + Page int + PageSize int +} + +// ListAdminUsers returns paginated users for the admin console. +func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []model.User, error) { + query := db.DB(ctx).Model(&model.User{}) + if filter.UserID != nil { + query = query.Where("id = ?", *filter.UserID) + } + if filter.Username != "" { + query = query.Where("username LIKE ?", filter.Username+"%") + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return 0, nil, err + } + + var users []model.User + offset := (filter.Page - 1) * filter.PageSize + if err := query. + Select("id, username, nickname, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at"). + Order("id DESC"). + Offset(offset). + Limit(filter.PageSize). + Find(&users).Error; err != nil { + return 0, nil, err + } + return total, users, nil +} + +// GetAdminUserDetail loads full user profile fields for admin detail view. +func GetAdminUserDetail(ctx context.Context, id uint64) (model.User, error) { + var user model.User + if err := db.DB(ctx). + Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at"). + Where("id = ?", id). + First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// UserAdminFlags stores minimal user authorization flags. +type UserAdminFlags struct { + ID uint64 + IsAdmin bool +} + +// GetUserAdminFlags loads id and is_admin for authorization checks. +func GetUserAdminFlags(ctx context.Context, id uint64) (UserAdminFlags, error) { + var flags UserAdminFlags + if err := db.DB(ctx). + Model(&model.User{}). + Select("id, is_admin"). + Where("id = ?", id). + First(&flags).Error; err != nil { + return UserAdminFlags{}, err + } + return flags, nil +} + +// UpdateUserActive updates the is_active flag for a user. +func UpdateUserActive(ctx context.Context, id uint64, active bool) error { + return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).Update("is_active", active).Error +} + +// DeleteUserWithRelations removes a user and related access tokens / external accounts. +func DeleteUserWithRelations(ctx context.Context, id uint64) error { + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("user_id = ?", id).Delete(&model.AccessToken{}).Error; err != nil { + return err + } + if err := tx.Where("user_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil { + return err + } + return tx.Where("id = ?", id).Delete(&model.User{}).Error + }) +} + +// CountUsersByUsername returns how many users share the username. +func CountUsersByUsername(ctx context.Context, username string) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", username).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CountUsersByEmail returns how many users share the email. +func CountUsersByEmail(ctx context.Context, email string) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CreateUser persists a new user record. +func CreateUser(ctx context.Context, user *model.User) error { + return db.DB(ctx).Create(user).Error +} + +// ListUsersByIDs loads users matching the given IDs. +func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) { + if len(ids) == 0 { + return []model.User{}, nil + } + var users []model.User + if err := db.DB(ctx).Where("id IN ?", ids).Find(&users).Error; err != nil { + return nil, err + } + return users, nil +} \ No newline at end of file diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go index 73fb89b3..2ef6d7f9 100644 --- a/internal/router/middlewares.go +++ b/internal/router/middlewares.go @@ -15,6 +15,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" "github.com/gin-gonic/gin" @@ -72,8 +73,8 @@ func loggerMiddleware() gin.HandlerFunc { } func isOriginAllowed(ctx context.Context, origin string) bool { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || sc.Value == "" { + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) + if err != nil || sc.Value == "" { return false } allowedOrigins := strings.Split(sc.Value, ",") diff --git a/internal/router/middlewares_test.go b/internal/router/middlewares_test.go index 7352ee0d..c801fd5e 100644 --- a/internal/router/middlewares_test.go +++ b/internal/router/middlewares_test.go @@ -10,6 +10,7 @@ import ( "testing" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/gin-gonic/gin" ) @@ -21,7 +22,7 @@ func TestCORSMiddleware(t *testing.T) { gin.SetMode(gin.TestMode) clearConfigCache := func() { - if err := model.InvalidateAllSystemConfigCaches(context.Background()); err != nil { + if err := repository.InvalidateAllSystemConfigCaches(context.Background()); err != nil { t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err) } } diff --git a/internal/storage/config.go b/internal/storage/config.go index c9deb840..2fe84811 100644 --- a/internal/storage/config.go +++ b/internal/storage/config.go @@ -14,6 +14,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" ) @@ -109,8 +110,8 @@ func LoadConfig(ctx context.Context) (Config, error) { } func loadConfigByKey(ctx context.Context, key string, fallback Config) (Config, error) { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, key); err != nil { + sc, err := repository.GetSystemConfigByKey(ctx, key) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return fallback, nil } @@ -208,7 +209,7 @@ func upsertSystemConfig(ctx context.Context, tx *gorm.DB, key string, value any, FirstOrCreate(&sc).Error; err != nil { return err } - return model.InvalidateSystemConfigCache(ctx, key) + return repository.InvalidateSystemConfigCache(ctx, key) } // MergeMaskedSecrets restores unchanged secrets from the current configuration. diff --git a/internal/storage/storage.go b/internal/storage/storage.go index 51de6a2f..f2e1de9f 100644 --- a/internal/storage/storage.go +++ b/internal/storage/storage.go @@ -15,6 +15,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" ) @@ -123,8 +124,7 @@ func Active(ctx context.Context) (Driver, Backend, error) { return activeDriver, activeBackend, nil } - var sc model.SystemConfig - err := sc.GetByKey(ctx, model.ConfigKeyStorageConfig) + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyStorageConfig) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { return "", nil, err } diff --git a/internal/testhelper/test_helper.go b/internal/testhelper/test_helper.go index 23b4092f..c1b0fdad 100644 --- a/internal/testhelper/test_helper.go +++ b/internal/testhelper/test_helper.go @@ -11,6 +11,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/alicebob/miniredis/v2" "github.com/glebarez/sqlite" "github.com/redis/go-redis/v9" @@ -76,7 +77,7 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) // Cleanup function cleanup := func() { runExtraCleanups() - model.ResetSystemConfigRAMCacheForTest() + repository.ResetSystemConfigRAMCacheForTest() _ = redisClient.Close() mr.Close() // Reset database and Redis references @@ -316,6 +317,6 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) { if _, ok := publicKeys[config.Key]; ok { config.Visibility = model.ConfigVisibilityVisible } - _ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, config.Key, &config) + _ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, config.Key, &config) } } diff --git a/internal/util/custom_types.go b/internal/util/custom_types.go index 27f43580..f7dd38ff 100644 --- a/internal/util/custom_types.go +++ b/internal/util/custom_types.go @@ -2,6 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 +// Package util provides framework-agnostic helper types and HTTP utilities. package util import (