mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 15:06:37 +08:00
wavelet init
This commit is contained in:
@@ -0,0 +1,218 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package user 提供用户认证与帐户管理功能
|
||||
package user
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
type createTokenRequest struct {
|
||||
Name string `json:"name"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
}
|
||||
|
||||
type tokenResponse struct {
|
||||
Token string `json:"token"`
|
||||
Record model.AccessToken `json:"record"`
|
||||
}
|
||||
|
||||
// ListAccessTokens 获取当前用户的 AccessToken 列表
|
||||
// @Summary 获取当前用户的 AccessToken 列表
|
||||
// @Description 返回当前登录用户的所有 active access tokens(脱敏后)
|
||||
// @Tags user
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.AccessToken} "令牌列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Router /api/v1/user/access-tokens [get]
|
||||
// ListAccessTokens 获取当前用户的 AccessToken 列表
|
||||
func ListAccessTokens(c *gin.Context) {
|
||||
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
var tokens []model.AccessToken
|
||||
if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tokens))
|
||||
}
|
||||
|
||||
// CreateAccessToken 创建一个新的 AccessToken
|
||||
// @Summary 创建一个新的 AccessToken
|
||||
// @Description 为当前用户新建一个 API 访问令牌,仅在此接口返回一次明文令牌值,请妥善保存。可通过 is_admin 字段赋予令牌管理员权限(仅管理员用户可设置)。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.createTokenRequest true "令牌名称"
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=user.tokenResponse} "新建令牌成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或超限"
|
||||
// @Router /api/v1/user/access-tokens [post]
|
||||
func CreateAccessToken(c *gin.Context) {
|
||||
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
var req createTokenRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, errBindParamsFailed)
|
||||
return
|
||||
}
|
||||
|
||||
req.Name = strings.TrimSpace(req.Name)
|
||||
if req.Name == "" {
|
||||
response.AbortBadRequest(c, errTokenNameRequired)
|
||||
return
|
||||
}
|
||||
|
||||
// 只有管理员才能创建具有管理员权限的令牌
|
||||
if req.IsAdmin && !currUser.IsAdmin {
|
||||
response.AbortBadRequest(c, errAdminTokenRequiresAdmin)
|
||||
return
|
||||
}
|
||||
|
||||
// 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5)
|
||||
maxLimit := 5
|
||||
if val, err := repository.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
|
||||
maxLimit = val
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if int(count) >= maxLimit {
|
||||
response.AbortBadRequest(c, errAccessTokenLimitReached)
|
||||
return
|
||||
}
|
||||
|
||||
// 生成 Token
|
||||
tokenStr, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errGenerateTokenFailed)
|
||||
return
|
||||
}
|
||||
|
||||
tokenHash := model.HashToken(tokenStr)
|
||||
maskedToken := model.MaskTokenString(tokenStr)
|
||||
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: currUser.ID,
|
||||
Name: req.Name,
|
||||
TokenHash: tokenHash,
|
||||
MaskedToken: maskedToken,
|
||||
IsAdmin: req.IsAdmin,
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tokenResponse{
|
||||
Token: tokenStr,
|
||||
Record: tokenRecord,
|
||||
}))
|
||||
}
|
||||
|
||||
// DeleteAccessToken 删除一个 AccessToken
|
||||
// @Summary 删除一个 AccessToken
|
||||
// @Description 撤销并删除一个属于当前用户的 API 访问令牌
|
||||
// @Tags user
|
||||
// @Produce json
|
||||
// @Param id path string true "令牌ID"
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Router /api/v1/user/access-tokens/{id} [delete]
|
||||
func DeleteAccessToken(c *gin.Context) {
|
||||
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
idStr := c.Param("id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errInvalidTokenID)
|
||||
return
|
||||
}
|
||||
|
||||
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{})
|
||||
if tx.Error != nil {
|
||||
response.AbortBadRequest(c, tx.Error.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if tx.RowsAffected == 0 {
|
||||
response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK("删除成功"))
|
||||
}
|
||||
|
||||
// RotateAccessToken 轮换一个 AccessToken
|
||||
// @Summary 轮换一个 AccessToken
|
||||
// @Description 轮换(重新生成)一个属于当前用户的 API 访问令牌的密钥,旧令牌将立即失效
|
||||
// @Tags user
|
||||
// @Produce json
|
||||
// @Param id path string true "令牌ID"
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=user.tokenResponse} "令牌轮换成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Router /api/v1/user/access-tokens/{id}/rotate [post]
|
||||
func RotateAccessToken(c *gin.Context) {
|
||||
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
idStr := c.Param("id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errInvalidTokenID)
|
||||
return
|
||||
}
|
||||
|
||||
var tokenRecord model.AccessToken
|
||||
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil {
|
||||
response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
// 生成新的 Token
|
||||
newTokenStr, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errGenerateTokenFailed)
|
||||
return
|
||||
}
|
||||
|
||||
newTokenHash := model.HashToken(newTokenStr)
|
||||
newMaskedToken := model.MaskTokenString(newTokenStr)
|
||||
|
||||
tokenRecord.TokenHash = newTokenHash
|
||||
tokenRecord.MaskedToken = newMaskedToken
|
||||
|
||||
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tokenResponse{
|
||||
Token: newTokenStr,
|
||||
Record: tokenRecord,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
verificationCodeRange = 900000 // 验证码随机范围
|
||||
verificationCodeOffset = 100000 // 验证码偏移量(保证 6 位)
|
||||
emailCodeExpiry = 5 * time.Minute // 验证码有效期
|
||||
emailCodeCooldown = 60 * time.Second // 验证码发送冷却时间
|
||||
minPasswordLength = 8 // 密码最小长度
|
||||
)
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
const (
|
||||
errBindParamsFailed = "参数绑定失败"
|
||||
errInvalidParams = "无效的参数"
|
||||
errPasswordLoginDisabled = "管理员关闭了密码登录" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errUsernameOrPasswordWrong = "用户名或密码错误" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginEmailMissing = "该账号未绑定邮箱,请联系管理员绑定邮箱后再登录"
|
||||
errNeedEmailCodePrefix = "need_email_code:"
|
||||
errSMTPInvalidUseTempCodePrefix = "smtp_invalid:"
|
||||
errSMTPInvalidUseTempCode = "smtp 配置无效,使用临时码登录"
|
||||
errEmailCodeInvalidOrExpired = "验证码错误或已过期"
|
||||
errSaveSessionFailed = "无法保存会话信息,请重试"
|
||||
errRegistrationDisabled = "管理员关闭了注册"
|
||||
errPasswordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailOrCodeRequired = "邮箱或验证码未填写"
|
||||
errNewPasswordTooShort = "新密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errLoginRequired = "请先登录"
|
||||
errUserNotFound = "未找到该用户"
|
||||
errOldPasswordIncorrect = "原密码不正确" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errPasswordEncryptFailed = "密码加密失败,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errEmailRequired = "邮箱地址不能为空"
|
||||
errUnsupportedEmailScene = "不支持的验证场景"
|
||||
errEmailAlreadyRegistered = "该邮箱已被注册"
|
||||
errEmailCodeCooldown = "验证码发送频繁,请稍后再试"
|
||||
errEmailFormatInvalid = "邮箱格式不正确"
|
||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||
errRenderEmailTemplateFailed = "渲染验证邮件模板失败:%w"
|
||||
errGenerateEmailCodeFailed = "生成验证码失败,请重试"
|
||||
errDispatchEmailTaskFailed = "投递验证邮件发送任务失败,请重试"
|
||||
errTokenNameRequired = "令牌名称不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAccessTokenLimitReached = "已达到访问令牌最大创建数量限制" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errGenerateTokenFailed = "生成令牌失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errInvalidTokenID = "无效的令牌ID" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTokenNotFoundOrForbidden = "令牌不存在或无权操作" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAdminTokenRequiresAdmin = "只有管理员才能创建具有管理员权限的令牌" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTaskPayloadRequired = "任务参数不能为空"
|
||||
errInvalidJSONFormat = "无效的 JSON 格式: %w"
|
||||
errEmailTaskFieldsRequired = "to、subject、body 不能为空"
|
||||
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
||||
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
||||
errSendMailFailed = "发送邮件失败: %w"
|
||||
)
|
||||
@@ -0,0 +1,287 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// LoginEmailVerificationStatus 登录邮箱验证的处理结果。
|
||||
type LoginEmailVerificationStatus int
|
||||
|
||||
const (
|
||||
// LoginEmailVerificationPassed 验证通过,可继续登录流程。
|
||||
LoginEmailVerificationPassed LoginEmailVerificationStatus = iota
|
||||
// LoginEmailVerificationPending 需要用户输入邮箱验证码。
|
||||
LoginEmailVerificationPending
|
||||
// LoginEmailVerificationRejected 验证被拒绝(验证码错误、临时码提示等)。
|
||||
LoginEmailVerificationRejected
|
||||
)
|
||||
|
||||
// LoginEmailVerificationResult 登录邮箱验证的业务结果。
|
||||
type LoginEmailVerificationResult struct {
|
||||
Status LoginEmailVerificationStatus
|
||||
Message string
|
||||
}
|
||||
|
||||
type updateProfileInput struct {
|
||||
Nickname string
|
||||
Email string
|
||||
AvatarURL string
|
||||
Bio string
|
||||
Phone string
|
||||
Gender string
|
||||
Website string
|
||||
Location string
|
||||
}
|
||||
|
||||
func isPasswordLoginEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isPasswordRegisterEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isRegistrationEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isEmailLoginVerificationEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isSMTPConfigured(ctx context.Context) bool {
|
||||
scHost, errHost := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
|
||||
scPort, errPort := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
|
||||
scUser, errUser := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
|
||||
scPass, errPass := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
|
||||
if errHost != nil || errPort != nil || errUser != nil || errPass != nil {
|
||||
return false
|
||||
}
|
||||
return scHost.Value != "" && scPort.Value != "" && scUser.Value != "" && scPass.Value != ""
|
||||
}
|
||||
|
||||
func generateVerificationCode() (string, error) {
|
||||
n, err := rand.Int(rand.Reader, big.NewInt(verificationCodeRange))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("%06d", n.Int64()+verificationCodeOffset), nil
|
||||
}
|
||||
|
||||
func getEmailCodeKey(scene, email string) string {
|
||||
return fmt.Sprintf("email_code:%s:%s", scene, email)
|
||||
}
|
||||
|
||||
func getEmailCooldownKey(scene, email string) string {
|
||||
return fmt.Sprintf("email_code:cooldown:%s:%s", scene, email)
|
||||
}
|
||||
|
||||
func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error {
|
||||
if !isSMTPConfigured(ctx) {
|
||||
return errors.New(errSMTPConfigIncomplete)
|
||||
}
|
||||
|
||||
code, err := generateVerificationCode()
|
||||
if err != nil {
|
||||
return errors.New(errGenerateEmailCodeFailed)
|
||||
}
|
||||
codeKey := getEmailCodeKey(scene, email)
|
||||
cooldownKey := getEmailCooldownKey(scene, email)
|
||||
|
||||
tmpl, err := repository.GetTemplateByKey(ctx, templateName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("模板 %s 不存在或不可用: %w", templateName, err)
|
||||
}
|
||||
emailSubject, emailBody, err := tmpl.Render(map[string]any{"Code": code})
|
||||
if err != nil {
|
||||
return fmt.Errorf(errRenderEmailTemplateFailed, err)
|
||||
}
|
||||
|
||||
if err := db.SetJSON(ctx, codeKey, code, emailCodeExpiry); err != nil {
|
||||
return errors.New(errGenerateEmailCodeFailed)
|
||||
}
|
||||
_ = db.SetJSON(ctx, cooldownKey, "1", emailCodeCooldown)
|
||||
|
||||
payload := SendEmailPayload{
|
||||
To: email,
|
||||
Subject: emailSubject,
|
||||
Body: emailBody,
|
||||
}
|
||||
payloadBytes, _ := json.Marshal(payload)
|
||||
_, err = task.DispatchTask(ctx, TaskTypeSendEmail, payloadBytes, "system")
|
||||
if err != nil {
|
||||
return errors.New(errDispatchEmailTaskFailed)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
|
||||
codeKey := getEmailCodeKey(scene, email)
|
||||
var storedCode string
|
||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||
return false
|
||||
}
|
||||
if storedCode != code {
|
||||
return false
|
||||
}
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
|
||||
return true
|
||||
}
|
||||
|
||||
func processLoginEmailVerification(ctx context.Context, code string, user *model.User) (LoginEmailVerificationResult, error) {
|
||||
if code != "" {
|
||||
if !verifyEmailCode(ctx, user.Email, "login", code) {
|
||||
return LoginEmailVerificationResult{
|
||||
Status: LoginEmailVerificationRejected,
|
||||
Message: errEmailCodeInvalidOrExpired,
|
||||
}, nil
|
||||
}
|
||||
return LoginEmailVerificationResult{Status: LoginEmailVerificationPassed}, nil
|
||||
}
|
||||
|
||||
// 如果 SMTP 未配置,或者用户没有绑定邮箱(无法发送验证码),则使用临时码 888888
|
||||
if !isSMTPConfigured(ctx) || user.Email == "" {
|
||||
codeKey := getEmailCodeKey("login", user.Email)
|
||||
if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil {
|
||||
return LoginEmailVerificationResult{}, errors.New(errGenerateEmailCodeFailed)
|
||||
}
|
||||
var msg string
|
||||
if !isSMTPConfigured(ctx) {
|
||||
msg = errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode
|
||||
} else {
|
||||
msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
|
||||
}
|
||||
return LoginEmailVerificationResult{
|
||||
Status: LoginEmailVerificationRejected,
|
||||
Message: msg,
|
||||
}, nil
|
||||
}
|
||||
|
||||
cooldownKey := getEmailCooldownKey("login", user.Email)
|
||||
var temp string
|
||||
if err := db.GetJSON(ctx, cooldownKey, &temp); err != nil {
|
||||
if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil {
|
||||
return LoginEmailVerificationResult{}, err
|
||||
}
|
||||
}
|
||||
|
||||
maskedEmail := pkgu.MaskEmail(user.Email)
|
||||
return LoginEmailVerificationResult{
|
||||
Status: LoginEmailVerificationPending,
|
||||
Message: errNeedEmailCodePrefix + maskedEmail,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func sendRegisterEmailCode(ctx context.Context, email string) error {
|
||||
email = strings.TrimSpace(email)
|
||||
if email == "" {
|
||||
return errors.New(errEmailRequired)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errEmailAlreadyRegistered)
|
||||
}
|
||||
|
||||
cooldownKey := getEmailCooldownKey("register", email)
|
||||
var temp string
|
||||
if err := db.GetJSON(ctx, cooldownKey, &temp); err == nil {
|
||||
return errors.New(errEmailCodeCooldown)
|
||||
}
|
||||
|
||||
return sendEmailVerificationCode(ctx, email, "register", "register_email")
|
||||
}
|
||||
|
||||
func validateRegisterEmailVerification(ctx context.Context, email, code string) error {
|
||||
if !isEmailRegisterVerificationEnabled(ctx) {
|
||||
return nil
|
||||
}
|
||||
if email == "" || code == "" {
|
||||
return errors.New(errEmailOrCodeRequired)
|
||||
}
|
||||
if !verifyEmailCode(ctx, email, "register", code) {
|
||||
return errors.New(errEmailCodeInvalidOrExpired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) {
|
||||
var dbUser model.User
|
||||
if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil {
|
||||
return nil, errors.New(errUserNotFound)
|
||||
}
|
||||
|
||||
input.Email = strings.TrimSpace(input.Email)
|
||||
if input.Email != "" && input.Email != dbUser.Email {
|
||||
if !strings.Contains(input.Email, "@") || !strings.Contains(input.Email, ".") {
|
||||
return nil, errors.New(errEmailFormatInvalid)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", input.Email, dbUser.ID).Count(&count).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count > 0 {
|
||||
return nil, errors.New(errEmailAlreadyBound)
|
||||
}
|
||||
}
|
||||
|
||||
dbUser.Nickname = strings.TrimSpace(input.Nickname)
|
||||
if dbUser.Nickname == "" {
|
||||
dbUser.Nickname = dbUser.Username
|
||||
}
|
||||
dbUser.Email = input.Email
|
||||
dbUser.AvatarURL = input.AvatarURL
|
||||
dbUser.Bio = input.Bio
|
||||
dbUser.Phone = strings.TrimSpace(input.Phone)
|
||||
dbUser.Gender = strings.TrimSpace(input.Gender)
|
||||
dbUser.Website = strings.TrimSpace(input.Website)
|
||||
dbUser.Location = strings.TrimSpace(input.Location)
|
||||
|
||||
if err := db.DB(ctx).Save(&dbUser).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dbUser, nil
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
)
|
||||
|
||||
func TestProcessLoginEmailVerificationSMTPFallback(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const email = "smtpuser@example.com"
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: 222,
|
||||
Username: "smtpuser",
|
||||
Nickname: "SMTP User",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword("newpassword123"); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySMTPHost).
|
||||
Update("value", "").Error; err != nil {
|
||||
t.Fatalf("clear SMTP host failed: %v", err)
|
||||
}
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
result, err := processLoginEmailVerification(ctx, "", &user)
|
||||
if err != nil {
|
||||
t.Fatalf("processLoginEmailVerification() error = %v, want nil", err)
|
||||
}
|
||||
expected := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode
|
||||
if result.Status != LoginEmailVerificationRejected || result.Message != expected {
|
||||
t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected)
|
||||
}
|
||||
|
||||
codeKey := getEmailCodeKey("login", email)
|
||||
var storedCode string
|
||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||
t.Fatalf("get stored verification code failed: %v", err)
|
||||
}
|
||||
if storedCode != "888888" {
|
||||
t.Errorf("stored verification code = %q, want %q", storedCode, "888888")
|
||||
}
|
||||
|
||||
passed, err := processLoginEmailVerification(ctx, "888888", &user)
|
||||
if err != nil {
|
||||
t.Fatalf("processLoginEmailVerification(valid code) error = %v, want nil", err)
|
||||
}
|
||||
if passed.Status != LoginEmailVerificationPassed {
|
||||
t.Fatalf("processLoginEmailVerification(valid code) status = %v, want passed", passed.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessLoginEmailVerificationEmptyEmailFallback(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: 223,
|
||||
Username: "emptyemailuser",
|
||||
Nickname: "Empty Email User",
|
||||
Email: "",
|
||||
IsActive: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword("newpassword123"); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
for _, cfg := range []struct {
|
||||
key string
|
||||
value string
|
||||
}{
|
||||
{model.ConfigKeySMTPHost, "smtp.example.com"},
|
||||
{model.ConfigKeySMTPPort, "587"},
|
||||
{model.ConfigKeySMTPUsername, "smtpuser"},
|
||||
{model.ConfigKeySMTPPassword, "smtppassword"},
|
||||
} {
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", cfg.key).Update("value", cfg.value).Error; err != nil {
|
||||
t.Fatalf("set %s failed: %v", cfg.key, err)
|
||||
}
|
||||
}
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
result, err := processLoginEmailVerification(ctx, "", &user)
|
||||
if err != nil {
|
||||
t.Fatalf("processLoginEmailVerification() error = %v, want nil", err)
|
||||
}
|
||||
expected := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
|
||||
if result.Status != LoginEmailVerificationRejected || result.Message != expected {
|
||||
t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessLoginEmailVerificationInvalidCode(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const email = "codeduser@example.com"
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: 224,
|
||||
Username: "codeduser",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := db.SetJSON(ctx, getEmailCodeKey("login", email), "123456", emailCodeExpiry); err != nil {
|
||||
t.Fatalf("seed verification code failed: %v", err)
|
||||
}
|
||||
|
||||
result, err := processLoginEmailVerification(ctx, "000000", &user)
|
||||
if err != nil {
|
||||
t.Fatalf("processLoginEmailVerification() error = %v, want nil", err)
|
||||
}
|
||||
if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired {
|
||||
t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,439 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
type registerRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Nickname string `json:"nickname"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Email string `json:"email"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
type sendEmailCodeRequest struct {
|
||||
Email string `json:"email" binding:"required,email"`
|
||||
Scene string `json:"scene" binding:"required"`
|
||||
}
|
||||
|
||||
type updateProfileRequest struct {
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
Bio string `json:"bio"`
|
||||
Phone string `json:"phone"`
|
||||
Gender string `json:"gender"`
|
||||
Website string `json:"website"`
|
||||
Location string `json:"location"`
|
||||
}
|
||||
|
||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, user.ID)
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
session.Set(oauth.PasswordHashKey, user.Password)
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||
if err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
session.Options(oauth.GetSessionOptions(maxAge))
|
||||
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if isSessionCookie {
|
||||
oauth.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Login 用户密码登录
|
||||
// @Summary 用户密码登录
|
||||
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.loginRequest true "登录请求参数"
|
||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "登录成功,返回用户信息"
|
||||
// @Failure 400 {object} response.Any "用户名或密码错误、帐号已禁用等"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/user/login [post]
|
||||
func Login(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
if !isPasswordLoginEnabled(ctx) {
|
||||
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
||||
return
|
||||
}
|
||||
var req loginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
if req.Username == "" || req.Password == "" {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return
|
||||
}
|
||||
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
|
||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||
return
|
||||
}
|
||||
if !user.IsActive {
|
||||
logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||
response.AbortBadRequest(c, common.BannedAccount)
|
||||
return
|
||||
}
|
||||
|
||||
// 判定是否是明文密码存储
|
||||
isPlaintext := !user.IsPasswordEncrypted()
|
||||
|
||||
if !user.CheckPassword(req.Password) {
|
||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||
return
|
||||
}
|
||||
|
||||
if isEmailLoginVerificationEnabled(ctx) {
|
||||
result, err := processLoginEmailVerification(ctx, req.Code, &user)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if result.Status != LoginEmailVerificationPassed {
|
||||
response.AbortBadRequest(c, result.Message)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
needChangePassword := isPlaintext
|
||||
|
||||
if isPlaintext {
|
||||
session.Set("need_change_password", true)
|
||||
} else {
|
||||
session.Delete("need_change_password")
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||
|
||||
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, needChangePassword)))
|
||||
}
|
||||
|
||||
// Register 用户注册
|
||||
// @Summary 用户注册
|
||||
// @Description 使用用户名和密码注册新账号,注册成功后自动登录并建立 Session。密码长度不能少于 8 位。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.registerRequest true "注册请求参数"
|
||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "注册并登录成功,返回用户信息"
|
||||
// @Failure 400 {object} response.Any "参数错误、用户名已存在或注册已关闭"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/user/register [post]
|
||||
func Register(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
if !isRegistrationEnabled(ctx) || !isPasswordRegisterEnabled(ctx) {
|
||||
response.AbortBadRequest(c, errRegistrationDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
var req registerRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
req.Password = strings.TrimSpace(req.Password)
|
||||
req.Nickname = strings.TrimSpace(req.Nickname)
|
||||
req.DisplayName = strings.TrimSpace(req.DisplayName)
|
||||
req.Email = strings.TrimSpace(req.Email)
|
||||
req.Code = strings.TrimSpace(req.Code)
|
||||
|
||||
if req.Username == "" || req.Password == "" {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return
|
||||
}
|
||||
if req.Email == "" {
|
||||
response.AbortBadRequest(c, errEmailRequired)
|
||||
return
|
||||
}
|
||||
if len(req.Password) < minPasswordLength {
|
||||
response.AbortBadRequest(c, errPasswordTooShort)
|
||||
return
|
||||
}
|
||||
|
||||
// 邮箱注册验证校验
|
||||
if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
user := model.User{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: req.Username,
|
||||
Nickname: req.Nickname,
|
||||
Email: req.Email,
|
||||
AvatarURL: "",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = req.DisplayName
|
||||
}
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = req.Username
|
||||
}
|
||||
if err := user.SetEncryptedPassword(req.Password); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, false)))
|
||||
}
|
||||
|
||||
// Logout 用户退出登录
|
||||
// @Summary 用户退出登录
|
||||
// @Description 清除用户登录 Session,完成退出
|
||||
// @Tags user
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=string} "退出成功"
|
||||
// @Failure 500 {object} response.Any "Session 清除失败"
|
||||
// @Router /api/v1/user/logout [get]
|
||||
func Logout(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
userID := session.Get(oauth.UserIDKey)
|
||||
username := session.Get(oauth.UserNameKey)
|
||||
if userID != nil {
|
||||
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
|
||||
}
|
||||
session.Options(oauth.GetSessionOptions(-1))
|
||||
session.Clear()
|
||||
if err := session.Save(); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(""))
|
||||
}
|
||||
|
||||
type changePasswordRequest struct {
|
||||
OldPassword string `json:"old_password"`
|
||||
NewPassword string `json:"new_password"`
|
||||
}
|
||||
|
||||
// ChangePassword 修改用户密码
|
||||
// @Summary 修改用户密码
|
||||
// @Description 修改当前登录用户的密码。修改成功后,如果是首次明文登录的升级提示,则清除修改密码的提示状态。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.changePasswordRequest true "修改密码请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "修改密码成功"
|
||||
// @Failure 400 {object} response.Any "原密码错误或新密码不符合要求"
|
||||
// @Failure 401 {object} response.Any "请先登录"
|
||||
// @Router /api/v1/user/change-password [post]
|
||||
func ChangePassword(c *gin.Context) {
|
||||
var req changePasswordRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
req.OldPassword = strings.TrimSpace(req.OldPassword)
|
||||
req.NewPassword = strings.TrimSpace(req.NewPassword)
|
||||
|
||||
if req.OldPassword == "" || req.NewPassword == "" {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return
|
||||
}
|
||||
if len(req.NewPassword) < minPasswordLength {
|
||||
response.AbortBadRequest(c, errNewPasswordTooShort)
|
||||
return
|
||||
}
|
||||
|
||||
userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
if userObj == nil {
|
||||
response.AbortUnauthorized(c, errLoginRequired)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
var dbUser model.User
|
||||
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil {
|
||||
response.AbortBadRequest(c, errUserNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验旧密码
|
||||
if !dbUser.CheckPassword(req.OldPassword) {
|
||||
response.AbortBadRequest(c, errOldPasswordIncorrect)
|
||||
return
|
||||
}
|
||||
|
||||
// 加密并更新为新密码
|
||||
if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil {
|
||||
response.AbortBadRequest(c, errPasswordEncryptFailed)
|
||||
return
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 吊销该用户所有的 Access Token
|
||||
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil {
|
||||
response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 销毁当前活跃会话以强制重新登录
|
||||
session := sessions.Default(c)
|
||||
session.Clear()
|
||||
_ = session.Save()
|
||||
|
||||
c.JSON(http.StatusOK, response.OK("密码修改成功"))
|
||||
}
|
||||
|
||||
// SendEmailCode 发送邮箱验证码
|
||||
// @Summary 发送邮箱验证码
|
||||
// @Description 向指定邮箱发送验证码(用于注册场景)
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数"
|
||||
// @Success 200 {object} response.Any "发送成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Router /api/v1/user/send-email-code [post]
|
||||
func SendEmailCode(c *gin.Context) {
|
||||
var req sendEmailCodeRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
req.Email = strings.TrimSpace(req.Email)
|
||||
if req.Email == "" {
|
||||
response.AbortBadRequest(c, errEmailRequired)
|
||||
return
|
||||
}
|
||||
|
||||
if req.Scene != "register" {
|
||||
response.AbortBadRequest(c, errUnsupportedEmailScene)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
if err := sendRegisterEmailCode(ctx, req.Email); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// UpdateProfile 修改当前登录用户的个人资料
|
||||
// @Summary 修改当前登录用户的个人资料
|
||||
// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.updateProfileRequest true "更新请求参数"
|
||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息"
|
||||
// @Failure 400 {object} response.Any "邮箱已被占用或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Router /api/v1/user/profile [put]
|
||||
func UpdateProfile(c *gin.Context) {
|
||||
var req updateProfileRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
if userObj == nil {
|
||||
response.AbortUnauthorized(c, errLoginRequired)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
dbUser, err := updateUserProfile(ctx, userObj.ID, updateProfileInput(req))
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
needChange := session.Get("need_change_password") == true
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(dbUser, needChange)))
|
||||
}
|
||||
@@ -0,0 +1,684 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
func setupUserTestRouter(t *testing.T) *gin.Engine {
|
||||
t.Helper()
|
||||
|
||||
oldCookieName := config.Config.App.SessionCookieName
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
oldDomain := config.Config.App.SessionDomain
|
||||
oldSecure := config.Config.App.SessionSecure
|
||||
oldHTTPOnly := config.Config.App.SessionHTTPOnly
|
||||
t.Cleanup(func() {
|
||||
config.Config.App.SessionCookieName = oldCookieName
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
config.Config.App.SessionDomain = oldDomain
|
||||
config.Config.App.SessionSecure = oldSecure
|
||||
config.Config.App.SessionHTTPOnly = oldHTTPOnly
|
||||
})
|
||||
|
||||
config.Config.App.SessionCookieName = "test_session_id"
|
||||
config.Config.App.SessionSecret = "test_session_secret"
|
||||
config.Config.App.SessionDomain = ""
|
||||
config.Config.App.SessionSecure = false
|
||||
config.Config.App.SessionHTTPOnly = true
|
||||
|
||||
store := cookie.NewStore([]byte(config.Config.App.SessionSecret))
|
||||
store.Options(oauth.GetSessionOptions(3600))
|
||||
r := testhelper.NewTestGinEngine(sessions.Sessions(config.Config.App.SessionCookieName, store))
|
||||
|
||||
api := r.Group("/api/v1")
|
||||
api.POST("/user/register", Register)
|
||||
api.POST("/user/login", Login)
|
||||
api.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo)
|
||||
return r
|
||||
}
|
||||
|
||||
func performUserRequest(r http.Handler, method, path string, body []byte, cookies []*http.Cookie) *httptest.ResponseRecorder {
|
||||
var reader *bytes.Reader
|
||||
if body != nil {
|
||||
reader = bytes.NewReader(body)
|
||||
} else {
|
||||
reader = bytes.NewReader(nil)
|
||||
}
|
||||
|
||||
req, _ := http.NewRequest(method, path, reader)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
for _, c := range cookies {
|
||||
req.AddCookie(c)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
func sessionCookieFromResponse(t *testing.T, w *httptest.ResponseRecorder) *http.Cookie {
|
||||
t.Helper()
|
||||
|
||||
for _, c := range w.Result().Cookies() {
|
||||
if c.Name == config.Config.App.SessionCookieName {
|
||||
return c
|
||||
}
|
||||
}
|
||||
t.Fatalf("sessionCookieFromResponse() did not find %q cookie", config.Config.App.SessionCookieName)
|
||||
return nil
|
||||
}
|
||||
|
||||
func basicUserInfoFromResponse(t *testing.T, w *httptest.ResponseRecorder) oauth.BasicUserInfo {
|
||||
t.Helper()
|
||||
|
||||
var resp response.Any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("basicUserInfoFromResponse() decode response failed: %v", err)
|
||||
}
|
||||
if resp.ErrorMsg != "" {
|
||||
t.Fatalf("basicUserInfoFromResponse() error_msg = %q, want empty", resp.ErrorMsg)
|
||||
}
|
||||
data, _ := json.Marshal(resp.Data)
|
||||
var info oauth.BasicUserInfo
|
||||
if err := json.Unmarshal(data, &info); err != nil {
|
||||
t.Fatalf("basicUserInfoFromResponse() decode data failed: %v", err)
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
func TestEmailCooldownKeyIncludesScene(t *testing.T) {
|
||||
email := "user@example.com"
|
||||
|
||||
loginKey := getEmailCooldownKey("login", email)
|
||||
registerKey := getEmailCooldownKey("register", email)
|
||||
if loginKey == registerKey {
|
||||
t.Errorf("getEmailCooldownKey(%q, %q) = %q, want different key from register scene", "login", email, loginKey)
|
||||
}
|
||||
if want := "email_code:cooldown:login:user@example.com"; loginKey != want {
|
||||
t.Errorf("getEmailCooldownKey(%q, %q) = %q, want %q", "login", email, loginKey, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateVerificationCode(t *testing.T) {
|
||||
code, err := generateVerificationCode()
|
||||
if err != nil {
|
||||
t.Fatalf("generateVerificationCode() error = %v, want nil", err)
|
||||
}
|
||||
if len(code) != 6 {
|
||||
t.Fatalf("generateVerificationCode() length = %d, want 6. Code: %q", len(code), code)
|
||||
}
|
||||
for _, r := range code {
|
||||
if r < '0' || r > '9' {
|
||||
t.Fatalf("generateVerificationCode() = %q, want only digits", code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterCreatesAuthenticatedEncryptedUser(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
payload := registerRequest{
|
||||
Username: "newuser",
|
||||
Password: "newpassword123",
|
||||
Nickname: "New User",
|
||||
Email: "newuser@example.com",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/register", body, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
info := basicUserInfoFromResponse(t, w)
|
||||
if info.NeedChangePassword {
|
||||
t.Errorf("Register(%q) need_change_password = true, want false", payload.Username)
|
||||
}
|
||||
|
||||
var dbUser model.User
|
||||
if err := dbConn.Where("username = ?", payload.Username).First(&dbUser).Error; err != nil {
|
||||
t.Fatalf("Register(%q) db lookup failed: %v", payload.Username, err)
|
||||
}
|
||||
if dbUser.ID < 1000 {
|
||||
t.Errorf("Register(%q) user ID = %d, want generated snowflake ID", payload.Username, dbUser.ID)
|
||||
}
|
||||
if !dbUser.IsPasswordEncrypted() {
|
||||
t.Errorf("Register(%q) stored plaintext password, want encrypted password", payload.Username)
|
||||
}
|
||||
if !dbUser.CheckPassword(payload.Password) {
|
||||
t.Errorf("Register(%q) stored password does not match original password", payload.Username)
|
||||
}
|
||||
|
||||
sessionCookie := sessionCookieFromResponse(t, w)
|
||||
w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie})
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("UserInfo() after Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginRequiresPasswordChangeForInitialPlaintextAdmin(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
adminID = uint64(1)
|
||||
adminUsername = "admin"
|
||||
adminPassword = "12345678"
|
||||
)
|
||||
now := time.Now()
|
||||
if err := dbConn.Exec(
|
||||
`INSERT INTO w_users (id, username, password, nickname, is_active, is_admin, last_login_at, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
adminID,
|
||||
adminUsername,
|
||||
adminPassword,
|
||||
"Administrator",
|
||||
true,
|
||||
true,
|
||||
now,
|
||||
now,
|
||||
now,
|
||||
).Error; err != nil {
|
||||
t.Fatalf("seed initial admin failed: %v", err)
|
||||
}
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
payload := loginRequest{
|
||||
Username: adminUsername,
|
||||
Password: adminPassword,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
info := basicUserInfoFromResponse(t, w)
|
||||
if !info.NeedChangePassword {
|
||||
t.Errorf("Login(%q) need_change_password = false, want true", adminUsername)
|
||||
}
|
||||
|
||||
var dbUser model.User
|
||||
if err := dbConn.Where("username = ?", adminUsername).First(&dbUser).Error; err != nil {
|
||||
t.Fatalf("Login(%q) db lookup failed: %v", adminUsername, err)
|
||||
}
|
||||
if dbUser.ID != adminID {
|
||||
t.Errorf("Login(%q) user ID = %d, want %d", adminUsername, dbUser.ID, adminID)
|
||||
}
|
||||
if dbUser.IsPasswordEncrypted() {
|
||||
t.Errorf("Login(%q) encrypted password during login, want plaintext until password change", adminUsername)
|
||||
}
|
||||
if !dbUser.CheckPassword(adminPassword) {
|
||||
t.Errorf("Login(%q) stored password does not match original password", adminUsername)
|
||||
}
|
||||
|
||||
sessionCookie := sessionCookieFromResponse(t, w)
|
||||
w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie})
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("UserInfo() after Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
info = basicUserInfoFromResponse(t, w)
|
||||
if !info.NeedChangePassword {
|
||||
t.Errorf("UserInfo() after Login(%q) need_change_password = false, want true", adminUsername)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
userID = uint64(222)
|
||||
username = "smtpuser"
|
||||
password = "newpassword123"
|
||||
email = "smtpuser@example.com"
|
||||
)
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "SMTP User",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword(password); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
// 1. Enable email login verification
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil {
|
||||
t.Fatalf("enable email login verification failed: %v", err)
|
||||
}
|
||||
// 2. Clear SMTP host to simulate unconfigured SMTP
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "").Error; err != nil {
|
||||
t.Fatalf("clear SMTP host failed: %v", err)
|
||||
}
|
||||
|
||||
// 2.5 Invalidate the system config cache in Redis
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
|
||||
// 3. Perform login request without verification code
|
||||
payload := loginRequest{
|
||||
Username: username,
|
||||
Password: password,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String())
|
||||
}
|
||||
|
||||
// Check response error msg
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response failed: %v", err)
|
||||
}
|
||||
expectedError := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode
|
||||
if resp.ErrorMsg != expectedError {
|
||||
t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Check that verification code stored in Redis is "888888"
|
||||
ctx := context.Background()
|
||||
codeKey := getEmailCodeKey("login", email)
|
||||
var storedCode string
|
||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||
t.Fatalf("get stored verification code failed: %v", err)
|
||||
}
|
||||
if storedCode != "888888" {
|
||||
t.Errorf("expected verification code '888888', got %q", storedCode)
|
||||
}
|
||||
|
||||
// 5. Retry login with code "888888"
|
||||
payload.Code = "888888"
|
||||
bodyWithCode, _ := json.Marshal(payload)
|
||||
w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
var successResp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil {
|
||||
t.Fatalf("unmarshal success response failed: %v", err)
|
||||
}
|
||||
if successResp.ErrorMsg != "" {
|
||||
t.Errorf("expected login success, got error %q", successResp.ErrorMsg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
userID = uint64(223)
|
||||
username = "emptyemailuser"
|
||||
password = "newpassword123"
|
||||
email = ""
|
||||
)
|
||||
now := time.Now()
|
||||
user := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "Empty Email User",
|
||||
Email: email,
|
||||
IsActive: true,
|
||||
IsAdmin: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := user.SetEncryptedPassword(password); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
// 1. Enable email login verification
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil {
|
||||
t.Fatalf("enable email login verification failed: %v", err)
|
||||
}
|
||||
// 2. Make sure SMTP is configured so we only trigger empty email check
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "smtp.example.com").Error; err != nil {
|
||||
t.Fatalf("set SMTP host failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPort).Update("value", "587").Error; err != nil {
|
||||
t.Fatalf("set SMTP port failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPUsername).Update("value", "smtpuser").Error; err != nil {
|
||||
t.Fatalf("set SMTP username failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPassword).Update("value", "smtppassword").Error; err != nil {
|
||||
t.Fatalf("set SMTP password failed: %v", err)
|
||||
}
|
||||
|
||||
// Invalidate the system config cache in Redis
|
||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||
}
|
||||
|
||||
router := setupUserTestRouter(t)
|
||||
|
||||
// 3. Perform login request without verification code
|
||||
payload := loginRequest{
|
||||
Username: username,
|
||||
Password: password,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
|
||||
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String())
|
||||
}
|
||||
|
||||
// Check response error msg
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response failed: %v", err)
|
||||
}
|
||||
expectedError := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录"
|
||||
if resp.ErrorMsg != expectedError {
|
||||
t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Check that verification code stored in Redis is "888888"
|
||||
ctx := context.Background()
|
||||
codeKey := getEmailCodeKey("login", email)
|
||||
var storedCode string
|
||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||
t.Fatalf("get stored verification code failed: %v", err)
|
||||
}
|
||||
if storedCode != "888888" {
|
||||
t.Errorf("expected verification code '888888', got %q", storedCode)
|
||||
}
|
||||
|
||||
// 5. Retry login with code "888888"
|
||||
payload.Code = "888888"
|
||||
bodyWithCode, _ := json.Marshal(payload)
|
||||
w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
var successResp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil {
|
||||
t.Fatalf("unmarshal success response failed: %v", err)
|
||||
}
|
||||
if successResp.ErrorMsg != "" {
|
||||
t.Errorf("expected login success, got error %q", successResp.ErrorMsg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// 1. Seed a user
|
||||
const (
|
||||
userID = uint64(500)
|
||||
username = "tokenuser"
|
||||
password = "tokenpassword123"
|
||||
)
|
||||
now := time.Now()
|
||||
userRecord := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "Token User",
|
||||
Email: "tokenuser@example.com",
|
||||
IsActive: true,
|
||||
IsAdmin: true, // Make them an admin so we can test with is_admin=true token requests if needed
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := userRecord.SetEncryptedPassword(password); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&userRecord).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
// Seed an active AccessToken for this user
|
||||
tokenStr, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
t.Fatalf("generate token string failed: %v", err)
|
||||
}
|
||||
tokenHash := model.HashToken(tokenStr)
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: userID,
|
||||
Name: "Test Token",
|
||||
TokenHash: tokenHash,
|
||||
MaskedToken: model.MaskTokenString(tokenStr),
|
||||
IsAdmin: false,
|
||||
}
|
||||
if err := dbConn.Create(&tokenRecord).Error; err != nil {
|
||||
t.Fatalf("create test access token failed: %v", err)
|
||||
}
|
||||
|
||||
// 2. Set up router with access-token routes and oauth middlewares
|
||||
store := cookie.NewStore([]byte("test_session_secret"))
|
||||
r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store))
|
||||
|
||||
apiV1Router := r.Group("/api/v1")
|
||||
userRouter := apiV1Router.Group("/user")
|
||||
tokenRouter := userRouter.Group("/access-tokens")
|
||||
tokenRouter.Use(oauth.LoginRequired(), oauth.DisallowTokenAuth())
|
||||
{
|
||||
tokenRouter.GET("", ListAccessTokens)
|
||||
tokenRouter.POST("", CreateAccessToken)
|
||||
tokenRouter.DELETE("/:id", DeleteAccessToken)
|
||||
tokenRouter.POST("/:id/rotate", RotateAccessToken)
|
||||
}
|
||||
|
||||
// 3. Test that accessing using an Access Token fails with 403 Forbidden
|
||||
req, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
req.Header.Set("X-Access-Token", tokenStr)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Errorf("expected status 403 Forbidden when accessing with Access Token, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("decode response failed: %v", err)
|
||||
}
|
||||
if resp.ErrorMsg != oauth.ErrTokenAuthNotAllowed {
|
||||
t.Errorf("expected error message %q, got %q", oauth.ErrTokenAuthNotAllowed, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Test that accessing using a Session succeeds
|
||||
sessionCookieStore := cookie.NewStore([]byte("test_session_secret"))
|
||||
rSession := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", sessionCookieStore))
|
||||
rSession.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth(), ListAccessTokens)
|
||||
|
||||
// We can login/register or just mock the session handler to set user ID
|
||||
rSession.GET("/mock-login", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, userID)
|
||||
session.Set(oauth.UserNameKey, username)
|
||||
session.Set(oauth.PasswordHashKey, userRecord.Password)
|
||||
_ = session.Save()
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
wMock := httptest.NewRecorder()
|
||||
reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil)
|
||||
rSession.ServeHTTP(wMock, reqMock)
|
||||
cookieVal := wMock.Header().Get("Set-Cookie")
|
||||
|
||||
reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqSession.Header.Set("Cookie", cookieVal)
|
||||
wSession := httptest.NewRecorder()
|
||||
rSession.ServeHTTP(wSession, reqSession)
|
||||
|
||||
if wSession.Code != http.StatusOK {
|
||||
t.Errorf("expected status 200 OK when accessing with Session, got %d. Body: %s", wSession.Code, wSession.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestChangePasswordRevocation(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// 1. Seed a user with a password
|
||||
user := model.User{
|
||||
ID: uint64(888),
|
||||
Username: "revoketest",
|
||||
Nickname: "Revoke Test",
|
||||
IsActive: true,
|
||||
}
|
||||
if err := user.SetEncryptedPassword("oldpassword123"); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
// 2. Seed an active AccessToken for this user
|
||||
tokenStr, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
t.Fatalf("generate token string failed: %v", err)
|
||||
}
|
||||
tokenHash := model.HashToken(tokenStr)
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: user.ID,
|
||||
Name: "Test Token",
|
||||
TokenHash: tokenHash,
|
||||
MaskedToken: model.MaskTokenString(tokenStr),
|
||||
}
|
||||
if err := dbConn.Create(&tokenRecord).Error; err != nil {
|
||||
t.Fatalf("create test access token failed: %v", err)
|
||||
}
|
||||
|
||||
// 3. Set up router
|
||||
store := cookie.NewStore([]byte("test_session_secret"))
|
||||
r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store))
|
||||
|
||||
r.GET("/mock-login", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, user.ID)
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
session.Set(oauth.PasswordHashKey, user.Password)
|
||||
_ = session.Save()
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
r.GET("/mock-old-session-login", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, user.ID)
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
session.Set(oauth.PasswordHashKey, "invalid_old_password_hash")
|
||||
_ = session.Save()
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
r.POST("/api/v1/user/change-password", oauth.LoginRequired(), ChangePassword)
|
||||
r.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), ListAccessTokens)
|
||||
|
||||
// 4. Perform mock login to get cookie
|
||||
wMock := httptest.NewRecorder()
|
||||
reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil)
|
||||
r.ServeHTTP(wMock, reqMock)
|
||||
cookieVal := wMock.Header().Get("Set-Cookie")
|
||||
|
||||
// 5. Test that session and token work initially
|
||||
reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqSession.Header.Set("Cookie", cookieVal)
|
||||
wSession := httptest.NewRecorder()
|
||||
r.ServeHTTP(wSession, reqSession)
|
||||
if wSession.Code != http.StatusOK {
|
||||
t.Errorf("expected 200, got %d", wSession.Code)
|
||||
}
|
||||
|
||||
reqToken, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqToken.Header.Set("X-Access-Token", tokenStr)
|
||||
wToken := httptest.NewRecorder()
|
||||
r.ServeHTTP(wToken, reqToken)
|
||||
if wToken.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 for token, got %d", wToken.Code)
|
||||
}
|
||||
|
||||
// 6. Change password using the active session
|
||||
reqBody := `{"old_password": "oldpassword123", "new_password": "newpassword12345"}`
|
||||
reqChange, _ := http.NewRequest(http.MethodPost, "/api/v1/user/change-password", strings.NewReader(reqBody))
|
||||
reqChange.Header.Set("Content-Type", "application/json")
|
||||
reqChange.Header.Set("Cookie", cookieVal)
|
||||
wChange := httptest.NewRecorder()
|
||||
r.ServeHTTP(wChange, reqChange)
|
||||
if wChange.Code != http.StatusOK {
|
||||
t.Fatalf("expected change password to return 200, got %d. Body: %s", wChange.Code, wChange.Body.String())
|
||||
}
|
||||
|
||||
// 7. Verification: The active session that performed change-password is now cleared (401)
|
||||
reqSessionAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqSessionAfter.Header.Set("Cookie", cookieVal)
|
||||
wSessionAfter := httptest.NewRecorder()
|
||||
r.ServeHTTP(wSessionAfter, reqSessionAfter)
|
||||
if wSessionAfter.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected session to be revoked (401), got %d", wSessionAfter.Code)
|
||||
}
|
||||
|
||||
// 8. Verification: An old session (holding an outdated password hash) should be rejected (401)
|
||||
wMockOld := httptest.NewRecorder()
|
||||
reqMockOld, _ := http.NewRequest(http.MethodGet, "/mock-old-session-login", nil)
|
||||
r.ServeHTTP(wMockOld, reqMockOld)
|
||||
oldCookieVal := wMockOld.Header().Get("Set-Cookie")
|
||||
|
||||
reqOldSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqOldSession.Header.Set("Cookie", oldCookieVal)
|
||||
wOldSession := httptest.NewRecorder()
|
||||
r.ServeHTTP(wOldSession, reqOldSession)
|
||||
if wOldSession.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected old session with invalid hash to return 401, got %d", wOldSession.Code)
|
||||
}
|
||||
|
||||
// 9. Verification: The Access Token should be deleted from DB and rejected (401)
|
||||
reqTokenAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqTokenAfter.Header.Set("X-Access-Token", tokenStr)
|
||||
wTokenAfter := httptest.NewRecorder()
|
||||
r.ServeHTTP(wTokenAfter, reqTokenAfter)
|
||||
if wTokenAfter.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected access token to be revoked (401), got %d", wTokenAfter.Code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/mail"
|
||||
)
|
||||
|
||||
// 异步任务名称与管理类型定义
|
||||
const (
|
||||
// SendEmailTask 发送邮件任务标识
|
||||
SendEmailTask = "mail:send"
|
||||
// TaskTypeSendEmail 发送邮件管理类型
|
||||
TaskTypeSendEmail = "send_email"
|
||||
)
|
||||
|
||||
// SendEmailMeta represents the task metadata.
|
||||
var SendEmailMeta = task.TaskMeta{
|
||||
Type: TaskTypeSendEmail,
|
||||
AsynqTask: SendEmailTask,
|
||||
Name: "发送邮件",
|
||||
Description: "异步发送系统邮件",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
{
|
||||
Name: "to",
|
||||
Label: "接收邮箱 (To)",
|
||||
Type: "string",
|
||||
Required: true,
|
||||
Placeholder: "receiver@example.com",
|
||||
Description: "接收邮件的目标邮箱地址",
|
||||
},
|
||||
{
|
||||
Name: "subject",
|
||||
Label: "邮件主题 (Subject)",
|
||||
Type: "string",
|
||||
Required: true,
|
||||
Placeholder: "请输入邮件主题",
|
||||
Description: "发送邮件的主题标题",
|
||||
},
|
||||
{
|
||||
Name: "body",
|
||||
Label: "邮件内容 (Body)",
|
||||
Type: "text",
|
||||
Required: true,
|
||||
Placeholder: "请输入邮件内容(支持 HTML 格式)",
|
||||
Description: "发送邮件的内容主体",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// SendEmailPayload 邮件发送任务载荷
|
||||
type SendEmailPayload struct {
|
||||
To string `json:"to"`
|
||||
Subject string `json:"subject"`
|
||||
Body string `json:"body"`
|
||||
}
|
||||
|
||||
// SendEmailHandler 发送验证码邮件的异步任务处理器
|
||||
type SendEmailHandler struct{}
|
||||
|
||||
// ValidatePayload 实现 task.PayloadValidator 接口
|
||||
// 校验并标准化邮件发送参数,框架在 Admin 下发时自动调用
|
||||
func (h *SendEmailHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
if len(payload) == 0 {
|
||||
return nil, errors.New(errTaskPayloadRequired)
|
||||
}
|
||||
|
||||
var req SendEmailPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
return nil, fmt.Errorf(errInvalidJSONFormat, err)
|
||||
}
|
||||
|
||||
req.To = strings.TrimSpace(req.To)
|
||||
req.Subject = strings.TrimSpace(req.Subject)
|
||||
req.Body = strings.TrimSpace(req.Body)
|
||||
|
||||
if req.To == "" || req.Subject == "" || req.Body == "" {
|
||||
return nil, errors.New(errEmailTaskFieldsRequired)
|
||||
}
|
||||
|
||||
return json.Marshal(req)
|
||||
}
|
||||
|
||||
// Execute 执行邮件异步发送逻辑
|
||||
func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
var req SendEmailPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
task.AppendLog(ctx, "解析邮件发送参数失败: %v", err)
|
||||
return nil, fmt.Errorf(errParseEmailPayloadFailed, err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始准备发送邮件到: %s, 主题: %s", req.To, req.Subject)
|
||||
|
||||
// 从数据库读取最新的 SMTP 系统配置
|
||||
var smtpHost string
|
||||
var smtpPortVal string
|
||||
var smtpUsername string
|
||||
var smtpPassword string
|
||||
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
||||
smtpHost = sc.Value
|
||||
}
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
||||
smtpPortVal = sc.Value
|
||||
}
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
||||
smtpUsername = sc.Value
|
||||
}
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword); err == nil {
|
||||
smtpPassword = sc.Value
|
||||
}
|
||||
|
||||
if smtpHost == "" || smtpPortVal == "" || smtpUsername == "" {
|
||||
err := errors.New(errSMTPConfigIncomplete)
|
||||
task.AppendLog(ctx, "发送失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
smtpPort, err := strconv.Atoi(smtpPortVal)
|
||||
if err != nil {
|
||||
smtpPort = 587
|
||||
}
|
||||
|
||||
cfg := mail.Config{
|
||||
Host: smtpHost,
|
||||
Port: smtpPort,
|
||||
Username: smtpUsername,
|
||||
Password: smtpPassword,
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "连接 SMTP 服务器: %s:%d, 用户名: %s", smtpHost, smtpPort, smtpUsername)
|
||||
|
||||
// 调用 SendMailHTML 执行邮件发送,这里会有 5s 拨号超时和 10s 读写限制
|
||||
err = mail.SendMailHTML(ctx, cfg, req.To, req.Subject, req.Body)
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "邮件发送失败: %v", err)
|
||||
return nil, fmt.Errorf(errSendMailFailed, err)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("邮件成功发送至: %s", req.To)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
|
||||
return &task.TaskResult{
|
||||
Message: msg,
|
||||
}, nil
|
||||
}
|
||||
Reference in New Issue
Block a user