diff --git a/internal/cmd/reset_passwd.go b/internal/cmd/reset_passwd.go new file mode 100644 index 00000000..68b4ce74 --- /dev/null +++ b/internal/cmd/reset_passwd.go @@ -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) +} diff --git a/internal/cmd/reset_passwd_test.go b/internal/cmd/reset_passwd_test.go new file mode 100644 index 00000000..2acae0e8 --- /dev/null +++ b/internal/cmd/reset_passwd_test.go @@ -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") + } +}