From e06f76436e512ac6b1c876cbdb36c7d8915e7de0 Mon Sep 17 00:00:00 2001 From: ryan Date: Tue, 9 Jun 2026 13:44:29 +0800 Subject: [PATCH] refactor: extract magic numbers to named constants for mnd lint compliance --- internal/apps/admin/logs/routers.go | 282 ++++++++++++++------------ internal/apps/admin/status/routers.go | 32 ++- internal/apps/admin/user/routers.go | 9 +- internal/apps/upload/routers.go | 257 +++++++++++++---------- internal/apps/user/controllers.go | 192 ++++++++++-------- internal/config/config.go | 19 +- internal/db/clickhouse.go | 11 +- internal/db/postgres_logger.go | 15 +- internal/logger/logger.go | 9 +- internal/logger/ringbuffer.go | 5 +- internal/logger/utils.go | 5 +- internal/model/access_token.go | 11 +- internal/router/middlewares.go | 2 +- internal/storage/cache.go | 8 +- internal/task/constants.go | 12 +- internal/task/scheduler/scheduler.go | 9 +- internal/task/worker/worker.go | 5 +- internal/util/cap/cap.go | 25 ++- internal/util/cap/manager.go | 40 ++-- internal/util/crypto.go | 7 +- internal/util/http_clients.go | 18 +- internal/util/mail/mail.go | 102 +++++----- internal/util/strings.go | 10 +- internal/util/uuid.go | 5 +- 24 files changed, 650 insertions(+), 440 deletions(-) diff --git a/internal/apps/admin/logs/routers.go b/internal/apps/admin/logs/routers.go index 3e3f3fdb..052f16e1 100644 --- a/internal/apps/admin/logs/routers.go +++ b/internal/apps/admin/logs/routers.go @@ -15,9 +15,11 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package logs 提供日志查询与分析功能 package logs import ( + "context" "encoding/json" "fmt" "net/http" @@ -35,7 +37,14 @@ import ( "github.com/gin-gonic/gin" ) -const defaultLimit = 200 +const ( + defaultLimit = 200 + maxLimit = 500 + maxPageSize = 100 + hoursInDay = 24 + analyticsDays = 7 + queryExtraArgs = 2 // pageSize + offset +) // logsResponse 历史日志查询响应 type logsResponse struct { @@ -68,8 +77,8 @@ func GetLogs(c *gin.Context) { if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 { limit = defaultLimit } - if limit > 500 { - limit = 500 + if limit > maxLimit { + limit = maxLimit } entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit) @@ -162,69 +171,26 @@ type accessLogsResponse struct { List []accessLogItem `json:"list"` } -// GetAccessLogs 获取 ClickHouse 异步采集的访问日志 -// @Summary 获取用户访问日志 -// @Description 分页并按照用户、接口路径、时间范围等维度检索 ClickHouse 用户访问日志列表(需要管理员权限,ClickHouse 未启用时报错) -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param page query int false "页码" default(1) -// @Param page_size query int false "每页条数" default(20) -// @Param username query string false "用户名模糊搜索" -// @Param path query string false "接口路径模糊搜索" -// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" -// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" -// @Success 200 {object} util.ResponseAny{data=logs.accessLogsResponse} "访问日志列表" -// @Failure 400 {object} util.ResponseAny "ClickHouse 未启用或参数错误" -// @Failure 401 {object} util.ResponseAny "未登录" -// @Failure 403 {object} util.ResponseAny "无管理员权限" -// @Router /api/v1/admin/logs/access [get] -func GetAccessLogs(c *gin.Context) { - // 1. 检查 ClickHouse 是否启用 - if !config.Config.ClickHouse.Enabled || db.ChConn == nil { - c.JSON(http.StatusBadRequest, util.Err("ClickHouse 存储服务未启用,无法检索访问日志")) - return - } - - // 2. 解析分页参数 - page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) - if page < 1 { - page = 1 - } - pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) - if pageSize < 1 { - pageSize = 20 - } - if pageSize > 100 { - pageSize = 100 - } - offset := (page - 1) * pageSize - - // 3. 按用户名过滤(预查 Postgres 映射 UserID) +// buildAccessLogFilters 构建 ClickHouse 访问日志查询过滤条件 +func buildAccessLogFilters(ctx context.Context, c *gin.Context) ([]string, []interface{}, []uint64, error) { + var conditions []string + var args []interface{} var userIDs []uint64 + + // 按用户名过滤 username := c.Query("username") if username != "" { - err := db.DB(c.Request.Context()).Model(&model.User{}). + err := db.DB(ctx).Model(&model.User{}). Where("username LIKE ?", "%"+username+"%"). Pluck("id", &userIDs).Error if err != nil { - c.JSON(http.StatusInternalServerError, util.Err("查询用户信息失败: "+err.Error())) - return + return nil, nil, nil, fmt.Errorf("查询用户信息失败: %w", err) } - // 如果指定了用户名搜索,但在 Postgres 中没匹配到任何用户,则直接返回空结果 if len(userIDs) == 0 { - c.JSON(http.StatusOK, util.OK(accessLogsResponse{ - Total: 0, - List: []accessLogItem{}, - })) - return + return nil, nil, nil, nil // 无匹配用户 } } - // 4. 构建 ClickHouse 条件查询子句与参数 - var conditions []string - var args []interface{} - if len(userIDs) > 0 { placeholders := make([]string, len(userIDs)) for i := range userIDs { @@ -259,29 +225,11 @@ func GetAccessLogs(c *gin.Context) { } } - whereClause := "" - if len(conditions) > 0 { - whereClause = "WHERE " + strings.Join(conditions, " AND ") - } + return conditions, args, userIDs, nil +} - // 5. 查询日志总数 - var total uint64 - countQuery := fmt.Sprintf("SELECT count() FROM user_access_logs %s", whereClause) - err := db.ChConn.QueryRow(c.Request.Context(), countQuery, args...).Scan(&total) - if err != nil { - c.JSON(http.StatusInternalServerError, util.Err("查询 ClickHouse 日志统计失败: "+err.Error())) - return - } - - if total == 0 { - c.JSON(http.StatusOK, util.OK(accessLogsResponse{ - Total: 0, - List: []accessLogItem{}, - })) - return - } - - // 6. 分页查询明细数据 +// fetchAccessLogDetails 查询 ClickHouse 访问日志明细并填充用户名 +func fetchAccessLogDetails(ctx context.Context, whereClause string, args []interface{}, pageSize int, offset int) ([]accessLogItem, error) { dataQuery := fmt.Sprintf(` SELECT id, user_id, path, method, ip, user_agent, headers, status, latency, created_at FROM user_access_logs @@ -290,11 +238,13 @@ func GetAccessLogs(c *gin.Context) { LIMIT ? OFFSET ? `, whereClause) - selectArgs := append(args, pageSize, offset) - rows, err := db.ChConn.Query(c.Request.Context(), dataQuery, selectArgs...) + selectArgs := make([]interface{}, len(args), len(args)+queryExtraArgs) + copy(selectArgs, args) + selectArgs = append(selectArgs, pageSize, offset) + + rows, err := db.ChConn.Query(ctx, dataQuery, selectArgs...) if err != nil { - c.JSON(http.StatusInternalServerError, util.Err("查询 ClickHouse 日志明细失败: "+err.Error())) - return + return nil, fmt.Errorf("查询 ClickHouse 日志明细失败: %w", err) } defer func() { _ = rows.Close() }() @@ -304,53 +254,105 @@ func GetAccessLogs(c *gin.Context) { for rows.Next() { var item accessLogItem var createdAt time.Time - err := rows.Scan( - &item.ID, - &item.UserID, - &item.Path, - &item.Method, - &item.IP, - &item.UserAgent, - &item.Headers, - &item.Status, - &item.Latency, - &createdAt, - ) - if err != nil { - c.JSON(http.StatusInternalServerError, util.Err("读取 ClickHouse 结果失败: "+err.Error())) - return + if err := rows.Scan(&item.ID, &item.UserID, &item.Path, &item.Method, &item.IP, &item.UserAgent, &item.Headers, &item.Status, &item.Latency, &createdAt); err != nil { + return nil, fmt.Errorf("读取 ClickHouse 结果失败: %w", err) } item.CreatedAt = createdAt.Format(time.RFC3339) list = append(list, item) fetchUserIDs = append(fetchUserIDs, item.UserID) } - // 7. 反查 Postgres 关联 Username 和 Nickname - userMap := make(map[uint64]struct { - Username string - Nickname string - }) - + // 反查 Postgres 关联 Username 和 Nickname if len(fetchUserIDs) > 0 { + userMap := make(map[uint64]struct{ Username, Nickname string }) var users []model.User - if err := db.DB(c.Request.Context()).Where("id IN ?", fetchUserIDs).Find(&users).Error; err == nil { + if err := db.DB(ctx).Where("id IN ?", fetchUserIDs).Find(&users).Error; err == nil { for _, u := range users { - userMap[u.ID] = struct { - Username string - Nickname string - }{ - Username: u.Username, - Nickname: u.Nickname, - } + userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} + } + } + for i := range list { + if info, ok := userMap[list[i].UserID]; ok { + list[i].Username = info.Username + list[i].Nickname = info.Nickname } } } - for i := range list { - if info, ok := userMap[list[i].UserID]; ok { - list[i].Username = info.Username - list[i].Nickname = info.Nickname - } + return list, nil +} + +// GetAccessLogs 获取 ClickHouse 异步采集的访问日志 +// @Summary 获取用户访问日志 +// @Description 分页并按照用户、接口路径、时间范围等维度检索 ClickHouse 用户访问日志列表(需要管理员权限,ClickHouse 未启用时报错) +// @Tags admin +// @Produce json +// @Security SessionCookie +// @Param page query int false "页码" default(1) +// @Param page_size query int false "每页条数" default(20) +// @Param username query string false "用户名模糊搜索" +// @Param path query string false "接口路径模糊搜索" +// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" +// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" +// @Success 200 {object} util.ResponseAny{data=logs.accessLogsResponse} "访问日志列表" +// @Failure 400 {object} util.ResponseAny "ClickHouse 未启用或参数错误" +// @Failure 401 {object} util.ResponseAny "未登录" +// @Failure 403 {object} util.ResponseAny "无管理员权限" +// @Router /api/v1/admin/logs/access [get] +func GetAccessLogs(c *gin.Context) { + // 1. 检查 ClickHouse 是否启用 + if !config.Config.ClickHouse.Enabled || db.ChConn == nil { + c.JSON(http.StatusBadRequest, util.Err("ClickHouse 存储服务未启用,无法检索访问日志")) + return + } + + // 2. 解析分页参数 + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + if page < 1 { + page = 1 + } + pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) + if pageSize < 1 { + pageSize = 20 + } + if pageSize > maxPageSize { + pageSize = maxPageSize + } + offset := (page - 1) * pageSize + + // 3. 构建过滤条件 + conditions, args, userIDs, err := buildAccessLogFilters(c.Request.Context(), c) + if err != nil { + c.JSON(http.StatusInternalServerError, util.Err(err.Error())) + return + } + if userIDs != nil && len(userIDs) == 0 { + c.JSON(http.StatusOK, util.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}})) + return + } + + whereClause := "" + if len(conditions) > 0 { + whereClause = "WHERE " + strings.Join(conditions, " AND ") + } + + // 4. 查询日志总数 + var total uint64 + countQuery := fmt.Sprintf("SELECT count() FROM user_access_logs %s", whereClause) + if err := db.ChConn.QueryRow(c.Request.Context(), countQuery, args...).Scan(&total); err != nil { + c.JSON(http.StatusInternalServerError, util.Err("查询 ClickHouse 日志统计失败: "+err.Error())) + return + } + if total == 0 { + c.JSON(http.StatusOK, util.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}})) + return + } + + // 5. 分页查询明细数据 + list, err := fetchAccessLogDetails(c.Request.Context(), whereClause, args, pageSize, offset) + if err != nil { + c.JSON(http.StatusInternalServerError, util.Err(err.Error())) + return } c.JSON(http.StatusOK, util.OK(accessLogsResponse{ @@ -404,11 +406,24 @@ func GetLogsAnalytics(c *gin.Context) { return } + ctx := c.Request.Context() // 7 天前 00:00:00 - startTime := time.Now().AddDate(0, 0, -6).Truncate(24 * time.Hour) + startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour) - // 2. 查询 7 天访问趋势 - trendRows, err := db.ChConn.Query(c.Request.Context(), ` + trendList := queryAccessTrend(ctx, startTime) + browserList := queryBrowserDistribution(ctx, startTime) + topUsers := queryTopActiveUsers(ctx, startTime) + + c.JSON(http.StatusOK, util.OK(logsAnalyticsResponse{ + Trend: trendList, + Browsers: browserList, + TopUsers: topUsers, + })) +} + +// queryAccessTrend 查询最近 7 天的访问趋势 +func queryAccessTrend(ctx context.Context, startTime time.Time) []trendItem { + trendRows, err := db.ChConn.Query(ctx, ` SELECT toDate(created_at) as date, count() as count FROM user_access_logs WHERE created_at >= ? @@ -417,8 +432,7 @@ func GetLogsAnalytics(c *gin.Context) { `, startTime) trendMap := make(map[string]uint64) - // 初始化最近 7 天的数据为 0,防止某天没有访问数据时导致日期断裂 - for i := 0; i < 7; i++ { + for i := 0; i < analyticsDays; i++ { dStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02") trendMap[dStr] = 0 } @@ -436,16 +450,19 @@ func GetLogsAnalytics(c *gin.Context) { } var trendList []trendItem - for i := 6; i >= 0; i-- { + for i := analyticsDays - 1; i >= 0; i-- { dStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02") trendList = append(trendList, trendItem{ Date: dStr, Count: trendMap[dStr], }) } + return trendList +} - // 3. 查询浏览器分布排行 - uaRows, err := db.ChConn.Query(c.Request.Context(), ` +// queryBrowserDistribution 查询浏览器分布排行 +func queryBrowserDistribution(ctx context.Context, startTime time.Time) []browserItem { + uaRows, err := db.ChConn.Query(ctx, ` SELECT user_agent, count() as count FROM user_access_logs WHERE created_at >= ? @@ -472,14 +489,15 @@ func GetLogsAnalytics(c *gin.Context) { Count: cnt, }) } - - // 排序:按访问次数降序 sort.Slice(browserList, func(i, j int) bool { return browserList[i].Count > browserList[j].Count }) + return browserList +} - // 4. 查询活跃用户 Top 10 (user_id > 0 代表已登录用户) - userRows, err := db.ChConn.Query(c.Request.Context(), ` +// queryTopActiveUsers 查询活跃用户 Top 10 +func queryTopActiveUsers(ctx context.Context, startTime time.Time) []topUserItem { + userRows, err := db.ChConn.Query(ctx, ` SELECT user_id, count() as count FROM user_access_logs WHERE created_at >= ? AND user_id > 0 @@ -509,10 +527,9 @@ func GetLogsAnalytics(c *gin.Context) { Username string Nickname string }) - if len(userIDs) > 0 { var users []model.User - if errProfile := db.DB(c.Request.Context()).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil { + if errProfile := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil { for _, u := range users { userProfileMap[u.ID] = struct { Username string @@ -534,12 +551,7 @@ func GetLogsAnalytics(c *gin.Context) { Count: userCountMap[uid], }) } - - c.JSON(http.StatusOK, util.OK(logsAnalyticsResponse{ - Trend: trendList, - Browsers: browserList, - TopUsers: topUsers, - })) + return topUsers } // parseBrowserName 简易的 User-Agent 浏览器类型识别 diff --git a/internal/apps/admin/status/routers.go b/internal/apps/admin/status/routers.go index 08f71b6f..92ee1e25 100644 --- a/internal/apps/admin/status/routers.go +++ b/internal/apps/admin/status/routers.go @@ -15,6 +15,7 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package status 提供系统状态查询接口 package status import ( @@ -31,6 +32,17 @@ import ( // startTime 记录服务启动时间 var startTime = time.Now() +const ( + hoursInDay = 24 + minutesInHour = 60 + secondsInMinute = 60 + nanosPerSecond = 1e9 + binaryKB = 2 + binaryMB = 3 + binaryGB = 4 + valueThreshold = 10 // 格式化时区分整数显示的阈值 +) + // SystemStatusResponse 系统状态响应结构体 type SystemStatusResponse struct { Uptime string `json:"uptime"` @@ -77,11 +89,11 @@ func formatBytes(bytes uint64) string { value := float64(bytes) / float64(div) var suffix string switch exp { - case 0: + case binaryKB: suffix = "KiB" - case 1: + case binaryMB: suffix = "MiB" - case 2: + case binaryGB: suffix = "GiB" default: suffix = "TiB" @@ -93,7 +105,7 @@ func formatBytes(bytes uint64) string { // - 如果 < 10,则格式化为 "%.1f" (e.g. "9.0 KiB") // - 如果不是整数(如 5.8, 9.1, 7.6, 4.8):格式化为 "%.1f" if value == math.Trunc(value) { - if value >= 10 { + if value >= valueThreshold { return fmt.Sprintf("%.0f %s", value, suffix) } return fmt.Sprintf("%.1f %s", value, suffix) @@ -103,10 +115,10 @@ func formatBytes(bytes uint64) string { // formatDuration 格式化时间持续时间 func formatDuration(d time.Duration) string { - days := int(d.Hours()) / 24 - hours := int(d.Hours()) % 24 - minutes := int(d.Minutes()) % 60 - seconds := int(d.Seconds()) % 60 + days := int(d.Hours()) / hoursInDay + hours := int(d.Hours()) % hoursInDay + minutes := int(d.Minutes()) % minutesInHour + seconds := int(d.Seconds()) % secondsInMinute var res string if days > 0 { @@ -150,7 +162,7 @@ func GetSystemStatus(c *gin.Context) { var lastPause string if m.NumGC > 0 { - lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/1e9) + lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond) } else { lastPause = "0.000s" } @@ -181,7 +193,7 @@ func GetSystemStatus(c *gin.Context) { OtherSys: formatBytes(m.OtherSys), NextGC: formatBytes(m.NextGC), LastGCTime: lastGCTime, - PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/1e9), + PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond), LastPause: lastPause, NumGC: m.NumGC, } diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index 03c18aec..f4f79254 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -30,6 +30,9 @@ import ( "gorm.io/gorm" ) +// minPasswordLength 密码最小长度 +const minPasswordLength = 8 + // listUsersRequest 用户列表查询请求 type listUsersRequest struct { Page int `form:"page" binding:"min=1"` @@ -42,7 +45,7 @@ type user struct { ID uint64 `json:"id"` Username string `json:"username"` Nickname string `json:"nickname"` - AvatarUrl string `json:"avatar_url"` + AvatarURL string `json:"avatar_url"` IsActive bool `json:"is_active"` IsAdmin bool `json:"is_admin"` LastLoginAt time.Time `json:"last_login_at"` @@ -216,7 +219,7 @@ func CreateUser(c *gin.Context) { c.JSON(http.StatusBadRequest, util.Err(usernameRequired)) return } - if len(req.Password) < 8 { + if len(req.Password) < minPasswordLength { c.JSON(http.StatusBadRequest, util.Err(passwordTooShort)) return } @@ -258,7 +261,7 @@ func CreateUser(c *gin.Context) { ID: newUser.ID, Username: newUser.Username, Nickname: newUser.Nickname, - AvatarUrl: newUser.AvatarUrl, + AvatarURL: newUser.AvatarURL, IsActive: newUser.IsActive, IsAdmin: newUser.IsAdmin, LastLoginAt: newUser.LastLoginAt, diff --git a/internal/apps/upload/routers.go b/internal/apps/upload/routers.go index 3773fb78..eec890c8 100644 --- a/internal/apps/upload/routers.go +++ b/internal/apps/upload/routers.go @@ -20,12 +20,14 @@ package upload import ( "archive/zip" "bytes" + "context" "crypto/sha256" "encoding/hex" "encoding/json" "errors" "fmt" "io" + "mime/multipart" "net/http" "net/url" "os" @@ -46,7 +48,12 @@ import ( "gorm.io/gorm" ) -const maxUploadSize = 32 * 1024 * 1024 // 32MB +const ( + maxUploadSize = 32 * 1024 * 1024 // 32MB + detectContentBytes = 512 // http.DetectContentType 需要的最小字节数 + uploadDirPerm = 0755 // 上传目录权限 + uploadFilePerm = 0644 // 上传文件权限 +) type batchDownloadRequest struct { IDs []string `json:"ids" binding:"required,min=1"` @@ -67,6 +74,8 @@ type batchDownloadRequest struct { // @Failure 401 {object} util.ResponseAny "未登录" // @Failure 500 {object} util.ResponseAny "内部错误" // @Router /api/v1/upload [post] +// +//nolint:revive func UploadFile(c *gin.Context) { c.Header("X-Content-Type-Options", "nosniff") c.Header("Content-Security-Policy", "sandbox") @@ -104,20 +113,9 @@ func UploadFile(c *gin.Context) { } // 3. 校验文件后缀是否在允许的系统配置列表中 - var sc model.SystemConfig - if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" { - allowedExts := strings.Split(strings.ToLower(sc.Value), ",") - allowed := false - for _, allowedExt := range allowedExts { - if strings.TrimSpace(allowedExt) == ext { - allowed = true - break - } - } - if !allowed { - c.JSON(http.StatusOK, util.Err(ErrUnsupportedFormat)) - return - } + if errMsg := validateUploadExtension(ctx, ext); errMsg != "" { + c.JSON(http.StatusOK, util.Err(errMsg)) + return } // 4. 读取文件并计算 Hash @@ -130,107 +128,39 @@ func UploadFile(c *gin.Context) { } fileHash := hex.EncodeToString(hashWriter.Sum(nil)) - - mimeType := http.DetectContentType(buf.Bytes()[:min(512, int(size))]) - if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" { - mimeType = header.Header.Get("Content-Type") - } + mimeType := detectMimeType(&buf, header, size) // 校验真实 MIME Type 是否与常见图片扩展名匹配,防止 Polyglot / HTML 注入攻击 - isImageExt := false - for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} { - if ext == imgExt { - isImageExt = true - break - } - } - if isImageExt && !strings.HasPrefix(mimeType, "image/") { + if isImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") { c.JSON(http.StatusOK, util.Err(ErrFileContentExtensionMismatch)) return } // 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件 - 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 { - // 命中了相同文件,直接生成新记录指向已有的存储路径(实现秒传) - id := idgen.NextUint64ID() - newUpload := model.Upload{ - ID: id, - UserID: currUser.ID, - FileName: origName, - FilePath: existing.FilePath, - FileSize: size, - MimeType: mimeType, - Extension: ext, - Hash: fileHash, - StorageDriver: existing.StorageDriver, - Type: c.DefaultPostForm("type", "generic"), - Status: model.UploadStatusUsed, - Metadata: existing.Metadata, - } - - if err := db.DB(ctx).Create(&newUpload).Error; err != nil { - c.JSON(http.StatusOK, util.Err(ErrSaveUploadRecordFailed)) - return - } - - logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath) - c.JSON(http.StatusOK, util.OK(newUpload)) + handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName) + if handled { return - } else if !errors.Is(err, gorm.ErrRecordNotFound) { + } + if lookupErr != nil && !errors.Is(lookupErr, gorm.ErrRecordNotFound) { c.JSON(http.StatusOK, util.Err(ErrFileValidationFailed)) return } // 7. 解析可选元数据字段 - metadataStr := c.DefaultPostForm("metadata", "") - var meta model.UploadMetadata - if metadataStr != "" { - if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil { - c.JSON(http.StatusOK, util.Err(ErrInvalidMetadataJSON)) - return - } + meta, errMsg := parseUploadMetadata(c, mimeType) + if errMsg != "" { + c.JSON(http.StatusOK, util.Err(errMsg)) + return } - meta.OriginalMime = mimeType - meta.UserAgent = c.Request.UserAgent() - meta.ClientIP = c.ClientIP() - id := idgen.NextUint64ID() subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext) - var storageDriver string - // 8. 写入底层存储驱动 (优先 S3 驱动,无配置或未开启则 fallback 至本地文件) - if storage.IsEnabled() { - storageDriver = "s3" - meta.Bucket = config.Config.S3.Bucket - fullKey := storage.BuildKey(subPath) - - err = storage.PutObject(ctx, fullKey, bytes.NewReader(buf.Bytes()), size, mimeType) - if err != nil { - logger.ErrorF(ctx, "S3 存储上传失败: %v", err) - c.JSON(http.StatusOK, util.Err(ErrSaveFileFailed)) - return - } - } else { - storageDriver = "local" - localDir := filepath.Join("uploads", time.Now().Format("2006/01/02")) - if err := os.MkdirAll(localDir, 0755); err != nil { - logger.ErrorF(ctx, "创建本地上传目录失败: %v", err) - c.JSON(http.StatusOK, util.Err(ErrSaveFileFailed)) - return - } - - localPath := filepath.Join(localDir, fmt.Sprintf("%d.%s", id, ext)) - if err := os.WriteFile(localPath, buf.Bytes(), 0644); err != nil { - logger.ErrorF(ctx, "本地磁盘写入文件失败: %v", err) - c.JSON(http.StatusOK, util.Err(ErrSaveFileFailed)) - return - } - // 统一使用相对路径,方便将来环境移植或备份 - subPath = localPath + storageDriver, subPath, errMsg := storeUploadFile(ctx, id, ext, subPath, size, mimeType, &buf, &meta) + if errMsg != "" { + c.JSON(http.StatusOK, util.Err(errMsg)) + return } // 9. 保存文件记录至数据库 @@ -249,12 +179,8 @@ func UploadFile(c *gin.Context) { Metadata: meta, } - if err := db.DB(ctx).Create(&newUpload).Error; err != nil { - // 失败时若为本地存储,可以尝试清理已保存的垃圾文件 - if storageDriver == "local" { - _ = os.Remove(subPath) - } - c.JSON(http.StatusOK, util.Err(ErrSaveUploadRecordFailed)) + if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" { + c.JSON(http.StatusOK, util.Err(err)) return } @@ -548,9 +474,126 @@ func DeleteFile(c *gin.Context) { c.JSON(http.StatusOK, util.OKNil()) } -func min(a, b int) int { - if a < b { - return a +// validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中 +func validateUploadExtension(ctx context.Context, ext string) string { + var sc model.SystemConfig + if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" { + allowedExts := strings.Split(strings.ToLower(sc.Value), ",") + allowed := false + for _, allowedExt := range allowedExts { + if strings.TrimSpace(allowedExt) == ext { + allowed = true + break + } + } + if !allowed { + return ErrUnsupportedFormat + } } - return b + return "" +} + +// tryInstantUpload 尝试秒传:若数据库已存在相同 Hash 且大小一致的可用文件,直接生成新记录 +func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string) (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 + } + + id := idgen.NextUint64ID() + newUpload := model.Upload{ + ID: id, + UserID: currUser.ID, + FileName: origName, + FilePath: existing.FilePath, + FileSize: size, + MimeType: mimeType, + Extension: ext, + Hash: fileHash, + StorageDriver: existing.StorageDriver, + Type: c.DefaultPostForm("type", "generic"), + Status: model.UploadStatusUsed, + Metadata: existing.Metadata, + } + + if err := db.DB(ctx).Create(&newUpload).Error; err != nil { + c.JSON(http.StatusOK, util.Err(ErrSaveUploadRecordFailed)) + return true, nil + } + + logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath) + c.JSON(http.StatusOK, util.OK(newUpload)) + return true, nil +} + +// storeUploadFile 将文件写入底层存储驱动(S3 或本地磁盘) +func storeUploadFile(ctx context.Context, id uint64, ext, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) { + if storage.IsEnabled() { + meta.Bucket = config.Config.S3.Bucket + fullKey := storage.BuildKey(subPath) + if err := storage.PutObject(ctx, fullKey, bytes.NewReader(buf.Bytes()), size, mimeType); err != nil { + logger.ErrorF(ctx, "S3 存储上传失败: %v", err) + return "", "", ErrSaveFileFailed + } + return "s3", subPath, "" + } + + localDir := filepath.Join("uploads", time.Now().Format("2006/01/02")) + if err := os.MkdirAll(localDir, uploadDirPerm); err != nil { + logger.ErrorF(ctx, "创建本地上传目录失败: %v", err) + return "", "", ErrSaveFileFailed + } + + localPath := filepath.Join(localDir, fmt.Sprintf("%d.%s", id, ext)) + if err := os.WriteFile(localPath, buf.Bytes(), uploadFilePerm); err != nil { + logger.ErrorF(ctx, "本地磁盘写入文件失败: %v", err) + return "", "", ErrSaveFileFailed + } + return "local", localPath, "" +} + +// 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 + } + } + meta.OriginalMime = mimeType + meta.UserAgent = c.Request.UserAgent() + meta.ClientIP = c.ClientIP() + 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))]) + 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 { + if storageDriver == "local" { + _ = os.Remove(filePath) + } + return ErrSaveUploadRecordFailed + } + return "" } diff --git a/internal/apps/user/controllers.go b/internal/apps/user/controllers.go index 31a301da..2fbc0a61 100644 --- a/internal/apps/user/controllers.go +++ b/internal/apps/user/controllers.go @@ -37,6 +37,14 @@ import ( "github.com/gin-gonic/gin" ) +const ( + verificationCodeRange = 900000 // 验证码随机范围 + verificationCodeOffset = 100000 // 验证码偏移量(保证 6 位) + emailCodeExpiry = 5 * time.Minute // 验证码有效期 + emailCodeCooldown = 60 * time.Second // 验证码发送冷却时间 + minPasswordLength = 8 // 密码最小长度 +) + type loginRequest struct { Username string `json:"username"` Password string `json:"password"` @@ -91,8 +99,8 @@ func isSMTPConfigured(ctx context.Context) bool { } func generateVerificationCode() string { - n, _ := rand.Int(rand.Reader, big.NewInt(900000)) - return fmt.Sprintf("%06d", n.Int64()+100000) + n, _ := rand.Int(rand.Reader, big.NewInt(verificationCodeRange)) + return fmt.Sprintf("%06d", n.Int64()+verificationCodeOffset) } func getEmailCodeKey(scene, email string) string { @@ -124,11 +132,11 @@ func sendEmailVerificationCode(ctx context.Context, email, scene, templateName s } // 存验证码,5分钟有效 - if err := db.SetJSON(ctx, codeKey, code, 5*time.Minute); err != nil { + if err := db.SetJSON(ctx, codeKey, code, emailCodeExpiry); err != nil { return errors.New(errGenerateEmailCodeFailed) } // 存冷却,60秒有效 - _ = db.SetJSON(ctx, cooldownKey, "1", 60*time.Second) + _ = db.SetJSON(ctx, cooldownKey, "1", emailCodeCooldown) // 构建异步邮件发送任务 payload := SendEmailPayload{ @@ -192,6 +200,40 @@ func setLoginSession(c *gin.Context, user *model.User) error { return nil } +// handleLoginEmailVerification 处理登录时的邮箱验证码校验流程 +func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { + if user.Email == "" { + c.JSON(http.StatusOK, util.Err(errLoginEmailMissing)) + return errors.New("handled") + } + + if req.Code == "" { + // 校验 Redis 发送冷却时间 + cooldownKey := getEmailCooldownKey(user.Email) + var temp string + err := db.GetJSON(ctx, cooldownKey, &temp) + if err != nil { + // 没有冷却,触发验证码发送 + if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return errors.New("handled") + } + } + + // 脱敏邮箱并返回错误,提示前端需要输入验证码 + maskedEmail := util.MaskEmail(user.Email) + c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail)) + return errors.New("handled") + } + + // 校验验证码 + if !verifyEmailCode(ctx, user.Email, "login", req.Code) { + c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired)) + return errors.New("handled") + } + return nil +} + // Login 用户密码登录 // @Summary 用户密码登录 // @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 @@ -239,33 +281,7 @@ func Login(c *gin.Context) { } if isEmailLoginVerificationEnabled() { - if user.Email == "" { - c.JSON(http.StatusOK, util.Err(errLoginEmailMissing)) - return - } - - if req.Code == "" { - // 校验 Redis 发送冷却时间 - cooldownKey := getEmailCooldownKey(user.Email) - var temp string - err := db.GetJSON(ctx, cooldownKey, &temp) - if err != nil { - // 没有冷却,触发验证码发送 - if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - } - - // 脱敏邮箱并返回错误,提示前端需要输入验证码 - maskedEmail := util.MaskEmail(user.Email) - c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail)) - return - } - - // 校验验证码 - if !verifyEmailCode(ctx, user.Email, "login", req.Code) { - c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired)) + if emailErr := handleLoginEmailVerification(ctx, c, &req, &user); emailErr != nil { return } } @@ -297,41 +313,7 @@ func Login(c *gin.Context) { } // 检查是否有未完成 of OAuth/OIDC 绑定 - pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey) - pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey) - pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey) - pendingEmail := session.Get(oauth.PendingOAuthEmailKey) - - if pendingSourceID != nil && pendingExternalID != nil { - var sourceID uint64 - switch v := pendingSourceID.(type) { - case uint64: - sourceID = v - case int: - sourceID = uint64(v) - case float64: - sourceID = uint64(v) - } - externalID, _ := pendingExternalID.(string) - externalUsername, _ := pendingExternalUsername.(string) - email, _ := pendingEmail.(string) - - if sourceID != 0 && externalID != "" { - _ = model.BindExternalAccount(&model.ExternalAccount{ - AuthSourceID: sourceID, - UserID: user.ID, - ExternalID: externalID, - ExternalUsername: externalUsername, - Email: email, - }) - } - // 清除 pending 信息 - session.Delete(oauth.PendingOAuthSourceIDKey) - session.Delete(oauth.PendingOAuthExternalIDKey) - session.Delete(oauth.PendingOAuthExternalUsernameKey) - session.Delete(oauth.PendingOAuthEmailKey) - _ = session.Save() - } + completePendingOAuthBinding(session, &user) c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&user, needChangePassword))) } @@ -370,7 +352,7 @@ func Register(c *gin.Context) { c.JSON(http.StatusOK, util.Err(errInvalidParams)) return } - if len(req.Password) < 8 { + if len(req.Password) < minPasswordLength { c.JSON(http.StatusOK, util.Err(errPasswordTooShort)) return } @@ -378,23 +360,16 @@ func Register(c *gin.Context) { ctx := c.Request.Context() // 邮箱注册验证校验 - if isEmailRegisterVerificationEnabled() { - if req.Email == "" || req.Code == "" { - c.JSON(http.StatusOK, util.Err(errEmailOrCodeRequired)) - return - } - - if !verifyEmailCode(ctx, req.Email, "register", req.Code) { - c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired)) - return - } + if err := validateRegisterEmailVerification(ctx, &req); err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return } user := model.User{ Username: req.Username, Nickname: req.Nickname, Email: req.Email, - AvatarUrl: "", + AvatarURL: "", IsActive: true, IsAdmin: false, LastLoginAt: time.Now(), @@ -473,7 +448,7 @@ func ChangePassword(c *gin.Context) { c.JSON(http.StatusOK, util.Err(errInvalidParams)) return } - if len(req.NewPassword) < 8 { + if len(req.NewPassword) < minPasswordLength { c.JSON(http.StatusOK, util.Err(errNewPasswordTooShort)) return } @@ -578,7 +553,7 @@ func SendEmailCode(c *gin.Context) { type updateProfileRequest struct { Nickname string `json:"nickname"` Email string `json:"email"` - AvatarUrl string `json:"avatar_url"` + AvatarURL string `json:"avatar_url"` Bio string `json:"bio"` Phone string `json:"phone"` Gender string `json:"gender"` @@ -642,7 +617,7 @@ func UpdateProfile(c *gin.Context) { dbUser.Nickname = dbUser.Username } dbUser.Email = req.Email - dbUser.AvatarUrl = req.AvatarUrl + dbUser.AvatarURL = req.AvatarURL dbUser.Bio = req.Bio dbUser.Phone = strings.TrimSpace(req.Phone) dbUser.Gender = strings.TrimSpace(req.Gender) @@ -659,3 +634,58 @@ func UpdateProfile(c *gin.Context) { c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&dbUser, needChange))) } + +// validateRegisterEmailVerification 校验注册时的邮箱验证码 +func validateRegisterEmailVerification(ctx context.Context, req *registerRequest) error { + if !isEmailRegisterVerificationEnabled() { + return nil + } + if req.Email == "" || req.Code == "" { + return errors.New(errEmailOrCodeRequired) + } + if !verifyEmailCode(ctx, req.Email, "register", req.Code) { + return errors.New(errEmailCodeInvalidOrExpired) + } + return nil +} + +// completePendingOAuthBinding 完成登录后的 OAuth 待绑定绑定流程 +func completePendingOAuthBinding(session sessions.Session, user *model.User) { + pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey) + pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey) + pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey) + pendingEmail := session.Get(oauth.PendingOAuthEmailKey) + + if pendingSourceID == nil || pendingExternalID == nil { + return + } + + var sourceID uint64 + switch v := pendingSourceID.(type) { + case uint64: + sourceID = v + case int: + sourceID = uint64(v) + case float64: + sourceID = uint64(v) + } + externalID, _ := pendingExternalID.(string) + externalUsername, _ := pendingExternalUsername.(string) + email, _ := pendingEmail.(string) + + if sourceID != 0 && externalID != "" { + _ = model.BindExternalAccount(&model.ExternalAccount{ + AuthSourceID: sourceID, + UserID: user.ID, + ExternalID: externalID, + ExternalUsername: externalUsername, + Email: email, + }) + } + // 清除 pending 信息 + session.Delete(oauth.PendingOAuthSourceIDKey) + session.Delete(oauth.PendingOAuthExternalIDKey) + session.Delete(oauth.PendingOAuthExternalUsernameKey) + session.Delete(oauth.PendingOAuthEmailKey) + _ = session.Save() +} diff --git a/internal/config/config.go b/internal/config/config.go index 53e6a8b2..28f3e85b 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -15,6 +15,7 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package config 负责应用配置的加载、解析与环境变量覆盖。 package config import ( @@ -28,6 +29,14 @@ import ( "github.com/spf13/viper" ) +// 默认队列优先级 +const ( + webhookQueuePriority = 10 + whitelistQueuePriority = 5 + defaultQueuePriority = 3 +) + +// Config 全局配置单例,初始化后不可变 var Config *configModel // findConfigPath searches upward for the config file to handle tests running in subdirectories. @@ -37,7 +46,7 @@ func findConfigPath(configPath string) string { } dir := "." for i := 0; i < 5; i++ { - dir = dir + "/.." + dir += "/.." path := dir + "/" + configPath if _, err := os.Stat(path); err == nil { return path @@ -177,7 +186,7 @@ func applyEnvOverrides(c *configModel) { c.App.SessionSecret = envStr("APP_SESSION_SECRET", c.App.SessionSecret) c.App.SessionDomain = envStr("APP_SESSION_DOMAIN", c.App.SessionDomain) c.App.SessionAge = envInt("APP_SESSION_AGE", c.App.SessionAge) - c.App.SessionHttpOnly = envBool("APP_SESSION_HTTP_ONLY", c.App.SessionHttpOnly) + c.App.SessionHTTPOnly = envBool("APP_SESSION_HTTP_ONLY", c.App.SessionHTTPOnly) c.App.SessionSecure = envBool("APP_SESSION_SECURE", c.App.SessionSecure) // ─── Database ─── @@ -245,9 +254,9 @@ func applyEnvOverrides(c *configModel) { // 无 yaml 且无环境变量时,使用代码级默认队列 if len(c.Worker.Queues) == 0 { c.Worker.Queues = []QueueConfig{ - {Name: "webhook", Priority: 10}, - {Name: "whitelist_only", Priority: 5}, - {Name: "default", Priority: 3}, + {Name: "webhook", Priority: webhookQueuePriority}, + {Name: "whitelist_only", Priority: whitelistQueuePriority}, + {Name: "default", Priority: defaultQueuePriority}, } } diff --git a/internal/db/clickhouse.go b/internal/db/clickhouse.go index 25a8bf85..ec4a1ec7 100644 --- a/internal/db/clickhouse.go +++ b/internal/db/clickhouse.go @@ -15,6 +15,7 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package db 提供数据库连接与基础设施 package db import ( @@ -27,7 +28,13 @@ import ( "github.com/Rain-kl/Wavelet/internal/config" ) +const ( + clickhouseMaxExecTime = 60 // ClickHouse 最大执行时间(秒) + clickhouseReadTimeoutFactor = 2 // ReadTimeout 为 DialTimeout 的倍数 +) + var ( + // ChConn ClickHouse 连接实例 ChConn driver.Conn ) @@ -48,7 +55,7 @@ func init() { Password: cfg.Password, }, Settings: clickhouse.Settings{ - "max_execution_time": 60, + "max_execution_time": clickhouseMaxExecTime, }, Compression: &clickhouse.Compression{ Method: clickhouse.CompressionLZ4, @@ -57,7 +64,7 @@ func init() { MaxOpenConns: cfg.MaxOpenConn, MaxIdleConns: cfg.MaxIdleConn, ConnMaxLifetime: time.Duration(cfg.ConnMaxLifetime) * time.Second, - ReadTimeout: time.Duration(cfg.DialTimeout*2) * time.Second, + ReadTimeout: time.Duration(cfg.DialTimeout*clickhouseReadTimeoutFactor) * time.Second, BlockBufferSize: cfg.BlockBufferSize, }) diff --git a/internal/db/postgres_logger.go b/internal/db/postgres_logger.go index 3fbb8df3..ca880ada 100644 --- a/internal/db/postgres_logger.go +++ b/internal/db/postgres_logger.go @@ -29,6 +29,9 @@ import ( gormLogger "gorm.io/gorm/logger" ) +// nanoToMilli 纳秒转毫秒的除数 +const nanoToMilli = 1e6 + type gormZapLogger struct { logLevel gormLogger.LogLevel ignoreRecordNotFoundError bool @@ -65,24 +68,24 @@ func (l *gormZapLogger) Trace(ctx context.Context, begin time.Time, fc func() (s case err != nil && l.logLevel >= gormLogger.Error && (!errors.Is(err, gorm.ErrRecordNotFound) || !l.ignoreRecordNotFoundError): sql, rows := fc() if rows == -1 { - logger.ErrorF(ctx, "%s\n[%.3fms] [rows:%v] %s", err, float64(elapsed.Nanoseconds())/1e6, "-", sql) + logger.ErrorF(ctx, "%s\n[%.3fms] [rows:%v] %s", err, float64(elapsed.Nanoseconds())/nanoToMilli, "-", sql) } else { - logger.ErrorF(ctx, "%s\n[%.3fms] [rows:%v] %s", err, float64(elapsed.Nanoseconds())/1e6, rows, sql) + logger.ErrorF(ctx, "%s\n[%.3fms] [rows:%v] %s", err, float64(elapsed.Nanoseconds())/nanoToMilli, rows, sql) } case elapsed > l.slowThreshold && l.slowThreshold != 0 && l.logLevel >= gormLogger.Warn: sql, rows := fc() slowLog := fmt.Sprintf("SLOW SQL >= %v", l.slowThreshold) if rows == -1 { - logger.WarnF(ctx, "%s\n[%.3fms] [rows:%v] %s", slowLog, float64(elapsed.Nanoseconds())/1e6, "-", sql) + logger.WarnF(ctx, "%s\n[%.3fms] [rows:%v] %s", slowLog, float64(elapsed.Nanoseconds())/nanoToMilli, "-", sql) } else { - logger.WarnF(ctx, "%s\n[%.3fms] [rows:%v] %s", slowLog, float64(elapsed.Nanoseconds())/1e6, rows, sql) + logger.WarnF(ctx, "%s\n[%.3fms] [rows:%v] %s", slowLog, float64(elapsed.Nanoseconds())/nanoToMilli, rows, sql) } case l.logLevel == gormLogger.Info: sql, rows := fc() if rows == -1 { - logger.InfoF(ctx, "[%.3fms] [rows:%v] %s", float64(elapsed.Nanoseconds())/1e6, "-", sql) + logger.InfoF(ctx, "[%.3fms] [rows:%v] %s", float64(elapsed.Nanoseconds())/nanoToMilli, "-", sql) } else { - logger.InfoF(ctx, "[%.3fms] [rows:%v] %s", float64(elapsed.Nanoseconds())/1e6, rows, sql) + logger.InfoF(ctx, "[%.3fms] [rows:%v] %s", float64(elapsed.Nanoseconds())/nanoToMilli, rows, sql) } } } diff --git a/internal/logger/logger.go b/internal/logger/logger.go index 747eb6d8..51e4c577 100644 --- a/internal/logger/logger.go +++ b/internal/logger/logger.go @@ -29,6 +29,9 @@ import ( var logger *otelzap.Logger +// ringBufferCapacity 环形缓冲区容量 +const ringBufferCapacity = 5000 + // GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用 var GlobalRingBuffer *LogRingBuffer @@ -39,7 +42,7 @@ func init() { } // 初始化 ring buffer(保留最近 5000 行日志) - GlobalRingBuffer = NewLogRingBuffer(5000) + GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity) // 使用 multi writer 同时写入原始输出和 ring buffer multiWriter := zapcore.NewMultiWriteSyncer( @@ -60,21 +63,25 @@ func init() { fmt.Printf("[Logger] %s\n", logger.Level()) } +// DebugF 输出 Debug 级别日志 func DebugF(ctx context.Context, format string, args ...interface{}) { msg := fmt.Sprintf(format, args...) logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...) } +// InfoF 输出 Info 级别日志 func InfoF(ctx context.Context, format string, args ...interface{}) { msg := fmt.Sprintf(format, args...) logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...) } +// WarnF 输出 Warn 级别日志 func WarnF(ctx context.Context, format string, args ...interface{}) { msg := fmt.Sprintf(format, args...) logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...) } +// ErrorF 输出 Error 级别日志 func ErrorF(ctx context.Context, format string, args ...interface{}) { msg := fmt.Sprintf(format, args...) logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...) diff --git a/internal/logger/ringbuffer.go b/internal/logger/ringbuffer.go index 9fbad256..f9e9c1cf 100644 --- a/internal/logger/ringbuffer.go +++ b/internal/logger/ringbuffer.go @@ -163,10 +163,13 @@ func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) { return ordered[start:cut], hasMore } +// subscribeChanSize 订阅者 channel 缓冲区大小 +const subscribeChanSize = 64 + // Subscribe 订阅实时日志推送 // 返回一个 channel,调用者应 defer Unsubscribe func (r *LogRingBuffer) Subscribe() chan LogEntry { - ch := make(chan LogEntry, 64) + ch := make(chan LogEntry, subscribeChanSize) r.subMu.Lock() r.subscribers[ch] = struct{}{} r.subMu.Unlock() diff --git a/internal/logger/utils.go b/internal/logger/utils.go index d2b1f91d..0e6b8738 100644 --- a/internal/logger/utils.go +++ b/internal/logger/utils.go @@ -47,6 +47,9 @@ func GetLogWriter() (zapcore.WriteSyncer, error) { return logWriter, initLogWriterErr } +// logDirPerm 日志目录权限 +const logDirPerm = 0750 + func initWriter() (zapcore.WriteSyncer, error) { logConfig := config.Config.Log @@ -54,7 +57,7 @@ func initWriter() (zapcore.WriteSyncer, error) { // 初始化日志目录 logPath := logConfig.FilePath logDir := filepath.Dir(logPath) - if err := os.MkdirAll(logDir, 0750); err != nil { + if err := os.MkdirAll(logDir, logDirPerm); err != nil { return nil, fmt.Errorf(errCreateLogFileDirFailed, err) } diff --git a/internal/model/access_token.go b/internal/model/access_token.go index 1b0d6695..8a8ffc9d 100644 --- a/internal/model/access_token.go +++ b/internal/model/access_token.go @@ -15,6 +15,7 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package model 定义数据模型与 GORM 实体 package model import ( @@ -25,6 +26,12 @@ import ( "time" ) +const ( + tokenByteLength = 24 // Token 随机字节长度 + maskThreshold = 8 // 脱敏显示阈值 +) + +// AccessToken 个人访问令牌实体 type AccessToken struct { ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` UserID uint64 `json:"user_id" gorm:"index;not null"` @@ -38,7 +45,7 @@ type AccessToken struct { // GenerateTokenString 生成加密安全的随机 Token 值 func GenerateTokenString() (string, error) { - bytes := make([]byte, 24) + bytes := make([]byte, tokenByteLength) if _, err := rand.Read(bytes); err != nil { return "", err } @@ -54,7 +61,7 @@ func HashToken(token string) string { // MaskTokenString 生成脱敏显示的 Token,仅保留前缀和最后四位 func MaskTokenString(token string) string { - if len(token) <= 8 { + if len(token) <= maskThreshold { return "at_****" } return fmt.Sprintf("%s...%s", token[:7], token[len(token)-4:]) diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go index bfd72a7b..e44afc35 100644 --- a/internal/router/middlewares.go +++ b/internal/router/middlewares.go @@ -73,7 +73,7 @@ func loggerMiddleware() gin.HandlerFunc { } // 设置 Span 状态 - if c.Writer.Status() >= 400 { + if c.Writer.Status() >= http.StatusBadRequest { span := trace.SpanFromContext(ctx) span.SetStatus(codes.Error, strconv.Itoa(c.Writer.Status())) } diff --git a/internal/storage/cache.go b/internal/storage/cache.go index 45f4bc6c..8eabd01c 100644 --- a/internal/storage/cache.go +++ b/internal/storage/cache.go @@ -47,17 +47,21 @@ type metaInfo struct { ContentLength int64 `json:"content_length"` } +// cacheDirPerm 缓存目录权限 +const cacheDirPerm = 0755 + func init() { cfg := config.Config.S3.LocalCache localCacheEnabled = cfg.Enabled && cfg.CacheDir != "" localCacheDir = strings.TrimSuffix(cfg.CacheDir, "/") if localCacheEnabled { - if err := os.MkdirAll(cfg.CacheDir, 0755); err != nil { + if err := os.MkdirAll(cfg.CacheDir, cacheDirPerm); err != nil { log.Fatalf("[Storage] failed to create local cache directory: %v\n", err) } } } +// GetObjectViaCache 通过本地缓存获取对象,缓存未命中时从 S3/CDN 拉取 func GetObjectViaCache(ctx context.Context, key string) (*ObjectInfo, error) { // 没有开启本地缓存 if !localCacheEnabled { @@ -156,7 +160,7 @@ func saveToLocalCache(ctx context.Context, localPath, metaPath string, objInfo * // 创建目录 localDir := filepath.Dir(localPath) - if err := os.MkdirAll(localDir, 0755); err != nil { + if err := os.MkdirAll(localDir, cacheDirPerm); err != nil { span.SetStatus(codes.Error, err.Error()) return err } diff --git a/internal/task/constants.go b/internal/task/constants.go index 111fa098..ff6196af 100644 --- a/internal/task/constants.go +++ b/internal/task/constants.go @@ -15,13 +15,16 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package task 定义异步任务类型与调度常量 package task +// 异步任务类型标识 const ( CleanupUnusedUploadsTask = "upload:cleanup_unused" SendEmailTask = "mail:send" ) +// 任务队列名称 const ( QueueDefault = "default" ) @@ -32,7 +35,11 @@ const ( TaskTypeSendEmail = "send_email" ) +// defaultMaxRetry 任务默认最大重试次数 +const defaultMaxRetry = 3 + // TaskParam 任务参数定义 +//nolint:revive // TaskParam 保留完整名称以避免与通用 Param 混淆 type TaskParam struct { Name string `json:"Name"` // 参数键名 Label string `json:"Label"` // 显示名称 @@ -43,6 +50,7 @@ type TaskParam struct { } // TaskMeta 任务元数据 +//nolint:revive // TaskMeta 保留完整名称以避免与通用 Meta 混淆 type TaskMeta struct { Type string AsynqTask string @@ -63,7 +71,7 @@ var DispatchableTasks = []TaskMeta{ Name: "清理未使用上传", Description: "清理超过1小时未使用的上传文件", SupportsTime: false, - MaxRetry: 3, + MaxRetry: defaultMaxRetry, Queue: QueueDefault, Retryable: true, }, @@ -73,7 +81,7 @@ var DispatchableTasks = []TaskMeta{ Name: "发送邮件", Description: "异步发送系统邮件", SupportsTime: false, - MaxRetry: 3, + MaxRetry: defaultMaxRetry, Queue: QueueDefault, Retryable: true, Params: []TaskParam{ diff --git a/internal/task/scheduler/scheduler.go b/internal/task/scheduler/scheduler.go index 0a5669ba..59bf82a4 100644 --- a/internal/task/scheduler/scheduler.go +++ b/internal/task/scheduler/scheduler.go @@ -28,6 +28,11 @@ import ( "github.com/hibiken/asynq" ) +const ( + cleanupDedupWindow = 23 * time.Hour // 清理任务去重窗口 + cleanupMaxRetry = 3 // 清理任务最大重试次数 +) + var ( scheduler *asynq.Scheduler schedulerOnce sync.Once @@ -62,8 +67,8 @@ func StartScheduler() error { if _, err = scheduler.Register( config.Config.Scheduler.CleanupUnusedUploadsTaskCron, asynq.NewTask(task.CleanupUnusedUploadsTask, nil), - asynq.Unique(23*time.Hour), - asynq.MaxRetry(3), + asynq.Unique(cleanupDedupWindow), + asynq.MaxRetry(cleanupMaxRetry), ); err != nil { return } diff --git a/internal/task/worker/worker.go b/internal/task/worker/worker.go index 78053ee7..7572baa0 100644 --- a/internal/task/worker/worker.go +++ b/internal/task/worker/worker.go @@ -26,6 +26,9 @@ import ( "github.com/hibiken/asynq" ) +// workerShutdownTimeout Worker 优雅关闭超时时间 +const workerShutdownTimeout = 3 * time.Minute + func init() { // 注册所有任务处理器 taskhandlers.Register() @@ -37,7 +40,7 @@ func StartWorker() error { task.RedisOpt, asynq.Config{ Concurrency: config.Config.Worker.Concurrency, - ShutdownTimeout: 3 * time.Minute, + ShutdownTimeout: workerShutdownTimeout, Queues: buildQueuesFromConfig(), StrictPriority: config.Config.Worker.StrictPriority, }, diff --git a/internal/util/cap/cap.go b/internal/util/cap/cap.go index d2cbd362..813c9e7f 100644 --- a/internal/util/cap/cap.go +++ b/internal/util/cap/cap.go @@ -14,6 +14,7 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package cap 提供人机验证(CAPTCHA)功能 package cap import ( @@ -29,7 +30,15 @@ import ( "time" ) -const jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9" +const ( + jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9" + jwtPartsCount = 3 // JWT 三段结构 + defaultChallengeCount = 50 // 默认 PoW 难题数 + defaultChallengeSize = 32 // 默认盐值长度 + defaultDifficulty = 4 // 默认难度 + defaultNonceLength = 25 // 随机 Nonce 字节长度 + defaultExpires = 10 * time.Minute // 默认过期时间 +) // ChallengeConfig holds parameters for the PoW challenge type ChallengeConfig struct { @@ -104,7 +113,7 @@ func jwtSign(payload []byte, secret []byte) string { func jwtVerify(token string, secret []byte) ([]byte, error) { parts := strings.Split(token, ".") - if len(parts) != 3 { + if len(parts) != jwtPartsCount { return nil, errors.New(errInvalidTokenFormat) } if parts[0] != jwtHeaderB64 { @@ -135,7 +144,7 @@ func jwtVerify(token string, secret []byte) ([]byte, error) { func jwtSigHex(token string) string { parts := strings.Split(token, ".") - if len(parts) != 3 { + if len(parts) != jwtPartsCount { return "" } sigBytes, err := b64urlDecode(parts[2]) @@ -148,23 +157,23 @@ func jwtSigHex(token string) string { // GenerateChallenge produces a new challenge and signed token func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) { if conf.Count <= 0 { - conf.Count = 50 + conf.Count = defaultChallengeCount } if conf.Size <= 0 { - conf.Size = 32 + conf.Size = defaultChallengeSize } if conf.Difficulty <= 0 { - conf.Difficulty = 4 + conf.Difficulty = defaultDifficulty } if conf.Expires <= 0 { - conf.Expires = 10 * time.Minute + conf.Expires = defaultExpires } now := time.Now().UnixNano() / int64(time.Millisecond) expires := now + int64(conf.Expires/time.Millisecond) payload := ChallengePayload{ - Nonce: randomHex(25), + Nonce: randomHex(defaultNonceLength), Count: conf.Count, Size: conf.Size, Difficulty: conf.Difficulty, diff --git a/internal/util/cap/manager.go b/internal/util/cap/manager.go index 26968182..35d304cd 100644 --- a/internal/util/cap/manager.go +++ b/internal/util/cap/manager.go @@ -30,6 +30,18 @@ import ( "github.com/Rain-kl/Wavelet/internal/model" ) +const ( + managerDefaultChallengeCount = 1 + managerDefaultChallengeSize = 32 + defaultChallengeDifficulty = 4 + defaultChallengeTTL = 10 * time.Minute + defaultTokenTTL = 20 * time.Minute + redeemTokenIDLength = 8 // 兑换 Token ID 字节长度 + redeemVerTokenLength = 15 // 兑换验证 Token 字节长度 + tokenPartsCount = 2 // 兑换 Token 由两部分组成 + valuePartsCount = 2 // 存储值由 scope 和过期时间组成 +) + // Config holds settings for the CAPTCHA manager type Config struct { Secret []byte // HMAC signing key @@ -49,19 +61,19 @@ type Manager struct { // NewManager creates a new CAPTCHA Manager func NewManager(conf Config, store Store) *Manager { if conf.ChallengeCount <= 0 { - conf.ChallengeCount = 1 + conf.ChallengeCount = managerDefaultChallengeCount } if conf.ChallengeSize <= 0 { - conf.ChallengeSize = 32 + conf.ChallengeSize = managerDefaultChallengeSize } if conf.ChallengeDifficulty <= 0 { - conf.ChallengeDifficulty = 4 + conf.ChallengeDifficulty = defaultChallengeDifficulty } if conf.ChallengeTTL <= 0 { - conf.ChallengeTTL = 10 * time.Minute + conf.ChallengeTTL = defaultChallengeTTL } if conf.TokenTTL <= 0 { - conf.TokenTTL = 20 * time.Minute + conf.TokenTTL = defaultTokenTTL } return &Manager{ conf: conf, @@ -116,8 +128,8 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco } // Generate a redeem token formatted as "id:verToken" - id := randomHex(8) - verToken := randomHex(15) + id := randomHex(redeemTokenIDLength) + verToken := randomHex(redeemVerTokenLength) verHashBytes := sha256.Sum256([]byte(verToken)) verHashHex := hex.EncodeToString(verHashBytes[:]) @@ -147,7 +159,7 @@ func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope s return false, nil } parts := strings.Split(token, ":") - if len(parts) != 2 { + if len(parts) != tokenPartsCount { return false, nil } id := parts[0] @@ -169,7 +181,7 @@ func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope s } valParts := strings.Split(val, "|") - if len(valParts) != 2 { + if len(valParts) != valuePartsCount { return false, nil } @@ -253,11 +265,11 @@ func GetDefaultManager() *Manager { secret = []byte("default-captcha-secret-key-at-least-16-bytes") } - challengeCount := 1 - challengeSize := 32 - challengeDifficulty := 4 - challengeTTL := 10 * time.Minute - tokenTTL := 20 * time.Minute + challengeCount := managerDefaultChallengeCount + challengeSize := managerDefaultChallengeSize + challengeDifficulty := defaultChallengeDifficulty + challengeTTL := defaultChallengeTTL + tokenTTL := defaultTokenTTL var store Store if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil { diff --git a/internal/util/crypto.go b/internal/util/crypto.go index 12894d13..eb357b75 100644 --- a/internal/util/crypto.go +++ b/internal/util/crypto.go @@ -29,6 +29,9 @@ import ( "io" ) +// aesKeyLength AES-256 密钥字节长度 +const aesKeyLength = 32 + // Encrypt 使用 SignKey 加密字符串数据 // signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256) // plaintext: 要加密的明文字符串 @@ -56,7 +59,7 @@ func encryptBytes(signKey string, plaintext []byte) (string, error) { if err != nil { return "", fmt.Errorf(errInvalidSignKey, err) } - if len(key) != 32 { + if len(key) != aesKeyLength { return "", errors.New(errSignKeyLengthInvalid) } @@ -92,7 +95,7 @@ func decryptBytes(signKey string, ciphertext string) ([]byte, error) { if err != nil { return nil, fmt.Errorf(errInvalidSignKey, err) } - if len(key) != 32 { + if len(key) != aesKeyLength { return nil, errors.New(errSignKeyLengthInvalid) } diff --git a/internal/util/http_clients.go b/internal/util/http_clients.go index 355cf6de..72f0fa56 100644 --- a/internal/util/http_clients.go +++ b/internal/util/http_clients.go @@ -38,20 +38,30 @@ func IsLocalhost(urlStr string) bool { return hostname == "localhost" || hostname == "127.0.0.1" || hostname == "::1" } +// HTTP 客户端配置常量 +const ( + httpClientTimeout = 10 // HTTP 客户端超时时间(秒) + httpMaxIdleConns = 100 + httpMaxIdleConnsPerHost = 20 + httpIdleConnTimeout = 60 // 空闲连接超时(秒) +) + // 配置HTTP客户端 使用 otelhttp 自动注入 trace span var httpClient = &http.Client{ - Timeout: 10 * time.Second, + Timeout: httpClientTimeout * time.Second, Transport: otelhttp.NewTransport(&http.Transport{ - MaxIdleConns: 100, - MaxIdleConnsPerHost: 20, - IdleConnTimeout: 60 * time.Second, + MaxIdleConns: httpMaxIdleConns, + MaxIdleConnsPerHost: httpMaxIdleConnsPerHost, + IdleConnTimeout: httpIdleConnTimeout * time.Second, }), } +// SetHTTPClient 替换全局 HTTP 客户端实例 func SetHTTPClient(c *http.Client) { httpClient = c } +// Request 发送 HTTP 请求,支持自定义 Headers 和 Cookies func Request(ctx context.Context, method, url string, body io.Reader, headers, cookies map[string]string) (*http.Response, error) { req, err := http.NewRequestWithContext(ctx, method, url, body) if err != nil { diff --git a/internal/util/mail/mail.go b/internal/util/mail/mail.go index 019db862..5c4675ef 100644 --- a/internal/util/mail/mail.go +++ b/internal/util/mail/mail.go @@ -27,6 +27,12 @@ import ( "time" ) +const ( + smtpSSLPort = 465 // SMTP SSL 端口 + smtpDialTimeout = 5 * time.Second // SMTP 连接超时 + smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间 +) + // Config represents SMTP mail configuration type Config struct { Host string @@ -61,49 +67,8 @@ func SendMailHTML(cfg Config, to string, subject, body string) error { auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) // If using SSL port 465, we connection via TLS dial - if cfg.Port == 465 { - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, - ServerName: cfg.Host, - } - dialer := &net.Dialer{Timeout: 5 * time.Second} - conn, err := tls.DialWithDialer(dialer, "tcp", addr, tlsConfig) - if err != nil { - return fmt.Errorf(errDialTLSFailed, err) - } - defer func() { _ = conn.Close() }() - _ = conn.SetDeadline(time.Now().Add(10 * time.Second)) - - client, err := smtp.NewClient(conn, cfg.Host) - if err != nil { - return fmt.Errorf(errSMTPClientCreationFailed, err) - } - defer func() { _ = client.Close() }() - - if err = client.Auth(auth); err != nil { - return fmt.Errorf(errSMTPAuthFailed, err) - } - - if err = client.Mail(cfg.Username); err != nil { - return fmt.Errorf(errSMTPMailCommandFailed, err) - } - - if err = client.Rcpt(to); err != nil { - return fmt.Errorf(errSMTPRcptCommandFailed, err) - } - - w, err := client.Data() - if err != nil { - return fmt.Errorf(errSMTPDataCommandFailed, err) - } - defer func() { _ = w.Close() }() - - _, err = w.Write([]byte(message)) - if err != nil { - return fmt.Errorf(errSMTPWritingBodyFailed, err) - } - - return nil + if cfg.Port == smtpSSLPort { + return sendMailViaSSL(addr, auth, cfg, to, message) } // For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it) @@ -115,6 +80,49 @@ func SendMailHTML(cfg Config, to string, subject, body string) error { return nil } +// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件 +func sendMailViaSSL(addr string, auth smtp.Auth, cfg Config, to, message string) error { + tlsConfig := &tls.Config{ + InsecureSkipVerify: true, + ServerName: cfg.Host, + } + dialer := &net.Dialer{Timeout: smtpDialTimeout} + conn, err := tls.DialWithDialer(dialer, "tcp", addr, tlsConfig) + if err != nil { + return fmt.Errorf(errDialTLSFailed, err) + } + defer func() { _ = conn.Close() }() + _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) + + client, err := smtp.NewClient(conn, cfg.Host) + if err != nil { + return fmt.Errorf(errSMTPClientCreationFailed, err) + } + defer func() { _ = client.Close() }() + + if err = client.Auth(auth); err != nil { + return fmt.Errorf(errSMTPAuthFailed, err) + } + if err = client.Mail(cfg.Username); err != nil { + return fmt.Errorf(errSMTPMailCommandFailed, err) + } + if err = client.Rcpt(to); err != nil { + return fmt.Errorf(errSMTPRcptCommandFailed, err) + } + + w, err := client.Data() + if err != nil { + return fmt.Errorf(errSMTPDataCommandFailed, err) + } + defer func() { _ = w.Close() }() + + _, err = w.Write([]byte(message)) + if err != nil { + return fmt.Errorf(errSMTPWritingBodyFailed, err) + } + return nil +} + // SendMailWithLog sends a test email and records a detailed SMTP connection log func SendMailWithLog(cfg Config, to string, subject, body string) (string, error) { var logBuf bytes.Buffer @@ -127,8 +135,8 @@ func SendMailWithLog(cfg Config, to string, subject, body string) (string, error var conn net.Conn var err error - dialer := &net.Dialer{Timeout: 5 * time.Second} - if cfg.Port == 465 { + dialer := &net.Dialer{Timeout: smtpDialTimeout} + if cfg.Port == smtpSSLPort { tlsConfig := &tls.Config{ InsecureSkipVerify: true, ServerName: cfg.Host, @@ -145,7 +153,7 @@ func SendMailWithLog(cfg Config, to string, subject, body string) (string, error logLine("System", "Connected successfully.") // Set a 10-second session deadline for read/write operations - _ = conn.SetDeadline(time.Now().Add(10 * time.Second)) + _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) client, err := smtp.NewClient(conn, cfg.Host) if err != nil { @@ -155,7 +163,7 @@ func SendMailWithLog(cfg Config, to string, subject, body string) (string, error defer func() { _ = client.Close() }() // If not 465, support STARTTLS if available - if cfg.Port != 465 { + if cfg.Port != smtpSSLPort { if ok, _ := client.Extension("STARTTLS"); ok { logLine("C", "STARTTLS") tlsConfig := &tls.Config{ diff --git a/internal/util/strings.go b/internal/util/strings.go index 91f894a7..2516ec1e 100644 --- a/internal/util/strings.go +++ b/internal/util/strings.go @@ -19,6 +19,12 @@ package util import "strings" +// emailPartsCount 邮箱地址由 @ 分割为两部分 +const ( + emailPartsCount = 2 + emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数 +) + // DerefString 安全地解引用字符串指针,nil 返回空字符串 func DerefString(s *string) string { if s == nil { @@ -30,12 +36,12 @@ func DerefString(s *string) string { // MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com) func MaskEmail(email string) string { parts := strings.Split(email, "@") - if len(parts) != 2 { + if len(parts) != emailPartsCount { return email } local := parts[0] domain := parts[1] - if len(local) <= 2 { + if len(local) <= emailLocalMinChars { return "**@" + domain } return local[:2] + "***" + local[len(local)-1:] + "@" + domain diff --git a/internal/util/uuid.go b/internal/util/uuid.go index eda96312..10fb314e 100644 --- a/internal/util/uuid.go +++ b/internal/util/uuid.go @@ -26,9 +26,12 @@ import ( "github.com/google/uuid" ) +// uniqueIDBytes 生成唯一 ID 所需的随机字节长度 +const uniqueIDBytes = 32 + // GenerateUniqueIDSimple 生成 64 位唯一标识符 func GenerateUniqueIDSimple() string { - randomBytes := make([]byte, 32) + randomBytes := make([]byte, uniqueIDBytes) if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil { // 如果随机数生成失败,使用 UUID 作为后备 uuidBytes := []byte(uuid.NewString())