refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps

- Merge all files inside openflare-server to the repository root directory.
- Relocate agent, relay, and flared subprojects from internal/ to internal/apps/.
- Combine docker-compose files and update build context paths to root.
- Update GitHub workflows and Dockerfiles to refer to new directories and package names.
- Rewrite Go package imports across all files.
- Resolve database renew test race condition and clean up docs.
This commit is contained in:
ryan
2026-06-19 14:23:29 +08:00
parent 19d476ed7f
commit 63cd906cfc
1064 changed files with 366 additions and 1397 deletions
+218
View File
@@ -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,
}))
}
+15
View File
@@ -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 // 密码最小长度
)
+47
View File
@@ -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"
)
+287
View File
@@ -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
}
+152
View File
@@ -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)
}
}
+439
View File
@@ -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)))
}
+689
View File
@@ -0,0 +1,689 @@
// 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()
// Enable registration for this test
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "true")
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyPasswordRegisterEnabled).Update("value", "true")
_ = repository.InvalidateAllSystemConfigCaches(context.Background())
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)
}
}
+161
View File
@@ -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
}