mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
feat(cmd): add reset-passwd command to reset user password
- Added ./wavelet reset-passwd subcommand to reset user passwords via CLI - Supported --user flag; if not specified, prompts for username interactively - Supported --password flag; if not specified, generates a secure random password - Handled access token deletion and cache invalidation - Added comprehensive unit tests
This commit is contained in:
@@ -0,0 +1,126 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/bootstrap"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db/migrator"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
usernameFlag string
|
||||||
|
passwordFlag string
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
passwdCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789!@#$%^&*"
|
||||||
|
defaultPasswordLength = 16
|
||||||
|
)
|
||||||
|
|
||||||
|
func generateRandomPassword(length int) (string, error) {
|
||||||
|
bytes := make([]byte, length)
|
||||||
|
if _, err := rand.Read(bytes); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
for i, b := range bytes {
|
||||||
|
bytes[i] = passwdCharset[int(b)%len(passwdCharset)]
|
||||||
|
}
|
||||||
|
return string(bytes), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var resetPasswdCmd = &cobra.Command{
|
||||||
|
Use: "reset-passwd",
|
||||||
|
Short: "重置指定账号密码",
|
||||||
|
PreRun: func(_ *cobra.Command, _ []string) {
|
||||||
|
migrator.Migrate()
|
||||||
|
},
|
||||||
|
Run: func(_ *cobra.Command, _ []string) {
|
||||||
|
ctx := context.Background()
|
||||||
|
runBootstrap(bootstrap.Options{})
|
||||||
|
|
||||||
|
var username string
|
||||||
|
if usernameFlag != "" {
|
||||||
|
username = usernameFlag
|
||||||
|
} else {
|
||||||
|
fmt.Print("请输入用户名: ")
|
||||||
|
reader := bufio.NewReader(os.Stdin)
|
||||||
|
input, err := reader.ReadString('\n')
|
||||||
|
if err != nil {
|
||||||
|
log.Fatalf("读取用户名失败: %v\n", err)
|
||||||
|
}
|
||||||
|
username = strings.TrimSpace(input)
|
||||||
|
if username == "" {
|
||||||
|
log.Fatal("用户名不能为空\n")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
user, err := repository.GetUserByUsername(ctx, username)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
log.Fatalf("错误: 用户 '%s' 不存在\n", username)
|
||||||
|
}
|
||||||
|
log.Fatalf("查询用户失败: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var password string
|
||||||
|
if passwordFlag != "" {
|
||||||
|
password = passwordFlag
|
||||||
|
} else {
|
||||||
|
password, err = generateRandomPassword(defaultPasswordLength)
|
||||||
|
if err != nil {
|
||||||
|
log.Fatalf("生成随机密码失败: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := user.SetEncryptedPassword(password); err != nil {
|
||||||
|
log.Fatalf("加密密码失败: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||||
|
if err := tx.Model(&user).Update("password", user.Password).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Invalidate existing tokens
|
||||||
|
var tokens []model.AccessToken
|
||||||
|
if err := tx.Where("user_id = ?", user.ID).Find(&tokens).Error; err == nil {
|
||||||
|
for _, token := range tokens {
|
||||||
|
oauth.InvalidateCachedToken(ctx, token.TokenHash)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return tx.Where("user_id = ?", user.ID).Delete(&model.AccessToken{}).Error
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatalf("重置密码失败: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
oauth.InvalidateCachedUser(ctx, user.ID)
|
||||||
|
|
||||||
|
fmt.Println("成功重置密码!")
|
||||||
|
fmt.Printf("用户名: %s\n", user.Username)
|
||||||
|
fmt.Printf("新密码: %s\n", password)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
resetPasswdCmd.Flags().StringVar(&usernameFlag, "user", "", "重置密码的目标用户名")
|
||||||
|
resetPasswdCmd.Flags().StringVar(&passwordFlag, "password", "", "新密码(若不指定,则自动生成随机密码)")
|
||||||
|
rootCmd.AddCommand(resetPasswdCmd)
|
||||||
|
}
|
||||||
@@ -0,0 +1,227 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestResetPasswdCmd_WithUserAndPassword(t *testing.T) {
|
||||||
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
// Seed test user
|
||||||
|
user := model.User{
|
||||||
|
ID: 1001,
|
||||||
|
Username: "testuser1",
|
||||||
|
Nickname: "Test User 1",
|
||||||
|
Email: "test1@example.com",
|
||||||
|
IsActive: true,
|
||||||
|
LastLoginAt: time.Now(),
|
||||||
|
}
|
||||||
|
_ = user.SetEncryptedPassword("oldpassword")
|
||||||
|
if err := dbConn.Create(&user).Error; err != nil {
|
||||||
|
t.Fatalf("failed to create test user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create access token to test invalidation/deletion
|
||||||
|
token := model.AccessToken{
|
||||||
|
ID: 1,
|
||||||
|
UserID: user.ID,
|
||||||
|
Name: "testtoken",
|
||||||
|
TokenHash: "somehash",
|
||||||
|
MaskedToken: "some...",
|
||||||
|
}
|
||||||
|
if err := dbConn.Create(&token).Error; err != nil {
|
||||||
|
t.Fatalf("failed to create access token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Override PreRun to bypass goose migrations in unit tests
|
||||||
|
oldPreRun := resetPasswdCmd.PreRun
|
||||||
|
resetPasswdCmd.PreRun = nil
|
||||||
|
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||||
|
|
||||||
|
// Execute command with args
|
||||||
|
rootCmd.SetArgs([]string{"reset-passwd", "--user", "testuser1", "--password", "newpassword123"})
|
||||||
|
|
||||||
|
// Capture output
|
||||||
|
oldStdout := os.Stdout
|
||||||
|
r, w, _ := os.Pipe()
|
||||||
|
os.Stdout = w
|
||||||
|
|
||||||
|
err := rootCmd.Execute()
|
||||||
|
w.Close()
|
||||||
|
os.Stdout = oldStdout
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("command execute failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
_, _ = io.Copy(&buf, r)
|
||||||
|
output := buf.String()
|
||||||
|
|
||||||
|
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||||
|
t.Errorf("expected output to contain success message, got: %s", output)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify password in DB
|
||||||
|
var dbUser model.User
|
||||||
|
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||||
|
t.Fatalf("failed to query user from DB: %v", err)
|
||||||
|
}
|
||||||
|
if !dbUser.CheckPassword("newpassword123") {
|
||||||
|
t.Errorf("password was not updated correctly in DB")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify token deleted
|
||||||
|
var count int64
|
||||||
|
dbConn.Model(&model.AccessToken{}).Where("user_id = ?", user.ID).Count(&count)
|
||||||
|
if count != 0 {
|
||||||
|
t.Errorf("expected access tokens to be deleted, got %d", count)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResetPasswdCmd_WithUserAndRandomPassword(t *testing.T) {
|
||||||
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
// Seed test user
|
||||||
|
user := model.User{
|
||||||
|
ID: 1002,
|
||||||
|
Username: "testuser2",
|
||||||
|
Nickname: "Test User 2",
|
||||||
|
Email: "test2@example.com",
|
||||||
|
IsActive: true,
|
||||||
|
LastLoginAt: time.Now(),
|
||||||
|
}
|
||||||
|
_ = user.SetEncryptedPassword("oldpassword")
|
||||||
|
if err := dbConn.Create(&user).Error; err != nil {
|
||||||
|
t.Fatalf("failed to create test user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Override PreRun
|
||||||
|
oldPreRun := resetPasswdCmd.PreRun
|
||||||
|
resetPasswdCmd.PreRun = nil
|
||||||
|
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||||
|
|
||||||
|
// Reset flags
|
||||||
|
usernameFlag = ""
|
||||||
|
passwordFlag = ""
|
||||||
|
|
||||||
|
// Execute command with args (no --password)
|
||||||
|
rootCmd.SetArgs([]string{"reset-passwd", "--user", "testuser2"})
|
||||||
|
|
||||||
|
// Capture output
|
||||||
|
oldStdout := os.Stdout
|
||||||
|
r, w, _ := os.Pipe()
|
||||||
|
os.Stdout = w
|
||||||
|
|
||||||
|
err := rootCmd.Execute()
|
||||||
|
w.Close()
|
||||||
|
os.Stdout = oldStdout
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("command execute failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
_, _ = io.Copy(&buf, r)
|
||||||
|
output := buf.String()
|
||||||
|
|
||||||
|
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||||
|
t.Errorf("expected output to contain success message, got: %s", output)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify password in DB (should be updated and not equal to old one)
|
||||||
|
var dbUser model.User
|
||||||
|
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||||
|
t.Fatalf("failed to query user from DB: %v", err)
|
||||||
|
}
|
||||||
|
if dbUser.CheckPassword("oldpassword") {
|
||||||
|
t.Errorf("expected password to change, but it matches the old one")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResetPasswdCmd_InteractiveMode(t *testing.T) {
|
||||||
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
// Seed test user
|
||||||
|
user := model.User{
|
||||||
|
ID: 1003,
|
||||||
|
Username: "testuser3",
|
||||||
|
Nickname: "Test User 3",
|
||||||
|
Email: "test3@example.com",
|
||||||
|
IsActive: true,
|
||||||
|
LastLoginAt: time.Now(),
|
||||||
|
}
|
||||||
|
_ = user.SetEncryptedPassword("oldpassword")
|
||||||
|
if err := dbConn.Create(&user).Error; err != nil {
|
||||||
|
t.Fatalf("failed to create test user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Override PreRun
|
||||||
|
oldPreRun := resetPasswdCmd.PreRun
|
||||||
|
resetPasswdCmd.PreRun = nil
|
||||||
|
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||||
|
|
||||||
|
// Mock stdin
|
||||||
|
inR, inW, err := os.Pipe()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer inR.Close()
|
||||||
|
defer inW.Close()
|
||||||
|
|
||||||
|
oldStdin := os.Stdin
|
||||||
|
os.Stdin = inR
|
||||||
|
defer func() { os.Stdin = oldStdin }()
|
||||||
|
|
||||||
|
// Write username to stdin
|
||||||
|
_, _ = inW.WriteString("testuser3\n")
|
||||||
|
inW.Close()
|
||||||
|
|
||||||
|
// Reset flags
|
||||||
|
usernameFlag = ""
|
||||||
|
passwordFlag = ""
|
||||||
|
rootCmd.SetArgs([]string{"reset-passwd"})
|
||||||
|
|
||||||
|
// Capture output
|
||||||
|
oldStdout := os.Stdout
|
||||||
|
r, w, _ := os.Pipe()
|
||||||
|
os.Stdout = w
|
||||||
|
|
||||||
|
err = rootCmd.Execute()
|
||||||
|
w.Close()
|
||||||
|
os.Stdout = oldStdout
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("command execute failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
_, _ = io.Copy(&buf, r)
|
||||||
|
output := buf.String()
|
||||||
|
|
||||||
|
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||||
|
t.Errorf("expected output to contain success message, got: %s", output)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify user password changed in DB
|
||||||
|
var dbUser model.User
|
||||||
|
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||||
|
t.Fatalf("failed to query user from DB: %v", err)
|
||||||
|
}
|
||||||
|
if dbUser.CheckPassword("oldpassword") {
|
||||||
|
t.Errorf("expected password to change, but it matches the old one")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user