diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go index b5b08159..d18e2c8e 100644 --- a/internal/apps/admin/auth_source/routers.go +++ b/internal/apps/admin/auth_source/routers.go @@ -50,7 +50,7 @@ type ToggleAuthSourceRequest struct { func ListAuthSources(c *gin.Context) { sources, err := model.GetAuthSources(c.Request.Context()) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(sources)) @@ -72,7 +72,7 @@ func ListAuthSources(c *gin.Context) { func CreateAuthSource(c *gin.Context) { var req AuthSourceRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -88,7 +88,7 @@ func CreateAuthSource(c *gin.Context) { IconURL: req.IconURL, } if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } source.Sanitize() @@ -113,13 +113,13 @@ func CreateAuthSource(c *gin.Context) { func UpdateAuthSource(c *gin.Context) { id, err := parseSourceID(c) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } var req AuthSourceRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -140,7 +140,7 @@ func UpdateAuthSource(c *gin.Context) { } keepSecret := source.ClientSecret == "" if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -153,7 +153,7 @@ func UpdateAuthSource(c *gin.Context) { updated, err := model.GetAuthSourceByID(c.Request.Context(), id) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } updated.Sanitize() @@ -177,18 +177,18 @@ func UpdateAuthSource(c *gin.Context) { func ToggleAuthSource(c *gin.Context) { id, err := parseSourceID(c) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } var req ToggleAuthSourceRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } c.JSON(http.StatusOK, response.OKNil()) @@ -209,11 +209,11 @@ func ToggleAuthSource(c *gin.Context) { func DeleteAuthSource(c *gin.Context) { id, err := parseSourceID(c) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } c.JSON(http.StatusOK, response.OKNil()) diff --git a/internal/apps/admin/cache/routers.go b/internal/apps/admin/cache/routers.go index 97ca546a..3b98a37b 100644 --- a/internal/apps/admin/cache/routers.go +++ b/internal/apps/admin/cache/routers.go @@ -56,7 +56,7 @@ func GetCacheStatus(c *gin.Context) { func UpdateCacheConfig(c *gin.Context) { var req updateCacheConfigRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -64,19 +64,19 @@ func UpdateCacheConfig(c *gin.Context) { // Update Max Size if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } // Update Default TTL if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } // Update LRU Enabled if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -99,7 +99,7 @@ func UpdateCacheConfig(c *gin.Context) { // @Router /api/v1/admin/cache/clear [post] func ClearCache(c *gin.Context) { if err := diskcache.GetGlobalCache().Clear(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OKNil()) diff --git a/internal/apps/admin/db_manage/routers.go b/internal/apps/admin/db_manage/routers.go index 897fe9e3..5d2b8f78 100644 --- a/internal/apps/admin/db_manage/routers.go +++ b/internal/apps/admin/db_manage/routers.go @@ -213,7 +213,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) { func GetDBOverview(c *gin.Context) { gormDB := db.DB(c.Request.Context()) if gormDB == nil { - c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化")) + response.AbortInternal(c, "数据库未初始化") return } @@ -227,7 +227,7 @@ func GetDBOverview(c *gin.Context) { } if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -248,7 +248,7 @@ func GetDBOverview(c *gin.Context) { func ListDBTables(c *gin.Context) { gormDB := db.DB(c.Request.Context()) if gormDB == nil { - c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化")) + response.AbortInternal(c, "数据库未初始化") return } @@ -262,7 +262,7 @@ func ListDBTables(c *gin.Context) { } if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -273,13 +273,13 @@ func ListDBTables(c *gin.Context) { func GetDBTableData(c *gin.Context) { var req GetTableDataRequest if err := c.ShouldBindQuery(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } gormDB := db.DB(c.Request.Context()) if gormDB == nil { - c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化")) + response.AbortInternal(c, "数据库未初始化") return } @@ -288,7 +288,7 @@ func GetDBTableData(c *gin.Context) { var total int64 if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -303,7 +303,7 @@ func GetDBTableData(c *gin.Context) { rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows() if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } defer func() { @@ -312,13 +312,13 @@ func GetDBTableData(c *gin.Context) { cols, err := rows.Columns() if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } results, err := scanTableRows(rows, cols) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -449,19 +449,19 @@ func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (Ex func ExecuteSQL(c *gin.Context) { var req ExecuteSQLRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } gormDB := db.DB(c.Request.Context()) if gormDB == nil { - c.JSON(http.StatusInternalServerError, response.Err("数据库未初始化")) + response.AbortInternal(c, "数据库未初始化") return } trimmedSQL := strings.TrimSpace(req.SQL) if trimmedSQL == "" { - c.JSON(http.StatusBadRequest, response.Err("SQL 语句不能为空")) + response.AbortBadRequest(c, "SQL 语句不能为空") return } @@ -488,7 +488,7 @@ func ExecuteSQL(c *gin.Context) { } if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/admin/middlewares.go b/internal/apps/admin/middlewares.go index 06856557..3bd1c25f 100644 --- a/internal/apps/admin/middlewares.go +++ b/internal/apps/admin/middlewares.go @@ -5,8 +5,7 @@ package admin import ( - "net/http" - + "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" @@ -29,13 +28,13 @@ func LoginAdminRequired() gin.HandlerFunc { if tokenAuth, _ := util.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth { tokenAdmin, _ := util.GetFromContext[bool](c, oauth.TokenAdminKey) if !tokenAdmin { - c.AbortWithStatusJSON(http.StatusNotFound, gin.H{"error_msg": TokenAdminRequired, "data": nil}) + response.AbortNotFound(c, TokenAdminRequired) return } } if !user.IsAdmin { - c.AbortWithStatusJSON(http.StatusNotFound, gin.H{"error_msg": AdminRequired, "data": nil}) + response.AbortNotFound(c, AdminRequired) return } diff --git a/internal/apps/admin/push/channels.go b/internal/apps/admin/push/channels.go index 34c32a7e..36ab1489 100644 --- a/internal/apps/admin/push/channels.go +++ b/internal/apps/admin/push/channels.go @@ -41,7 +41,7 @@ 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(channels)) @@ -71,7 +71,7 @@ type CreateChannelRequest struct { func CreateChannel(c *gin.Context) { var req CreateChannelRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -79,11 +79,11 @@ func CreateChannel(c *gin.Context) { var count int64 if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", req.Name).Count(&count).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if count > 0 { - c.JSON(http.StatusBadRequest, response.Err("channel name already exists")) + response.AbortBadRequest(c, "channel name already exists") return } @@ -98,12 +98,12 @@ func CreateChannel(c *gin.Context) { } if err := channel.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(ctx).Create(&channel).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -138,13 +138,13 @@ func UpdateChannel(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("invalid channel id")) + response.AbortBadRequest(c, "invalid channel id") return } var req UpdateChannelRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -153,10 +153,10 @@ func UpdateChannel(c *gin.Context) { var channel model.PushChannel if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err("channel not found")) + response.AbortNotFound(c, "channel not found") return } - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -168,12 +168,12 @@ func UpdateChannel(c *gin.Context) { channel.Enabled = req.Enabled if err := channel.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(ctx).Save(&channel).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -196,7 +196,7 @@ func DeleteChannel(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("invalid channel id")) + response.AbortBadRequest(c, "invalid channel id") return } @@ -204,15 +204,15 @@ func DeleteChannel(c *gin.Context) { var channel model.PushChannel if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err("channel not found")) + response.AbortNotFound(c, "channel not found") return } - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if err := db.DB(ctx).Delete(&channel).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -245,7 +245,7 @@ type TestChannelRequest struct { func TestChannel(c *gin.Context) { var req TestChannelRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -255,7 +255,7 @@ func TestChannel(c *gin.Context) { if req.Name != "" { var channel model.PushChannel if err := db.DB(ctx).Where("name = ?", req.Name).First(&channel).Error; err != nil { - c.JSON(http.StatusBadRequest, response.Err("channel not found")) + response.AbortBadRequest(c, "channel not found") return } url = channel.URL @@ -284,7 +284,7 @@ func TestChannel(c *gin.Context) { } if err := tempChannel.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } url = tempChannel.URL @@ -342,7 +342,7 @@ func TestChannel(c *gin.Context) { } if err := enqueuePushTask(ctx, payload); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } diff --git a/internal/apps/admin/push/routers.go b/internal/apps/admin/push/routers.go index 8858c5f7..1cc50d0c 100644 --- a/internal/apps/admin/push/routers.go +++ b/internal/apps/admin/push/routers.go @@ -76,7 +76,7 @@ func ListEvents(c *gin.Context) { var events []model.PushEvent if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(events)) @@ -166,7 +166,7 @@ func getEventInfo(req CreateEventRequest) (string, string, []byte, error) { func CreateEvent(c *gin.Context) { var req CreateEventRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -174,18 +174,18 @@ func CreateEvent(c *gin.Context) { eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if count > 0 { - c.JSON(http.StatusBadRequest, response.Err("this notification event is already configured")) + response.AbortBadRequest(c, "this notification event is already configured") return } @@ -196,7 +196,7 @@ func CreateEvent(c *gin.Context) { } else { var tempMap map[string]any if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { - c.JSON(http.StatusBadRequest, response.Err("custom template is not a valid JSON format")) + response.AbortBadRequest(c, "custom template is not a valid JSON format") return } } @@ -222,12 +222,12 @@ func CreateEvent(c *gin.Context) { } if err := event.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(ctx).Create(&event).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -250,7 +250,7 @@ func DeleteEvent(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("invalid event id")) + response.AbortBadRequest(c, "invalid event id") return } @@ -258,15 +258,15 @@ func DeleteEvent(c *gin.Context) { var event model.PushEvent if err := db.DB(ctx).First(&event, id).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err("notification event not found")) + response.AbortNotFound(c, "notification event not found") } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } if err := db.DB(ctx).Delete(&event).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -291,22 +291,22 @@ func UpdateEvent(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("invalid event id")) + response.AbortBadRequest(c, "invalid event id") return } var req UpdateEventRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } var event model.PushEvent if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err("notification event not found")) + response.AbortNotFound(c, "notification event not found") } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } @@ -317,12 +317,12 @@ func UpdateEvent(c *gin.Context) { event.Enabled = req.Enabled if err := event.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(c.Request.Context()).Save(&event).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -345,27 +345,27 @@ func ToggleEvent(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("invalid event id")) + response.AbortBadRequest(c, "invalid event id") return } var event model.PushEvent if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err("notification event not found")) + response.AbortNotFound(c, "notification event not found") } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } event.Enabled = !event.Enabled if event.Enabled && len(event.Channels) == 0 { - c.JSON(http.StatusBadRequest, response.Err("cannot enable event without any push channels configured")) + 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -420,14 +420,14 @@ func ListHistories(c *gin.Context) { var total int64 if err := query.Count(&total).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -450,19 +450,19 @@ func ListHistories(c *gin.Context) { func TestPush(c *gin.Context) { var req TestPushRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } pusher, err := push.GetPusher(req.Config.Channel) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 校验配置 if err := pusher.ValidateConfig(req.Config); err != nil { - c.JSON(http.StatusBadRequest, response.Err(fmt.Sprintf("validation failed: %v", err))) + response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err)) return } @@ -494,7 +494,7 @@ func TestPush(c *gin.Context) { err = pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil) if err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/admin/status/routers.go b/internal/apps/admin/status/routers.go index 2890ada8..513d701b 100644 --- a/internal/apps/admin/status/routers.go +++ b/internal/apps/admin/status/routers.go @@ -290,7 +290,7 @@ func exportSQLite(c *gin.Context) { f, err := os.Open(path) //nolint:gosec // path is loaded from server startup configuration, not user input if err != nil { - c.JSON(http.StatusInternalServerError, response.Err("无法打开数据库文件: "+err.Error())) + response.AbortInternal(c, "无法打开数据库文件: "+err.Error()) return } defer func() { @@ -301,7 +301,7 @@ func exportSQLite(c *gin.Context) { fi, err := f.Stat() if err != nil { - c.JSON(http.StatusInternalServerError, response.Err("无法读取数据库文件信息: "+err.Error())) + response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error()) return } @@ -319,7 +319,7 @@ func exportPostgres(c *gin.Context) { // 检查 pg_dump 是否可用 pgDumpPath, err := exec.LookPath("pg_dump") if err != nil { - c.JSON(http.StatusInternalServerError, response.Err("pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具")) + response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具") return } diff --git a/internal/apps/admin/system_config/routers.go b/internal/apps/admin/system_config/routers.go index 1efc0c6c..89f15096 100644 --- a/internal/apps/admin/system_config/routers.go +++ b/internal/apps/admin/system_config/routers.go @@ -60,17 +60,17 @@ type UpdateSystemConfigRequest struct { func CreateSystemConfig(c *gin.Context) { var req CreateSystemConfigRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 检查配置键是否已存在 var existing model.SystemConfig if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil { - c.JSON(http.StatusBadRequest, response.Err(ConfigKeyExists)) + response.AbortBadRequest(c, ConfigKeyExists) return } else if !errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -90,7 +90,7 @@ func CreateSystemConfig(c *gin.Context) { return nil }); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -124,7 +124,7 @@ func ListSystemConfigs(c *gin.Context) { var configs []model.SystemConfig if err := query.Find(&configs).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -152,9 +152,9 @@ 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 { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err(SystemConfigNotFound)) + response.AbortNotFound(c, SystemConfigNotFound) } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } @@ -183,7 +183,7 @@ func GetSystemConfig(c *gin.Context) { func UpdateSystemConfig(c *gin.Context) { var req UpdateSystemConfigRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -193,9 +193,9 @@ func UpdateSystemConfig(c *gin.Context) { var config model.SystemConfig if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, response.Err(SystemConfigNotFound)) + response.AbortNotFound(c, SystemConfigNotFound) } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } @@ -209,7 +209,7 @@ func UpdateSystemConfig(c *gin.Context) { validatedVal, err := validateAndMergeStorageConfig(c.Request.Context(), req.Value, config.Value) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } req.Value = validatedVal @@ -242,7 +242,7 @@ func UpdateSystemConfig(c *gin.Context) { return nil }); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -339,7 +339,7 @@ type TestSMTPResponse struct { func TestSMTP(c *gin.Context) { var req TestSMTPRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/admin/task/routers.go b/internal/apps/admin/task/routers.go index 580a2261..c2f98949 100644 --- a/internal/apps/admin/task/routers.go +++ b/internal/apps/admin/task/routers.go @@ -65,13 +65,13 @@ type DispatchTaskRequest struct { func DispatchTask(c *gin.Context) { var req DispatchTaskRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } meta := task.GetTaskMeta(req.TaskType) if meta == nil { - c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType)) + response.AbortBadRequest(c, InvalidTaskType) return } @@ -82,13 +82,13 @@ func DispatchTask(c *gin.Context) { validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual") if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", TaskDispatchFailed, err))) + response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err)) return } @@ -112,7 +112,7 @@ func DispatchTask(c *gin.Context) { func ListTaskExecutions(c *gin.Context) { var req model.ListTaskExecutionsRequest if err := c.ShouldBindQuery(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -124,7 +124,7 @@ func ListTaskExecutions(c *gin.Context) { executions, total, err := model.ListTaskExecutions(c.Request.Context(), req) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -152,13 +152,13 @@ func ListTaskExecutions(c *gin.Context) { func GetTaskExecution(c *gin.Context) { id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(admin.InvalidTaskExecutionID)) + response.AbortBadRequest(c, admin.InvalidTaskExecutionID) return } execution, err := model.GetTaskExecutionByID(c.Request.Context(), id) if err != nil { - c.JSON(http.StatusNotFound, response.Err(TaskNotFound)) + response.AbortNotFound(c, TaskNotFound) return } @@ -182,7 +182,7 @@ func GetTaskExecution(c *gin.Context) { func RetryTask(c *gin.Context) { id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(admin.InvalidTaskExecutionID)) + response.AbortBadRequest(c, admin.InvalidTaskExecutionID) return } @@ -191,11 +191,11 @@ func RetryTask(c *gin.Context) { errMsg := err.Error() switch { case strings.Contains(errMsg, "不存在"): - c.JSON(http.StatusNotFound, response.Err(errMsg)) + response.AbortNotFound(c, errMsg) case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"): - c.JSON(http.StatusBadRequest, response.Err(errMsg)) + response.AbortBadRequest(c, errMsg) default: - c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", TaskRetryFailed, err))) + response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err)) } return } @@ -216,7 +216,7 @@ func RetryTask(c *gin.Context) { func ListSchedules(c *gin.Context) { schedules, err := model.ListSchedules(c.Request.Context()) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(schedules)) @@ -248,20 +248,20 @@ type CreateScheduleRequest struct { func CreateSchedule(c *gin.Context) { var req CreateScheduleRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 校验 Cron 表达式 if _, err := cron.ParseStandard(req.Cron); err != nil { - c.JSON(http.StatusBadRequest, response.Err(InvalidCronExpression)) + response.AbortBadRequest(c, InvalidCronExpression) return } // 校验关联的异步任务类型 meta := task.GetTaskMeta(req.TaskType) if meta == nil { - c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType)) + response.AbortBadRequest(c, InvalidTaskType) return } @@ -272,7 +272,7 @@ func CreateSchedule(c *gin.Context) { } validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -285,7 +285,7 @@ func CreateSchedule(c *gin.Context) { } if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))) + response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) return } @@ -325,33 +325,33 @@ type UpdateScheduleRequest struct { func UpdateSchedule(c *gin.Context) { id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("无效的定时任务ID")) + response.AbortBadRequest(c, "无效的定时任务ID") return } var req UpdateScheduleRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 检查定时任务是否存在 schedule, err := model.GetScheduleByID(c.Request.Context(), id) if err != nil { - c.JSON(http.StatusNotFound, response.Err(ScheduleNotFound)) + response.AbortNotFound(c, ScheduleNotFound) return } // 校验 Cron 表达式 if _, err := cron.ParseStandard(req.Cron); err != nil { - c.JSON(http.StatusBadRequest, response.Err(InvalidCronExpression)) + response.AbortBadRequest(c, InvalidCronExpression) return } // 校验关联的异步任务类型 meta := task.GetTaskMeta(req.TaskType) if meta == nil { - c.JSON(http.StatusBadRequest, response.Err(InvalidTaskType)) + response.AbortBadRequest(c, InvalidTaskType) return } @@ -362,7 +362,7 @@ func UpdateSchedule(c *gin.Context) { } validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -373,7 +373,7 @@ func UpdateSchedule(c *gin.Context) { schedule.IsActive = *req.IsActive if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))) + response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) return } @@ -401,12 +401,12 @@ func UpdateSchedule(c *gin.Context) { func DeleteSchedule(c *gin.Context) { id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusBadRequest, response.Err("无效的定时任务ID")) + response.AbortBadRequest(c, "无效的定时任务ID") return } if err := model.DeleteSchedule(c.Request.Context(), id); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))) + response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err)) return } diff --git a/internal/apps/admin/template/routers.go b/internal/apps/admin/template/routers.go index 7580aa59..09d0dc0d 100644 --- a/internal/apps/admin/template/routers.go +++ b/internal/apps/admin/template/routers.go @@ -49,17 +49,17 @@ type UpdateTemplateRequest struct { func CreateTemplate(c *gin.Context) { var req CreateTemplateRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 检查模板 Key 是否已存在 var existing model.Template if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil { - c.JSON(http.StatusBadRequest, response.Err(TemplateKeyExists)) + response.AbortBadRequest(c, TemplateKeyExists) return } else if !errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -74,12 +74,12 @@ func CreateTemplate(c *gin.Context) { } if err := tmpl.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(c.Request.Context()).Create(&tmpl).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -100,7 +100,7 @@ func CreateTemplate(c *gin.Context) { 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -124,9 +124,9 @@ 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) { - c.JSON(http.StatusNotFound, response.Err(TemplateNotFound)) + response.AbortNotFound(c, TemplateNotFound) } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } @@ -153,7 +153,7 @@ func GetTemplate(c *gin.Context) { func UpdateTemplate(c *gin.Context) { var req UpdateTemplateRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -163,9 +163,9 @@ func UpdateTemplate(c *gin.Context) { 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) { - c.JSON(http.StatusNotFound, response.Err(TemplateNotFound)) + response.AbortNotFound(c, TemplateNotFound) } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } @@ -177,12 +177,12 @@ func UpdateTemplate(c *gin.Context) { tmpl.Description = req.Description if err := tmpl.Validate(); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := db.DB(c.Request.Context()).Save(&tmpl).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -210,21 +210,21 @@ func DeleteTemplate(c *gin.Context) { 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) { - c.JSON(http.StatusNotFound, response.Err(TemplateNotFound)) + response.AbortNotFound(c, TemplateNotFound) } else { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) } return } // 限制系统模板删除 if tmpl.IsSystem { - c.JSON(http.StatusBadRequest, response.Err(SystemTemplateCannotDelete)) + response.AbortBadRequest(c, SystemTemplateCannotDelete) return } if err := db.DB(c.Request.Context()).Delete(&tmpl).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } diff --git a/internal/apps/admin/updater/routers.go b/internal/apps/admin/updater/routers.go index 3e5ba0da..d96c447e 100644 --- a/internal/apps/admin/updater/routers.go +++ b/internal/apps/admin/updater/routers.go @@ -27,7 +27,7 @@ func GetUpdateStatus(c *gin.Context) { status, _, err := defaultManager.status(c.Request.Context()) if err != nil { logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err) - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(status)) @@ -49,7 +49,7 @@ func ApplyUpdate(c *gin.Context) { executable, stagedBinary, err := defaultManager.prepareUpgrade(c.Request.Context()) if err != nil { logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err) - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index e6ef0081..509882d9 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -59,7 +59,7 @@ type listUsersResponse struct { func parseUserID(c *gin.Context) (uint64, bool) { id, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil || id == 0 { - c.JSON(http.StatusBadRequest, response.Err(userNotFound)) + response.AbortBadRequest(c, userNotFound) return 0, false } return id, true @@ -102,7 +102,7 @@ func toUser(u model.User) user { func ListUsers(c *gin.Context) { var req listUsersRequest if err := c.ShouldBindQuery(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -122,7 +122,7 @@ func ListUsers(c *gin.Context) { } if err := query.Count(&total).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -134,7 +134,7 @@ func ListUsers(c *gin.Context) { Offset(offset). Limit(req.PageSize). Find(&modelUsers).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -176,10 +176,10 @@ func GetUser(c *gin.Context) { Where("id = ?", id). First(&targetUser).Error; err != nil { if err == gorm.ErrRecordNotFound { - c.JSON(http.StatusNotFound, response.Err(userNotFound)) + response.AbortNotFound(c, userNotFound) return } - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } @@ -210,7 +210,7 @@ type updateUserStatusRequest struct { func UpdateUserStatus(c *gin.Context) { var req updateUserStatusRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -229,15 +229,15 @@ func UpdateUserStatus(c *gin.Context) { Where("id = ?", id). First(&targetUser).Error; err != nil { if err == gorm.ErrRecordNotFound { - c.JSON(http.StatusNotFound, response.Err(userNotFound)) + response.AbortNotFound(c, userNotFound) return } - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if !req.IsActive && targetUser.IsAdmin { - c.JSON(http.StatusForbidden, response.Err(cannotDisable)) + response.AbortForbidden(c, cannotDisable) return } @@ -245,7 +245,7 @@ func UpdateUserStatus(c *gin.Context) { Model(&model.User{}). Where("id = ?", id). Update("is_active", req.IsActive).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(updateUserFailed)) + response.AbortInternal(c, updateUserFailed) return } @@ -274,7 +274,7 @@ func DeleteUser(c *gin.Context) { currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) if currUser != nil && currUser.ID == id { - c.JSON(http.StatusForbidden, response.Err(cannotDeleteSelf)) + response.AbortForbidden(c, cannotDeleteSelf) return } @@ -288,15 +288,15 @@ func DeleteUser(c *gin.Context) { Where("id = ?", id). First(&targetUser).Error; err != nil { if err == gorm.ErrRecordNotFound { - c.JSON(http.StatusNotFound, response.Err(userNotFound)) + response.AbortNotFound(c, userNotFound) return } - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if targetUser.IsAdmin { - c.JSON(http.StatusForbidden, response.Err(cannotDelete)) + response.AbortForbidden(c, cannotDelete) return } @@ -309,7 +309,7 @@ func DeleteUser(c *gin.Context) { } return tx.Where("id = ?", id).Delete(&model.User{}).Error }); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(deleteUserFailed)) + response.AbortInternal(c, deleteUserFailed) return } @@ -343,7 +343,7 @@ type createUserRequest struct { func CreateUser(c *gin.Context) { var req createUserRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -353,36 +353,36 @@ func CreateUser(c *gin.Context) { req.Email = strings.TrimSpace(req.Email) if req.Username == "" { - c.JSON(http.StatusBadRequest, response.Err(usernameRequired)) + response.AbortBadRequest(c, usernameRequired) return } if req.Email == "" { - c.JSON(http.StatusBadRequest, response.Err(emailRequired)) + response.AbortBadRequest(c, emailRequired) return } if len(req.Password) < minPasswordLength { - c.JSON(http.StatusBadRequest, response.Err(passwordTooShort)) + 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if count > 0 { - c.JSON(http.StatusBadRequest, response.Err(usernameExists)) + 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 { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if emailCount > 0 { - c.JSON(http.StatusBadRequest, response.Err(emailExists)) + response.AbortBadRequest(c, emailExists) return } @@ -400,12 +400,12 @@ func CreateUser(c *gin.Context) { } if err := newUser.SetEncryptedPassword(req.Password); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } if err := db.DB(ctx).Create(&newUser).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } diff --git a/internal/apps/cap/middleware.go b/internal/apps/cap/middleware.go index dfe97f78..e8039bf4 100644 --- a/internal/apps/cap/middleware.go +++ b/internal/apps/cap/middleware.go @@ -4,8 +4,6 @@ package cap import ( - "net/http" - "github.com/gin-gonic/gin" "github.com/Rain-kl/Wavelet/internal/common/response" @@ -21,13 +19,13 @@ func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc { token := c.GetHeader("X-Cap-Token") if token == "" { - c.AbortWithStatusJSON(http.StatusUnauthorized, response.Err(errCapTokenMissing)) + response.AbortUnauthorized(c, errCapTokenMissing) return } valid, err := mgr.VerifyToken(c.Request.Context(), token, scope) if err != nil || !valid { - c.AbortWithStatusJSON(http.StatusUnauthorized, response.Err(errCapTokenInvalidOrExpired)) + response.AbortUnauthorized(c, errCapTokenInvalidOrExpired) return } diff --git a/internal/apps/config/routers.go b/internal/apps/config/routers.go index 8f693cb0..7226e7dc 100644 --- a/internal/apps/config/routers.go +++ b/internal/apps/config/routers.go @@ -24,7 +24,7 @@ func GetPublicConfig(c *gin.Context) { ctx := c.Request.Context() configs, err := model.ListVisibleSystemConfigs(ctx) if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } diff --git a/internal/apps/oauth/auth_source_resolver.go b/internal/apps/oauth/auth_source_resolver.go new file mode 100644 index 00000000..58dc7218 --- /dev/null +++ b/internal/apps/oauth/auth_source_resolver.go @@ -0,0 +1,117 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "errors" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" +) + +func isOIDCLoginEnabled(ctx context.Context) bool { + enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + if err != nil { + return true + } + return enabled +} + +func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) { + name := strings.TrimSpace(strings.ToLower(sourceName)) + if name == "" { + sources, err := model.GetActiveAuthSources(ctx) + if err != nil { + return nil, err + } + if len(sources) == 0 { + return nil, errors.New(errNoActiveAuthSource) + } + return &sources[0], nil + } + return model.GetAuthSourceByName(ctx, name) +} + +func activeLoginSources(ctx context.Context) []AuthSourceView { + enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + if err == nil && !enabled { + return nil + } + + dbSources, err := model.GetActiveAuthSources(ctx) + if err != nil { + return nil + } + sources := make([]AuthSourceView, 0, len(dbSources)) + for _, source := range dbSources { + sources = append(sources, AuthSourceView{ + ID: source.ID, + Name: source.Name, + Type: source.Type, + DisplayName: source.DisplayName, + IsActive: source.IsActive, + IconURL: source.IconURL, + ClientSecretConfigured: source.ClientSecretConfigured, + }) + } + return sources +} + +func getFrontendLoginRedirectURL(ctx context.Context) (string, error) { + var sc model.SystemConfig + if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" { + return "", errors.New(errServerAddressMissing) + } + return strings.TrimRight(sc.Value, "/") + "/login", nil +} + +func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { + if source == nil { + return nil, nil, errors.New(errAuthSourceRequired) + } + + if source.OpenIDDiscoveryURL == "" { + return nil, nil, errors.New(errDiscoveryURLRequired) + } + + // Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake) + issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/") + issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") + issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") + + // 使用进程级缓存获取 provider,避免每次调用都向 issuer 发起 + // /.well-known/openid-configuration HTTP 请求。 + provider, err := globalOIDCProviderCache.get(ctx, issuer) + if err != nil { + return nil, nil, err + } + verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID}) + scopes := strings.Fields(source.Scopes) + if len(scopes) == 0 { + scopes = []string{oidc.ScopeOpenID, "profile", "email"} + } + if !containsScope(scopes, oidc.ScopeOpenID) { + scopes = append([]string{oidc.ScopeOpenID}, scopes...) + } + + return &oauth2.Config{ + ClientID: source.ClientID, + ClientSecret: source.ClientSecret, + RedirectURL: redirectURL, + Scopes: scopes, + Endpoint: provider.Endpoint(), + }, verifier, nil +} + +func containsScope(scopes []string, scope string) bool { + for _, item := range scopes { + if item == scope { + return true + } + } + return false +} diff --git a/internal/apps/oauth/handler_authorize.go b/internal/apps/oauth/handler_authorize.go new file mode 100644 index 00000000..42ac36b4 --- /dev/null +++ b/internal/apps/oauth/handler_authorize.go @@ -0,0 +1,173 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "fmt" + "net/http" + "strings" + + "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/model" + "github.com/coreos/go-oidc/v3/oidc" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "github.com/google/uuid" +) + +// GetLoginURL 获取登录授权地址 +// @Summary 获取登录授权地址 +// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 +// @Tags oauth +// @Produce json +// @Param source query string false "认证源名称,为空使用第一个启用的认证源" +// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL" +// @Failure 400 {object} response.Any "认证源不存在或未配置" +// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败" +// @Router /api/v1/oauth/login [get] +func GetLoginURL(c *gin.Context) { + ctx := c.Request.Context() + if !isOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + + source, err := resolveAuthSource(ctx, c.Query("source")) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + + session := sessions.Default(c) + token, isNew := ensureSessionToken(session) + if isNew { + if err := session.Save(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + userID := GetUserIDFromSession(session) + sessionHash := hashSessionToken(token) + + state := uuid.NewString() + payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ + SourceName: source.Name, + Purpose: OAuthPurposeLogin, + UserID: userID, + SessionHash: sessionHash, + }) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + + authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) +} + +func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) { + redirectURL, err := getFrontendLoginRedirectURL(ctx) + if err != nil { + return "", err + } + authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return "", err + } + if verifier != nil { + return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil + } + return authConfig.AuthCodeURL(state), nil +} + +// Authorize 发起指定认证源授权 +// @Summary 发起指定认证源授权 +// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。 +// @Tags oauth +// @Produce json +// @Param source path string true "认证源名称" +// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login" +// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL" +// @Failure 400 {object} response.Any "认证源不存在或未启用" +// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败" +// @Router /api/v1/oauth/{source}/authorize [get] +func Authorize(c *gin.Context) { + ctx := c.Request.Context() + if !isOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + + source, err := resolveAuthSource(ctx, c.Param("source")) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose"))) + if purpose != OAuthPurposeBind { + purpose = OAuthPurposeLogin + } + + session := sessions.Default(c) + userID := GetUserIDFromSession(session) + if purpose == OAuthPurposeBind && userID == 0 { + response.AbortUnauthorized(c, common.UnAuthorized) + return + } + + token, isNew := ensureSessionToken(session) + if isNew { + if err := session.Save(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + sessionHash := hashSessionToken(token) + + state := uuid.NewString() + payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ + SourceName: source.Name, + Purpose: purpose, + UserID: userID, + SessionHash: sessionHash, + }) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + + authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) +} \ No newline at end of file diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go new file mode 100644 index 00000000..04a29b3b --- /dev/null +++ b/internal/apps/oauth/handler_callback.go @@ -0,0 +1,226 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "errors" + "fmt" + "net/http" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events" + "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/model" + "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +// Callback OAuth 回调处理 +// @Summary OAuth 回调处理 +// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。 +// @Tags oauth +// @Accept json +// @Produce json +// @Param request body oauth.CallbackRequest true "回调请求参数" +// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功" +// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误" +// @Failure 401 {object} response.Any "绑定场景未登录" +// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误" +// @Router /api/v1/oauth/callback [post] +func Callback(c *gin.Context) { + var req CallbackRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + ctx := c.Request.Context() + stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)) + payloadRaw, err := db.Redis.Get(ctx, stateKey).Result() + if err != nil { + response.AbortBadRequest(c, errInvalidState) + return + } + _ = db.Redis.Del(ctx, stateKey) + + payload, err := decodeOAuthStatePayload(payloadRaw) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + session := sessions.Default(c) + currentUserID := GetUserIDFromSession(session) + + if payload.Purpose == OAuthPurposeBind && currentUserID == 0 { + response.AbortUnauthorized(c, common.UnAuthorized) + return + } + + token, ok := session.Get(SessionTokenKey).(string) + if !ok || token == "" { + response.AbortBadRequest(c, "invalid session context") + return + } + + if hashSessionToken(token) != payload.SessionHash { + response.AbortBadRequest(c, "session mismatch for oauth state") + return + } + + if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID { + response.AbortBadRequest(c, "user context mismatch for oauth binding") + return + } + + if !isOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + + source, err := resolveAuthSource(ctx, payload.SourceName) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, errAuthSourceDisabled) + return + } + + redirectURL, err := getFrontendLoginRedirectURL(ctx) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := normalizeOAuthUserInfo(userInfo); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + if userInfo.Sub == "" { + userInfo.Sub = userInfo.Username + } + + if payload.Purpose == OAuthPurposeBind { + handleCallbackBind(ctx, c, source, userInfo) + return + } + + handleCallbackLogin(ctx, c, source, userInfo) +} + +// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程 +func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { + userID := GetUserIDFromContext(c) + if userID == 0 { + response.AbortUnauthorized(c, common.UnAuthorized) + return + } + var user model.User + if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: source.ID, + UserID: user.ID, + ExternalID: userInfo.Sub, + ExternalUsername: userInfo.Username, + Email: userInfo.Email, + }); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + user.LastLoginAt = time.Now() + _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error + c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound"))) +} + +// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册) +func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { + var user model.User + + account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub) + switch { + case err == nil: + if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil { + response.AbortInternal(c, err.Error()) + return + } + case errors.Is(err, gorm.ErrRecordNotFound): + newUser, ok := handleCallbackRegister(ctx, c, source, userInfo) + if !ok { + return + } + user = newUser + default: + response.AbortInternal(c, err.Error()) + return + } + + user.LastLoginAt = time.Now() + _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error + if err := setLoginSession(ctx, c, &user); err != nil { + response.AbortInternal(c, err.Error()) + return + } + + logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) + + custom_events.TriggerAdminLoginEvent(ctx, &user, c.ClientIP()) + + c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in"))) +} + +// 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) + if regErr != nil { + registrationEnabled = true + } + + if !registrationEnabled { + c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind"))) + return model.User{}, false + } + + username, uniqueErr := uniqueUsername(ctx, userInfo.Username) + if uniqueErr != nil { + response.AbortInternal(c, uniqueErr.Error()) + return model.User{}, false + } + userInfo.Username = username + + var user model.User + if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil { + response.AbortInternal(c, err.Error()) + return model.User{}, false + } + if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: source.ID, + UserID: user.ID, + ExternalID: userInfo.Sub, + ExternalUsername: userInfo.Username, + Email: userInfo.Email, + }); err != nil { + response.AbortBadRequest(c, err.Error()) + return model.User{}, false + } + 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/handler_external_accounts.go b/internal/apps/oauth/handler_external_accounts.go new file mode 100644 index 00000000..7f766154 --- /dev/null +++ b/internal/apps/oauth/handler_external_accounts.go @@ -0,0 +1,65 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "net/http" + "strconv" + "strings" + + "github.com/Rain-kl/Wavelet/internal/common" + "github.com/Rain-kl/Wavelet/internal/common/response" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +// ListExternalAccounts 获取当前用户的外部帐号绑定列表 +// @Summary 获取外部帐号列表 +// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表" +// @Failure 401 {object} response.Any "未登录" +// @Failure 500 {object} response.Any "内部错误" +// @Router /api/v1/oauth/external-accounts [get] +func ListExternalAccounts(c *gin.Context) { + userID := GetUserIDFromContext(c) + accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(accounts)) +} + +// DeleteExternalAccount 解除外部帐号绑定 +// @Summary 解除外部帐号绑定 +// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Param id path uint64 true "外部帐号绑定记录 ID" +// @Success 200 {object} response.Any{data=string} "解除绑定成功" +// @Failure 400 {object} response.Any "ID 无效或解除失败" +// @Failure 401 {object} response.Any "未登录" +// @Router /api/v1/oauth/external-accounts/{id}/delete [post] +func DeleteExternalAccount(c *gin.Context) { + userID := GetUserIDFromContext(c) + if userID == 0 { + response.AbortUnauthorized(c, common.UnAuthorized) + return + } + rawID := strings.TrimSpace(c.Param("id")) + id, err := strconv.ParseUint(rawID, 10, 64) + if err != nil || id == 0 { + response.AbortBadRequest(c, errInvalidExternalAccountBindingID) + return + } + if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); 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/oauth/handler_sources.go b/internal/apps/oauth/handler_sources.go new file mode 100644 index 00000000..fa1c628a --- /dev/null +++ b/internal/apps/oauth/handler_sources.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "net/http" + + "github.com/Rain-kl/Wavelet/internal/common/response" + "github.com/gin-gonic/gin" +) + +// GetLoginSources 获取可用登录源列表 +// @Summary 获取可用登录源 +// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用 +// @Tags oauth +// @Produce json +// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表" +// @Router /api/v1/oauth/sources [get] +func GetLoginSources(c *gin.Context) { + c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context()))) +} diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 8f4a4f68..5eacbbda 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -7,9 +7,9 @@ package oauth import ( "context" "errors" - "net/http" "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/model" "github.com/Rain-kl/Wavelet/internal/util" @@ -111,7 +111,7 @@ func LoginRequired() gin.HandlerFunc { user, err := GetUserFromRequest(c) if err != nil { - c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil}) + response.AbortUnauthorized(c, common.UnAuthorized) return } @@ -130,7 +130,7 @@ func LoginRequired() gin.HandlerFunc { func DisallowTokenAuth() gin.HandlerFunc { return func(c *gin.Context) { if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth { - c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error_msg": ErrTokenAuthNotAllowed, "data": nil}) + response.AbortForbidden(c, ErrTokenAuthNotAllowed) return } c.Next() diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 5e2a9468..6d58fbf1 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/testhelper" "github.com/Rain-kl/Wavelet/internal/util" ) @@ -127,6 +128,12 @@ func (m *mockRedisClient) HGet(ctx context.Context, key string, field string) *r return cmd } +func (m *mockRedisClient) Subscribe(ctx context.Context, channels ...string) *redis.PubSub { + return redis.NewClient(&redis.Options{ + Addr: "127.0.0.1:0", + }).Subscribe(ctx, channels...) +} + type mockRoundTripper struct { roundTripFunc func(req *http.Request) (*http.Response, error) } @@ -332,9 +339,7 @@ func mockContextMiddleware(mockClient *http.Client) gin.HandlerFunc { } func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() - r.Use(gin.Recovery()) + r := testhelper.NewTestGinEngine(gin.Recovery()) // Inject context mock middleware r.Use(mockContextMiddleware(mockClient)) @@ -1199,7 +1204,7 @@ func TestSystemUserBlockedByMiddleware(t *testing.T) { // 3. 设置全局测试数据库连接并构建测试路由组 db.SetDB(dbConn) - rProtected := gin.New() + rProtected := testhelper.NewTestGinEngine() store := cookie.NewStore([]byte("secret")) rProtected.Use(sessions.Sessions("mysession", store)) rProtected.Use(LoginRequired()) diff --git a/internal/apps/oauth/oauth_types.go b/internal/apps/oauth/oauth_types.go new file mode 100644 index 00000000..17ab267c --- /dev/null +++ b/internal/apps/oauth/oauth_types.go @@ -0,0 +1,36 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +// AuthSourceView 登录源展示信息 +type AuthSourceView struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + IsActive bool `json:"is_active"` + IconURL string `json:"icon_url"` + ClientSecretConfigured bool `json:"client_secret_configured"` +} + +// OAuthAuthorizeResponse 授权 URL 响应 +// +//nolint:revive // OAuth 前缀保持包内语义清晰 +type OAuthAuthorizeResponse struct { + AuthorizeURL string `json:"authorize_url"` +} + +// OAuthCallbackResult 回调处理结果 +// +//nolint:revive // OAuth 前缀保持包内语义清晰 +type OAuthCallbackResult struct { + Status string `json:"status"` + User *BasicUserInfo `json:"user,omitempty"` +} + +// CallbackRequest OAuth 回调请求参数 +type CallbackRequest struct { + State string `json:"state" binding:"required"` + Code string `json:"code" binding:"required"` +} diff --git a/internal/apps/oauth/oauth_userinfo.go b/internal/apps/oauth/oauth_userinfo.go new file mode 100644 index 00000000..84ae7124 --- /dev/null +++ b/internal/apps/oauth/oauth_userinfo.go @@ -0,0 +1,141 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" +) + +func uniqueUsername(ctx context.Context, base string) (string, error) { + base = strings.TrimSpace(base) + if base == "" { + base = "user" + } + + var existingUsernames []string + if err := db.DB(ctx).Model(&model.User{}). + Where("username = ? OR username LIKE ?", base, base+"-%"). + Pluck("username", &existingUsernames).Error; err != nil { + return "", err + } + + // 将现有的用户名放入 map 中,以便 O(1) 查找 + exists := make(map[string]bool, len(existingUsernames)) + for _, u := range existingUsernames { + exists[strings.ToLower(u)] = true + } + + // 检查 base 是否被占用 + if !exists[strings.ToLower(base)] { + return base, nil + } + + // 顺序查找第一个可用的带后缀用户名 + for i := 1; i <= 1000; i++ { + candidate := fmt.Sprintf("%s-%d", base, i) + if !exists[strings.ToLower(candidate)] { + return candidate, nil + } + } + + return "", errors.New(errUsernameGenerateFailed) +} + +func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) { + authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return nil, err + } + + token, err := authConfig.Exchange(ctx, code) + if err != nil { + return nil, err + } + + userInfo := &model.OAuthUserInfo{Active: true} + if verifier != nil { + if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil { + return nil, verifyErr + } + } + + if userInfo.Username == "" && userInfo.PreferredUsername != "" { + userInfo.Username = userInfo.PreferredUsername + } + if userInfo.Username == "" && userInfo.Email != "" { + userInfo.Username = strings.Split(userInfo.Email, "@")[0] + } + if userInfo.Username == "" && userInfo.Sub != "" { + userInfo.Username = userInfo.Sub + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + + return userInfo, nil +} + +// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo +func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error { + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + return nil + } + idToken, verifyErr := verifier.Verify(ctx, rawIDToken) + if verifyErr != nil { + return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr) + } + if nonce != "" && idToken.Nonce != nonce { + return errors.New(errNonceMismatch) + } + if claimsErr := idToken.Claims(userInfo); claimsErr != nil { + return claimsErr + } + return nil +} + +func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error { + userInfo.Username = strings.TrimSpace(userInfo.Username) + userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername) + userInfo.Email = strings.TrimSpace(userInfo.Email) + userInfo.Name = strings.TrimSpace(userInfo.Name) + userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL) + + if userInfo.Username == "" && userInfo.PreferredUsername != "" { + userInfo.Username = userInfo.PreferredUsername + } + if userInfo.Username == "" && userInfo.Email != "" { + userInfo.Username = strings.Split(userInfo.Email, "@")[0] + } + if userInfo.Username == "" && userInfo.Sub != "" { + userInfo.Username = userInfo.Sub + } + if userInfo.Username == "" { + return errors.New(errUsernameFromSourceFailed) + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + if !userInfo.Active { + userInfo.Active = true + } + return nil +} + +func buildCallbackResult(user *model.User, status string) OAuthCallbackResult { + result := OAuthCallbackResult{Status: status} + if user != nil { + info := BuildBasicUserInfo(user, false) + result.User = &info + } + return result +} diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go index e5be215a..6eeba2f3 100644 --- a/internal/apps/oauth/routers.go +++ b/internal/apps/oauth/routers.go @@ -98,7 +98,7 @@ func Logout(c *gin.Context) { session.Options(GetSessionOptions(-1)) session.Clear() if err := session.Save(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } c.JSON(http.StatusOK, response.OKNil()) diff --git a/internal/apps/oauth/session_context.go b/internal/apps/oauth/session_context.go new file mode 100644 index 00000000..42be0c55 --- /dev/null +++ b/internal/apps/oauth/session_context.go @@ -0,0 +1,82 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "crypto/sha256" + "encoding/hex" + + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "github.com/google/uuid" +) + +// GetUserIDFromSession 从 Session 中提取用户 ID +func GetUserIDFromSession(s sessions.Session) uint64 { + userID, ok := s.Get(UserIDKey).(uint64) + if !ok { + return 0 + } + return userID +} + +// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID +func GetUserIDFromContext(c *gin.Context) uint64 { + session := sessions.Default(c) + return GetUserIDFromSession(session) +} + +func ensureSessionToken(s sessions.Session) (string, bool) { + token, ok := s.Get(SessionTokenKey).(string) + if !ok || token == "" { + token = uuid.NewString() + s.Set(SessionTokenKey, token) + return token, true + } + return token, false +} + +func hashSessionToken(token string) string { + h := sha256.New() + h.Write([]byte(token)) + return hex.EncodeToString(h.Sum(nil)) +} + +func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error { + session := sessions.Default(c) + session.Set(UserIDKey, user.ID) + session.Set(UserNameKey, user.Username) + session.Set(PasswordHashKey, user.Password) + + // 根据系统配置动态设置 Session 过期时间 + maxAge := config.Config.App.SessionAge + isSessionCookie := false + + ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) + if err == nil { + switch { + case ttlHours == -1: + // 永不过期,设置为 10 年 + maxAge = 10 * 365 * 24 * 3600 + case ttlHours > 0: + maxAge = ttlHours * 3600 + case ttlHours == 0: + isSessionCookie = true + } + } + session.Options(GetSessionOptions(maxAge)) + + if err := session.Save(); err != nil { + return err + } + + if isSessionCookie { + StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName) + } + + return nil +} diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go deleted file mode 100644 index 3194c0c8..00000000 --- a/internal/apps/oauth/sources.go +++ /dev/null @@ -1,774 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "crypto/sha256" - "encoding/hex" - "errors" - "fmt" - "net/http" - "strconv" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events" - "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/model" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/coreos/go-oidc/v3/oidc" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "github.com/google/uuid" - "golang.org/x/oauth2" - "gorm.io/gorm" -) - -// AuthSourceView 登录源展示信息 -type AuthSourceView struct { - ID uint64 `json:"id"` - Name string `json:"name"` - Type string `json:"type"` - DisplayName string `json:"display_name"` - IsActive bool `json:"is_active"` - IconURL string `json:"icon_url"` - ClientSecretConfigured bool `json:"client_secret_configured"` -} - -// OAuthAuthorizeResponse 授权 URL 响应 -// -//nolint:revive // OAuth 前缀保持包内语义清晰 -type OAuthAuthorizeResponse struct { - AuthorizeURL string `json:"authorize_url"` -} - -// OAuthCallbackResult 回调处理结果 -// -//nolint:revive // OAuth 前缀保持包内语义清晰 -type OAuthCallbackResult struct { - Status string `json:"status"` - User *BasicUserInfo `json:"user,omitempty"` -} - -// CallbackRequest OAuth 回调请求参数 -type CallbackRequest struct { - State string `json:"state" binding:"required"` - Code string `json:"code" binding:"required"` -} - -// GetUserIDFromSession 从 Session 中提取用户 ID -func GetUserIDFromSession(s sessions.Session) uint64 { - userID, ok := s.Get(UserIDKey).(uint64) - if !ok { - return 0 - } - return userID -} - -// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID -func GetUserIDFromContext(c *gin.Context) uint64 { - session := sessions.Default(c) - return GetUserIDFromSession(session) -} - -func ensureSessionToken(s sessions.Session) (string, bool) { - token, ok := s.Get(SessionTokenKey).(string) - if !ok || token == "" { - token = uuid.NewString() - s.Set(SessionTokenKey, token) - return token, true - } - return token, false -} - -func hashSessionToken(token string) string { - h := sha256.New() - h.Write([]byte(token)) - return hex.EncodeToString(h.Sum(nil)) -} - -func isOIDCLoginEnabled(ctx context.Context) bool { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) - if err != nil { - return true - } - return enabled -} - -func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) { - name := strings.TrimSpace(strings.ToLower(sourceName)) - if name == "" { - sources, err := model.GetActiveAuthSources(ctx) - if err != nil { - return nil, err - } - if len(sources) == 0 { - return nil, errors.New(errNoActiveAuthSource) - } - return &sources[0], nil - } - return model.GetAuthSourceByName(ctx, name) -} - -func activeLoginSources(ctx context.Context) []AuthSourceView { - enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) - if err == nil && !enabled { - return nil - } - - dbSources, err := model.GetActiveAuthSources(ctx) - if err != nil { - return nil - } - sources := make([]AuthSourceView, 0, len(dbSources)) - for _, source := range dbSources { - sources = append(sources, AuthSourceView{ - ID: source.ID, - Name: source.Name, - Type: source.Type, - DisplayName: source.DisplayName, - IsActive: source.IsActive, - IconURL: source.IconURL, - ClientSecretConfigured: source.ClientSecretConfigured, - }) - } - return sources -} - -func getFrontendLoginRedirectURL(ctx context.Context) (string, error) { - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" { - return "", errors.New(errServerAddressMissing) - } - return strings.TrimRight(sc.Value, "/") + "/login", nil -} - -func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { - if source == nil { - return nil, nil, errors.New(errAuthSourceRequired) - } - - if source.OpenIDDiscoveryURL == "" { - return nil, nil, errors.New(errDiscoveryURLRequired) - } - - // Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake) - issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/") - issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") - issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") - - // 使用进程级缓存获取 provider,避免每次调用都向 issuer 发起 - // /.well-known/openid-configuration HTTP 请求。 - provider, err := globalOIDCProviderCache.get(ctx, issuer) - if err != nil { - return nil, nil, err - } - verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID}) - scopes := strings.Fields(source.Scopes) - if len(scopes) == 0 { - scopes = []string{oidc.ScopeOpenID, "profile", "email"} - } - if !containsScope(scopes, oidc.ScopeOpenID) { - scopes = append([]string{oidc.ScopeOpenID}, scopes...) - } - - return &oauth2.Config{ - ClientID: source.ClientID, - ClientSecret: source.ClientSecret, - RedirectURL: redirectURL, - Scopes: scopes, - Endpoint: provider.Endpoint(), - }, verifier, nil -} - -func containsScope(scopes []string, scope string) bool { - for _, item := range scopes { - if item == scope { - return true - } - } - return false -} - -func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error { - session := sessions.Default(c) - session.Set(UserIDKey, user.ID) - session.Set(UserNameKey, user.Username) - session.Set(PasswordHashKey, user.Password) - - // 根据系统配置动态设置 Session 过期时间 - maxAge := config.Config.App.SessionAge - isSessionCookie := false - - ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) - if err == nil { - switch { - case ttlHours == -1: - // 永不过期,设置为 10 年 - maxAge = 10 * 365 * 24 * 3600 - case ttlHours > 0: - maxAge = ttlHours * 3600 - case ttlHours == 0: - isSessionCookie = true - } - } - session.Options(GetSessionOptions(maxAge)) - - if err := session.Save(); err != nil { - return err - } - - if isSessionCookie { - StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName) - } - - return nil -} - -func uniqueUsername(ctx context.Context, base string) (string, error) { - base = strings.TrimSpace(base) - if base == "" { - base = "user" - } - - var existingUsernames []string - if err := db.DB(ctx).Model(&model.User{}). - Where("username = ? OR username LIKE ?", base, base+"-%"). - Pluck("username", &existingUsernames).Error; err != nil { - return "", err - } - - // 将现有的用户名放入 map 中,以便 O(1) 查找 - exists := make(map[string]bool, len(existingUsernames)) - for _, u := range existingUsernames { - exists[strings.ToLower(u)] = true - } - - // 检查 base 是否被占用 - if !exists[strings.ToLower(base)] { - return base, nil - } - - // 顺序查找第一个可用的带后缀用户名 - for i := 1; i <= 1000; i++ { - candidate := fmt.Sprintf("%s-%d", base, i) - if !exists[strings.ToLower(candidate)] { - return candidate, nil - } - } - - return "", errors.New(errUsernameGenerateFailed) -} - -func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) { - authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) - if err != nil { - return nil, err - } - - token, err := authConfig.Exchange(ctx, code) - if err != nil { - return nil, err - } - - userInfo := &model.OAuthUserInfo{Active: true} - if verifier != nil { - if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil { - return nil, verifyErr - } - } - - if userInfo.Username == "" && userInfo.PreferredUsername != "" { - userInfo.Username = userInfo.PreferredUsername - } - if userInfo.Username == "" && userInfo.Email != "" { - userInfo.Username = strings.Split(userInfo.Email, "@")[0] - } - if userInfo.Username == "" && userInfo.Sub != "" { - userInfo.Username = userInfo.Sub - } - if userInfo.Name == "" { - userInfo.Name = userInfo.Username - } - - return userInfo, nil -} - -// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo -func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error { - rawIDToken, ok := token.Extra("id_token").(string) - if !ok { - return nil - } - idToken, verifyErr := verifier.Verify(ctx, rawIDToken) - if verifyErr != nil { - return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr) - } - if nonce != "" && idToken.Nonce != nonce { - return errors.New(errNonceMismatch) - } - if claimsErr := idToken.Claims(userInfo); claimsErr != nil { - return claimsErr - } - return nil -} - -func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error { - userInfo.Username = strings.TrimSpace(userInfo.Username) - userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername) - userInfo.Email = strings.TrimSpace(userInfo.Email) - userInfo.Name = strings.TrimSpace(userInfo.Name) - userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL) - - if userInfo.Username == "" && userInfo.PreferredUsername != "" { - userInfo.Username = userInfo.PreferredUsername - } - if userInfo.Username == "" && userInfo.Email != "" { - userInfo.Username = strings.Split(userInfo.Email, "@")[0] - } - if userInfo.Username == "" && userInfo.Sub != "" { - userInfo.Username = userInfo.Sub - } - if userInfo.Username == "" { - return errors.New(errUsernameFromSourceFailed) - } - if userInfo.Name == "" { - userInfo.Name = userInfo.Username - } - if !userInfo.Active { - userInfo.Active = true - } - return nil -} - -func buildCallbackResult(user *model.User, status string) OAuthCallbackResult { - result := OAuthCallbackResult{Status: status} - if user != nil { - info := BuildBasicUserInfo(user, false) - result.User = &info - } - return result -} - -// GetLoginSources 获取可用登录源列表 -// @Summary 获取可用登录源 -// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用 -// @Tags oauth -// @Produce json -// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表" -// @Router /api/v1/oauth/sources [get] -func GetLoginSources(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context()))) -} - -// GetLoginURL 获取登录授权地址 -// @Summary 获取登录授权地址 -// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 -// @Tags oauth -// @Produce json -// @Param source query string false "认证源名称,为空使用第一个启用的认证源" -// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未配置" -// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败" -// @Router /api/v1/oauth/login [get] -func GetLoginURL(c *gin.Context) { - ctx := c.Request.Context() - if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - - source, err := resolveAuthSource(ctx, c.Query("source")) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - if !source.IsActive { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - - session := sessions.Default(c) - token, isNew := ensureSessionToken(session) - if isNew { - if err := session.Save(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - } - - userID := GetUserIDFromSession(session) - sessionHash := hashSessionToken(token) - - state := uuid.NewString() - payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ - SourceName: source.Name, - Purpose: OAuthPurposeLogin, - UserID: userID, - SessionHash: sessionHash, - }) - if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - - authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) -} - -func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) { - redirectURL, err := getFrontendLoginRedirectURL(ctx) - if err != nil { - return "", err - } - authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) - if err != nil { - return "", err - } - if verifier != nil { - return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil - } - return authConfig.AuthCodeURL(state), nil -} - -// Authorize 发起指定认证源授权 -// @Summary 发起指定认证源授权 -// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。 -// @Tags oauth -// @Produce json -// @Param source path string true "认证源名称" -// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login" -// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未启用" -// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败" -// @Router /api/v1/oauth/{source}/authorize [get] -func Authorize(c *gin.Context) { - ctx := c.Request.Context() - if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - - source, err := resolveAuthSource(ctx, c.Param("source")) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - if !source.IsActive { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose"))) - if purpose != OAuthPurposeBind { - purpose = OAuthPurposeLogin - } - - session := sessions.Default(c) - userID := GetUserIDFromSession(session) - if purpose == OAuthPurposeBind && userID == 0 { - c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized)) - return - } - - token, isNew := ensureSessionToken(session) - if isNew { - if err := session.Save(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - } - - sessionHash := hashSessionToken(token) - - state := uuid.NewString() - payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ - SourceName: source.Name, - Purpose: purpose, - UserID: userID, - SessionHash: sessionHash, - }) - if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - - authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) -} - -// Callback OAuth 回调处理 -// @Summary OAuth 回调处理 -// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。 -// @Tags oauth -// @Accept json -// @Produce json -// @Param request body oauth.CallbackRequest true "回调请求参数" -// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功" -// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误" -// @Failure 401 {object} response.Any "绑定场景未登录" -// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误" -// @Router /api/v1/oauth/callback [post] -func Callback(c *gin.Context) { - var req CallbackRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - ctx := c.Request.Context() - stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)) - payloadRaw, err := db.Redis.Get(ctx, stateKey).Result() - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(errInvalidState)) - return - } - _ = db.Redis.Del(ctx, stateKey) - - payload, err := decodeOAuthStatePayload(payloadRaw) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - session := sessions.Default(c) - currentUserID := GetUserIDFromSession(session) - - if payload.Purpose == OAuthPurposeBind && currentUserID == 0 { - c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized)) - return - } - - token, ok := session.Get(SessionTokenKey).(string) - if !ok || token == "" { - c.JSON(http.StatusBadRequest, response.Err("invalid session context")) - return - } - - if hashSessionToken(token) != payload.SessionHash { - c.JSON(http.StatusBadRequest, response.Err("session mismatch for oauth state")) - return - } - - if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID { - c.JSON(http.StatusBadRequest, response.Err("user context mismatch for oauth binding")) - return - } - - if !isOIDCLoginEnabled(ctx) { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - - source, err := resolveAuthSource(ctx, payload.SourceName) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - if !source.IsActive { - c.JSON(http.StatusBadRequest, response.Err(errAuthSourceDisabled)) - return - } - - redirectURL, err := getFrontendLoginRedirectURL(ctx) - if err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - - userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL) - if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - if err := normalizeOAuthUserInfo(userInfo); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - if userInfo.Sub == "" { - userInfo.Sub = userInfo.Username - } - - if payload.Purpose == OAuthPurposeBind { - handleCallbackBind(ctx, c, source, userInfo) - return - } - - handleCallbackLogin(ctx, c, source, userInfo) -} - -// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程 -func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { - userID := GetUserIDFromContext(c) - if userID == 0 { - c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized)) - return - } - var user model.User - if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ - AuthSourceID: source.ID, - UserID: user.ID, - ExternalID: userInfo.Sub, - ExternalUsername: userInfo.Username, - Email: userInfo.Email, - }); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - user.LastLoginAt = time.Now() - _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error - c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound"))) -} - -// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册) -func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { - var user model.User - - account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub) - switch { - case err == nil: - if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - case errors.Is(err, gorm.ErrRecordNotFound): - newUser, ok := handleCallbackRegister(ctx, c, source, userInfo) - if !ok { - return - } - user = newUser - default: - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - - user.LastLoginAt = time.Now() - _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error - if err := setLoginSession(ctx, c, &user); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - - logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) - - custom_events.TriggerAdminLoginEvent(ctx, &user, c.ClientIP()) - - c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in"))) -} - -// 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) - if regErr != nil { - registrationEnabled = true - } - - if !registrationEnabled { - c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind"))) - return model.User{}, false - } - - username, uniqueErr := uniqueUsername(ctx, userInfo.Username) - if uniqueErr != nil { - c.JSON(http.StatusInternalServerError, response.Err(uniqueErr.Error())) - return model.User{}, false - } - userInfo.Username = username - - var user model.User - if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return model.User{}, false - } - if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ - AuthSourceID: source.ID, - UserID: user.ID, - ExternalID: userInfo.Sub, - ExternalUsername: userInfo.Username, - Email: userInfo.Email, - }); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return model.User{}, false - } - 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 -} - -// ListExternalAccounts 获取当前用户的外部帐号绑定列表 -// @Summary 获取外部帐号列表 -// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/oauth/external-accounts [get] -func ListExternalAccounts(c *gin.Context) { - userID := GetUserIDFromContext(c) - accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) - if err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) - return - } - c.JSON(http.StatusOK, response.OK(accounts)) -} - -// DeleteExternalAccount 解除外部帐号绑定 -// @Summary 解除外部帐号绑定 -// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "外部帐号绑定记录 ID" -// @Success 200 {object} response.Any{data=string} "解除绑定成功" -// @Failure 400 {object} response.Any "ID 无效或解除失败" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/oauth/external-accounts/{id}/delete [post] -func DeleteExternalAccount(c *gin.Context) { - userID := GetUserIDFromContext(c) - if userID == 0 { - c.JSON(http.StatusUnauthorized, response.Err(common.UnAuthorized)) - return - } - rawID := strings.TrimSpace(c.Param("id")) - id, err := strconv.ParseUint(rawID, 10, 64) - if err != nil || id == 0 { - c.JSON(http.StatusBadRequest, response.Err(errInvalidExternalAccountBindingID)) - return - } - if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/risk_control/middleware.go b/internal/apps/risk_control/middleware.go index 8b88afa5..ee5daf06 100644 --- a/internal/apps/risk_control/middleware.go +++ b/internal/apps/risk_control/middleware.go @@ -29,7 +29,7 @@ func RiskControlMiddleware() gin.HandlerFunc { // 1. 限流背压检测(检测本地缓冲队列是否已满) if IsBufferFull() { - c.AbortWithStatusJSON(http.StatusTooManyRequests, response.Err("系统繁忙,请稍后再试")) + response.AbortTooManyRequests(c, "系统繁忙,请稍后再试") return } diff --git a/internal/apps/risk_control/middleware_test.go b/internal/apps/risk_control/middleware_test.go index b4f5652a..58d1227f 100644 --- a/internal/apps/risk_control/middleware_test.go +++ b/internal/apps/risk_control/middleware_test.go @@ -13,6 +13,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "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" @@ -25,8 +26,7 @@ func TestRiskControlMiddleware(t *testing.T) { config.Config.ClickHouse.Enabled = false defer func() { config.Config.ClickHouse.Enabled = false }() - r := gin.New() - r.Use(RiskControlMiddleware()) + r := testhelper.NewTestGinEngine(RiskControlMiddleware()) r.GET("/test", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) @@ -91,8 +91,7 @@ func TestRiskControlMiddleware(t *testing.T) { logChan = nil }() - r := gin.New() - r.Use(RiskControlMiddleware()) + r := testhelper.NewTestGinEngine(RiskControlMiddleware()) r.GET("/test", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) @@ -126,8 +125,7 @@ func TestRiskControlMiddleware(t *testing.T) { logChan <- &UserAccessLog{} } - r := gin.New() - r.Use(RiskControlMiddleware()) + r := testhelper.NewTestGinEngine(RiskControlMiddleware()) r.GET("/test", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) diff --git a/internal/apps/upload/access_cache.go b/internal/apps/upload/cache/access_cache.go similarity index 61% rename from internal/apps/upload/access_cache.go rename to internal/apps/upload/cache/access_cache.go index efddb0ee..eff0681d 100644 --- a/internal/apps/upload/access_cache.go +++ b/internal/apps/upload/cache/access_cache.go @@ -1,7 +1,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +// Package cache provides in-process upload access-control caches. +package cache import ( "context" @@ -10,32 +11,19 @@ import ( "sync" "time" + "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/storage" ) -const accessCacheTTL = 5 * time.Second - const fileAccessInvalidationChannel = "upload:file_access_invalidation" -type migrationAccessState struct { - readOnly bool - target storage.Config - hasTarget bool - targetErr error - loadErr error -} - var ( accessCacheOnce sync.Once - migrationAccessMu sync.RWMutex - migrationAccessCached migrationAccessState - migrationAccessValid bool - migrationAccessCheckedAt time.Time - - fileAccessWhitelistMu sync.RWMutex + fileAccessWhitelistMu sync.RWMutex fileAccessWhitelistTypes map[string]struct{} fileAccessWhitelistValid bool fileAccessWhitelistCheckedAt time.Time @@ -43,9 +31,7 @@ var ( // ResetAccessCaches clears in-process upload access caches. func ResetAccessCaches() { - migrationAccessMu.Lock() - migrationAccessValid = false - migrationAccessMu.Unlock() + uploadstorage.ResetMigrationAccessCache() fileAccessWhitelistMu.Lock() fileAccessWhitelistValid = false @@ -85,62 +71,18 @@ func startAccessCacheInvalidationListener() { }() } -func loadMigrationAccessState(ctx context.Context) migrationAccessState { - ensureAccessCacheListener() - - migrationAccessMu.RLock() - if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL { - state := migrationAccessCached - migrationAccessMu.RUnlock() - return state - } - migrationAccessMu.RUnlock() - - migrationAccessMu.Lock() - defer migrationAccessMu.Unlock() - - if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL { - return migrationAccessCached - } - - migrationAccessCached = buildMigrationAccessState(ctx) - migrationAccessValid = true - migrationAccessCheckedAt = time.Now() - return migrationAccessCached -} - -func buildMigrationAccessState(ctx context.Context) migrationAccessState { - execution, ok, err := latestStorageMigrationExecution(ctx) - if err != nil { - return migrationAccessState{loadErr: err, readOnly: true} - } - if !ok { - return migrationAccessState{} - } - - state := migrationAccessState{ - readOnly: execution.Status != model.TaskExecutionStatusSucceeded, - } - if execution.Status == model.TaskExecutionStatusSucceeded { - return state - } - - target, err := parseMigrationTargetConfig(ctx, []byte(execution.Payload)) - if err != nil { - state.targetErr = err - return state - } - - state.target = target - state.hasTarget = true - return state +// IsFilePublic reports whether uploadType is in the public access whitelist. +func IsFilePublic(ctx context.Context, uploadType string) bool { + whitelist := loadFileAccessWhitelist(ctx) + _, ok := whitelist[strings.ToLower(uploadType)] + return ok } func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} { ensureAccessCacheListener() fileAccessWhitelistMu.RLock() - if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL { + if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second { types := fileAccessWhitelistTypes fileAccessWhitelistMu.RUnlock() return types @@ -150,7 +92,7 @@ func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} { fileAccessWhitelistMu.Lock() defer fileAccessWhitelistMu.Unlock() - if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL { + if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second { return fileAccessWhitelistTypes } @@ -172,7 +114,7 @@ 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 == "" { - return []string{defaultPublicUploadType} + return []string{shared.DefaultPublicUploadType} } var whitelist []string @@ -182,7 +124,7 @@ func parseFileAccessWhitelist(ctx context.Context) []string { whitelist = parseCommaSeparatedWhitelist(sc.Value) if len(whitelist) == 0 { - return []string{defaultPublicUploadType} + return []string{shared.DefaultPublicUploadType} } return whitelist } diff --git a/internal/apps/upload/access_cache_test.go b/internal/apps/upload/cache/access_cache_test.go similarity index 69% rename from internal/apps/upload/access_cache_test.go rename to internal/apps/upload/cache/access_cache_test.go index 155023b5..1e7fb3d5 100644 --- a/internal/apps/upload/access_cache_test.go +++ b/internal/apps/upload/cache/access_cache_test.go @@ -1,13 +1,15 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package cache import ( "context" "testing" "time" + uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" @@ -19,14 +21,14 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) { ResetAccessCaches() ctx := context.Background() - first := loadMigrationAccessState(ctx) - second := loadMigrationAccessState(ctx) + first := uploadstorage.LoadMigrationAccessState(ctx) + second := uploadstorage.LoadMigrationAccessState(ctx) - if first.readOnly != second.readOnly { - t.Fatalf("readOnly mismatch: first=%v second=%v", first.readOnly, second.readOnly) + if first.ReadOnly != second.ReadOnly { + t.Fatalf("readOnly mismatch: first=%v second=%v", first.ReadOnly, second.ReadOnly) } - if first.hasTarget != second.hasTarget { - t.Fatalf("hasTarget mismatch: first=%v second=%v", first.hasTarget, second.hasTarget) + if first.HasTarget != second.HasTarget { + t.Fatalf("hasTarget mismatch: first=%v second=%v", first.HasTarget, second.HasTarget) } } @@ -36,13 +38,13 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) { ResetAccessCaches() ctx := context.Background() - if !isFilePublic(ctx, "avatar") { + if !IsFilePublic(ctx, "avatar") { t.Fatal("expected avatar to be public by default") } - if isFilePublic(ctx, "attachment") { + if IsFilePublic(ctx, "attachment") { t.Fatal("expected attachment to be private by default") } - if !isFilePublic(ctx, "AVATAR") { + if !IsFilePublic(ctx, "AVATAR") { t.Fatal("expected whitelist lookup to be case-insensitive") } } @@ -53,7 +55,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) { ResetAccessCaches() ctx := context.Background() - if !isFilePublic(ctx, "avatar") { + if !IsFilePublic(ctx, "avatar") { t.Fatal("expected seeded avatar whitelist before reset") } @@ -68,12 +70,13 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) { if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil { t.Fatalf("refresh whitelist redis cache: %v", err) } + model.ResetSystemConfigRAMCacheForTest() ResetAccessCaches() - if !isFilePublic(ctx, "attachment") { + if !IsFilePublic(ctx, "attachment") { t.Fatal("expected attachment to be public after whitelist refresh") } - if isFilePublic(ctx, "avatar") { + if IsFilePublic(ctx, "avatar") { t.Fatal("expected avatar to be private after whitelist refresh") } } @@ -87,11 +90,11 @@ func TestAccessCacheTTLExpires(t *testing.T) { _ = loadFileAccessWhitelist(ctx) fileAccessWhitelistMu.Lock() - fileAccessWhitelistCheckedAt = time.Now().Add(-accessCacheTTL - time.Second) + fileAccessWhitelistCheckedAt = time.Now().Add(-time.Duration(shared.AccessCacheTTL)*time.Second - time.Second) fileAccessWhitelistMu.Unlock() // Should still work after TTL by reloading from config. - if !isFilePublic(ctx, "avatar") { + if !IsFilePublic(ctx, "avatar") { t.Fatal("expected whitelist reload after TTL expiration") } } \ No newline at end of file diff --git a/internal/apps/upload/constants.go b/internal/apps/upload/constants.go deleted file mode 100644 index 22bff1d6..00000000 --- a/internal/apps/upload/constants.go +++ /dev/null @@ -1,20 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package upload - -import "github.com/Rain-kl/Wavelet/internal/storage" - -const ( - maxUploadSize = 32 * 1024 * 1024 // 32MB - detectContentBytes = 512 // http.DetectContentType 需要的最小字节数 - uploadDirPerm = 0755 // 上传目录权限 - uploadFilePerm = 0644 // 上传文件权限 - imageQualityLow = "low" - imageQualityMedium = "medium" - imageQualityHigh = "high" - imageQualityOrigin = "origin" - storageDriverLocal = string(storage.DriverLocal) - defaultPublicUploadType = "avatar" - fileStatsTrendDays = 7 -) diff --git a/internal/apps/upload/errs.go b/internal/apps/upload/errs.go index 5b2cafd6..e146bc4a 100644 --- a/internal/apps/upload/errs.go +++ b/internal/apps/upload/errs.go @@ -5,37 +5,34 @@ // Package upload 提供文件上传与下载功能 package upload +import "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + // 文件管理常量 const ( - ErrNoFileSelected = "请选择要上传的文件" - ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片" - ErrProcessFileFailed = "处理文件失败" - ErrSaveFileFailed = "保存文件失败" - ErrOpenFileFailed = "打开文件失败" - ErrSaveUploadRecordFailed = "保存上传记录失败" - ErrGenericFileTooLarge = "文件大小不能超过 32MB" - ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险" - ErrFileValidationFailed = "文件校验失败" - ErrInvalidMetadataJSON = "元数据 JSON 格式不合法" - ErrInvalidFileID = "无效的文件 ID" - ErrQueryUploadRecordFailed = "查询文件记录失败" - ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组" - ErrInvalidIDValueFormat = "无效的 ID 值: %s" - ErrRetrieveUploadRecordsFailed = "检索文件记录失败" - ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包" - ErrInvalidParams = "参数错误" - ErrQueryFileCountFailed = "查询文件数量失败" - ErrQueryFileListFailed = "查询文件列表失败" - ErrDeleteFileFailed = "删除文件失败" - ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件" - ErrS3KeyRequired = "s3 key must not be empty" - ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d" - ErrS3KeyStartsWithSlash = "s3 key must not start with /" - ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes" - ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w" - errImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空" - errInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w" - errInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high" - errParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w" - errQueryImagesForCacheWarmup = "查询待预热图片失败: %w" -) + ErrNoFileSelected = shared.ErrNoFileSelected + ErrUnsupportedFormat = shared.ErrUnsupportedFormat + ErrProcessFileFailed = shared.ErrProcessFileFailed + ErrSaveFileFailed = shared.ErrSaveFileFailed + ErrOpenFileFailed = shared.ErrOpenFileFailed + ErrSaveUploadRecordFailed = shared.ErrSaveUploadRecordFailed + ErrGenericFileTooLarge = shared.ErrGenericFileTooLarge + ErrFileContentExtensionMismatch = shared.ErrFileContentExtensionMismatch + ErrFileValidationFailed = shared.ErrFileValidationFailed + ErrInvalidMetadataJSON = shared.ErrInvalidMetadataJSON + ErrInvalidFileID = shared.ErrInvalidFileID + ErrQueryUploadRecordFailed = shared.ErrQueryUploadRecordFailed + ErrInvalidBatchDownloadRequest = shared.ErrInvalidBatchDownloadRequest + ErrInvalidIDValueFormat = shared.ErrInvalidIDValueFormat + ErrRetrieveUploadRecordsFailed = shared.ErrRetrieveUploadRecordsFailed + ErrNoValidFilesForArchive = shared.ErrNoValidFilesForArchive + ErrInvalidParams = shared.ErrInvalidParams + ErrQueryFileCountFailed = shared.ErrQueryFileCountFailed + ErrQueryFileListFailed = shared.ErrQueryFileListFailed + ErrDeleteFileFailed = shared.ErrDeleteFileFailed + ErrStorageReadOnly = shared.ErrStorageReadOnly + ErrS3KeyRequired = shared.ErrS3KeyRequired + ErrS3KeyTooLongFormat = shared.ErrS3KeyTooLongFormat + ErrS3KeyStartsWithSlash = shared.ErrS3KeyStartsWithSlash + ErrS3KeyContainsNullBytes = shared.ErrS3KeyContainsNullBytes + ErrQueryUnusedUploadsFailed = shared.ErrQueryUnusedUploadsFailed +) \ No newline at end of file diff --git a/internal/apps/upload/exports.go b/internal/apps/upload/exports.go new file mode 100644 index 00000000..0ea8a489 --- /dev/null +++ b/internal/apps/upload/exports.go @@ -0,0 +1,86 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package upload + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" + "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" + "github.com/Rain-kl/Wavelet/internal/apps/upload/handler" + uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" + uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" + "github.com/Rain-kl/Wavelet/internal/apps/upload/util" + "github.com/Rain-kl/Wavelet/internal/task" +) + +// HTTP handlers +var ( + UploadFile = handler.UploadFile + DownloadFile = handler.DownloadFile + BatchDownloadFiles = handler.BatchDownloadFiles + ListFiles = handler.ListFiles + DeleteFile = handler.DeleteFile + GetDistinctUploadTypes = handler.GetDistinctUploadTypes + ListMyFiles = handler.ListMyFiles + DeleteMyFile = handler.DeleteMyFile + UpdateMyFile = handler.UpdateMyFile + GetFileStats = handler.GetFileStats + ServeFileByID = filesrv.ServeFileByID +) + +// Cache management +var ( + ResetAccessCaches = cache.ResetAccessCaches + PublishAccessCacheInvalidation = cache.PublishAccessCacheInvalidation +) + +// Stats +var ( + ApplyUploadStatsAdd = uploadstats.ApplyUploadStatsAdd + ApplyUploadStatsRemove = uploadstats.ApplyUploadStatsRemove + RebuildUploadStats = uploadstats.RebuildUploadStats +) + +// Utilities +var ( + CompressImageToWebP = util.CompressImageToWebP + ValidateS3Key = util.ValidateS3Key +) + +// Task identifiers and metadata +const ( + StorageMigrationTask = uploadtask.StorageMigrationTask + SystemCleanupTask = uploadtask.SystemCleanupTask + WarmImageCacheTask = uploadtask.WarmImageCacheTask +) + +var ( + // StorageMigrationMeta describes the storage migration async task. + StorageMigrationMeta = uploadtask.StorageMigrationMeta + // SystemCleanupMeta describes the orphaned upload cleanup task. + SystemCleanupMeta = uploadtask.SystemCleanupMeta + // WarmImageCacheMeta describes the image compression cache warmup task. + WarmImageCacheMeta = uploadtask.WarmImageCacheMeta +) + +// MigrationHandler executes storage migration tasks. +type MigrationHandler = uploadtask.MigrationHandler + +// SystemCleanupHandler removes orphaned upload files. +type SystemCleanupHandler = uploadtask.SystemCleanupHandler + +// WarmImageCacheHandler pre-warms compressed image caches. +type WarmImageCacheHandler = uploadtask.WarmImageCacheHandler + +// WarmImageCachePayload is the payload for image cache warmup tasks. +type WarmImageCachePayload = uploadtask.WarmImageCachePayload + +// Ensure task handler types implement required interfaces. +var ( + _ task.TaskHandler = (*MigrationHandler)(nil) + _ task.TaskHandler = (*SystemCleanupHandler)(nil) + _ interface { + task.TaskHandler + ValidatePayload([]byte) ([]byte, error) + } = (*WarmImageCacheHandler)(nil) +) \ No newline at end of file diff --git a/internal/apps/upload/file_server.go b/internal/apps/upload/filesrv/file_server.go similarity index 70% rename from internal/apps/upload/file_server.go rename to internal/apps/upload/filesrv/file_server.go index 4fd247ca..d4a3156a 100644 --- a/internal/apps/upload/file_server.go +++ b/internal/apps/upload/filesrv/file_server.go @@ -2,7 +2,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +// Package filesrv serves uploaded files with access control and image compression. +package filesrv import ( "bytes" @@ -15,11 +16,16 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" + "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/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/diskcache" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/util" + 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" @@ -34,6 +40,15 @@ type compressedImageCacheResult struct { err error } +type fileTypeCategory string + +const ( + fileTypeImage fileTypeCategory = "image" + fileTypeVideo fileTypeCategory = "video" + fileTypeAudio fileTypeCategory = "audio" + fileTypeOther fileTypeCategory = "other" +) + // ServeFileByID 根据 ID 获取并提供已上传的文件 // @Summary 获取已上传文件 // @Description 根据文件 ID 获取并提供已上传的临时或正式文件,若配置了缓存则优先走本地缓存,否则从 S3 等后端存储读取并流式返回 @@ -48,7 +63,7 @@ type compressedImageCacheResult struct { // @Failure 500 {object} response.Any "服务内部错误" // @Router /f/{id} [get] func ServeFileByID(c *gin.Context) { - upload, err := getUploadRecordByID(c) + upload, err := GetUploadRecordByID(c) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { c.AbortWithStatus(http.StatusNotFound) @@ -62,18 +77,16 @@ func ServeFileByID(c *gin.Context) { return } - // 校验业务白名单与访问权限 - if err := checkFileAccessPermission(c, upload); err != nil { - c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil}) + if err := CheckFileAccessPermission(c, upload); err != nil { + response.AbortUnauthorized(c, common.UnAuthorized) return } ServeUpload(c, upload) } -// getUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。 -// 同时会自动设置通用的安全响应头。 -func getUploadRecordByID(c *gin.Context) (*model.Upload, error) { +// GetUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。 +func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) { c.Header("X-Content-Type-Options", "nosniff") c.Header("Content-Security-Policy", "sandbox") @@ -93,22 +106,11 @@ func getUploadRecordByID(c *gin.Context) (*model.Upload, error) { return &upload, nil } -// fileTypeCategory 定义文件的大类,便于未来扩展不同的处理方式 -type fileTypeCategory string - -const ( - fileTypeImage fileTypeCategory = "image" - fileTypeVideo fileTypeCategory = "video" - fileTypeAudio fileTypeCategory = "audio" - fileTypeOther fileTypeCategory = "other" -) - -// getFileTypeCategory 判断并返回文件的大类 func getFileTypeCategory(upload *model.Upload) fileTypeCategory { mime := strings.ToLower(upload.MimeType) ext := strings.ToLower(upload.Extension) - if strings.HasPrefix(mime, "image/") || isImageExtension(ext) { + if strings.HasPrefix(mime, "image/") || util.IsImageExtension(ext) { return fileTypeImage } if strings.HasPrefix(mime, "video/") { @@ -120,32 +122,27 @@ func getFileTypeCategory(upload *model.Upload) fileTypeCategory { return fileTypeOther } -// ServeUpload 将已存在的文件内容读取并流式响应给客户端,支持本地和 S3/CDN 驱动,并可选支持 WebP 图片压缩与本地缓存。 +// ServeUpload 将已存在的文件内容读取并流式响应给客户端。 func ServeUpload(c *gin.Context, upload *model.Upload) { - // 设置通用的缓存控制响应头 setCacheHeaders(c, upload) category := getFileTypeCategory(upload) - quality := normalizeImageQuality(c.Query("quality")) + quality := util.NormalizeImageQuality(c.Query("quality")) switch category { case fileTypeImage: - // 如果是图片且不是原图质量,则提供压缩优化后的图片预览 - if quality != imageQualityOrigin { + if quality != shared.ImageQualityOrigin { serveCompressedImage(c, upload, quality) return } - // 请求原图质量时,退化到默认提供原文件 fallthrough - default: - // 默认提供原文件,并执行协商缓存校验 serveOriginalWithConditionalCheck(c, upload) } } func setCacheHeaders(c *gin.Context, upload *model.Upload) { - if isFilePublic(c.Request.Context(), upload.Type) { + if cache.IsFilePublic(c.Request.Context(), upload.Type) { c.Header("Cache-Control", "public, max-age=31536000") } else { c.Header("Cache-Control", "private, no-cache") @@ -173,7 +170,7 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string) return } - webpBytes, _, err := ensureCompressedImageCache(c.Request.Context(), upload, quality) + webpBytes, _, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality) if err != nil { if len(webpBytes) > 0 { logger.WarnF(c.Request.Context(), "failed to cache compressed image: %v", err) @@ -188,14 +185,15 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string) c.Data(http.StatusOK, "image/webp", webpBytes) } -func ensureCompressedImageCache( +// EnsureCompressedImageCache returns cached or freshly generated WebP bytes for an upload. +func EnsureCompressedImageCache( ctx context.Context, upload *model.Upload, quality string, ) ([]byte, bool, error) { - cache := diskcache.GetGlobalCache() - cacheKey := imageCompressionCacheKey(upload, quality) - webpBytes, err := cache.Get(cacheKey) + cacheStore := diskcache.GetGlobalCache() + cacheKey := ImageCompressionCacheKey(upload, quality) + webpBytes, err := cacheStore.Get(cacheKey) if err == nil { return webpBytes, true, nil } @@ -220,9 +218,9 @@ func generateCompressedImageCache( quality string, cacheKey string, ) (compressedImageCacheResult, error) { - cache := diskcache.GetGlobalCache() + cacheStore := diskcache.GetGlobalCache() - webpBytes, err := cache.Get(cacheKey) + webpBytes, err := cacheStore.Get(cacheKey) if err == nil { return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil } @@ -235,12 +233,12 @@ func generateCompressedImageCache( return compressedImageCacheResult{}, fmt.Errorf("read original image: %w", err) } - webpBytes, err = CompressImageToWebP(bytes.NewReader(origBytes), quality) + webpBytes, err = util.CompressImageToWebP(bytes.NewReader(origBytes), quality) if err != nil { return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err) } - if err := cache.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil { + if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil { return compressedImageCacheResult{ bytes: webpBytes, err: fmt.Errorf("write compressed image cache: %w", err), @@ -250,7 +248,8 @@ func generateCompressedImageCache( return compressedImageCacheResult{bytes: webpBytes}, nil } -func imageCompressionCacheKey(upload *model.Upload, quality string) string { +// ImageCompressionCacheKey returns the disk cache key for a compressed upload image. +func ImageCompressionCacheKey(upload *model.Upload, quality string) string { return fmt.Sprintf( "upload_webp_v1_%d_%d_%d_%s_%s", upload.ID, @@ -261,18 +260,8 @@ func imageCompressionCacheKey(upload *model.Upload, quality string) string { ) } -func normalizeImageQuality(quality string) string { - switch strings.ToLower(quality) { - case imageQualityLow, imageQualityMedium, imageQualityHigh: - return strings.ToLower(quality) - default: - return imageQualityOrigin - } -} - -// serveOriginal 原始文件的流式响应逻辑 func serveOriginal(c *gin.Context, upload *model.Upload) { - obj, err := openStoredObject(c.Request.Context(), upload) + obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload) if err != nil { c.AbortWithStatus(http.StatusNotFound) return @@ -281,9 +270,8 @@ func serveOriginal(c *gin.Context, upload *model.Upload) { c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil) } -// getOriginalFileBytes 获取原始文件所有字节 func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, error) { - obj, err := openStoredObject(ctx, upload) + obj, err := uploadstorage.OpenStoredObject(ctx, upload) if err != nil { return nil, err } @@ -291,17 +279,10 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er return io.ReadAll(obj.Body) } -// isFilePublic 校验文件类型是否在公开访问白名单中 -func isFilePublic(ctx context.Context, uploadType string) bool { - whitelist := loadFileAccessWhitelist(ctx) - _, ok := whitelist[strings.ToLower(uploadType)] - return ok -} - func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error { var currUser *model.User var err error - if u, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil { + if u, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil { currUser = u } else { currUser, err = oauth.GetUserFromRequest(c) @@ -318,21 +299,18 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error { return nil } -// checkFileAccessPermission 校验文件是否可以被当前请求访问 -func checkFileAccessPermission(c *gin.Context, upload *model.Upload) error { - // 1. 私有文件校验(优先级高于当前白名单逻辑) +// CheckFileAccessPermission 校验文件是否可以被当前请求访问 +func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error { if upload.AccessMode == 0 { return checkPrivateFileOwner(c, upload.UserID) } - // 2. 如果类型为公开的则再进行校验白名单 - if !isFilePublic(c.Request.Context(), upload.Type) { - // 必须进行鉴权 - if _, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); !ok { + if !cache.IsFilePublic(c.Request.Context(), upload.Type) { + if _, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); !ok { if _, err := oauth.GetUserFromRequest(c); err != nil { return err } } } return nil -} +} \ No newline at end of file diff --git a/internal/apps/upload/file_server_test.go b/internal/apps/upload/filesrv/file_server_test.go similarity index 83% rename from internal/apps/upload/file_server_test.go rename to internal/apps/upload/filesrv/file_server_test.go index 779fcb33..5dc426e2 100644 --- a/internal/apps/upload/file_server_test.go +++ b/internal/apps/upload/filesrv/file_server_test.go @@ -2,7 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package filesrv import ( "bytes" @@ -15,7 +15,11 @@ import ( "os" "testing" + "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + "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/diskcache" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" @@ -27,6 +31,7 @@ import ( func TestServeFileByIDAccessControl(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() + cache.ResetAccessCaches() // Ensure uploads dir is cleaned up defer func() { _ = os.RemoveAll("uploads") }() @@ -91,6 +96,7 @@ func TestServeFileByIDAccessControl(t *testing.T) { // Set up router gin.SetMode(gin.TestMode) r := gin.New() + r.Use(response.ErrorHandlerMiddleware()) store := cookie.NewStore([]byte("secret")) r.Use(sessions.Sessions("test_session", store)) r.GET("/f/:id", ServeFileByID) @@ -151,58 +157,6 @@ func TestServeFileByIDAccessControl(t *testing.T) { }) } -func TestGetDistinctUploadTypes(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed some uploads with new custom types - user := model.User{ID: 2222, Username: "test_user_2"} - dbConn.Create(&user) - - customUpload := model.Upload{ - ID: 9001, - UserID: user.ID, - FileName: "custom.txt", - FilePath: "uploads/custom.txt", - FileSize: 10, - MimeType: "text/plain", - Extension: "txt", - StorageDriver: "local", - Type: "custom_type_xyz", - Status: model.UploadStatusUsed, - } - dbConn.Create(&customUpload) - - gin.SetMode(gin.TestMode) - r := gin.New() - r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes) - - req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200, got %d", w.Code) - } - - var resp struct { - ErrorMsg string `json:"error_msg"` - Data []string `json:"data"` - } - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("failed to parse JSON: %v", err) - } - - if resp.ErrorMsg != "" { - t.Fatalf("unexpected error: %s", resp.ErrorMsg) - } - - // Verify that only custom_type_xyz is present - if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" { - t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data) - } -} - func TestImageCompression(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() @@ -295,7 +249,7 @@ func TestImageCompression(t *testing.T) { t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type")) } - cacheKey := imageCompressionCacheKey(&uploadRecord, imageQualityMedium) + cacheKey := ImageCompressionCacheKey(&uploadRecord, shared.ImageQualityMedium) cachedBytes, err := cache.Get(cacheKey) if err != nil { t.Fatalf("disk cache Get(%q) returned error: %v", cacheKey, err) @@ -376,19 +330,19 @@ func TestNormalizeImageQuality(t *testing.T) { quality string want string }{ - {name: imageQualityLow, quality: imageQualityLow, want: imageQualityLow}, - {name: imageQualityMedium, quality: imageQualityMedium, want: imageQualityMedium}, - {name: imageQualityHigh, quality: imageQualityHigh, want: imageQualityHigh}, + {name: shared.ImageQualityLow, quality: shared.ImageQualityLow, want: shared.ImageQualityLow}, + {name: shared.ImageQualityMedium, quality: shared.ImageQualityMedium, want: shared.ImageQualityMedium}, + {name: shared.ImageQualityHigh, quality: shared.ImageQualityHigh, want: shared.ImageQualityHigh}, {name: "origin", quality: "origin", want: "origin"}, - {name: "uppercase", quality: "LOW", want: imageQualityLow}, + {name: "uppercase", quality: "LOW", want: shared.ImageQualityLow}, {name: "empty", quality: "", want: "origin"}, {name: "invalid", quality: "maximum", want: "origin"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := normalizeImageQuality(tt.quality); got != tt.want { - t.Errorf("normalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want) + if got := util.NormalizeImageQuality(tt.quality); got != tt.want { + t.Errorf("NormalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want) } }) } diff --git a/internal/apps/upload/file_management.go b/internal/apps/upload/handler/file_management.go similarity index 84% rename from internal/apps/upload/file_management.go rename to internal/apps/upload/handler/file_management.go index af14f8cf..2f005c2d 100644 --- a/internal/apps/upload/file_management.go +++ b/internal/apps/upload/handler/file_management.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package handler import ( "errors" @@ -11,10 +11,13 @@ import ( "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" - "github.com/Rain-kl/Wavelet/internal/util" + apputil "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-gonic/gin" "gorm.io/gorm" ) @@ -56,7 +59,7 @@ func ListFiles(c *gin.Context) { var req listFilesRequest if err := c.ShouldBindQuery(&req); err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidParams)) + response.AbortBadRequest(c, shared.ErrInvalidParams) return } if req.Page <= 0 { @@ -84,14 +87,14 @@ func ListFiles(c *gin.Context) { var total int64 if err := query.Count(&total).Error; err != nil { - c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed)) + 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 { - c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed)) + response.AbortBadRequest(c, shared.ErrQueryFileListFailed) return } @@ -116,14 +119,14 @@ func ListFiles(c *gin.Context) { // @Router /api/v1/admin/uploads/{id} [delete] func DeleteFile(c *gin.Context) { ctx := c.Request.Context() - if StorageReadOnly(ctx) { - c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly)) + if uploadstorage.ReadOnly(ctx) { + response.AbortConflict(c, shared.ErrStorageReadOnly) return } uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidFileID)) + response.AbortBadRequest(c, shared.ErrInvalidFileID) return } @@ -133,14 +136,14 @@ func DeleteFile(c *gin.Context) { c.AbortWithStatus(http.StatusNotFound) return } - c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed)) + response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) return } if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil { - c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed)) + response.AbortBadRequest(c, shared.ErrDeleteFileFailed) return } - recordUploadStatsRemove(ctx, &upload) + uploadstats.RecordUploadStatsRemove(ctx, &upload) c.JSON(http.StatusOK, response.OKNil()) } @@ -161,7 +164,7 @@ func GetDistinctUploadTypes(c *gin.Context) { Where("type IS NOT NULL AND type != ''"). Distinct(). Pluck("type", &dbTypes).Error; err != nil { - c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + response.AbortInternal(c, err.Error()) return } sort.Strings(dbTypes) @@ -198,12 +201,12 @@ type listMyFilesResponse struct { // @Failure 401 {object} response.Any "未登录" // @Router /api/v1/upload/my [get] func ListMyFiles(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() var req listMyFilesRequest if err := c.ShouldBindQuery(&req); err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidParams)) + response.AbortBadRequest(c, shared.ErrInvalidParams) return } if req.Page <= 0 { @@ -228,14 +231,14 @@ func ListMyFiles(c *gin.Context) { var total int64 if err := query.Count(&total).Error; err != nil { - c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed)) + 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 { - c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed)) + response.AbortBadRequest(c, shared.ErrQueryFileListFailed) return } @@ -259,16 +262,16 @@ func ListMyFiles(c *gin.Context) { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [delete] func DeleteMyFile(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() - if StorageReadOnly(ctx) { - c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly)) + if uploadstorage.ReadOnly(ctx) { + response.AbortConflict(c, shared.ErrStorageReadOnly) return } uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidFileID)) + response.AbortBadRequest(c, shared.ErrInvalidFileID) return } @@ -278,7 +281,7 @@ func DeleteMyFile(c *gin.Context) { c.AbortWithStatus(http.StatusNotFound) return } - c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed)) + response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) return } @@ -288,10 +291,10 @@ func DeleteMyFile(c *gin.Context) { } if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil { - c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed)) + response.AbortBadRequest(c, shared.ErrDeleteFileFailed) return } - recordUploadStatsRemove(ctx, &upload) + uploadstats.RecordUploadStatsRemove(ctx, &upload) c.JSON(http.StatusOK, response.OKNil()) } @@ -314,22 +317,22 @@ type updateMyFileRequest struct { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [put] func UpdateMyFile(c *gin.Context) { - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() - if StorageReadOnly(ctx) { - c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly)) + if uploadstorage.ReadOnly(ctx) { + response.AbortConflict(c, shared.ErrStorageReadOnly) return } uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidFileID)) + response.AbortBadRequest(c, shared.ErrInvalidFileID) return } var req updateMyFileRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidParams)) + response.AbortBadRequest(c, shared.ErrInvalidParams) return } @@ -339,7 +342,7 @@ func UpdateMyFile(c *gin.Context) { c.AbortWithStatus(http.StatusNotFound) return } - c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed)) + response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) return } @@ -358,10 +361,10 @@ func UpdateMyFile(c *gin.Context) { if len(updates) > 0 { if err := db.DB(ctx).Model(&upload).Updates(updates).Error; err != nil { - c.JSON(http.StatusOK, response.Err("更新文件记录失败")) + response.AbortBadRequest(c, "更新文件记录失败") return } } c.JSON(http.StatusOK, response.OK(upload)) -} +} \ No newline at end of file diff --git a/internal/apps/upload/handler/file_management_test.go b/internal/apps/upload/handler/file_management_test.go new file mode 100644 index 00000000..7607b526 --- /dev/null +++ b/internal/apps/upload/handler/file_management_test.go @@ -0,0 +1,65 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package handler + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/gin-gonic/gin" +) + +func TestGetDistinctUploadTypes(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + user := model.User{ID: 2222, Username: "test_user_2"} + dbConn.Create(&user) + + customUpload := model.Upload{ + ID: 9001, + UserID: user.ID, + FileName: "custom.txt", + FilePath: "uploads/custom.txt", + FileSize: 10, + MimeType: "text/plain", + Extension: "txt", + StorageDriver: "local", + Type: "custom_type_xyz", + Status: model.UploadStatusUsed, + } + dbConn.Create(&customUpload) + + gin.SetMode(gin.TestMode) + r := gin.New() + r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes) + + req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", w.Code) + } + + var resp struct { + ErrorMsg string `json:"error_msg"` + Data []string `json:"data"` + } + if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { + t.Fatalf("failed to parse JSON: %v", err) + } + + if resp.ErrorMsg != "" { + t.Fatalf("unexpected error: %s", resp.ErrorMsg) + } + + if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" { + t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data) + } +} \ No newline at end of file diff --git a/internal/apps/upload/routers.go b/internal/apps/upload/handler/routers.go similarity index 71% rename from internal/apps/upload/routers.go rename to internal/apps/upload/handler/routers.go index 15d71e37..7f0953e7 100644 --- a/internal/apps/upload/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -2,9 +2,11 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +// Package handler provides upload HTTP API handlers. +package handler -import ("archive/zip" +import ( + "archive/zip" "bytes" "context" "crypto/sha256" @@ -22,16 +24,21 @@ import ("archive/zip" "time" "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" - "github.com/Rain-kl/Wavelet/internal/util" + apputil "github.com/Rain-kl/Wavelet/internal/util" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response" ) type batchDownloadRequest struct { @@ -59,59 +66,53 @@ func UploadFile(c *gin.Context) { c.Header("X-Content-Type-Options", "nosniff") c.Header("Content-Security-Policy", "sandbox") - // 限制请求体大小以防止 DoS - c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxUploadSize) + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize) - currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() header, err := c.FormFile("file") if err != nil { - c.JSON(http.StatusOK, response.Err(ErrNoFileSelected)) + response.AbortBadRequest(c, shared.ErrNoFileSelected) return } file, err := header.Open() if err != nil { - c.JSON(http.StatusOK, response.Err(ErrOpenFileFailed)) + response.AbortBadRequest(c, shared.ErrOpenFileFailed) return } defer func() { _ = file.Close() }() - // 校验大小 - if header.Size > maxUploadSize { - c.JSON(http.StatusOK, response.Err(ErrGenericFileTooLarge)) + if header.Size > shared.MaxUploadSize { + response.AbortBadRequest(c, shared.ErrGenericFileTooLarge) return } - // 2. 提取文件基本元数据 origName := header.Filename ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(origName), ".")) if ext == "" { ext = "bin" } - // 3. 校验文件后缀是否在允许的系统配置列表中 if errMsg := validateUploadExtension(ctx, ext); errMsg != "" { - c.JSON(http.StatusOK, response.Err(errMsg)) + response.AbortBadRequest(c, errMsg) return } - // 4. 读取文件并计算 Hash hashWriter := sha256.New() var buf bytes.Buffer size, err := io.Copy(&buf, io.TeeReader(file, hashWriter)) if err != nil { - c.JSON(http.StatusOK, response.Err(ErrProcessFileFailed)) + response.AbortBadRequest(c, shared.ErrProcessFileFailed) return } fileHash := hex.EncodeToString(hashWriter.Sum(nil)) mimeType := detectMimeType(&buf, header, size) - // 校验真实 MIME Type 是否与常见图片扩展名匹配,防止 Polyglot / HTML 注入攻击 - if isImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") { - c.JSON(http.StatusOK, response.Err(ErrFileContentExtensionMismatch)) + if util.IsImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") { + response.AbortBadRequest(c, shared.ErrFileContentExtensionMismatch) return } @@ -119,38 +120,34 @@ func UploadFile(c *gin.Context) { accessMode, errMsg := resolveUploadAccessMode(c, uploadType) if errMsg != "" { - c.JSON(http.StatusOK, response.Err(errMsg)) + response.AbortBadRequest(c, errMsg) return } - // 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件 handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName, accessMode) if handled { return } if lookupErr != nil && !errors.Is(lookupErr, gorm.ErrRecordNotFound) { - c.JSON(http.StatusOK, response.Err(ErrFileValidationFailed)) + response.AbortBadRequest(c, shared.ErrFileValidationFailed) return } - // 7. 解析可选元数据字段 meta, errMsg := parseUploadMetadata(c, mimeType) if errMsg != "" { - c.JSON(http.StatusOK, response.Err(errMsg)) + response.AbortBadRequest(c, errMsg) return } id := idgen.NextUint64ID() subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext) - // 8. 写入当前活动存储驱动。 storageDriver, subPath, errMsg := storeUploadFile(ctx, subPath, size, mimeType, &buf, &meta) if errMsg != "" { - c.JSON(http.StatusOK, response.Err(errMsg)) + response.AbortBadRequest(c, errMsg) return } - // 9. 保存文件记录至数据库 newUpload := model.Upload{ ID: id, UserID: currUser.ID, @@ -168,7 +165,7 @@ func UploadFile(c *gin.Context) { } if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" { - c.JSON(http.StatusOK, response.Err(err)) + response.AbortBadRequest(c, err) return } @@ -189,31 +186,30 @@ func UploadFile(c *gin.Context) { // @Failure 500 {object} response.Any "服务内部错误" // @Router /api/v1/admin/uploads/download/{id} [get] func DownloadFile(c *gin.Context) { - upload, err := getUploadRecordByID(c) + upload, err := filesrv.GetUploadRecordByID(c) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { c.AbortWithStatus(http.StatusNotFound) return } if _, ok := err.(*strconv.NumError); ok { - c.JSON(http.StatusOK, response.Err(ErrInvalidFileID)) + response.AbortBadRequest(c, shared.ErrInvalidFileID) return } - c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed)) + response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed) return } - // 校验文件访问权限 - if err := checkFileAccessPermission(c, upload); err != nil { - c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil}) + if err := filesrv.CheckFileAccessPermission(c, upload); err != nil { + response.AbortUnauthorized(c, common.UnAuthorized) return } fileName := upload.FileName - quality := normalizeImageQuality(c.Query("quality")) - isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || isImageExtension(strings.ToLower(upload.Extension)) + quality := util.NormalizeImageQuality(c.Query("quality")) + isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || util.IsImageExtension(strings.ToLower(upload.Extension)) - if quality != imageQualityOrigin && isImage { + if quality != shared.ImageQualityOrigin && isImage { ext := filepath.Ext(fileName) if ext != "" { fileName = strings.TrimSuffix(fileName, ext) + ".webp" @@ -222,9 +218,8 @@ func DownloadFile(c *gin.Context) { } } - // 设置下载 Attachment 响应头 (支持 UTF-8 中文文件名转义) c.Header("Content-Disposition", fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(fileName))) - ServeUpload(c, upload) + filesrv.ServeUpload(c, upload) } // BatchDownloadFiles 批量打包 ZIP 下载接口 @@ -233,7 +228,7 @@ func DownloadFile(c *gin.Context) { // @Tags admin // @Accept json // @Produce octet-stream -// @Param request body upload.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体" +// @Param request body handler.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体" // @Security SessionCookie // @Success 200 {file} file "成功下载打包后的 ZIP" // @Failure 400 {object} response.Any "参数错误" @@ -244,52 +239,45 @@ func BatchDownloadFiles(c *gin.Context) { var req batchDownloadRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusOK, response.Err(ErrInvalidBatchDownloadRequest)) + response.AbortBadRequest(c, shared.ErrInvalidBatchDownloadRequest) return } - // 转换 ID 列表 var ids []uint64 for _, idStr := range req.IDs { id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(fmt.Sprintf(ErrInvalidIDValueFormat, idStr))) + response.AbortBadRequest(c, fmt.Sprintf(shared.ErrInvalidIDValueFormat, idStr)) return } 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 { - c.JSON(http.StatusOK, response.Err(ErrRetrieveUploadRecordsFailed)) + response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed) return } if len(uploads) == 0 { - c.JSON(http.StatusOK, response.Err(ErrNoValidFilesForArchive)) + response.AbortBadRequest(c, shared.ErrNoValidFilesForArchive) return } - // 设置 ZIP 格式流的响应头 c.Header("Content-Type", "application/zip") c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"") - // 开启实时 ZIP 压缩器并直接输出给 Response Writer zipWriter := zip.NewWriter(c.Writer) defer func() { _ = zipWriter.Close() }() - // 用于解决 ZIP 内部文件名称发生碰撞冲突的问题 usedNames := make(map[string]int) for _, upload := range uploads { - // 校验文件访问权限 - if err := checkFileAccessPermission(c, &upload); err != nil { + if err := filesrv.CheckFileAccessPermission(c, &upload); err != nil { logger.WarnF(ctx, "Batch download: skip file %d due to permission denied: %v", upload.ID, err) continue } - // 校验防冲突重命名逻辑 fileName := upload.FileName if count, exists := usedNames[fileName]; exists { usedNames[fileName] = count + 1 @@ -300,23 +288,19 @@ func BatchDownloadFiles(c *gin.Context) { usedNames[fileName] = 1 } - // 在 ZIP 包内建新条目 zipFileEntry, err := zipWriter.Create(fileName) if err != nil { logger.ErrorF(ctx, "ZIP 添加条目失败 [%s]: %v", fileName, err) continue } - // 打开底层文件数据源 - var rc io.ReadCloser - obj, err := openStoredObject(ctx, &upload) + obj, err := uploadstorage.OpenStoredObject(ctx, &upload) if err != nil { logger.ErrorF(ctx, "打包时读取文件失败: %v", err) continue } - rc = obj.Body + rc := obj.Body - // 流式拷贝到 ZIP entry _, err = io.Copy(zipFileEntry, rc) _ = rc.Close() if err != nil { @@ -328,7 +312,7 @@ func BatchDownloadFiles(c *gin.Context) { func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) { accessModeStr := c.PostForm("access_mode") if accessModeStr == "" { - if uploadType == defaultPublicUploadType { + if uploadType == shared.DefaultPublicUploadType { return 1, "" } return 0, "" @@ -341,7 +325,6 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) { return accessMode, "" } -// validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中 func validateUploadExtension(ctx context.Context, ext string) string { var sc model.SystemConfig if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" { @@ -354,21 +337,20 @@ func validateUploadExtension(ctx context.Context, ext string) string { } } if !allowed { - return ErrUnsupportedFormat + return shared.ErrUnsupportedFormat } } return "" } -// tryInstantUpload 尝试秒传:若数据库已存在相同 Hash 且大小一致的可用文件,直接生成新记录 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 if err != nil { return false, err } - if StorageReadOnly(ctx) { - c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly)) + if uploadstorage.ReadOnly(ctx) { + response.AbortConflict(c, shared.ErrStorageReadOnly) return true, nil } @@ -390,52 +372,40 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, } if err := db.DB(ctx).Create(&newUpload).Error; err != nil { - c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed)) + response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed) return true, err } - recordUploadStatsAdd(ctx, &newUpload) + uploadstats.RecordUploadStatsAdd(ctx, &newUpload) logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath) c.JSON(http.StatusOK, response.OK(newUpload)) return true, nil } -// storeUploadFile 将文件写入当前活动存储驱动。 func storeUploadFile(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) { - if StorageReadOnly(ctx) { - return "", "", ErrStorageReadOnly + if uploadstorage.ReadOnly(ctx) { + return "", "", shared.ErrStorageReadOnly } driver, backend, err := storage.Active(ctx) if err != nil { logger.ErrorF(ctx, "初始化活动存储失败: %v", err) - return "", "", ErrSaveFileFailed + 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 "", "", ErrSaveFileFailed + return "", "", shared.ErrSaveFileFailed } meta.Bucket = result.Bucket return string(driver), result.Key, "" } -// isImageExtension 判断文件扩展名是否属于常见图片格式 -func isImageExtension(ext string) bool { - for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} { - if ext == imgExt { - return true - } - } - return false -} - -// parseUploadMetadata 解析上传元数据字段 func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) { var meta model.UploadMetadata metadataStr := c.DefaultPostForm("metadata", "") if metadataStr != "" { if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil { - return meta, ErrInvalidMetadataJSON + return meta, shared.ErrInvalidMetadataJSON } } meta.OriginalMime = mimeType @@ -444,16 +414,14 @@ func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, return meta, "" } -// detectMimeType 检测文件的 MIME 类型,优先使用 Content-Type 头部信息 func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) string { - mimeType := http.DetectContentType(buf.Bytes()[:min(detectContentBytes, int(size))]) + mimeType := http.DetectContentType(buf.Bytes()[:min(shared.DetectContentBytes, int(size))]) if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" { mimeType = header.Header.Get("Content-Type") } return mimeType } -// saveUploadRecord 保存上传记录到数据库,失败时清理本地垃圾文件 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)) @@ -462,8 +430,8 @@ func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr) } } - return ErrSaveUploadRecordFailed + return shared.ErrSaveUploadRecordFailed } - recordUploadStatsAdd(ctx, upload) + uploadstats.RecordUploadStatsAdd(ctx, upload) return "" -} +} \ No newline at end of file diff --git a/internal/apps/upload/routers_test.go b/internal/apps/upload/handler/routers_test.go similarity index 98% rename from internal/apps/upload/routers_test.go rename to internal/apps/upload/handler/routers_test.go index 98a6aeff..1dbf9d4e 100644 --- a/internal/apps/upload/routers_test.go +++ b/internal/apps/upload/handler/routers_test.go @@ -2,7 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package handler import ( "archive/zip" @@ -20,6 +20,9 @@ import ( "time" "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/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/storage" @@ -36,6 +39,7 @@ type testResponse struct { func setupTestRouter(authUser *model.User) *gin.Engine { gin.SetMode(gin.TestMode) r := gin.New() + r.Use(response.ErrorHandlerMiddleware()) authMiddleware := func(c *gin.Context) { if authUser != nil { @@ -211,13 +215,13 @@ func TestUploadFile(t *testing.T) { w := httptest.NewRecorder() router.ServeHTTP(w, req) - if w.Code != http.StatusOK { - t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String()) + if w.Code != http.StatusBadRequest { + t.Fatalf("expected status 400, got %d. Body: %s", w.Code, w.Body.String()) } var resp testResponse _ = json.Unmarshal(w.Body.Bytes(), &resp) - if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, ErrUnsupportedFormat) { + if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, shared.ErrUnsupportedFormat) { t.Errorf("expected unsupported format error, got: %v", resp) } }) @@ -294,6 +298,7 @@ func TestUploadFile(t *testing.T) { sc.Value = "jpg,png,webp,txt" dbConn.Save(&sc) _ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc) + model.ResetSystemConfigRAMCacheForTest() contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{ "type": "document", @@ -806,7 +811,7 @@ func TestGetFileStats(t *testing.T) { t.Fatalf("failed to create upload: %v", err) } } - if err := RebuildUploadStats(context.Background()); err != nil { + if err := uploadstats.RebuildUploadStats(context.Background()); err != nil { t.Fatalf("failed to rebuild upload stats: %v", err) } diff --git a/internal/apps/upload/stats.go b/internal/apps/upload/handler/stats.go similarity index 66% rename from internal/apps/upload/stats.go rename to internal/apps/upload/handler/stats.go index 46527730..88a0a753 100644 --- a/internal/apps/upload/stats.go +++ b/internal/apps/upload/handler/stats.go @@ -1,27 +1,17 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package handler import ( "net/http" - "strings" "time" + "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" - - "github.com/Rain-kl/Wavelet/internal/common/response" -) - -const ( - catImage = "图片" - catVideo = "视频" - catAudio = "音频" - catDocument = "文档" - catArchive = "压缩包" - catOther = "其他" ) type trendItem struct { @@ -60,15 +50,15 @@ func GetFileStats(c *gin.Context) { var stats []model.UploadStat if err := db.DB(ctx).Find(&stats).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } now := time.Now() - trendDates := make([]string, 0, fileStatsTrendDays) - trendCountMap := make(map[string]int64, fileStatsTrendDays) - trendSizeMap := make(map[string]int64, fileStatsTrendDays) - for i := fileStatsTrendDays - 1; i >= 0; i-- { + trendDates := make([]string, 0, shared.FileStatsTrendDays) + trendCountMap := make(map[string]int64, shared.FileStatsTrendDays) + trendSizeMap := make(map[string]int64, shared.FileStatsTrendDays) + for i := shared.FileStatsTrendDays - 1; i >= 0; i-- { date := now.AddDate(0, 0, -i).Format("2006-01-02") trendDates = append(trendDates, date) trendCountMap[date] = 0 @@ -82,7 +72,7 @@ func GetFileStats(c *gin.Context) { categories []distributionItem ) - categoriesList := []string{catImage, catVideo, catAudio, catDocument, catArchive, catOther} + categoriesList := []string{"图片", "视频", "音频", "文档", "压缩包", "其他"} categoryMap := make(map[string]distributionItem, len(categoriesList)) for _, cat := range categoriesList { categoryMap[cat] = distributionItem{Name: cat} @@ -134,44 +124,4 @@ func GetFileStats(c *gin.Context) { Categories: categories, Types: types, })) -} - -func getFileCategory(mimeType, ext string) string { - mimeType = strings.ToLower(mimeType) - ext = strings.ToLower(ext) - - if strings.HasPrefix(mimeType, "image/") || isImageExtension(ext) { - return catImage - } - if strings.HasPrefix(mimeType, "video/") { - return catVideo - } - if strings.HasPrefix(mimeType, "audio/") { - return catAudio - } - if isArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") { - return catArchive - } - if isDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" { - return catDocument - } - return catOther -} - -func isArchiveExtension(ext string) bool { - for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} { - if ext == e { - return true - } - } - return false -} - -func isDocumentExtension(ext string) bool { - for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} { - if ext == e { - return true - } - } - return false } \ No newline at end of file diff --git a/internal/apps/upload/shared/constants.go b/internal/apps/upload/shared/constants.go new file mode 100644 index 00000000..405bf08e --- /dev/null +++ b/internal/apps/upload/shared/constants.go @@ -0,0 +1,23 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package shared + +import "github.com/Rain-kl/Wavelet/internal/storage" + +// Upload size, path, media quality, and cache constants shared across subpackages. +const ( + MaxUploadSize = 32 * 1024 * 1024 // 32MB + DetectContentBytes = 512 // http.DetectContentType 需要的最小字节数 + UploadDirPerm = 0755 // 上传目录权限 + UploadFilePerm = 0644 // 上传文件权限 + ImageQualityLow = "low" + ImageQualityMedium = "medium" + ImageQualityHigh = "high" + ImageQualityOrigin = "origin" + StorageDriverLocal = string(storage.DriverLocal) + DefaultPublicUploadType = "avatar" + FileStatsTrendDays = 7 + MaxS3KeyLength = 1024 + AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site +) \ No newline at end of file diff --git a/internal/apps/upload/shared/errs.go b/internal/apps/upload/shared/errs.go new file mode 100644 index 00000000..e9a466e6 --- /dev/null +++ b/internal/apps/upload/shared/errs.go @@ -0,0 +1,41 @@ +// Copyright 2025 linux.do +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package shared holds upload error and configuration constants shared across subpackages. +package shared + +// 文件管理常量 +const ( + ErrNoFileSelected = "请选择要上传的文件" + ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片" + ErrProcessFileFailed = "处理文件失败" + ErrSaveFileFailed = "保存文件失败" + ErrOpenFileFailed = "打开文件失败" + ErrSaveUploadRecordFailed = "保存上传记录失败" + ErrGenericFileTooLarge = "文件大小不能超过 32MB" + ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险" + ErrFileValidationFailed = "文件校验失败" + ErrInvalidMetadataJSON = "元数据 JSON 格式不合法" + ErrInvalidFileID = "无效的文件 ID" + ErrQueryUploadRecordFailed = "查询文件记录失败" + ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组" + ErrInvalidIDValueFormat = "无效的 ID 值: %s" + ErrRetrieveUploadRecordsFailed = "检索文件记录失败" + ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包" + ErrInvalidParams = "参数错误" + ErrQueryFileCountFailed = "查询文件数量失败" + ErrQueryFileListFailed = "查询文件列表失败" + ErrDeleteFileFailed = "删除文件失败" + ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件" + ErrS3KeyRequired = "s3 key must not be empty" + ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d" + ErrS3KeyStartsWithSlash = "s3 key must not start with /" + ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes" + ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w" + ErrImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空" + ErrInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w" + ErrInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high" + ErrParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w" + ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w" +) \ No newline at end of file diff --git a/internal/apps/upload/stats/category.go b/internal/apps/upload/stats/category.go new file mode 100644 index 00000000..5ed059e6 --- /dev/null +++ b/internal/apps/upload/stats/category.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package stats maintains incremental upload statistics and aggregations. +package stats + +import ( + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/upload/util" +) + +const ( + catImage = "图片" + catVideo = "视频" + catAudio = "音频" + catDocument = "文档" + catArchive = "压缩包" + catOther = "其他" +) + +// GetFileCategory classifies a file by mime type and extension. +func GetFileCategory(mimeType, ext string) string { + mimeType = strings.ToLower(mimeType) + ext = strings.ToLower(ext) + + if strings.HasPrefix(mimeType, "image/") || util.IsImageExtension(ext) { + return catImage + } + if strings.HasPrefix(mimeType, "video/") { + return catVideo + } + if strings.HasPrefix(mimeType, "audio/") { + return catAudio + } + if util.IsArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") { + return catArchive + } + if util.IsDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" { + return catDocument + } + return catOther +} \ No newline at end of file diff --git a/internal/apps/upload/stats_counter.go b/internal/apps/upload/stats/stats_counter.go similarity index 91% rename from internal/apps/upload/stats_counter.go rename to internal/apps/upload/stats/stats_counter.go index 3f5c28d0..eb2170d2 100644 --- a/internal/apps/upload/stats_counter.go +++ b/internal/apps/upload/stats/stats_counter.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package stats import ( "context" @@ -72,7 +72,7 @@ func applyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) erro }{ {model.UploadStatDimensionTotal, ""}, {model.UploadStatDimensionType, typeKey}, - {model.UploadStatDimensionCategory, getFileCategory(upload.MimeType, upload.Extension)}, + {model.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)}, {model.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")}, } @@ -111,13 +111,15 @@ func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeD }).Error } -func recordUploadStatsAdd(ctx context.Context, upload *model.Upload) { +// RecordUploadStatsAdd logs and applies upload stats increment. +func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) { if err := ApplyUploadStatsAdd(ctx, upload); err != nil { logger.WarnF(ctx, "increment upload stats failed: %v", err) } } -func recordUploadStatsRemove(ctx context.Context, upload *model.Upload) { +// RecordUploadStatsRemove logs and applies upload stats decrement. +func RecordUploadStatsRemove(ctx context.Context, upload *model.Upload) { if err := ApplyUploadStatsRemove(ctx, upload); err != nil { logger.WarnF(ctx, "decrement upload stats failed: %v", err) } diff --git a/internal/apps/upload/stats_counter_test.go b/internal/apps/upload/stats/stats_counter_test.go similarity index 99% rename from internal/apps/upload/stats_counter_test.go rename to internal/apps/upload/stats/stats_counter_test.go index dad330a6..597c929a 100644 --- a/internal/apps/upload/stats_counter_test.go +++ b/internal/apps/upload/stats/stats_counter_test.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package stats import ( "context" diff --git a/internal/apps/upload/storage/access_state.go b/internal/apps/upload/storage/access_state.go new file mode 100644 index 00000000..2f8054c0 --- /dev/null +++ b/internal/apps/upload/storage/access_state.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package storage provides upload storage backend operations and migration state. +package storage + +import ( + "context" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/storage" +) + +// MigrationAccessState captures cached migration maintenance state. +type MigrationAccessState struct { + ReadOnly bool + Target storage.Config + HasTarget bool + TargetErr error + LoadErr error +} + +var ( + migrationAccessMu sync.RWMutex + migrationAccessCached MigrationAccessState + migrationAccessValid bool + migrationAccessCheckedAt time.Time +) + +// ResetMigrationAccessCache clears the in-process migration access cache. +func ResetMigrationAccessCache() { + migrationAccessMu.Lock() + migrationAccessValid = false + migrationAccessMu.Unlock() +} + +// LoadMigrationAccessState returns cached migration maintenance state. +func LoadMigrationAccessState(ctx context.Context) MigrationAccessState { + migrationAccessMu.RLock() + if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second { + state := migrationAccessCached + migrationAccessMu.RUnlock() + return state + } + migrationAccessMu.RUnlock() + + migrationAccessMu.Lock() + defer migrationAccessMu.Unlock() + + if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second { + return migrationAccessCached + } + + migrationAccessCached = buildMigrationAccessState(ctx) + migrationAccessValid = true + migrationAccessCheckedAt = time.Now() + return migrationAccessCached +} + +func buildMigrationAccessState(ctx context.Context) MigrationAccessState { + execution, ok, err := LatestMigrationExecution(ctx) + if err != nil { + return MigrationAccessState{LoadErr: err, ReadOnly: true} + } + if !ok { + return MigrationAccessState{} + } + + state := MigrationAccessState{ + ReadOnly: execution.Status != model.TaskExecutionStatusSucceeded, + } + if execution.Status == model.TaskExecutionStatusSucceeded { + return state + } + + target, err := ParseMigrationTargetConfig(ctx, []byte(execution.Payload)) + if err != nil { + state.TargetErr = err + return state + } + + state.Target = target + state.HasTarget = true + return state +} \ No newline at end of file diff --git a/internal/apps/upload/storage/migration.go b/internal/apps/upload/storage/migration.go new file mode 100644 index 00000000..1f8536e0 --- /dev/null +++ b/internal/apps/upload/storage/migration.go @@ -0,0 +1,93 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package storage + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/storage" + "gorm.io/gorm" +) + +// StorageMigrationTask is the Asynq task name for storage migration. +const StorageMigrationTask = "storage:migrate" + +// LatestMigrationExecution returns the most recent storage migration task execution. +func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) { + var execution model.TaskExecution + err := db.DB(ctx). + Where("task_type = ?", StorageMigrationTask). + Order("id DESC"). + First(&execution).Error + if err == nil { + return &execution, true, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, false, err + } + return nil, false, nil +} + +// ParseMigrationTargetConfig parses and validates a storage migration target payload. +func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) { + if strings.TrimSpace(string(payload)) == "" { + return storage.Config{}, errors.New("storage migration target payload is required") + } + + var raw struct { + Target json.RawMessage `json:"target"` + } + if err := json.Unmarshal(payload, &raw); err != nil { + return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err) + } + + if len(raw.Target) == 0 { + return storage.Config{}, errors.New("storage migration target payload is required") + } + + var targetBytes []byte + var targetStr string + if err := json.Unmarshal(raw.Target, &targetStr); err == nil { + targetBytes = []byte(targetStr) + } else { + targetBytes = raw.Target + } + + var target storage.Config + if err := json.Unmarshal(targetBytes, &target); err != nil { + return storage.Config{}, fmt.Errorf("parse target storage config: %w", err) + } + + current, err := storage.LoadConfig(ctx) + if err != nil { + return storage.Config{}, fmt.Errorf("load active storage config: %w", err) + } + target = storage.MergeMaskedSecrets(target, current) + if err := storage.ValidateConfig(target); err != nil { + return storage.Config{}, fmt.Errorf("validate target storage config: %w", err) + } + return target, nil +} + +// NormalizeMigrationPayload validates and normalizes a storage migration payload. +func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) { + target, err := ParseMigrationTargetConfig(ctx, payload) + if err != nil { + return nil, storage.Config{}, err + } + type storageMigrationPayload struct { + Target storage.Config `json:"target"` + } + normalized, err := json.Marshal(storageMigrationPayload{Target: target}) + if err != nil { + return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err) + } + return normalized, target, nil +} \ No newline at end of file diff --git a/internal/apps/upload/storage_ops.go b/internal/apps/upload/storage/storage_ops.go similarity index 55% rename from internal/apps/upload/storage_ops.go rename to internal/apps/upload/storage/storage_ops.go index 369f9f43..8e4e36e7 100644 --- a/internal/apps/upload/storage_ops.go +++ b/internal/apps/upload/storage/storage_ops.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package storage import ( "context" @@ -12,17 +12,18 @@ import ( "github.com/Rain-kl/Wavelet/pkg/logger" ) -// StorageReadOnly checks if the storage system is in read-only maintenance mode. -func StorageReadOnly(ctx context.Context) bool { - state := loadMigrationAccessState(ctx) - if state.loadErr != nil { - logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.loadErr) +// ReadOnly checks if the storage system is in read-only maintenance mode. +func ReadOnly(ctx context.Context) bool { + state := LoadMigrationAccessState(ctx) + if state.LoadErr != nil { + logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.LoadErr) return true } - return state.readOnly + return state.ReadOnly } -func openStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) { +// OpenStoredObject opens a stored upload object from its configured backend. +func OpenStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) { driver := storage.Driver(upload.StorageDriver) if driver == "" { driver = storage.DriverLocal @@ -40,7 +41,7 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage return backend, nil } - target, ok, targetErr := currentMigrationTargetConfig(ctx) + target, ok, targetErr := CurrentMigrationTargetConfig(ctx) if targetErr != nil { return nil, targetErr } @@ -50,16 +51,17 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage return nil, fmt.Errorf("storage configuration for driver %q is unavailable", driver) } -func currentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) { - state := loadMigrationAccessState(ctx) - if state.loadErr != nil { - return storage.Config{}, false, state.loadErr +// CurrentMigrationTargetConfig returns the pending migration target config when available. +func CurrentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) { + state := LoadMigrationAccessState(ctx) + if state.LoadErr != nil { + return storage.Config{}, false, state.LoadErr } - if state.targetErr != nil { - return storage.Config{}, false, state.targetErr + if state.TargetErr != nil { + return storage.Config{}, false, state.TargetErr } - if !state.hasTarget { + if !state.HasTarget { return storage.Config{}, false, nil } - return state.target, true, nil -} + return state.Target, true, nil +} \ No newline at end of file diff --git a/internal/apps/upload/cleanup.go b/internal/apps/upload/task/cleanup.go similarity index 78% rename from internal/apps/upload/cleanup.go rename to internal/apps/upload/task/cleanup.go index adfa80c1..82ac2fc1 100644 --- a/internal/apps/upload/cleanup.go +++ b/internal/apps/upload/task/cleanup.go @@ -1,8 +1,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package upload implements upload tasks and file cleanup services. -package upload +// Package task provides upload-related async background task handlers. +package task import ( "context" @@ -10,6 +10,9 @@ import ( "fmt" "time" + "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" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/storage" @@ -18,16 +21,11 @@ import ( "gorm.io/gorm" ) -// 异步任务名称与管理类型定义 const ( // SystemCleanupTask 系统定期垃圾清理任务标识 SystemCleanupTask = "system:cleanup" // TaskTypeSystemCleanup 系统定期垃圾清理管理类型 TaskTypeSystemCleanup = "system_cleanup" - - // 错误描述常量 - errStorageReadOnly = "存储迁移维护中,当前仅允许读取文件" - errQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w" ) // SystemCleanupMeta represents the task metadata. @@ -47,21 +45,19 @@ type SystemCleanupHandler struct{} // Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理) func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) { - if storageReadOnly(ctx) { - return nil, errors.New(errStorageReadOnly) + if uploadstorage.ReadOnly(ctx) { + return nil, errors.New(shared.ErrStorageReadOnly) } - const batchSize = 100 // 每批处理100个文件 + const batchSize = 100 var lastID uint64 var totalProcessed int var totalDeleted int - // 计算1小时前的时间 oneHourAgo := time.Now().Add(-1 * time.Hour) task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339)) for { - // 使用游标分页查询未使用且超过1小时的上传记录 var unusedUploads []model.Upload if err := db.DB(ctx). Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo). @@ -69,22 +65,19 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas Limit(batchSize). Find(&unusedUploads).Error; err != nil { task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err) - return nil, fmt.Errorf(errQueryUnusedUploadsFailed, err) + return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err) } - // 没有更多数据,退出循环 if len(unusedUploads) == 0 { break } task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads)) - // 处理每个未使用的上传文件 for _, u := range unusedUploads { totalProcessed++ if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - // 更新上传记录状态 if err := tx.Model(&model.Upload{}). Where("id = ? AND status = ?", u.ID, model.UploadStatusPending). Update("status", model.UploadStatusDeleted).Error; err != nil { @@ -110,13 +103,12 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas continue } - recordUploadStatsRemove(ctx, &u) + uploadstats.RecordUploadStatsRemove(ctx, &u) totalDeleted++ lastID = u.ID } } - // 2. 清理超过7天的历史推送日志 task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...") cutoff := time.Now().AddDate(0, 0, -7) var pushHistoryCount int64 @@ -132,7 +124,6 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05")) } - // 3. 清理任务执行日志:高频任务保留3天,低频任务保留30天。 task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...") taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now()) if err != nil { @@ -154,17 +145,4 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas ) task.AppendLog(ctx, "%s", msg) return &task.TaskResult{Message: msg}, nil -} - -func storageReadOnly(ctx context.Context) bool { - var execution model.TaskExecution - err := db.DB(ctx).Where("task_type = ?", "storage:migrate").Order("id DESC").First(&execution).Error - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return false - } - logger.ErrorF(ctx, "读取存储维护状态失败: %v", err) - return true - } - return execution.Status != model.TaskExecutionStatusSucceeded -} +} \ No newline at end of file diff --git a/internal/apps/upload/storage_migration_task.go b/internal/apps/upload/task/storage_migration.go similarity index 79% rename from internal/apps/upload/storage_migration_task.go rename to internal/apps/upload/task/storage_migration.go index 5b089d74..1a54f8c3 100644 --- a/internal/apps/upload/storage_migration_task.go +++ b/internal/apps/upload/task/storage_migration.go @@ -1,13 +1,12 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package task import ( "context" "crypto/sha256" "encoding/hex" - "encoding/json" "errors" "fmt" "io" @@ -16,17 +15,18 @@ import ( "sync/atomic" "time" + 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" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/task" "golang.org/x/sync/errgroup" - "gorm.io/gorm" ) const ( // StorageMigrationTask is the Asynq task name for storage migration. - StorageMigrationTask = "storage:migrate" + StorageMigrationTask = uploadstorage.StorageMigrationTask // TaskTypeStorageMigration is the task metadata type for storage migration. TaskTypeStorageMigration = "storage_migration" @@ -58,13 +58,9 @@ var StorageMigrationMeta = task.TaskMeta{ // MigrationHandler copies stored objects and activates the target backend. type MigrationHandler struct{} -type storageMigrationPayload struct { - Target storage.Config `json:"target"` -} - // ValidatePayload rejects duplicate active migrations through the task framework. func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) { - normalized, _, err := normalizeStorageMigrationPayload(context.Background(), payload) + normalized, _, err := uploadstorage.NormalizeMigrationPayload(context.Background(), payload) if err != nil { return payload, err } @@ -95,7 +91,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T return nil, errors.New("另一个存储迁移任务正在运行中") } - // 任务结束时清理锁,使用 Background context 避免受任务 context 取消的影响 stopRenewal := make(chan struct{}) //nolint:contextcheck defer func() { @@ -105,7 +100,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T _ = db.Redis.Del(cleanupCtx, lockKey) }() - // 启动看门狗续租协程,每 10 分钟将锁的 TTL 自动延长为 1 小时 //nolint:contextcheck,gosec go func() { ticker := time.NewTicker(renewalInterval) @@ -129,7 +123,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T if err != nil { return nil, fmt.Errorf("load active storage config: %w", err) } - target, err := parseMigrationTargetConfig(ctx, payload) + target, err := uploadstorage.ParseMigrationTargetConfig(ctx, payload) if err != nil { return nil, err } @@ -178,62 +172,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T return &task.TaskResult{Message: message}, nil } -func normalizeStorageMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) { - target, err := parseMigrationTargetConfig(ctx, payload) - if err != nil { - return nil, storage.Config{}, err - } - normalized, err := json.Marshal(storageMigrationPayload{Target: target}) - if err != nil { - return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err) - } - return normalized, target, nil -} - -func parseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) { - if strings.TrimSpace(string(payload)) == "" { - return storage.Config{}, errors.New("storage migration target payload is required") - } - - // Try to parse using raw JSON message to handle both struct and string payload formats - var raw struct { - Target json.RawMessage `json:"target"` - } - if err := json.Unmarshal(payload, &raw); err != nil { - return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err) - } - - if len(raw.Target) == 0 { - return storage.Config{}, errors.New("storage migration target payload is required") - } - - var targetBytes []byte - var targetStr string - // Check if Target is a JSON string - if err := json.Unmarshal(raw.Target, &targetStr); err == nil { - // It is a string (e.g. from dynamic form input), parse its content as JSON - targetBytes = []byte(targetStr) - } else { - // It is a JSON object, use directly - targetBytes = raw.Target - } - - var target storage.Config - if err := json.Unmarshal(targetBytes, &target); err != nil { - return storage.Config{}, fmt.Errorf("parse target storage config: %w", err) - } - - current, err := storage.LoadConfig(ctx) - if err != nil { - return storage.Config{}, fmt.Errorf("load active storage config: %w", err) - } - target = storage.MergeMaskedSecrets(target, current) - if err := storage.ValidateConfig(target); err != nil { - return storage.Config{}, fmt.Errorf("validate target storage config: %w", err) - } - return target, nil -} - func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, error) { var count int64 err := db.DB(ctx).Model(&model.Upload{}). @@ -244,28 +182,13 @@ func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, err } func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) { - execution, ok, err := latestStorageMigrationExecution(ctx) + execution, ok, err := uploadstorage.LatestMigrationExecution(ctx) if err != nil || !ok { return false, err } return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil } -func latestStorageMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) { - var execution model.TaskExecution - err := db.DB(ctx). - Where("task_type = ?", StorageMigrationTask). - Order("id DESC"). - First(&execution).Error - if err == nil { - return &execution, true, nil - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return nil, false, err - } - return nil, false, nil -} - func migrateObjects( ctx context.Context, sourceBackend storage.Backend, @@ -311,7 +234,7 @@ func migrateObjects( g.SetLimit(migrationConcurrency) for _, object := range objects { - obj := object // Capture range variable + obj := object g.Go(func() error { if err := migrateSingleObject(ctx, sourceBackend, targetBackend, sourceDriver, targetDriver, obj, sha256HexLength); err != nil { return err @@ -344,7 +267,6 @@ func migrateSingleObject( }, sha256HexLength int, ) error { - // Check if the file already exists in target storage and has matching size if shouldSkipMigration(ctx, targetBackend, obj) { task.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件且校验一致: %s", obj.FilePath) if err := db.DB(ctx).Model(&model.Upload{}). @@ -375,7 +297,6 @@ func migrateSingleObject( return fmt.Errorf("close source object %q: %w", obj.FilePath, closeErr) } - // Data integrity check (SHA-256 hash verification) if len(obj.Hash) == sha256HexLength { task.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key) targetObj, getErr := targetBackend.Get(ctx, targetResult.Key) @@ -450,13 +371,13 @@ func markMissingMigrationObjectDeleted( if err := db.DB(ctx).Model(&model.Upload{}). Where("storage_driver = ? AND file_path = ?", sourceDriver, filePath). Updates(map[string]any{ - "status": model.UploadStatusDeleted, - colStorageDriver: targetDriver, + "status": model.UploadStatusDeleted, + colStorageDriver: targetDriver, }).Error; err != nil { return fmt.Errorf("update missing object %q: %w", filePath, err) } for i := range affectedUploads { - recordUploadStatsRemove(ctx, &affectedUploads[i]) + uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i]) } return nil } @@ -475,4 +396,4 @@ func isNotFoundError(err error) bool { } } return false -} +} \ No newline at end of file diff --git a/internal/apps/upload/storage_migration_task_test.go b/internal/apps/upload/task/storage_migration_task_test.go similarity index 96% rename from internal/apps/upload/storage_migration_task_test.go rename to internal/apps/upload/task/storage_migration_task_test.go index 86d4177f..f7216971 100644 --- a/internal/apps/upload/storage_migration_task_test.go +++ b/internal/apps/upload/task/storage_migration_task_test.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package task import ( "bytes" @@ -52,7 +52,9 @@ func TestMigrationHandlerExecute(t *testing.T) { AccessKeyID: "key", SecretAccessKey: "secret", } - payload, err := json.Marshal(storageMigrationPayload{Target: target}) + payload, err := json.Marshal(struct { + Target storage.Config `json:"target"` + }{Target: target}) if err != nil { t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err) } @@ -150,7 +152,9 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { AccessKeyID: "key", SecretAccessKey: "secret", } - payload, err := json.Marshal(storageMigrationPayload{Target: target}) + payload, err := json.Marshal(struct { + Target storage.Config `json:"target"` + }{Target: target}) if err != nil { t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err) } @@ -259,7 +263,9 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) { t.Fatalf("SaveActiveConfig() returned error: %v", err) } - payload, err := json.Marshal(storageMigrationPayload{Target: active}) + payload, err := json.Marshal(struct { + Target storage.Config `json:"target"` + }{Target: active}) if err != nil { t.Fatalf("Marshal payload failed: %v", err) } diff --git a/internal/apps/upload/tasks.go b/internal/apps/upload/task/tasks.go similarity index 85% rename from internal/apps/upload/tasks.go rename to internal/apps/upload/task/tasks.go index f203a7cd..bcfec085 100644 --- a/internal/apps/upload/tasks.go +++ b/internal/apps/upload/task/tasks.go @@ -2,7 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package task import ( "context" @@ -12,12 +12,13 @@ import ( "strings" "sync" + "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" ) -// 异步任务名称与管理类型定义 const ( // WarmImageCacheTask 图片压缩缓存预热任务标识 WarmImageCacheTask = "upload:warm_image_cache" @@ -54,29 +55,25 @@ type WarmImageCachePayload struct { Quality string `json:"quality"` } -// SystemCleanupHandler 系统定期垃圾清理异步任务处理器 - // WarmImageCacheHandler serially warms compressed image cache entries. type WarmImageCacheHandler struct{} -// Execute 执行系统清理(包含文件清理和历史消息推送日志清理) - // ValidatePayload validates and normalizes image cache warmup parameters. func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error) { if len(payload) == 0 { - return nil, errors.New(errImageCacheWarmupPayloadRequired) + return nil, errors.New(shared.ErrImageCacheWarmupPayloadRequired) } var req WarmImageCachePayload if err := json.Unmarshal(payload, &req); err != nil { - return nil, fmt.Errorf(errInvalidImageCacheWarmupPayload, err) + return nil, fmt.Errorf(shared.ErrInvalidImageCacheWarmupPayload, err) } req.Quality = strings.ToLower(strings.TrimSpace(req.Quality)) - if req.Quality != imageQualityLow && - req.Quality != imageQualityMedium && - req.Quality != imageQualityHigh { - return nil, errors.New(errInvalidImageCacheWarmupQuality) + if req.Quality != shared.ImageQualityLow && + req.Quality != shared.ImageQualityMedium && + req.Quality != shared.ImageQualityHigh { + return nil, errors.New(shared.ErrInvalidImageCacheWarmupQuality) } return json.Marshal(req) @@ -92,7 +89,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t var req WarmImageCachePayload if err := json.Unmarshal(normalizedPayload, &req); err != nil { - return nil, fmt.Errorf(errParseImageCacheWarmupPayload, err) + return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err) } task.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality) @@ -128,7 +125,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t Limit(batchSize). Find(&uploads).Error; err != nil { task.AppendLog(ctx, "查询图片上传记录失败: %v", err) - return nil, fmt.Errorf(errQueryImagesForCacheWarmup, err) + return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err) } if len(uploads) == 0 { @@ -147,7 +144,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t totalProcessed++ lastID = upload.ID - _, cacheHit, err := ensureCompressedImageCache(ctx, upload, req.Quality) + _, cacheHit, err := filesrv.EnsureCompressedImageCache(ctx, upload, req.Quality) if err != nil { totalFailed++ batchFailed++ @@ -184,4 +181,4 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t ) task.AppendLog(ctx, "%s", msg) return &task.TaskResult{Message: msg}, nil -} +} \ No newline at end of file diff --git a/internal/apps/upload/tasks_test.go b/internal/apps/upload/task/tasks_test.go similarity index 96% rename from internal/apps/upload/tasks_test.go rename to internal/apps/upload/task/tasks_test.go index 27bb5439..79b3293b 100644 --- a/internal/apps/upload/tasks_test.go +++ b/internal/apps/upload/task/tasks_test.go @@ -2,7 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +package task import ( "bytes" @@ -17,6 +17,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" @@ -201,7 +203,7 @@ func TestWarmImageCacheHandlerValidatePayload(t *testing.T) { { name: "normalizes quality", payload: []byte(`{"quality":" HIGH "}`), - wantQuality: imageQualityHigh, + wantQuality: shared.ImageQualityHigh, }, { name: "empty payload", @@ -281,7 +283,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { FilePath: firstPath, MimeType: "image/png", Extension: "png", - StorageDriver: storageDriverLocal, + StorageDriver: shared.StorageDriverLocal, Status: model.UploadStatusUsed, }, { @@ -291,7 +293,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { FilePath: secondPath, MimeType: "application/octet-stream", Extension: "jpg", - StorageDriver: storageDriverLocal, + StorageDriver: shared.StorageDriverLocal, Status: model.UploadStatusPending, }, { @@ -301,7 +303,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { FilePath: filepath.Join(testDir, "notes.txt"), MimeType: "text/plain", Extension: "txt", - StorageDriver: storageDriverLocal, + StorageDriver: shared.StorageDriverLocal, Status: model.UploadStatusUsed, }, { @@ -311,7 +313,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { FilePath: firstPath, MimeType: "image/png", Extension: "png", - StorageDriver: storageDriverLocal, + StorageDriver: shared.StorageDriverLocal, Status: model.UploadStatusDeleted, }, } @@ -339,7 +341,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { } for i := range records[:2] { - key := imageCompressionCacheKey(&records[i], imageQualityLow) + key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow) got, err := cache.Get(key) if err != nil { t.Errorf("cache.Get(%q) returned error: %v", key, err) diff --git a/internal/apps/upload/util/media.go b/internal/apps/upload/util/media.go new file mode 100644 index 00000000..0550422c --- /dev/null +++ b/internal/apps/upload/util/media.go @@ -0,0 +1,50 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" +) + +// IsImageExtension reports whether ext is a common image format. +func IsImageExtension(ext string) bool { + for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} { + if ext == imgExt { + return true + } + } + return false +} + +// IsArchiveExtension reports whether ext is a common archive format. +func IsArchiveExtension(ext string) bool { + for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} { + if ext == e { + return true + } + } + return false +} + +// IsDocumentExtension reports whether ext is a common document format. +func IsDocumentExtension(ext string) bool { + for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} { + if ext == e { + return true + } + } + return false +} + +// NormalizeImageQuality normalizes the requested image quality query parameter. +func NormalizeImageQuality(quality string) string { + switch strings.ToLower(quality) { + case shared.ImageQualityLow, shared.ImageQualityMedium, shared.ImageQualityHigh: + return strings.ToLower(quality) + default: + return shared.ImageQualityOrigin + } +} \ No newline at end of file diff --git a/internal/apps/upload/utils.go b/internal/apps/upload/util/utils.go similarity index 73% rename from internal/apps/upload/utils.go rename to internal/apps/upload/util/utils.go index 5f65aa4c..177f7365 100644 --- a/internal/apps/upload/utils.go +++ b/internal/apps/upload/util/utils.go @@ -2,7 +2,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package upload +// Package util provides upload media helpers and image utilities. +package util import ( "bytes" @@ -15,28 +16,27 @@ import ( "io" "strings" + "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/deepteams/webp" _ "golang.org/x/image/webp" // Register WebP decoder for image.Decode ) -const maxS3KeyLength = 1024 - // ValidateS3Key validates an S3 object key for safety. func ValidateS3Key(key string) error { if key == "" { - return errors.New(ErrS3KeyRequired) + return errors.New(shared.ErrS3KeyRequired) } - if len(key) > maxS3KeyLength { - return fmt.Errorf(ErrS3KeyTooLongFormat, maxS3KeyLength) + if len(key) > shared.MaxS3KeyLength { + return fmt.Errorf(shared.ErrS3KeyTooLongFormat, shared.MaxS3KeyLength) } if strings.HasPrefix(key, "/") { - return errors.New(ErrS3KeyStartsWithSlash) + return errors.New(shared.ErrS3KeyStartsWithSlash) } if strings.Contains(key, "\x00") { - return errors.New(ErrS3KeyContainsNullBytes) + return errors.New(shared.ErrS3KeyContainsNullBytes) } return nil @@ -45,34 +45,31 @@ func ValidateS3Key(key string) error { // CompressImageToWebP decodes an image from srcReader and encodes it into WebP format // using the specified quality (low -> 60, medium -> 75, high -> 85). func CompressImageToWebP(srcReader io.Reader, quality string) ([]byte, error) { - // Decode the image img, format, err := image.Decode(srcReader) if err != nil { return nil, fmt.Errorf("failed to decode image (format: %s): %w", format, err) } - // Determine quality var qualityScore float32 switch strings.ToLower(quality) { - case imageQualityLow: + case shared.ImageQualityLow: qualityScore = 60 - case imageQualityMedium: + case shared.ImageQualityMedium: qualityScore = 75 - case imageQualityHigh, "": + case shared.ImageQualityHigh, "": qualityScore = 85 default: qualityScore = 85 } - // Encode to WebP var buf bytes.Buffer err = webp.Encode(&buf, img, &webp.EncoderOptions{ Quality: qualityScore, - Method: 4, // Default method + Method: 4, }) if err != nil { return nil, fmt.Errorf("failed to encode WebP: %w", err) } return buf.Bytes(), nil -} +} \ No newline at end of file diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go index e128a981..28e98644 100644 --- a/internal/apps/user/access_tokens.go +++ b/internal/apps/user/access_tokens.go @@ -43,7 +43,7 @@ func ListAccessTokens(c *gin.Context) { var tokens []model.AccessToken if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -67,19 +67,19 @@ func CreateAccessToken(c *gin.Context) { var req createTokenRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusOK, response.Err(errBindParamsFailed)) + response.AbortBadRequest(c, errBindParamsFailed) return } req.Name = strings.TrimSpace(req.Name) if req.Name == "" { - c.JSON(http.StatusOK, response.Err(errTokenNameRequired)) + response.AbortBadRequest(c, errTokenNameRequired) return } // 只有管理员才能创建具有管理员权限的令牌 if req.IsAdmin && !currUser.IsAdmin { - c.JSON(http.StatusOK, response.Err(errAdminTokenRequiresAdmin)) + response.AbortBadRequest(c, errAdminTokenRequiresAdmin) return } @@ -91,19 +91,19 @@ func CreateAccessToken(c *gin.Context) { var count int64 if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if int(count) >= maxLimit { - c.JSON(http.StatusOK, response.Err(errAccessTokenLimitReached)) + response.AbortBadRequest(c, errAccessTokenLimitReached) return } // 生成 Token tokenStr, err := model.GenerateTokenString() if err != nil { - c.JSON(http.StatusOK, response.Err(errGenerateTokenFailed)) + response.AbortBadRequest(c, errGenerateTokenFailed) return } @@ -119,7 +119,7 @@ func CreateAccessToken(c *gin.Context) { } if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -146,18 +146,18 @@ func DeleteAccessToken(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(errInvalidTokenID)) + response.AbortBadRequest(c, errInvalidTokenID) return } tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{}) if tx.Error != nil { - c.JSON(http.StatusOK, response.Err(tx.Error.Error())) + response.AbortBadRequest(c, tx.Error.Error()) return } if tx.RowsAffected == 0 { - c.JSON(http.StatusOK, response.Err(errTokenNotFoundOrForbidden)) + response.AbortBadRequest(c, errTokenNotFoundOrForbidden) return } @@ -181,20 +181,20 @@ func RotateAccessToken(c *gin.Context) { idStr := c.Param("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { - c.JSON(http.StatusOK, response.Err(errInvalidTokenID)) + response.AbortBadRequest(c, errInvalidTokenID) return } var tokenRecord model.AccessToken if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil { - c.JSON(http.StatusOK, response.Err(errTokenNotFoundOrForbidden)) + response.AbortBadRequest(c, errTokenNotFoundOrForbidden) return } // 生成新的 Token newTokenStr, err := model.GenerateTokenString() if err != nil { - c.JSON(http.StatusOK, response.Err(errGenerateTokenFailed)) + response.AbortBadRequest(c, errGenerateTokenFailed) return } @@ -205,7 +205,7 @@ func RotateAccessToken(c *gin.Context) { tokenRecord.MaskedToken = newMaskedToken if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 02923f41..c95796c0 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -139,7 +139,7 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool { func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { if req.Code != "" { if !verifyEmailCode(ctx, user.Email, "login", req.Code) { - c.JSON(http.StatusOK, response.Err(errEmailCodeInvalidOrExpired)) + response.AbortBadRequest(c, errEmailCodeInvalidOrExpired) return errors.New("handled") } return nil @@ -149,7 +149,7 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi if !isSMTPConfigured(ctx) || user.Email == "" { codeKey := getEmailCodeKey("login", user.Email) if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil { - c.JSON(http.StatusOK, response.Err(errGenerateEmailCodeFailed)) + response.AbortBadRequest(c, errGenerateEmailCodeFailed) return errors.New("handled") } var msg string @@ -158,7 +158,7 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi } else { msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" } - c.JSON(http.StatusOK, response.Err(msg)) + response.AbortBadRequest(c, msg) return errors.New("handled") } @@ -167,13 +167,13 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi err := db.GetJSON(ctx, cooldownKey, &temp) if err != nil { if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return errors.New("handled") } } maskedEmail := pkgu.MaskEmail(user.Email) - c.JSON(http.StatusOK, response.Err(errNeedEmailCodePrefix+maskedEmail)) + response.AbortBadRequest(c, errNeedEmailCodePrefix+maskedEmail) return errors.New("handled") } @@ -190,18 +190,18 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi func SendEmailCode(c *gin.Context) { var req sendEmailCodeRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } req.Email = strings.TrimSpace(req.Email) if req.Email == "" { - c.JSON(http.StatusOK, response.Err(errEmailRequired)) + response.AbortBadRequest(c, errEmailRequired) return } if req.Scene != "register" { - c.JSON(http.StatusOK, response.Err(errUnsupportedEmailScene)) + response.AbortBadRequest(c, errUnsupportedEmailScene) return } @@ -209,11 +209,11 @@ func SendEmailCode(c *gin.Context) { var count int64 if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if count > 0 { - c.JSON(http.StatusOK, response.Err(errEmailAlreadyRegistered)) + response.AbortBadRequest(c, errEmailAlreadyRegistered) return } @@ -221,12 +221,12 @@ func SendEmailCode(c *gin.Context) { var temp string err := db.GetJSON(ctx, cooldownKey, &temp) if err == nil { - c.JSON(http.StatusOK, response.Err(errEmailCodeCooldown)) + response.AbortBadRequest(c, errEmailCodeCooldown) return } if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -271,37 +271,37 @@ type updateProfileRequest struct { func UpdateProfile(c *gin.Context) { var req updateProfileRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) if userObj == nil { - c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired)) + response.AbortUnauthorized(c, errLoginRequired) return } ctx := c.Request.Context() var dbUser model.User if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, response.Err(errUserNotFound)) + response.AbortBadRequest(c, errUserNotFound) return } req.Email = strings.TrimSpace(req.Email) if req.Email != "" && req.Email != dbUser.Email { if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") { - c.JSON(http.StatusOK, response.Err(errEmailFormatInvalid)) + response.AbortBadRequest(c, errEmailFormatInvalid) return } var count int64 if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", req.Email, dbUser.ID).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if count > 0 { - c.JSON(http.StatusOK, response.Err(errEmailAlreadyBound)) + response.AbortBadRequest(c, errEmailAlreadyBound) return } } @@ -319,7 +319,7 @@ func UpdateProfile(c *gin.Context) { dbUser.Location = strings.TrimSpace(req.Location) if err := db.DB(ctx).Save(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index fad35df2..70e85d49 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -109,17 +109,17 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro // @Router /api/v1/user/login [post] func Login(c *gin.Context) { if !isPasswordLoginEnabled() { - c.JSON(http.StatusOK, response.Err(errPasswordLoginDisabled)) + response.AbortBadRequest(c, errPasswordLoginDisabled) return } var req loginRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } req.Username = strings.TrimSpace(req.Username) if req.Username == "" || req.Password == "" { - c.JSON(http.StatusOK, response.Err(errInvalidParams)) + response.AbortBadRequest(c, errInvalidParams) return } @@ -127,12 +127,12 @@ func Login(c *gin.Context) { 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()) - c.JSON(http.StatusOK, response.Err(errUsernameOrPasswordWrong)) + response.AbortBadRequest(c, errUsernameOrPasswordWrong) return } if !user.IsActive { logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - c.JSON(http.StatusOK, response.Err(common.BannedAccount)) + response.AbortBadRequest(c, common.BannedAccount) return } @@ -141,7 +141,7 @@ func Login(c *gin.Context) { if !user.CheckPassword(req.Password) { logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - c.JSON(http.StatusOK, response.Err(errUsernameOrPasswordWrong)) + response.AbortBadRequest(c, errUsernameOrPasswordWrong) return } @@ -162,11 +162,11 @@ func Login(c *gin.Context) { user.LastLoginAt = time.Now() if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := setLoginSession(ctx, c, &user); err != nil { - c.JSON(http.StatusOK, response.Err(errSaveSessionFailed)) + response.AbortBadRequest(c, errSaveSessionFailed) return } @@ -190,13 +190,13 @@ func Login(c *gin.Context) { // @Router /api/v1/user/register [post] func Register(c *gin.Context) { if !isRegistrationEnabled() || !isPasswordRegisterEnabled() { - c.JSON(http.StatusOK, response.Err(errRegistrationDisabled)) + response.AbortBadRequest(c, errRegistrationDisabled) return } var req registerRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -208,15 +208,15 @@ func Register(c *gin.Context) { req.Code = strings.TrimSpace(req.Code) if req.Username == "" || req.Password == "" { - c.JSON(http.StatusOK, response.Err(errInvalidParams)) + response.AbortBadRequest(c, errInvalidParams) return } if req.Email == "" { - c.JSON(http.StatusOK, response.Err(errEmailRequired)) + response.AbortBadRequest(c, errEmailRequired) return } if len(req.Password) < minPasswordLength { - c.JSON(http.StatusOK, response.Err(errPasswordTooShort)) + response.AbortBadRequest(c, errPasswordTooShort) return } @@ -224,7 +224,7 @@ func Register(c *gin.Context) { // 邮箱注册验证校验 if err := validateRegisterEmailVerification(ctx, &req); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -245,17 +245,17 @@ func Register(c *gin.Context) { user.Nickname = req.Username } if err := user.SetEncryptedPassword(req.Password); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } if err := setLoginSession(ctx, c, &user); err != nil { - c.JSON(http.StatusOK, response.Err(errSaveSessionFailed)) + response.AbortBadRequest(c, errSaveSessionFailed) return } @@ -281,7 +281,7 @@ func Logout(c *gin.Context) { session.Options(oauth.GetSessionOptions(-1)) session.Clear() if err := session.Save(); err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } c.JSON(http.StatusOK, response.OK("")) @@ -306,7 +306,7 @@ type changePasswordRequest struct { func ChangePassword(c *gin.Context) { var req changePasswordRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } @@ -314,47 +314,47 @@ func ChangePassword(c *gin.Context) { req.NewPassword = strings.TrimSpace(req.NewPassword) if req.OldPassword == "" || req.NewPassword == "" { - c.JSON(http.StatusOK, response.Err(errInvalidParams)) + response.AbortBadRequest(c, errInvalidParams) return } if len(req.NewPassword) < minPasswordLength { - c.JSON(http.StatusOK, response.Err(errNewPasswordTooShort)) + response.AbortBadRequest(c, errNewPasswordTooShort) return } userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) if userObj == nil { - c.JSON(http.StatusUnauthorized, response.Err(errLoginRequired)) + response.AbortUnauthorized(c, errLoginRequired) return } ctx := c.Request.Context() var dbUser model.User if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, response.Err(errUserNotFound)) + response.AbortBadRequest(c, errUserNotFound) return } // 校验旧密码 if !dbUser.CheckPassword(req.OldPassword) { - c.JSON(http.StatusOK, response.Err(errOldPasswordIncorrect)) + response.AbortBadRequest(c, errOldPasswordIncorrect) return } // 加密并更新为新密码 if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil { - c.JSON(http.StatusOK, response.Err(errPasswordEncryptFailed)) + response.AbortBadRequest(c, errPasswordEncryptFailed) return } if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { - c.JSON(http.StatusOK, response.Err(err.Error())) + response.AbortBadRequest(c, err.Error()) return } // 吊销该用户所有的 Access Token if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { - c.JSON(http.StatusOK, response.Err("吊销 Access Token 失败: "+err.Error())) + response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error()) return } diff --git a/internal/common/response/abort.go b/internal/common/response/abort.go new file mode 100644 index 00000000..581629ab --- /dev/null +++ b/internal/common/response/abort.go @@ -0,0 +1,45 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package response + +import ( + "net/http" + + "github.com/gin-gonic/gin" +) + +// AbortBadRequest 以 400 中断请求并将错误挂载到 Gin Error 链,供全局中间件统一记录 Trace 并响应。 +func AbortBadRequest(c *gin.Context, msg string) { + AbortWithError(c, http.StatusBadRequest, msg) +} + +// AbortUnauthorized 以 401 中断请求并将错误挂载到 Gin Error 链。 +func AbortUnauthorized(c *gin.Context, msg string) { + AbortWithError(c, http.StatusUnauthorized, msg) +} + +// AbortForbidden 以 403 中断请求并将错误挂载到 Gin Error 链。 +func AbortForbidden(c *gin.Context, msg string) { + AbortWithError(c, http.StatusForbidden, msg) +} + +// AbortNotFound 以 404 中断请求并将错误挂载到 Gin Error 链。 +func AbortNotFound(c *gin.Context, msg string) { + AbortWithError(c, http.StatusNotFound, msg) +} + +// AbortInternal 以 500 中断请求并将错误挂载到 Gin Error 链。 +func AbortInternal(c *gin.Context, msg string) { + AbortWithError(c, http.StatusInternalServerError, msg) +} + +// AbortTooManyRequests 以 429 中断请求并将错误挂载到 Gin Error 链。 +func AbortTooManyRequests(c *gin.Context, msg string) { + AbortWithError(c, http.StatusTooManyRequests, msg) +} + +// AbortConflict 以 409 中断请求并将错误挂载到 Gin Error 链。 +func AbortConflict(c *gin.Context, msg string) { + AbortWithError(c, http.StatusConflict, msg) +} \ No newline at end of file diff --git a/internal/common/response/middleware.go b/internal/common/response/middleware.go new file mode 100644 index 00000000..5f895085 --- /dev/null +++ b/internal/common/response/middleware.go @@ -0,0 +1,40 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package response + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" + "go.opentelemetry.io/otel/codes" + "go.opentelemetry.io/otel/trace" +) + +// ErrorHandlerMiddleware 捕获 c.Errors 并统一格式化为 JSON 返回给客户端,同时将其记录到 Span 异常中。 +// 与 AbortWithError / AbortBadRequest 等配合使用,是全局 OTel 友好错误响应的唯一出口。 +func ErrorHandlerMiddleware() gin.HandlerFunc { + return func(c *gin.Context) { + c.Next() + + if len(c.Errors) == 0 || c.Writer.Written() { + return + } + + err := c.Errors.Last().Err + span := trace.SpanFromContext(c.Request.Context()) + if span.IsRecording() { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + + var apiErr *APIError + if errors.As(err, &apiErr) { + c.JSON(apiErr.Code, Err(apiErr.Msg)) + return + } + + c.JSON(http.StatusInternalServerError, Err("内部系统错误")) + } +} \ No newline at end of file diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go index 77f05633..73fb89b3 100644 --- a/internal/router/middlewares.go +++ b/internal/router/middlewares.go @@ -7,7 +7,6 @@ package router import ( "context" - "errors" "net/http" "strconv" "strings" @@ -106,29 +105,7 @@ func corsMiddleware() gin.HandlerFunc { } } -// errorHandlerMiddleware 捕获 c.Errors 并统一格式化为 JSON 返回给客户端,同时将其记录到 Span 异常中 +// errorHandlerMiddleware 委托给 response.ErrorHandlerMiddleware,保持路由层单一入口。 func errorHandlerMiddleware() gin.HandlerFunc { - return func(c *gin.Context) { - c.Next() - - if len(c.Errors) > 0 { - err := c.Errors.Last().Err - span := trace.SpanFromContext(c.Request.Context()) - - // 1. 如果有活跃的 Span,将错误信息记录到 Trace 中,并把 Span 状态置为 Error - if span.IsRecording() { - span.RecordError(err) - span.SetStatus(codes.Error, err.Error()) - } - - // 2. 将错误转化为统一的 JSON 格式响应给客户端 - var apiErr *response.APIError - if errors.As(err, &apiErr) { - c.JSON(apiErr.Code, response.Err(apiErr.Msg)) - } else { - // 兜底策略:未知的系统级错误 - c.JSON(http.StatusInternalServerError, response.Err("内部系统错误")) - } - } - } + return response.ErrorHandlerMiddleware() } diff --git a/internal/testhelper/gin.go b/internal/testhelper/gin.go new file mode 100644 index 00000000..11b881b4 --- /dev/null +++ b/internal/testhelper/gin.go @@ -0,0 +1,20 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package testhelper + +import ( + "github.com/Rain-kl/Wavelet/internal/common/response" + "github.com/gin-gonic/gin" +) + +// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。 +func NewTestGinEngine(middlewares ...gin.HandlerFunc) *gin.Engine { + gin.SetMode(gin.TestMode) + r := gin.New() + r.Use(response.ErrorHandlerMiddleware()) + for _, middleware := range middlewares { + r.Use(middleware) + } + return r +} \ No newline at end of file