Merge pull request #1072 from trheyi/main

Implement MFA functionality in user management
This commit is contained in:
Max 2025-08-02 21:09:21 +08:00 committed by GitHub
commit 1d5de18eb9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 1241 additions and 124 deletions

File diff suppressed because one or more lines are too long

2
go.mod
View file

@ -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
View file

@ -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=

View file

@ -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,
}
}

View file

@ -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
}

View 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")
})
}

View file

@ -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)

View file

@ -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"`

View file

@ -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",