Merge pull request #1072 from trheyi/main
Implement MFA functionality in user management
This commit is contained in:
commit
1d5de18eb9
9 changed files with 1241 additions and 124 deletions
212
data/bindata.go
212
data/bindata.go
File diff suppressed because one or more lines are too long
2
go.mod
2
go.mod
|
|
@ -58,6 +58,7 @@ require (
|
|||
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.15 // indirect
|
||||
github.com/aws/smithy-go v1.22.3 // indirect
|
||||
github.com/blang/semver/v4 v4.0.0 // indirect
|
||||
github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc // indirect
|
||||
github.com/bytedance/sonic v1.13.2 // indirect
|
||||
github.com/bytedance/sonic/loader v0.2.4 // indirect
|
||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||
|
|
@ -116,6 +117,7 @@ require (
|
|||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||
github.com/pkg/errors v0.9.1 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/pquerna/otp v1.5.0 // indirect
|
||||
github.com/qdrant/go-client v1.14.0 // indirect
|
||||
github.com/richardlehane/mscfb v1.0.4 // indirect
|
||||
github.com/richardlehane/msoleps v1.0.4 // indirect
|
||||
|
|
|
|||
4
go.sum
4
go.sum
|
|
@ -34,6 +34,8 @@ github.com/blang/semver v3.5.1+incompatible h1:cQNTCjp13qL8KC3Nbxr/y2Bqb63oX6wdn
|
|||
github.com/blang/semver v3.5.1+incompatible/go.mod h1:kRBLl5iJ+tD4TcOOxsy/0fnwebNt5EWlYSAyrTnjyyk=
|
||||
github.com/blang/semver/v4 v4.0.0 h1:1PFHFE6yCCTv8C1TeyNNarDzntLi7wMI5i/pzqYIsAM=
|
||||
github.com/blang/semver/v4 v4.0.0/go.mod h1:IbckMUScFkM3pff0VJDNKRiT6TG/YpiHIM2yvyW5YoQ=
|
||||
github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc h1:biVzkmvwrH8WK8raXaxBx6fRVTlJILwEwQGL1I/ByEI=
|
||||
github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc/go.mod h1:paBWMcWSl3LHKBqUq+rly7CNSldXjb2rDl3JlRe0mD8=
|
||||
github.com/bufbuild/protocompile v0.4.0 h1:LbFKd2XowZvQ/kajzguUp2DC9UEIQhIq77fZZlaQsNA=
|
||||
github.com/bufbuild/protocompile v0.4.0/go.mod h1:3v93+mbWn/v3xzN+31nwkJfrEpAUwp+BagBSZWx+TP8=
|
||||
github.com/bytedance/sonic v1.13.2 h1:8/H1FempDZqC4VqjptGo14QQlJx8VdZJegxs6wwfqpQ=
|
||||
|
|
@ -237,6 +239,8 @@ github.com/pkoukk/tiktoken-go v0.1.7 h1:qOBHXX4PHtvIvmOtyg1EeKlwFRiMKAcoMp4Q+bLQ
|
|||
github.com/pkoukk/tiktoken-go v0.1.7/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/pquerna/otp v1.5.0 h1:NMMR+WrmaqXU4EzdGJEE1aUUI0AMRzsp96fFFWNPwxs=
|
||||
github.com/pquerna/otp v1.5.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg=
|
||||
github.com/qdrant/go-client v1.14.0 h1:cyz9OOooAexudw5w69LRe9vKCQFYJvaFvt9icOciI1U=
|
||||
github.com/qdrant/go-client v1.14.0/go.mod h1:iO8ts78jL4x6LDHFOViyYWELVtIBDTjOykBmiOTHLnQ=
|
||||
github.com/rhysd/go-github-selfupdate v1.2.3 h1:iaa+J202f+Nc+A8zi75uccC8Wg3omaM7HDeimXA22Ag=
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ package user
|
|||
|
||||
import (
|
||||
"github.com/yaoapp/gou/store"
|
||||
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||
)
|
||||
|
||||
// Error messages
|
||||
|
|
@ -32,6 +33,17 @@ const (
|
|||
ErrFailedToDeleteRole = "failed to delete role: %w"
|
||||
ErrFailedToDeleteType = "failed to delete type: %w"
|
||||
ErrFailedToDeleteOAuth = "failed to delete oauth account: %w"
|
||||
|
||||
// MFA related errors
|
||||
ErrMFANotEnabled = "MFA is not enabled for this user"
|
||||
ErrMFAAlreadyEnabled = "MFA is already enabled for this user"
|
||||
ErrInvalidMFACode = "invalid MFA code"
|
||||
ErrInvalidRecoveryCode = "invalid recovery code"
|
||||
ErrFailedToGenerateMFASecret = "failed to generate MFA secret: %w"
|
||||
ErrFailedToGenerateQRCode = "failed to generate QR code: %w"
|
||||
ErrFailedToVerifyMFACode = "failed to verify MFA code: %w"
|
||||
ErrFailedToUpdateMFAStatus = "failed to update MFA status: %w"
|
||||
ErrRecoveryCodeNotFound = "recovery code not found or already used"
|
||||
)
|
||||
|
||||
// Default field lists - used when not configured
|
||||
|
|
@ -103,6 +115,17 @@ var (
|
|||
"is_active", "is_default", "sort_order", "max_sessions", "session_timeout",
|
||||
"password_policy", "features", "limits", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// DefaultMFAOptions contains default MFA configuration
|
||||
DefaultMFAOptions = &types.MFAOptions{
|
||||
Issuer: "Yao App Engine",
|
||||
Algorithm: "SHA256",
|
||||
Digits: 6,
|
||||
Period: 30,
|
||||
SecretSize: 32,
|
||||
RecoveryCount: 16, // 16 codes (~960 bytes, under 1024 char limit)
|
||||
RecoveryLength: 12, // 12-character codes for better security
|
||||
}
|
||||
)
|
||||
|
||||
// DefaultUser provides a default implementation of UserProvider
|
||||
|
|
@ -135,6 +158,9 @@ type DefaultUser struct {
|
|||
// Type Field lists
|
||||
typeFields []interface{} // configurable
|
||||
typeDetailFields []interface{} // configurable
|
||||
|
||||
// MFA Configuration
|
||||
mfaOptions *types.MFAOptions // configurable MFA settings
|
||||
}
|
||||
|
||||
// IDStrategy defines the strategy for generating user IDs
|
||||
|
|
@ -175,6 +201,9 @@ type DefaultUserOptions struct {
|
|||
// Type field lists (use defaults if not specified)
|
||||
TypeFields []interface{} // basic type fields
|
||||
TypeDetailFields []interface{} // detailed type fields including configuration and metadata
|
||||
|
||||
// MFA configuration (use defaults if not specified)
|
||||
MFAOptions *types.MFAOptions // MFA settings
|
||||
}
|
||||
|
||||
// NewDefaultUser creates a new DefaultUser
|
||||
|
|
@ -253,6 +282,12 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
|||
typeDetailFields = DefaultTypeDetailFields
|
||||
}
|
||||
|
||||
// Set MFA options with defaults if not specified
|
||||
mfaOptions := options.MFAOptions
|
||||
if mfaOptions == nil {
|
||||
mfaOptions = DefaultMFAOptions
|
||||
}
|
||||
|
||||
return &DefaultUser{
|
||||
prefix: options.Prefix,
|
||||
model: model,
|
||||
|
|
@ -278,5 +313,8 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
|||
// Type field lists
|
||||
typeFields: typeFields,
|
||||
typeDetailFields: typeDetailFields,
|
||||
|
||||
// MFA Configuration
|
||||
mfaOptions: mfaOptions,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,56 +2,637 @@ package user
|
|||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/pquerna/otp"
|
||||
"github.com/pquerna/otp/totp"
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// User MFA Management
|
||||
|
||||
// GenerateMFASecret generates a new TOTP secret for user
|
||||
func (u *DefaultUser) GenerateMFASecret(ctx context.Context, userID string, issuer string, accountName string) (string, string, error) {
|
||||
// TODO: implement
|
||||
return "", "", nil
|
||||
func (u *DefaultUser) GenerateMFASecret(ctx context.Context, userID string, options *types.MFAOptions) (string, string, error) {
|
||||
// Verify user exists
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return "", "", fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
// Use provided options or fallback to instance defaults
|
||||
if options == nil {
|
||||
options = u.mfaOptions
|
||||
}
|
||||
|
||||
// Apply defaults for individual fields if not specified
|
||||
issuer := options.Issuer
|
||||
if issuer == "" {
|
||||
issuer = u.mfaOptions.Issuer
|
||||
}
|
||||
|
||||
accountName := options.AccountName
|
||||
if accountName == "" {
|
||||
accountName = userID // Default to userID
|
||||
}
|
||||
|
||||
algorithm := options.Algorithm
|
||||
if algorithm == "" {
|
||||
algorithm = u.mfaOptions.Algorithm
|
||||
}
|
||||
|
||||
digits := options.Digits
|
||||
if digits == 0 {
|
||||
digits = u.mfaOptions.Digits
|
||||
}
|
||||
|
||||
period := options.Period
|
||||
if period == 0 {
|
||||
period = u.mfaOptions.Period
|
||||
}
|
||||
|
||||
secretSize := options.SecretSize
|
||||
if secretSize == 0 {
|
||||
secretSize = u.mfaOptions.SecretSize
|
||||
}
|
||||
|
||||
// Convert algorithm string to otp.Algorithm
|
||||
var otpAlgorithm otp.Algorithm
|
||||
switch algorithm {
|
||||
case "SHA1":
|
||||
otpAlgorithm = otp.AlgorithmSHA1
|
||||
case "SHA256":
|
||||
otpAlgorithm = otp.AlgorithmSHA256
|
||||
case "SHA512":
|
||||
otpAlgorithm = otp.AlgorithmSHA512
|
||||
default:
|
||||
otpAlgorithm = otp.AlgorithmSHA256 // Default fallback
|
||||
}
|
||||
|
||||
// Convert digits to otp.Digits
|
||||
var otpDigits otp.Digits
|
||||
if digits == 8 {
|
||||
otpDigits = otp.DigitsEight
|
||||
} else {
|
||||
otpDigits = otp.DigitsSix // Default
|
||||
}
|
||||
|
||||
// Generate TOTP key
|
||||
key, err := totp.Generate(totp.GenerateOpts{
|
||||
Issuer: issuer,
|
||||
AccountName: accountName,
|
||||
SecretSize: uint(secretSize),
|
||||
Algorithm: otpAlgorithm,
|
||||
Digits: otpDigits,
|
||||
Period: uint(period),
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(ErrFailedToGenerateMFASecret, err)
|
||||
}
|
||||
|
||||
secret := key.Secret()
|
||||
qrCodeURL := key.URL()
|
||||
|
||||
// Store the secret temporarily (not enabled yet until user verifies)
|
||||
updateData := maps.MapStrAny{
|
||||
"mfa_secret": secret,
|
||||
"mfa_issuer": issuer,
|
||||
"mfa_algorithm": algorithm,
|
||||
"mfa_digits": digits,
|
||||
"mfa_period": period,
|
||||
}
|
||||
|
||||
_, err = m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(ErrFailedToUpdateMFAStatus, err)
|
||||
}
|
||||
|
||||
return secret, qrCodeURL, nil
|
||||
}
|
||||
|
||||
// EnableMFA enables multi-factor authentication for user
|
||||
func (u *DefaultUser) EnableMFA(ctx context.Context, userID string, secret string, code string) error {
|
||||
// TODO: implement
|
||||
// Get user and current MFA status
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled", "mfa_secret"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is already enabled
|
||||
if mfaEnabled, ok := user["mfa_enabled"].(bool); ok && mfaEnabled {
|
||||
return fmt.Errorf(ErrMFAAlreadyEnabled)
|
||||
}
|
||||
// Handle different boolean types from database
|
||||
if mfaEnabledInt, ok := user["mfa_enabled"].(int64); ok && mfaEnabledInt != 0 {
|
||||
return fmt.Errorf(ErrMFAAlreadyEnabled)
|
||||
}
|
||||
|
||||
// Use stored secret if not provided
|
||||
if secret == "" {
|
||||
if storedSecret, ok := user["mfa_secret"].(string); ok && storedSecret != "" {
|
||||
secret = storedSecret
|
||||
} else {
|
||||
return fmt.Errorf("no MFA secret found, please generate one first")
|
||||
}
|
||||
}
|
||||
|
||||
// Verify the provided code
|
||||
valid := totp.Validate(code, secret)
|
||||
if !valid {
|
||||
return fmt.Errorf(ErrInvalidMFACode)
|
||||
}
|
||||
|
||||
// Enable MFA
|
||||
updateData := maps.MapStrAny{
|
||||
"mfa_enabled": true,
|
||||
"mfa_secret": secret, // Store the verified secret
|
||||
"mfa_enabled_at": time.Now(),
|
||||
}
|
||||
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateMFAStatus, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DisableMFA disables multi-factor authentication for user
|
||||
func (u *DefaultUser) DisableMFA(ctx context.Context, userID string, code string) error {
|
||||
// TODO: implement
|
||||
// Get user and current MFA status
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled", "mfa_secret"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is enabled
|
||||
mfaEnabled := false
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
mfaEnabled = enabled
|
||||
} else if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
mfaEnabled = enabledInt != 0
|
||||
}
|
||||
|
||||
if !mfaEnabled {
|
||||
return fmt.Errorf(ErrMFANotEnabled)
|
||||
}
|
||||
|
||||
// Get stored secret
|
||||
secret, ok := user["mfa_secret"].(string)
|
||||
if !ok || secret == "" {
|
||||
return fmt.Errorf("no MFA secret found")
|
||||
}
|
||||
|
||||
// Verify the provided code
|
||||
valid := totp.Validate(code, secret)
|
||||
if !valid {
|
||||
return fmt.Errorf(ErrInvalidMFACode)
|
||||
}
|
||||
|
||||
// Disable MFA and clear sensitive data
|
||||
updateData := maps.MapStrAny{
|
||||
"mfa_enabled": false,
|
||||
"mfa_secret": nil, // Clear the secret
|
||||
"mfa_recovery_hash": nil, // Clear recovery codes
|
||||
"mfa_enabled_at": nil,
|
||||
"mfa_last_verified_at": nil,
|
||||
}
|
||||
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateMFAStatus, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// VerifyMFACode verifies a TOTP code for user
|
||||
func (u *DefaultUser) VerifyMFACode(ctx context.Context, userID string, code string) (bool, error) {
|
||||
// TODO: implement
|
||||
return false, nil
|
||||
// Get user and MFA status
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled", "mfa_secret"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return false, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is enabled
|
||||
mfaEnabled := false
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
mfaEnabled = enabled
|
||||
} else if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
mfaEnabled = enabledInt != 0
|
||||
}
|
||||
|
||||
if !mfaEnabled {
|
||||
return false, fmt.Errorf(ErrMFANotEnabled)
|
||||
}
|
||||
|
||||
// Get stored secret
|
||||
secret, ok := user["mfa_secret"].(string)
|
||||
if !ok || secret == "" {
|
||||
return false, fmt.Errorf("no MFA secret found")
|
||||
}
|
||||
|
||||
// Verify the code
|
||||
valid := totp.Validate(code, secret)
|
||||
if valid {
|
||||
// Update last verified timestamp
|
||||
_, err = m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, maps.MapStrAny{
|
||||
"mfa_last_verified_at": time.Now(),
|
||||
})
|
||||
// Don't fail verification if timestamp update fails
|
||||
if err != nil {
|
||||
// Log the error but continue
|
||||
}
|
||||
}
|
||||
|
||||
return valid, nil
|
||||
}
|
||||
|
||||
// GenerateRecoveryCodes generates new recovery codes for user and stores their hash
|
||||
func (u *DefaultUser) GenerateRecoveryCodes(ctx context.Context, userID string) ([]string, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// Verify user exists and MFA is enabled
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return nil, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is enabled
|
||||
mfaEnabled := false
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
mfaEnabled = enabled
|
||||
} else if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
mfaEnabled = enabledInt != 0
|
||||
}
|
||||
|
||||
if !mfaEnabled {
|
||||
return nil, fmt.Errorf(ErrMFANotEnabled)
|
||||
}
|
||||
|
||||
// Generate multiple recovery codes and store bcrypt hashes (512 char limit)
|
||||
recoveryCount := u.mfaOptions.RecoveryCount
|
||||
recoveryLength := u.mfaOptions.RecoveryLength
|
||||
|
||||
recoveryCodes := make([]string, recoveryCount)
|
||||
|
||||
for i := 0; i < recoveryCount; i++ {
|
||||
code, err := generateRecoveryCode(recoveryLength)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate recovery code: %w", err)
|
||||
}
|
||||
recoveryCodes[i] = code
|
||||
}
|
||||
|
||||
// Hash each recovery code with bcrypt and store all hashes
|
||||
recoveryHashes := make([]string, recoveryCount)
|
||||
for i, code := range recoveryCodes {
|
||||
hashedCode, err := bcrypt.GenerateFromPassword([]byte(code), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to hash recovery code: %w", err)
|
||||
}
|
||||
recoveryHashes[i] = string(hashedCode)
|
||||
}
|
||||
|
||||
// Join all bcrypt hashes (~60 bytes each, 8 hashes = ~480 bytes, under 512 limit)
|
||||
allHashesStr := strings.Join(recoveryHashes, "|||")
|
||||
|
||||
updateData := maps.MapStrAny{
|
||||
"mfa_recovery_hash": allHashesStr, // Store bcrypt hashes (~480 bytes, under 512 limit)
|
||||
}
|
||||
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToUpdateMFAStatus, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return nil, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
// Return all generated recovery codes
|
||||
return recoveryCodes, nil
|
||||
}
|
||||
|
||||
// VerifyRecoveryCode verifies and consumes a recovery code
|
||||
func (u *DefaultUser) VerifyRecoveryCode(ctx context.Context, userID string, code string) (bool, error) {
|
||||
// TODO: implement
|
||||
return false, nil
|
||||
// Get user and MFA status
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled", "mfa_recovery_hash"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return false, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is enabled
|
||||
mfaEnabled := false
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
mfaEnabled = enabled
|
||||
} else if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
mfaEnabled = enabledInt != 0
|
||||
}
|
||||
|
||||
if !mfaEnabled {
|
||||
return false, fmt.Errorf(ErrMFANotEnabled)
|
||||
}
|
||||
|
||||
// Get stored bcrypt hashes string (512 char limit - no problem!)
|
||||
recoveryHashesStr, ok := user["mfa_recovery_hash"].(string)
|
||||
if !ok || recoveryHashesStr == "" {
|
||||
return false, fmt.Errorf("no recovery codes found")
|
||||
}
|
||||
|
||||
// Split into hashes list
|
||||
recoveryHashes := strings.Split(recoveryHashesStr, "|||")
|
||||
if len(recoveryHashes) == 0 {
|
||||
return false, fmt.Errorf("no recovery codes found")
|
||||
}
|
||||
|
||||
// Check if user input code matches any stored bcrypt hash
|
||||
matchIndex := -1
|
||||
for i, storedHash := range recoveryHashes {
|
||||
if storedHash == "" {
|
||||
continue // Skip already used codes
|
||||
}
|
||||
// Verify user input against stored bcrypt hash
|
||||
err := bcrypt.CompareHashAndPassword([]byte(storedHash), []byte(code))
|
||||
if err == nil {
|
||||
matchIndex = i
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if matchIndex == -1 {
|
||||
return false, nil // Invalid code
|
||||
}
|
||||
|
||||
// Mark the code as used by clearing its hash
|
||||
recoveryHashes[matchIndex] = ""
|
||||
updatedHashesStr := strings.Join(recoveryHashes, "|||")
|
||||
|
||||
// Update recovery codes in database and mark verification time
|
||||
updateData := maps.MapStrAny{
|
||||
"mfa_recovery_hash": updatedHashesStr,
|
||||
"mfa_last_verified_at": time.Now(),
|
||||
}
|
||||
|
||||
_, err = m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf(ErrFailedToUpdateMFAStatus, err)
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// IsMFAEnabled checks if MFA is enabled for a user
|
||||
func (u *DefaultUser) IsMFAEnabled(ctx context.Context, userID string) (bool, error) {
|
||||
// TODO: implement
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "mfa_enabled"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return false, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check MFA status
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
return enabled, nil
|
||||
}
|
||||
// Handle different boolean types from database
|
||||
if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
return enabledInt != 0, nil
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// GetMFAConfig retrieves MFA configuration for a user
|
||||
func (u *DefaultUser) GetMFAConfig(ctx context.Context, userID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{
|
||||
"user_id", "mfa_enabled", "mfa_issuer", "mfa_algorithm",
|
||||
"mfa_digits", "mfa_period", "mfa_enabled_at", "mfa_last_verified_at",
|
||||
"mfa_recovery_hash", // Include recovery hash field
|
||||
},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetUser, err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return nil, fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
user := users[0]
|
||||
|
||||
// Check if MFA is enabled
|
||||
mfaEnabled := false
|
||||
if enabled, ok := user["mfa_enabled"].(bool); ok {
|
||||
mfaEnabled = enabled
|
||||
} else if enabledInt, ok := user["mfa_enabled"].(int64); ok {
|
||||
mfaEnabled = enabledInt != 0
|
||||
}
|
||||
|
||||
config := maps.MapStrAny{
|
||||
"user_id": userID,
|
||||
"mfa_enabled": mfaEnabled,
|
||||
}
|
||||
|
||||
if mfaEnabled {
|
||||
// Include MFA configuration details (but not the secret)
|
||||
config["mfa_issuer"] = user["mfa_issuer"]
|
||||
config["mfa_algorithm"] = user["mfa_algorithm"]
|
||||
config["mfa_digits"] = user["mfa_digits"]
|
||||
config["mfa_period"] = user["mfa_period"]
|
||||
config["mfa_enabled_at"] = user["mfa_enabled_at"]
|
||||
config["mfa_last_verified_at"] = user["mfa_last_verified_at"]
|
||||
|
||||
// Check how many recovery codes are available (bcrypt hash storage)
|
||||
if recoveryHashesStr, ok := user["mfa_recovery_hash"].(string); ok && recoveryHashesStr != "" {
|
||||
hashes := strings.Split(recoveryHashesStr, "|||")
|
||||
remainingCodes := 0
|
||||
for _, hash := range hashes {
|
||||
if hash != "" {
|
||||
remainingCodes++
|
||||
}
|
||||
}
|
||||
config["recovery_codes_available"] = remainingCodes
|
||||
} else {
|
||||
config["recovery_codes_available"] = 0
|
||||
}
|
||||
}
|
||||
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// Helper function to generate recovery codes
|
||||
func generateRecoveryCode(length int) (string, error) {
|
||||
// Use alphanumeric charset (excluding similar-looking characters for better UX)
|
||||
// Excludes: 0, O, 1, I, l to avoid confusion
|
||||
const charset = "23456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz"
|
||||
b := make([]byte, length)
|
||||
_, err := rand.Read(b)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
for i := range b {
|
||||
b[i] = charset[b[i]%byte(len(charset))]
|
||||
}
|
||||
|
||||
// Format with dashes for better readability (like GitHub)
|
||||
// For 8-character codes: XXXX-XXXX
|
||||
// For 12-character codes: XXXX-XXXX-XXXX
|
||||
result := string(b)
|
||||
if length == 8 {
|
||||
return fmt.Sprintf("%s-%s", result[:4], result[4:]), nil
|
||||
} else if length >= 12 {
|
||||
return fmt.Sprintf("%s-%s-%s", result[:4], result[4:8], result[8:]), nil
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
|
|
|||
481
openapi/oauth/providers/user/user_mfa_test.go
Normal file
481
openapi/oauth/providers/user/user_mfa_test.go
Normal file
|
|
@ -0,0 +1,481 @@
|
|||
package user_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/pquerna/otp/totp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||
)
|
||||
|
||||
func TestMFAOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Create test user data dynamically
|
||||
testUser := createTestUserData(testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
var mfaSecret string
|
||||
var recoveryCodes []string
|
||||
|
||||
// Test complete MFA setup and usage flow
|
||||
t.Run("CompleteFlow", func(t *testing.T) {
|
||||
// Step 1: Generate MFA Secret
|
||||
secret, qrURL, err := testProvider.GenerateMFASecret(ctx, testUserID, nil)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, secret)
|
||||
assert.NotEmpty(t, qrURL)
|
||||
assert.Contains(t, qrURL, "otpauth://totp/")
|
||||
assert.Contains(t, qrURL, testUserID) // Account name should default to userID
|
||||
|
||||
mfaSecret = secret
|
||||
|
||||
// Verify MFA is not enabled yet
|
||||
enabled, err := testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, enabled)
|
||||
|
||||
// Step 2: Enable MFA
|
||||
code, err := totp.GenerateCode(mfaSecret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = testProvider.EnableMFA(ctx, testUserID, mfaSecret, code)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify MFA is now enabled
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, enabled)
|
||||
|
||||
// Step 3: Verify MFA Code
|
||||
validCode, err := totp.GenerateCode(mfaSecret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
valid, err := testProvider.VerifyMFACode(ctx, testUserID, validCode)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
|
||||
// Test invalid code
|
||||
valid, err = testProvider.VerifyMFACode(ctx, testUserID, "000000")
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid)
|
||||
|
||||
// Step 4: Get MFA Config
|
||||
config, err := testProvider.GetMFAConfig(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, config)
|
||||
|
||||
assert.Equal(t, testUserID, config["user_id"])
|
||||
assert.Equal(t, true, config["mfa_enabled"])
|
||||
assert.Equal(t, "Yao App Engine", config["mfa_issuer"]) // Default issuer
|
||||
assert.Equal(t, "SHA256", config["mfa_algorithm"]) // Default algorithm
|
||||
// Handle database type variations for integers
|
||||
if digits, ok := config["mfa_digits"].(int64); ok {
|
||||
assert.Equal(t, int64(6), digits) // Default digits
|
||||
} else {
|
||||
assert.Equal(t, 6, config["mfa_digits"])
|
||||
}
|
||||
if period, ok := config["mfa_period"].(int64); ok {
|
||||
assert.Equal(t, int64(30), period) // Default period
|
||||
} else {
|
||||
assert.Equal(t, 30, config["mfa_period"])
|
||||
}
|
||||
assert.NotNil(t, config["mfa_enabled_at"])
|
||||
assert.NotNil(t, config["mfa_last_verified_at"]) // Should be set after VerifyMFACode
|
||||
|
||||
// Step 5: Generate Recovery Codes
|
||||
codes, err := testProvider.GenerateRecoveryCodes(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, codes)
|
||||
assert.Len(t, codes, 16) // 16 recovery codes following GitHub standard
|
||||
|
||||
recoveryCodes = codes
|
||||
|
||||
// Verify code format (should be 12 characters with dashes: XXXX-XXXX-XXXX)
|
||||
for _, code := range codes {
|
||||
assert.Len(t, code, 14) // 12 chars + 2 dashes
|
||||
assert.Regexp(t, `^[23456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz]{4}-[23456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz]{4}-[23456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz]{4}$`, code)
|
||||
}
|
||||
|
||||
// Verify all codes are unique
|
||||
codeSet := make(map[string]bool)
|
||||
for _, code := range codes {
|
||||
assert.False(t, codeSet[code], "Duplicate recovery code: %s", code)
|
||||
codeSet[code] = true
|
||||
}
|
||||
|
||||
// Verify MFA config now shows recovery codes available
|
||||
config, err = testProvider.GetMFAConfig(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 16, config["recovery_codes_available"])
|
||||
|
||||
// Step 6: Test Recovery Code Verification
|
||||
testCode := recoveryCodes[0]
|
||||
valid, err = testProvider.VerifyRecoveryCode(ctx, testUserID, testCode)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
|
||||
// Test same code again (should be consumed/invalid)
|
||||
valid, err = testProvider.VerifyRecoveryCode(ctx, testUserID, testCode)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid)
|
||||
|
||||
// Verify recovery codes available decreased
|
||||
config, err = testProvider.GetMFAConfig(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 15, config["recovery_codes_available"]) // One less (16-1=15)
|
||||
|
||||
// Test invalid recovery code
|
||||
valid, err = testProvider.VerifyRecoveryCode(ctx, testUserID, "invalid-code")
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid)
|
||||
|
||||
// Step 7: Disable MFA
|
||||
disableCode, err := totp.GenerateCode(mfaSecret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = testProvider.DisableMFA(ctx, testUserID, disableCode)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify MFA is now disabled
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, enabled)
|
||||
|
||||
// Verify MFA config reflects disabled state
|
||||
config, err = testProvider.GetMFAConfig(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, false, config["mfa_enabled"])
|
||||
// Should not contain MFA-specific fields when disabled
|
||||
assert.NotContains(t, config, "mfa_issuer")
|
||||
assert.NotContains(t, config, "mfa_algorithm")
|
||||
assert.NotContains(t, config, "recovery_codes_available")
|
||||
})
|
||||
}
|
||||
|
||||
func TestMFAErrorHandling(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
nonExistentUserID := "non-existent-user-" + testUUID
|
||||
|
||||
// Create test user for some error tests
|
||||
testUser := createTestUserData("mfaerror" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
t.Run("GenerateMFASecret_UserNotFound", func(t *testing.T) {
|
||||
_, _, err := testProvider.GenerateMFASecret(ctx, nonExistentUserID, nil)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("EnableMFA_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.EnableMFA(ctx, nonExistentUserID, "testsecret", "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("EnableMFA_InvalidCode", func(t *testing.T) {
|
||||
// Generate MFA secret first
|
||||
secret, _, err := testProvider.GenerateMFASecret(ctx, testUserID, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to enable with invalid code
|
||||
err = testProvider.EnableMFA(ctx, testUserID, secret, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid MFA code")
|
||||
})
|
||||
|
||||
t.Run("EnableMFA_NoSecret", func(t *testing.T) {
|
||||
// Try to enable MFA without generating secret first (use new user)
|
||||
newUser := createTestUserData("nomfasecret" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
err := testProvider.EnableMFA(ctx, newUserID, "", "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "no MFA secret found")
|
||||
})
|
||||
|
||||
t.Run("DisableMFA_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.DisableMFA(ctx, nonExistentUserID, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("DisableMFA_NotEnabled", func(t *testing.T) {
|
||||
// Use user without MFA enabled
|
||||
newUser := createTestUserData("nomfauser" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
err := testProvider.DisableMFA(ctx, newUserID, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
})
|
||||
|
||||
t.Run("VerifyMFACode_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.VerifyMFACode(ctx, nonExistentUserID, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("VerifyMFACode_NotEnabled", func(t *testing.T) {
|
||||
// Use user without MFA enabled
|
||||
newUser := createTestUserData("nomfaverify" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
_, err := testProvider.VerifyMFACode(ctx, newUserID, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
})
|
||||
|
||||
t.Run("GenerateRecoveryCodes_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GenerateRecoveryCodes(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("GenerateRecoveryCodes_NotEnabled", func(t *testing.T) {
|
||||
// Use user without MFA enabled
|
||||
newUser := createTestUserData("norecovery" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
_, err := testProvider.GenerateRecoveryCodes(ctx, newUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
})
|
||||
|
||||
t.Run("VerifyRecoveryCode_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.VerifyRecoveryCode(ctx, nonExistentUserID, "test-code")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("VerifyRecoveryCode_NotEnabled", func(t *testing.T) {
|
||||
// Use user without MFA enabled
|
||||
newUser := createTestUserData("noverifyrecov" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
_, err := testProvider.VerifyRecoveryCode(ctx, newUserID, "test-code")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
})
|
||||
|
||||
t.Run("VerifyRecoveryCode_NoRecoveryCodes", func(t *testing.T) {
|
||||
// Create user, enable MFA, but don't generate recovery codes
|
||||
newUser := createTestUserData("norecodes" + testUUID)
|
||||
_, newUserID := setupTestUser(t, ctx, newUser)
|
||||
|
||||
// Generate and enable MFA
|
||||
secret, _, err := testProvider.GenerateMFASecret(ctx, newUserID, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
code, err := totp.GenerateCode(secret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = testProvider.EnableMFA(ctx, newUserID, secret, code)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to verify recovery code without generating them
|
||||
_, err = testProvider.VerifyRecoveryCode(ctx, newUserID, "test-code")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "no recovery codes found")
|
||||
})
|
||||
|
||||
t.Run("IsMFAEnabled_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.IsMFAEnabled(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("GetMFAConfig_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetMFAConfig(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestMFACustomOptions(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
t.Run("CustomMFAOptions", func(t *testing.T) {
|
||||
// Create test user
|
||||
testUser := createTestUserData("mfaoptions" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
// Test with custom options
|
||||
customOptions := &types.MFAOptions{
|
||||
Issuer: "Custom MFA Test",
|
||||
Algorithm: "SHA1",
|
||||
Digits: 8,
|
||||
Period: 60,
|
||||
SecretSize: 16,
|
||||
AccountName: "custom@test.com",
|
||||
}
|
||||
|
||||
// Generate MFA secret with custom options
|
||||
secret, qrURL, err := testProvider.GenerateMFASecret(ctx, testUserID, customOptions)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, secret)
|
||||
assert.Contains(t, qrURL, "Custom%20MFA%20Test")
|
||||
assert.Contains(t, qrURL, "custom@test.com")
|
||||
assert.Contains(t, qrURL, "algorithm=SHA1")
|
||||
assert.Contains(t, qrURL, "digits=8")
|
||||
assert.Contains(t, qrURL, "period=60")
|
||||
|
||||
// Enable MFA
|
||||
code, err := totp.GenerateCode(secret, time.Now())
|
||||
require.NoError(t, err)
|
||||
err = testProvider.EnableMFA(ctx, testUserID, secret, code)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify MFA config shows custom settings
|
||||
config, err := testProvider.GetMFAConfig(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "Custom MFA Test", config["mfa_issuer"])
|
||||
assert.Equal(t, "SHA1", config["mfa_algorithm"])
|
||||
// Handle database type variations
|
||||
if digits, ok := config["mfa_digits"].(int64); ok {
|
||||
assert.Equal(t, int64(8), digits)
|
||||
} else {
|
||||
assert.Equal(t, 8, config["mfa_digits"])
|
||||
}
|
||||
if period, ok := config["mfa_period"].(int64); ok {
|
||||
assert.Equal(t, int64(60), period)
|
||||
} else {
|
||||
assert.Equal(t, 60, config["mfa_period"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("PartialCustomOptions", func(t *testing.T) {
|
||||
// Create another user for partial options test
|
||||
testUser := createTestUserData("mfapartial" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
// Test with partial custom options (some fields empty)
|
||||
partialOptions := &types.MFAOptions{
|
||||
Issuer: "Partial Test",
|
||||
AccountName: "partial@test.com",
|
||||
// Other fields empty, should use defaults
|
||||
}
|
||||
|
||||
secret, qrURL, err := testProvider.GenerateMFASecret(ctx, testUserID, partialOptions)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, secret)
|
||||
assert.Contains(t, qrURL, "Partial%20Test")
|
||||
assert.Contains(t, qrURL, "partial@test.com")
|
||||
// Should use defaults for other parameters
|
||||
assert.Contains(t, qrURL, "algorithm=SHA256") // Default
|
||||
assert.Contains(t, qrURL, "digits=6") // Default
|
||||
assert.Contains(t, qrURL, "period=30") // Default
|
||||
})
|
||||
}
|
||||
|
||||
func TestMFAStateMachine(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Create test user
|
||||
testUser := createTestUserData("mfastate" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
t.Run("StateTransitions", func(t *testing.T) {
|
||||
// State 1: No MFA configured
|
||||
enabled, err := testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, enabled)
|
||||
|
||||
// Should fail to verify code when MFA not enabled
|
||||
_, err = testProvider.VerifyMFACode(ctx, testUserID, "000000")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
|
||||
// State 2: Generate secret (but not enabled yet)
|
||||
secret, _, err := testProvider.GenerateMFASecret(ctx, testUserID, nil)
|
||||
assert.NoError(t, err)
|
||||
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, enabled) // Still not enabled
|
||||
|
||||
// State 3: Enable MFA
|
||||
code, err := totp.GenerateCode(secret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = testProvider.EnableMFA(ctx, testUserID, secret, code)
|
||||
assert.NoError(t, err)
|
||||
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, enabled) // Now enabled
|
||||
|
||||
// Should be able to verify codes now
|
||||
newCode, err := totp.GenerateCode(secret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
valid, err := testProvider.VerifyMFACode(ctx, testUserID, newCode)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
|
||||
// State 4: Regenerate secret (MFA remains enabled but with new secret)
|
||||
newSecret, _, err := testProvider.GenerateMFASecret(ctx, testUserID, nil)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEqual(t, secret, newSecret)
|
||||
|
||||
// MFA should still be enabled
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, enabled) // Still enabled
|
||||
|
||||
// Old codes should not work after regenerating secret, but new codes should work
|
||||
newCode2, err := totp.GenerateCode(newSecret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
valid, err = testProvider.VerifyMFACode(ctx, testUserID, newCode2)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Should work with new secret
|
||||
|
||||
// Old codes should not work
|
||||
oldCode, err := totp.GenerateCode(secret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
valid, err = testProvider.VerifyMFACode(ctx, testUserID, oldCode)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should fail with old secret
|
||||
|
||||
// State 5: Disable MFA
|
||||
disableCode, err := totp.GenerateCode(newSecret, time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = testProvider.DisableMFA(ctx, testUserID, disableCode)
|
||||
assert.NoError(t, err)
|
||||
|
||||
enabled, err = testProvider.IsMFAEnabled(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, enabled) // Back to disabled
|
||||
|
||||
// Should not be able to verify codes after disabling
|
||||
_, err = testProvider.VerifyMFACode(ctx, testUserID, disableCode)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "MFA is not enabled")
|
||||
})
|
||||
}
|
||||
|
|
@ -181,7 +181,7 @@ type UserProvider interface {
|
|||
ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error)
|
||||
|
||||
// User MFA Management
|
||||
GenerateMFASecret(ctx context.Context, userID string, issuer string, accountName string) (string, string, error)
|
||||
GenerateMFASecret(ctx context.Context, userID string, options *MFAOptions) (string, string, error)
|
||||
EnableMFA(ctx context.Context, userID string, secret string, code string) error
|
||||
DisableMFA(ctx context.Context, userID string, code string) error
|
||||
VerifyMFACode(ctx context.Context, userID string, code string) (bool, error)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,18 @@ import (
|
|||
"github.com/golang-jwt/jwt/v4"
|
||||
)
|
||||
|
||||
// MFAOptions contains configuration for MFA operations
|
||||
type MFAOptions struct {
|
||||
Issuer string // Issuer name displayed in authenticator app
|
||||
Algorithm string // TOTP algorithm: "SHA1", "SHA256", "SHA512"
|
||||
Digits int // Number of digits in TOTP code (6 or 8)
|
||||
Period int // TOTP time period in seconds (usually 30)
|
||||
SecretSize int // Secret key size in bytes (usually 32)
|
||||
RecoveryCount int // Number of recovery codes to generate
|
||||
RecoveryLength int // Length of each recovery code
|
||||
AccountName string // Optional account name (defaults to userID)
|
||||
}
|
||||
|
||||
// ErrorResponse represents an OAuth 2.1 error response
|
||||
type ErrorResponse struct {
|
||||
Code string `json:"error"`
|
||||
|
|
|
|||
|
|
@ -316,9 +316,8 @@
|
|||
"type": "string",
|
||||
"label": "MFA Recovery Hash",
|
||||
"comment": "Hashed recovery code for MFA backup authentication",
|
||||
"length": 255,
|
||||
"nullable": true,
|
||||
"crypt": "PASSWORD"
|
||||
"length": 1024,
|
||||
"nullable": true
|
||||
},
|
||||
{
|
||||
"name": "mfa_enabled_at",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue