refactor: extract magic numbers to named constants for mnd lint compliance

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