fix: purge authentication data on hard user deletion (#6168)
* fix: purge authentication data on hard user deletion * fix: fail closed when 2FA status lookup fails * fix: reject stale Telegram login callbacks * fix(twofa): prevent concurrent backup code and lockout bypasses * fix(auth): harden user deletion and Telegram verification
This commit is contained in:
+52
-21
@@ -4,9 +4,13 @@ import (
|
|||||||
"crypto/hmac"
|
"crypto/hmac"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"io"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/url"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/common"
|
"github.com/QuantumNous/new-api/common"
|
||||||
"github.com/QuantumNous/new-api/model"
|
"github.com/QuantumNous/new-api/model"
|
||||||
@@ -15,6 +19,13 @@ import (
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// The legacy Telegram widget has no nonce. Keep its signed assertion short-lived
|
||||||
|
// so captured callbacks cannot be reused indefinitely.
|
||||||
|
telegramAuthorizationMaxAge = 5 * time.Minute
|
||||||
|
telegramAuthorizationFutureSkew = 2 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
func TelegramBind(c *gin.Context) {
|
func TelegramBind(c *gin.Context) {
|
||||||
if !common.TelegramOAuthEnabled {
|
if !common.TelegramOAuthEnabled {
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
@@ -24,14 +35,15 @@ func TelegramBind(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
params := c.Request.URL.Query()
|
params := c.Request.URL.Query()
|
||||||
if !checkTelegramAuthorization(params, common.TelegramBotToken) {
|
telegramId, err := verifyTelegramAuthorization(params, common.TelegramBotToken, time.Now())
|
||||||
|
if err != nil {
|
||||||
|
common.SysLog("TelegramBind authorization failed: " + err.Error())
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
"message": "无效的请求",
|
"message": "无效的请求",
|
||||||
"success": false,
|
"success": false,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
telegramId := params["id"][0]
|
|
||||||
if model.IsTelegramIdAlreadyTaken(telegramId) {
|
if model.IsTelegramIdAlreadyTaken(telegramId) {
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
"message": "该 Telegram 账户已被绑定",
|
"message": "该 Telegram 账户已被绑定",
|
||||||
@@ -78,7 +90,9 @@ func TelegramLogin(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
params := c.Request.URL.Query()
|
params := c.Request.URL.Query()
|
||||||
if !checkTelegramAuthorization(params, common.TelegramBotToken) {
|
telegramId, err := verifyTelegramAuthorization(params, common.TelegramBotToken, time.Now())
|
||||||
|
if err != nil {
|
||||||
|
common.SysLog("TelegramLogin authorization failed: " + err.Error())
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
"message": "无效的请求",
|
"message": "无效的请求",
|
||||||
"success": false,
|
"success": false,
|
||||||
@@ -86,7 +100,6 @@ func TelegramLogin(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
telegramId := params["id"][0]
|
|
||||||
user := model.User{TelegramId: telegramId}
|
user := model.User{TelegramId: telegramId}
|
||||||
if err := user.FillUserByTelegramId(); err != nil {
|
if err := user.FillUserByTelegramId(); err != nil {
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
@@ -98,28 +111,46 @@ func TelegramLogin(c *gin.Context) {
|
|||||||
setupLogin(&user, c)
|
setupLogin(&user, c)
|
||||||
}
|
}
|
||||||
|
|
||||||
func checkTelegramAuthorization(params map[string][]string, token string) bool {
|
func verifyTelegramAuthorization(params url.Values, token string, now time.Time) (string, error) {
|
||||||
strs := []string{}
|
if token == "" {
|
||||||
var hash = ""
|
return "", errors.New("telegram bot token is empty")
|
||||||
|
}
|
||||||
|
for _, values := range params {
|
||||||
|
if len(values) != 1 {
|
||||||
|
return "", errors.New("telegram authorization contains duplicate parameters")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
telegramID := params.Get("id")
|
||||||
|
hash := params.Get("hash")
|
||||||
|
authDateText := params.Get("auth_date")
|
||||||
|
if telegramID == "" || hash == "" || authDateText == "" {
|
||||||
|
return "", errors.New("telegram authorization is incomplete")
|
||||||
|
}
|
||||||
|
authDate, err := strconv.ParseInt(authDateText, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return "", errors.New("telegram authorization date is invalid")
|
||||||
|
}
|
||||||
|
if authDate < now.Add(-telegramAuthorizationMaxAge).Unix() ||
|
||||||
|
authDate > now.Add(telegramAuthorizationFutureSkew).Unix() {
|
||||||
|
return "", errors.New("telegram authorization has expired")
|
||||||
|
}
|
||||||
|
|
||||||
|
strs := make([]string, 0, len(params)-1)
|
||||||
for k, v := range params {
|
for k, v := range params {
|
||||||
if k == "hash" {
|
if k == "hash" {
|
||||||
hash = v[0]
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
strs = append(strs, k+"="+v[0])
|
strs = append(strs, k+"="+v[0])
|
||||||
}
|
}
|
||||||
sort.Strings(strs)
|
sort.Strings(strs)
|
||||||
var imploded = ""
|
secret := sha256.Sum256([]byte(token))
|
||||||
for _, s := range strs {
|
mac := hmac.New(sha256.New, secret[:])
|
||||||
if imploded != "" {
|
_, _ = mac.Write([]byte(strings.Join(strs, "\n")))
|
||||||
imploded += "\n"
|
providedHash, err := hex.DecodeString(hash)
|
||||||
}
|
if err != nil || !hmac.Equal(providedHash, mac.Sum(nil)) {
|
||||||
imploded += s
|
return "", errors.New("telegram authorization signature is invalid")
|
||||||
}
|
}
|
||||||
sha256hash := sha256.New()
|
|
||||||
io.WriteString(sha256hash, token)
|
return telegramID, nil
|
||||||
hmachash := hmac.New(sha256.New, sha256hash.Sum(nil))
|
|
||||||
io.WriteString(hmachash, imploded)
|
|
||||||
ss := hex.EncodeToString(hmachash.Sum(nil))
|
|
||||||
return hash == ss
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,78 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/hmac"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"net/url"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestVerifyTelegramAuthorization(t *testing.T) {
|
||||||
|
const token = "telegram-test-token"
|
||||||
|
now := time.Unix(1_700_000_000, 0)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
authDate time.Time
|
||||||
|
mutate func(url.Values)
|
||||||
|
wantID string
|
||||||
|
wantErr string
|
||||||
|
}{
|
||||||
|
{name: "valid", authDate: now, wantID: "123456"},
|
||||||
|
{name: "small future clock skew", authDate: now.Add(90 * time.Second), wantID: "123456"},
|
||||||
|
{name: "expired", authDate: now.Add(-telegramAuthorizationMaxAge - time.Second), wantErr: "expired"},
|
||||||
|
{name: "too far in future", authDate: now.Add(telegramAuthorizationFutureSkew + time.Second), wantErr: "expired"},
|
||||||
|
{name: "invalid signature", authDate: now, mutate: func(values url.Values) { values.Set("hash", "00") }, wantErr: "signature"},
|
||||||
|
{name: "duplicate parameter", authDate: now, mutate: func(values url.Values) { values["id"] = append(values["id"], "654321") }, wantErr: "duplicate"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
params := signedTelegramAuthorization(token, tt.authDate)
|
||||||
|
if tt.mutate != nil {
|
||||||
|
tt.mutate(params)
|
||||||
|
}
|
||||||
|
|
||||||
|
telegramID, err := verifyTelegramAuthorization(params, token, now)
|
||||||
|
if tt.wantErr != "" {
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.ErrorContains(t, err, tt.wantErr)
|
||||||
|
assert.Empty(t, telegramID)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, tt.wantID, telegramID)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func signedTelegramAuthorization(token string, authDate time.Time) url.Values {
|
||||||
|
params := url.Values{
|
||||||
|
"auth_date": {strconv.FormatInt(authDate.Unix(), 10)},
|
||||||
|
"first_name": {"Test"},
|
||||||
|
"id": {"123456"},
|
||||||
|
}
|
||||||
|
keys := make([]string, 0, len(params))
|
||||||
|
for key := range params {
|
||||||
|
keys = append(keys, key)
|
||||||
|
}
|
||||||
|
sort.Strings(keys)
|
||||||
|
dataCheck := make([]string, 0, len(keys))
|
||||||
|
for _, key := range keys {
|
||||||
|
dataCheck = append(dataCheck, key+"="+params.Get(key))
|
||||||
|
}
|
||||||
|
secret := sha256.Sum256([]byte(token))
|
||||||
|
mac := hmac.New(sha256.New, secret[:])
|
||||||
|
_, _ = mac.Write([]byte(strings.Join(dataCheck, "\n")))
|
||||||
|
params.Set("hash", hex.EncodeToString(mac.Sum(nil)))
|
||||||
|
return params
|
||||||
|
}
|
||||||
+7
-1
@@ -73,7 +73,13 @@ func Login(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 检查是否启用2FA
|
// 检查是否启用2FA
|
||||||
if model.IsTwoFAEnabled(user.Id) {
|
twoFAEnabled, err := model.IsTwoFAEnabled(user.Id)
|
||||||
|
if err != nil {
|
||||||
|
common.SysLog(fmt.Sprintf("Login failed to load 2FA status for user %d: %v", user.Id, err))
|
||||||
|
common.ApiErrorI18n(c, i18n.MsgDatabaseError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if twoFAEnabled {
|
||||||
// 设置pending session,等待2FA验证
|
// 设置pending session,等待2FA验证
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
session.Set("pending_username", user.Username)
|
session.Set("pending_username", user.Username)
|
||||||
|
|||||||
@@ -38,6 +38,9 @@ func TestMain(m *testing.M) {
|
|||||||
&Task{},
|
&Task{},
|
||||||
&User{},
|
&User{},
|
||||||
&Token{},
|
&Token{},
|
||||||
|
&PasskeyCredential{},
|
||||||
|
&TwoFA{},
|
||||||
|
&TwoFABackupCode{},
|
||||||
&Log{},
|
&Log{},
|
||||||
&Channel{},
|
&Channel{},
|
||||||
&QuotaData{},
|
&QuotaData{},
|
||||||
@@ -62,8 +65,12 @@ func truncateTables(t *testing.T) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
DB.Exec("DELETE FROM tasks")
|
DB.Exec("DELETE FROM tasks")
|
||||||
DB.Exec("DELETE FROM users")
|
DB.Exec("DELETE FROM passkey_credentials")
|
||||||
|
DB.Exec("DELETE FROM two_fa_backup_codes")
|
||||||
|
DB.Exec("DELETE FROM two_fas")
|
||||||
DB.Exec("DELETE FROM tokens")
|
DB.Exec("DELETE FROM tokens")
|
||||||
|
DB.Exec("DELETE FROM user_oauth_bindings")
|
||||||
|
DB.Exec("DELETE FROM users")
|
||||||
DB.Exec("DELETE FROM logs")
|
DB.Exec("DELETE FROM logs")
|
||||||
DB.Exec("DELETE FROM channels")
|
DB.Exec("DELETE FROM channels")
|
||||||
DB.Exec("DELETE FROM quota_data")
|
DB.Exec("DELETE FROM quota_data")
|
||||||
@@ -72,7 +79,6 @@ func truncateTables(t *testing.T) {
|
|||||||
DB.Exec("DELETE FROM subscription_orders")
|
DB.Exec("DELETE FROM subscription_orders")
|
||||||
DB.Exec("DELETE FROM subscription_plans")
|
DB.Exec("DELETE FROM subscription_plans")
|
||||||
DB.Exec("DELETE FROM user_subscriptions")
|
DB.Exec("DELETE FROM user_subscriptions")
|
||||||
DB.Exec("DELETE FROM user_oauth_bindings")
|
|
||||||
DB.Exec("DELETE FROM perf_metrics")
|
DB.Exec("DELETE FROM perf_metrics")
|
||||||
DB.Exec("DELETE FROM system_instances")
|
DB.Exec("DELETE FROM system_instances")
|
||||||
DB.Exec("DELETE FROM system_task_locks")
|
DB.Exec("DELETE FROM system_task_locks")
|
||||||
|
|||||||
@@ -505,6 +505,13 @@ func InvalidateUserTokensCache(userId int) error {
|
|||||||
Find(&tokens).Error; err != nil {
|
Find(&tokens).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
return invalidateTokensCache(tokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
func invalidateTokensCache(tokens []Token) error {
|
||||||
|
if !common.RedisEnabled {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
var firstErr error
|
var firstErr error
|
||||||
for _, t := range tokens {
|
for _, t := range tokens {
|
||||||
if t.Key == "" {
|
if t.Key == "" {
|
||||||
|
|||||||
+55
-19
@@ -54,12 +54,12 @@ func GetTwoFAByUserId(userId int) (*TwoFA, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// IsTwoFAEnabled 检查用户是否启用了2FA
|
// IsTwoFAEnabled 检查用户是否启用了2FA
|
||||||
func IsTwoFAEnabled(userId int) bool {
|
func IsTwoFAEnabled(userId int) (bool, error) {
|
||||||
twoFA, err := GetTwoFAByUserId(userId)
|
twoFA, err := GetTwoFAByUserId(userId)
|
||||||
if err != nil || twoFA == nil {
|
if err != nil {
|
||||||
return false
|
return false, err
|
||||||
}
|
}
|
||||||
return twoFA.IsEnabled
|
return twoFA != nil && twoFA.IsEnabled, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// CreateTwoFA 创建2FA设置
|
// CreateTwoFA 创建2FA设置
|
||||||
@@ -120,15 +120,50 @@ func (t *TwoFA) ResetFailedAttempts() error {
|
|||||||
|
|
||||||
// IncrementFailedAttempts 增加失败尝试次数
|
// IncrementFailedAttempts 增加失败尝试次数
|
||||||
func (t *TwoFA) IncrementFailedAttempts() error {
|
func (t *TwoFA) IncrementFailedAttempts() error {
|
||||||
t.FailedAttempts++
|
if t.Id == 0 {
|
||||||
|
return errors.New("2FA记录ID不能为空")
|
||||||
// 检查是否需要锁定
|
|
||||||
if t.FailedAttempts >= common.MaxFailAttempts {
|
|
||||||
lockUntil := time.Now().Add(time.Duration(common.LockoutDuration) * time.Second)
|
|
||||||
t.LockedUntil = &lockUntil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return t.Update()
|
const maxUpdateRetries = 5
|
||||||
|
for range maxUpdateRetries {
|
||||||
|
var current TwoFA
|
||||||
|
if err := DB.Select("id", "failed_attempts", "locked_until").First(¤t, t.Id).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
if current.LockedUntil != nil && now.Before(*current.LockedUntil) {
|
||||||
|
t.FailedAttempts = current.FailedAttempts
|
||||||
|
t.LockedUntil = current.LockedUntil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
nextFailedAttempts := current.FailedAttempts + 1
|
||||||
|
nextLockedUntil := current.LockedUntil
|
||||||
|
if nextFailedAttempts >= common.MaxFailAttempts {
|
||||||
|
lockUntil := now.Add(time.Duration(common.LockoutDuration) * time.Second)
|
||||||
|
nextLockedUntil = &lockUntil
|
||||||
|
}
|
||||||
|
|
||||||
|
result := DB.Model(&TwoFA{}).
|
||||||
|
Where("id = ? AND failed_attempts = ? AND (locked_until IS NULL OR locked_until <= ?)", current.Id, current.FailedAttempts, now).
|
||||||
|
Updates(map[string]interface{}{
|
||||||
|
"failed_attempts": nextFailedAttempts,
|
||||||
|
"locked_until": nextLockedUntil,
|
||||||
|
})
|
||||||
|
if result.Error != nil {
|
||||||
|
return result.Error
|
||||||
|
}
|
||||||
|
if result.RowsAffected == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
t.FailedAttempts = nextFailedAttempts
|
||||||
|
t.LockedUntil = nextLockedUntil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.New("更新2FA失败次数冲突,请重试")
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsLocked 检查账户是否被锁定
|
// IsLocked 检查账户是否被锁定
|
||||||
@@ -186,16 +221,17 @@ func ValidateBackupCode(userId int, code string) (bool, error) {
|
|||||||
// 验证备用码
|
// 验证备用码
|
||||||
for _, bc := range backupCodes {
|
for _, bc := range backupCodes {
|
||||||
if common.ValidatePasswordAndHash(normalizedCode, bc.CodeHash) {
|
if common.ValidatePasswordAndHash(normalizedCode, bc.CodeHash) {
|
||||||
// 标记为已使用
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
bc.IsUsed = true
|
result := DB.Model(&TwoFABackupCode{}).
|
||||||
bc.UsedAt = &now
|
Where("id = ? AND is_used = ?", bc.Id, false).
|
||||||
|
Updates(map[string]interface{}{
|
||||||
if err := DB.Save(&bc).Error; err != nil {
|
"is_used": true,
|
||||||
return false, err
|
"used_at": now,
|
||||||
|
})
|
||||||
|
if result.Error != nil {
|
||||||
|
return false, result.Error
|
||||||
}
|
}
|
||||||
|
return result.RowsAffected == 1, nil
|
||||||
return true, nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+34
-8
@@ -423,12 +423,8 @@ func HardDeleteUserById(id int) error {
|
|||||||
if id == 0 {
|
if id == 0 {
|
||||||
return errors.New("id 为空!")
|
return errors.New("id 为空!")
|
||||||
}
|
}
|
||||||
return DB.Transaction(func(tx *gorm.DB) error {
|
user := User{Id: id}
|
||||||
if err := deleteUserOAuthBindingsByUserId(tx, id); err != nil {
|
return user.HardDelete()
|
||||||
return err
|
|
||||||
}
|
|
||||||
return tx.Unscoped().Delete(&User{}, "id = ?", id).Error
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func inviteUser(inviterId int) (err error) {
|
func inviteUser(inviterId int) (err error) {
|
||||||
@@ -754,12 +750,42 @@ func (user *User) HardDelete() error {
|
|||||||
if user.Id == 0 {
|
if user.Id == 0 {
|
||||||
return errors.New("id 为空!")
|
return errors.New("id 为空!")
|
||||||
}
|
}
|
||||||
return DB.Transaction(func(tx *gorm.DB) error {
|
var tokens []Token
|
||||||
if err := deleteUserOAuthBindingsByUserId(tx, user.Id); err != nil {
|
err := DB.Transaction(func(tx *gorm.DB) error {
|
||||||
|
if common.RedisEnabled {
|
||||||
|
if err := tx.Unscoped().Select("id", commonKeyCol).Where("user_id = ?", user.Id).Find(&tokens).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := deleteUserAuthenticationData(tx, user.Id); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return tx.Unscoped().Delete(user).Error
|
return tx.Unscoped().Delete(user).Error
|
||||||
})
|
})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := invalidateTokensCache(tokens); err != nil {
|
||||||
|
common.SysError(fmt.Sprintf("failed to invalidate token cache after hard deleting user %d: %v", user.Id, err))
|
||||||
|
}
|
||||||
|
if err := invalidateUserCache(user.Id); err != nil {
|
||||||
|
common.SysError(fmt.Sprintf("failed to invalidate user cache after hard deleting user %d: %v", user.Id, err))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func deleteUserAuthenticationData(tx *gorm.DB, userId int) error {
|
||||||
|
for _, authenticationData := range []any{
|
||||||
|
&TwoFABackupCode{},
|
||||||
|
&TwoFA{},
|
||||||
|
&PasskeyCredential{},
|
||||||
|
&Token{},
|
||||||
|
} {
|
||||||
|
if err := tx.Unscoped().Where("user_id = ?", userId).Delete(authenticationData).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return deleteUserOAuthBindingsByUserId(tx, userId)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ValidateAndFill check password & user status
|
// ValidateAndFill check password & user status
|
||||||
|
|||||||
@@ -0,0 +1,130 @@
|
|||||||
|
package model
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/go-redis/redis/v8"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHardDeleteUserPurgesAuthenticationDataWhenRedisFails(t *testing.T) {
|
||||||
|
truncateTables(t)
|
||||||
|
|
||||||
|
user := User{Username: "hard-delete-user", Password: "password"}
|
||||||
|
require.NoError(t, DB.Create(&user).Error)
|
||||||
|
require.NoError(t, DB.Create(&Token{UserId: user.Id, Key: "hard-delete-token"}).Error)
|
||||||
|
require.NoError(t, DB.Create(&TwoFA{UserId: user.Id, Secret: "secret", IsEnabled: true}).Error)
|
||||||
|
require.NoError(t, DB.Create(&TwoFABackupCode{UserId: user.Id, CodeHash: "hash"}).Error)
|
||||||
|
require.NoError(t, DB.Create(&PasskeyCredential{UserID: user.Id, CredentialID: "credential", PublicKey: "public-key"}).Error)
|
||||||
|
require.NoError(t, DB.Create(&UserOAuthBinding{UserId: user.Id, ProviderId: 1, ProviderUserId: "provider-user"}).Error)
|
||||||
|
|
||||||
|
oldRedisEnabled, oldRDB := common.RedisEnabled, common.RDB
|
||||||
|
common.RedisEnabled = true
|
||||||
|
var cacheInvalidatedAfterCommit atomic.Bool
|
||||||
|
common.RDB = redis.NewClient(&redis.Options{
|
||||||
|
Dialer: func(context.Context, string, string) (net.Conn, error) {
|
||||||
|
var count int64
|
||||||
|
if err := DB.Unscoped().Model(&User{}).Where("id = ?", user.Id).Count(&count).Error; err == nil && count == 0 {
|
||||||
|
cacheInvalidatedAfterCommit.Store(true)
|
||||||
|
}
|
||||||
|
return nil, errors.New("forced redis failure")
|
||||||
|
},
|
||||||
|
MaxRetries: -1,
|
||||||
|
})
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = common.RDB.Close()
|
||||||
|
common.RedisEnabled, common.RDB = oldRedisEnabled, oldRDB
|
||||||
|
})
|
||||||
|
|
||||||
|
require.NoError(t, HardDeleteUserById(user.Id))
|
||||||
|
assert.True(t, cacheInvalidatedAfterCommit.Load())
|
||||||
|
|
||||||
|
var count int64
|
||||||
|
require.NoError(t, DB.Unscoped().Model(&User{}).Where("id = ?", user.Id).Count(&count).Error)
|
||||||
|
assert.Zero(t, count)
|
||||||
|
for _, record := range []any{
|
||||||
|
&Token{},
|
||||||
|
&TwoFA{},
|
||||||
|
&TwoFABackupCode{},
|
||||||
|
&PasskeyCredential{},
|
||||||
|
&UserOAuthBinding{},
|
||||||
|
} {
|
||||||
|
require.NoError(t, DB.Unscoped().Model(record).Where("user_id = ?", user.Id).Count(&count).Error)
|
||||||
|
assert.Zero(t, count)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIncrementFailedAttemptsCountsConcurrentFailures(t *testing.T) {
|
||||||
|
truncateTables(t)
|
||||||
|
|
||||||
|
user := User{Username: "twofa-cas-user", Password: "password"}
|
||||||
|
require.NoError(t, DB.Create(&user).Error)
|
||||||
|
twoFA := TwoFA{UserId: user.Id, Secret: "secret", IsEnabled: true}
|
||||||
|
require.NoError(t, DB.Create(&twoFA).Error)
|
||||||
|
|
||||||
|
const attempts = 4
|
||||||
|
errs := make(chan error, attempts)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range attempts {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
errs <- (&TwoFA{Id: twoFA.Id}).IncrementFailedAttempts()
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(errs)
|
||||||
|
for err := range errs {
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var reloaded TwoFA
|
||||||
|
require.NoError(t, DB.First(&reloaded, twoFA.Id).Error)
|
||||||
|
assert.Equal(t, attempts, reloaded.FailedAttempts)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateBackupCodeCanOnlySucceedOnce(t *testing.T) {
|
||||||
|
truncateTables(t)
|
||||||
|
|
||||||
|
const code = "ABCD-1234"
|
||||||
|
require.NoError(t, CreateBackupCodes(123, []string{code}))
|
||||||
|
|
||||||
|
const attempts = 2
|
||||||
|
results := make(chan bool, attempts)
|
||||||
|
errs := make(chan error, attempts)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range attempts {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
valid, err := ValidateBackupCode(123, code)
|
||||||
|
results <- valid
|
||||||
|
errs <- err
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(results)
|
||||||
|
close(errs)
|
||||||
|
|
||||||
|
for err := range errs {
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
wins := 0
|
||||||
|
for valid := range results {
|
||||||
|
if valid {
|
||||||
|
wins++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert.Equal(t, 1, wins)
|
||||||
|
|
||||||
|
remaining, err := GetUnusedBackupCodeCount(123)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Zero(t, remaining)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user