mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 07:06:36 +08:00
refactor: extract magic numbers to named constants for mnd lint compliance
This commit is contained in:
+147
-135
@@ -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 浏览器类型识别
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
+150
-107
@@ -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 ""
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user