Merge pull request #1071 from trheyi/main
Implement user type and user role management
This commit is contained in:
commit
7b5777e728
12 changed files with 3379 additions and 1916 deletions
|
|
@ -76,6 +76,33 @@ var (
|
|||
"website", "gender", "birthdate", "zoneinfo", "locale", "phone_number", "phone_number_verified",
|
||||
"address", "raw", "last_login_at", "is_active", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// DefaultRoleFields contains basic role fields
|
||||
DefaultRoleFields = []interface{}{
|
||||
"id", "role_id", "name", "description", "is_active", "is_default", "is_system",
|
||||
"level", "sort_order", "color", "icon", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// DefaultRoleDetailFields contains all role fields including permissions and metadata
|
||||
DefaultRoleDetailFields = []interface{}{
|
||||
"id", "role_id", "name", "description", "permissions", "restricted_permissions",
|
||||
"parent_role_id", "level", "is_active", "is_default", "is_system", "sort_order",
|
||||
"color", "icon", "max_users", "requires_approval", "auto_revoke_days",
|
||||
"metadata", "conditions", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// DefaultTypeFields contains basic type fields
|
||||
DefaultTypeFields = []interface{}{
|
||||
"id", "type_id", "name", "description", "is_active", "is_default", "sort_order",
|
||||
"default_role_id", "max_sessions", "session_timeout", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// DefaultTypeDetailFields contains all type fields including configuration and metadata
|
||||
DefaultTypeDetailFields = []interface{}{
|
||||
"id", "type_id", "name", "description", "default_role_id", "schema", "metadata",
|
||||
"is_active", "is_default", "sort_order", "max_sessions", "session_timeout",
|
||||
"password_policy", "features", "limits", "created_at", "updated_at",
|
||||
}
|
||||
)
|
||||
|
||||
// DefaultUser provides a default implementation of UserProvider
|
||||
|
|
@ -100,6 +127,14 @@ type DefaultUser struct {
|
|||
// OAuth Account Field lists
|
||||
oauthAccountFields []interface{} // configurable
|
||||
oauthAccountDetailFields []interface{} // configurable
|
||||
|
||||
// Role Field lists
|
||||
roleFields []interface{} // configurable
|
||||
roleDetailFields []interface{} // configurable
|
||||
|
||||
// Type Field lists
|
||||
typeFields []interface{} // configurable
|
||||
typeDetailFields []interface{} // configurable
|
||||
}
|
||||
|
||||
// IDStrategy defines the strategy for generating user IDs
|
||||
|
|
@ -132,6 +167,14 @@ type DefaultUserOptions struct {
|
|||
// OAuth Account field lists (use defaults if not specified)
|
||||
OAuthAccountFields []interface{} // basic OAuth account fields
|
||||
OAuthAccountDetailFields []interface{} // detailed OAuth account fields with OIDC claims
|
||||
|
||||
// Role field lists (use defaults if not specified)
|
||||
RoleFields []interface{} // basic role fields
|
||||
RoleDetailFields []interface{} // detailed role fields including permissions and metadata
|
||||
|
||||
// Type field lists (use defaults if not specified)
|
||||
TypeFields []interface{} // basic type fields
|
||||
TypeDetailFields []interface{} // detailed type fields including configuration and metadata
|
||||
}
|
||||
|
||||
// NewDefaultUser creates a new DefaultUser
|
||||
|
|
@ -188,6 +231,28 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
|||
oauthAccountDetailFields = DefaultOAuthAccountDetailFields
|
||||
}
|
||||
|
||||
// Set role field lists with defaults if not specified
|
||||
roleFields := options.RoleFields
|
||||
if roleFields == nil {
|
||||
roleFields = DefaultRoleFields
|
||||
}
|
||||
|
||||
roleDetailFields := options.RoleDetailFields
|
||||
if roleDetailFields == nil {
|
||||
roleDetailFields = DefaultRoleDetailFields
|
||||
}
|
||||
|
||||
// Set type field lists with defaults if not specified
|
||||
typeFields := options.TypeFields
|
||||
if typeFields == nil {
|
||||
typeFields = DefaultTypeFields
|
||||
}
|
||||
|
||||
typeDetailFields := options.TypeDetailFields
|
||||
if typeDetailFields == nil {
|
||||
typeDetailFields = DefaultTypeDetailFields
|
||||
}
|
||||
|
||||
return &DefaultUser{
|
||||
prefix: options.Prefix,
|
||||
model: model,
|
||||
|
|
@ -205,5 +270,13 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
|||
// OAuth Account field lists
|
||||
oauthAccountFields: oauthAccountFields,
|
||||
oauthAccountDetailFields: oauthAccountDetailFields,
|
||||
|
||||
// Role field lists
|
||||
roleFields: roleFields,
|
||||
roleDetailFields: roleDetailFields,
|
||||
|
||||
// Type field lists
|
||||
typeFields: typeFields,
|
||||
typeDetailFields: typeDetailFields,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,864 +0,0 @@
|
|||
package removedfuncref
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha1"
|
||||
"crypto/sha256"
|
||||
"crypto/sha512"
|
||||
"encoding/base32"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"hash"
|
||||
"math"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/gou/store"
|
||||
)
|
||||
|
||||
// Safe user fields that can be displayed to users
|
||||
var (
|
||||
// PublicUserFields contains fields that can be safely returned to users
|
||||
PublicUserFields = []interface{}{
|
||||
"id", "subject", "username", "email", "first_name", "last_name",
|
||||
"full_name", "avatar_url", "mobile", "address", "scopes", "status",
|
||||
"email_verified", "mobile_verified", "two_factor_enabled",
|
||||
"last_login_at", "metadata", "preferences", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
// BasicUserFields contains minimal fields for basic user info
|
||||
BasicUserFields = []interface{}{
|
||||
"id", "subject", "username", "email", "first_name", "last_name",
|
||||
"full_name", "avatar_url", "status", "email_verified", "mobile_verified",
|
||||
}
|
||||
|
||||
// AuthUserFields contains fields needed for authentication
|
||||
AuthUserFields = []interface{}{
|
||||
"id", "subject", "username", "email", "password_hash", "scopes", "status",
|
||||
"email_verified", "mobile_verified", "two_factor_enabled", "last_login_at",
|
||||
}
|
||||
|
||||
// TwoFactorUserFields contains fields needed for two-factor authentication
|
||||
TwoFactorUserFields = []interface{}{
|
||||
"id", "two_factor_enabled", "two_factor_secret", "two_factor_algorithm",
|
||||
"two_factor_digits", "two_factor_period", "two_factor_recovery_codes",
|
||||
}
|
||||
)
|
||||
|
||||
// DefaultUser provides a default implementation of UserProvider
|
||||
type DefaultUser struct {
|
||||
prefix string
|
||||
model string
|
||||
cache store.Store
|
||||
tokenStore store.Store
|
||||
}
|
||||
|
||||
// DefaultUserOptions provides options for the DefaultUser
|
||||
type DefaultUserOptions struct {
|
||||
Prefix string
|
||||
Model string // bind to a specific user model
|
||||
Cache store.Store
|
||||
TokenStore store.Store // store for OAuth tokens
|
||||
}
|
||||
|
||||
// NewDefaultUser creates a new DefaultUser
|
||||
func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
||||
// Set default model name if not specified
|
||||
modelName := options.Model
|
||||
if modelName == "" {
|
||||
modelName = "__yao.user"
|
||||
}
|
||||
|
||||
return &DefaultUser{
|
||||
prefix: options.Prefix,
|
||||
model: modelName,
|
||||
cache: options.Cache,
|
||||
tokenStore: options.TokenStore,
|
||||
}
|
||||
}
|
||||
|
||||
// Key generation methods
|
||||
|
||||
func (u *DefaultUser) tokenKey(accessToken string) string {
|
||||
return fmt.Sprintf("%s:token:%s", u.prefix, accessToken)
|
||||
}
|
||||
|
||||
func (u *DefaultUser) cacheKey(userID string) string {
|
||||
return fmt.Sprintf("%s:user:%s", u.prefix, userID)
|
||||
}
|
||||
|
||||
func (u *DefaultUser) subjectCacheKey(subject string) string {
|
||||
return fmt.Sprintf("%s:user:subject:%s", u.prefix, subject)
|
||||
}
|
||||
|
||||
func (u *DefaultUser) usernameCacheKey(username string) string {
|
||||
return fmt.Sprintf("%s:user:username:%s", u.prefix, username)
|
||||
}
|
||||
|
||||
func (u *DefaultUser) emailCacheKey(email string) string {
|
||||
return fmt.Sprintf("%s:user:email:%s", u.prefix, email)
|
||||
}
|
||||
|
||||
// GetUserByAccessToken retrieves user information using an access token
|
||||
func (u *DefaultUser) GetUserByAccessToken(ctx context.Context, accessToken string) (interface{}, error) {
|
||||
// Get token information from tokenStore
|
||||
tokenData, exists := u.tokenStore.Get(u.tokenKey(accessToken))
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("token not found")
|
||||
}
|
||||
|
||||
// Parse token data to get user subject
|
||||
var tokenInfo map[string]interface{}
|
||||
var ok bool
|
||||
|
||||
// Try to convert to map[string]interface{} directly
|
||||
if tokenInfo, ok = tokenData.(map[string]interface{}); !ok {
|
||||
// If direct conversion fails, try to handle other possible types
|
||||
switch v := tokenData.(type) {
|
||||
case map[interface{}]interface{}:
|
||||
// Convert map[interface{}]interface{} to map[string]interface{}
|
||||
tokenInfo = make(map[string]interface{})
|
||||
for key, val := range v {
|
||||
if keyStr, ok := key.(string); ok {
|
||||
tokenInfo[keyStr] = val
|
||||
}
|
||||
}
|
||||
default:
|
||||
// Try to convert using map[string]interface{} casting
|
||||
// This handles primitive.M and other MongoDB types
|
||||
if reflect.TypeOf(v).Kind() == reflect.Map {
|
||||
tokenInfo = make(map[string]interface{})
|
||||
rv := reflect.ValueOf(v)
|
||||
for _, key := range rv.MapKeys() {
|
||||
if keyStr, ok := key.Interface().(string); ok {
|
||||
tokenInfo[keyStr] = rv.MapIndex(key).Interface()
|
||||
}
|
||||
}
|
||||
if len(tokenInfo) == 0 {
|
||||
return nil, fmt.Errorf("invalid token data format: %T", tokenData)
|
||||
}
|
||||
} else {
|
||||
return nil, fmt.Errorf("invalid token data format: %T", tokenData)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
subject, ok := tokenInfo["subject"].(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid subject in token")
|
||||
}
|
||||
|
||||
// Get user by subject
|
||||
return u.GetUserBySubject(ctx, subject)
|
||||
}
|
||||
|
||||
// GetUserBySubject retrieves user information using a subject identifier
|
||||
func (u *DefaultUser) GetUserBySubject(ctx context.Context, subject string) (interface{}, error) {
|
||||
// Try cache first if available
|
||||
if u.cache != nil {
|
||||
if cached, ok := u.cache.Get(u.subjectCacheKey(subject)); ok {
|
||||
return cached, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Get user from database using the model
|
||||
m := model.Select(u.model)
|
||||
|
||||
user, err := m.Get(model.QueryParam{
|
||||
Select: PublicUserFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "subject", Value: subject},
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get user by subject: %w", err)
|
||||
}
|
||||
|
||||
if len(user) == 0 {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
userData := user[0]
|
||||
|
||||
// Cache the result if cache is available
|
||||
if u.cache != nil {
|
||||
u.cache.Set(u.subjectCacheKey(subject), userData, 5*time.Minute)
|
||||
}
|
||||
|
||||
return userData, nil
|
||||
}
|
||||
|
||||
// ValidateUserScope validates if a user has access to requested scopes
|
||||
func (u *DefaultUser) ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error) {
|
||||
var user interface{}
|
||||
var err error
|
||||
|
||||
// Try cache first if available
|
||||
if u.cache != nil {
|
||||
if cached, ok := u.cache.Get(u.cacheKey(userID)); ok {
|
||||
user = cached
|
||||
}
|
||||
}
|
||||
|
||||
// If not in cache, get from database
|
||||
if user == nil {
|
||||
m := model.Select(u.model)
|
||||
user, err = m.Find(userID, model.QueryParam{
|
||||
Select: []interface{}{"scopes", "status"},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to get user: %w", err)
|
||||
}
|
||||
|
||||
// Cache the result if cache is available
|
||||
if u.cache != nil {
|
||||
u.cache.Set(u.cacheKey(userID), user, 5*time.Minute)
|
||||
}
|
||||
}
|
||||
|
||||
// Check if user data is valid
|
||||
if user == nil {
|
||||
return false, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
// Convert user to map for indexing
|
||||
var userMap map[string]interface{}
|
||||
switch v := user.(type) {
|
||||
case map[string]interface{}:
|
||||
userMap = v
|
||||
default:
|
||||
// Try to convert using reflection if it's a map-like type
|
||||
if reflect.TypeOf(v).Kind() == reflect.Map {
|
||||
userMap = make(map[string]interface{})
|
||||
rv := reflect.ValueOf(v)
|
||||
for _, key := range rv.MapKeys() {
|
||||
if keyStr, ok := key.Interface().(string); ok {
|
||||
userMap[keyStr] = rv.MapIndex(key).Interface()
|
||||
}
|
||||
}
|
||||
} else {
|
||||
return false, fmt.Errorf("invalid user data format")
|
||||
}
|
||||
}
|
||||
|
||||
// Check if user is active
|
||||
if status, ok := userMap["status"].(string); ok && status != "active" {
|
||||
return false, fmt.Errorf("user is not active")
|
||||
}
|
||||
|
||||
// Get user scopes
|
||||
userScopes, ok := userMap["scopes"].([]interface{})
|
||||
if !ok {
|
||||
// If no scopes defined, deny access
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// Convert user scopes to string slice
|
||||
userScopeStrings := make([]string, len(userScopes))
|
||||
for i, scope := range userScopes {
|
||||
if scopeStr, ok := scope.(string); ok {
|
||||
userScopeStrings[i] = scopeStr
|
||||
}
|
||||
}
|
||||
|
||||
// Check if user has all requested scopes
|
||||
for _, requestedScope := range scopes {
|
||||
hasScope := false
|
||||
for _, userScope := range userScopeStrings {
|
||||
if userScope == requestedScope {
|
||||
hasScope = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !hasScope {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// // StoreToken stores a token in the token store with expiration time
|
||||
// func (u *DefaultUser) StoreToken(accessToken string, tokenData map[string]interface{}, expiration time.Duration) error {
|
||||
// return u.tokenStore.Set(u.tokenKey(accessToken), tokenData, expiration)
|
||||
// }
|
||||
|
||||
// // RevokeToken revokes a token by removing it from the token store
|
||||
// func (u *DefaultUser) RevokeToken(accessToken string) error {
|
||||
// u.tokenStore.Del(u.tokenKey(accessToken))
|
||||
// return nil
|
||||
// }
|
||||
|
||||
// // TokenExists checks if a token exists in the token store
|
||||
// func (u *DefaultUser) TokenExists(accessToken string) bool {
|
||||
// _, exists := u.tokenStore.Get(u.tokenKey(accessToken))
|
||||
// return exists
|
||||
// }
|
||||
|
||||
// // GetTokenData retrieves token data from the token store
|
||||
// func (u *DefaultUser) GetTokenData(accessToken string) (map[string]interface{}, error) {
|
||||
// tokenData, exists := u.tokenStore.Get(u.tokenKey(accessToken))
|
||||
// if !exists {
|
||||
// return nil, fmt.Errorf("token not found")
|
||||
// }
|
||||
|
||||
// // Try to convert to map[string]interface{} directly
|
||||
// if tokenInfo, ok := tokenData.(map[string]interface{}); ok {
|
||||
// return tokenInfo, nil
|
||||
// }
|
||||
|
||||
// // If direct conversion fails, try to handle other possible types
|
||||
// // This handles cases where MongoDB might return different types
|
||||
// switch v := tokenData.(type) {
|
||||
// case map[string]interface{}:
|
||||
// return v, nil
|
||||
// case map[interface{}]interface{}:
|
||||
// // Convert map[interface{}]interface{} to map[string]interface{}
|
||||
// result := make(map[string]interface{})
|
||||
// for key, val := range v {
|
||||
// if keyStr, ok := key.(string); ok {
|
||||
// result[keyStr] = val
|
||||
// }
|
||||
// }
|
||||
// return result, nil
|
||||
// default:
|
||||
// // Try to convert using map[string]interface{} casting
|
||||
// // This handles primitive.M and other MongoDB types
|
||||
// if reflect.TypeOf(v).Kind() == reflect.Map {
|
||||
// result := make(map[string]interface{})
|
||||
// rv := reflect.ValueOf(v)
|
||||
// for _, key := range rv.MapKeys() {
|
||||
// if keyStr, ok := key.Interface().(string); ok {
|
||||
// result[keyStr] = rv.MapIndex(key).Interface()
|
||||
// }
|
||||
// }
|
||||
// if len(result) > 0 {
|
||||
// return result, nil
|
||||
// }
|
||||
// }
|
||||
// return nil, fmt.Errorf("invalid token data format: %T", tokenData)
|
||||
// }
|
||||
// }
|
||||
|
||||
// CreateUser creates a new user in the database
|
||||
func (u *DefaultUser) CreateUser(userData map[string]interface{}) (interface{}, error) {
|
||||
m := model.Select(u.model)
|
||||
userID, err := m.Create(userData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Note: No need to cache newly created user data since it will be cached
|
||||
// when accessed for the first time through other methods
|
||||
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
// UpdateUserLastLogin updates the user's last login timestamp
|
||||
func (u *DefaultUser) UpdateUserLastLogin(userID interface{}) error {
|
||||
m := model.Select(u.model)
|
||||
err := m.Update(userID, map[string]interface{}{
|
||||
"last_login_at": time.Now(),
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Clear cache for this user since data has changed
|
||||
if u.cache != nil {
|
||||
userIDStr := fmt.Sprintf("%v", userID)
|
||||
u.cache.Del(u.cacheKey(userIDStr))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUserByUsername retrieves user by username
|
||||
func (u *DefaultUser) GetUserByUsername(username string) (interface{}, error) {
|
||||
// Try cache first if available
|
||||
if u.cache != nil {
|
||||
if cached, ok := u.cache.Get(u.usernameCacheKey(username)); ok {
|
||||
return cached, nil
|
||||
}
|
||||
}
|
||||
|
||||
m := model.Select(u.model)
|
||||
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: PublicUserFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "username", Value: username},
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get user by username: %w", err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
userData := users[0]
|
||||
|
||||
// Cache the result if cache is available
|
||||
if u.cache != nil {
|
||||
u.cache.Set(u.usernameCacheKey(username), userData, 5*time.Minute)
|
||||
}
|
||||
|
||||
return userData, nil
|
||||
}
|
||||
|
||||
// GetUserByEmail retrieves user by email
|
||||
func (u *DefaultUser) GetUserByEmail(email string) (interface{}, error) {
|
||||
// Try cache first if available
|
||||
if u.cache != nil {
|
||||
if cached, ok := u.cache.Get(u.emailCacheKey(email)); ok {
|
||||
return cached, nil
|
||||
}
|
||||
}
|
||||
|
||||
m := model.Select(u.model)
|
||||
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: PublicUserFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "email", Value: email},
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get user by email: %w", err)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
userData := users[0]
|
||||
|
||||
// Cache the result if cache is available
|
||||
if u.cache != nil {
|
||||
u.cache.Set(u.emailCacheKey(email), userData, 5*time.Minute)
|
||||
}
|
||||
|
||||
return userData, nil
|
||||
}
|
||||
|
||||
// GenerateTOTPSecret generates a new TOTP secret for user
|
||||
func (u *DefaultUser) GenerateTOTPSecret(ctx context.Context, userID string, issuer string, accountName string) (string, string, error) {
|
||||
// Generate a random 20-byte secret
|
||||
secret := make([]byte, 20)
|
||||
if _, err := rand.Read(secret); err != nil {
|
||||
return "", "", fmt.Errorf("failed to generate secret: %w", err)
|
||||
}
|
||||
|
||||
// Encode secret as Base32
|
||||
secretBase32 := base32.StdEncoding.EncodeToString(secret)
|
||||
secretBase32 = strings.TrimRight(secretBase32, "=") // Remove padding
|
||||
|
||||
// Set default values
|
||||
if issuer == "" {
|
||||
issuer = "YAO OAuth"
|
||||
}
|
||||
if accountName == "" {
|
||||
accountName = userID
|
||||
}
|
||||
|
||||
// Generate QR code URL
|
||||
qrURL := u.generateQRCodeURL(secretBase32, issuer, accountName)
|
||||
|
||||
return secretBase32, qrURL, nil
|
||||
}
|
||||
|
||||
// EnableTwoFactor enables two-factor authentication for user
|
||||
func (u *DefaultUser) EnableTwoFactor(ctx context.Context, userID string, secret string, code string) error {
|
||||
// Verify the provided code with the secret
|
||||
if !u.verifyTOTPWithSecret(secret, code, "SHA1", 6, 30) {
|
||||
return fmt.Errorf("invalid verification code")
|
||||
}
|
||||
|
||||
// Generate recovery codes
|
||||
recoveryCodes, err := u.generateRecoveryCodesList()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate recovery codes: %w", err)
|
||||
}
|
||||
|
||||
// Update user record
|
||||
m := model.Select(u.model)
|
||||
now := time.Now()
|
||||
err = m.Update(userID, map[string]interface{}{
|
||||
"two_factor_enabled": true,
|
||||
"two_factor_secret": secret,
|
||||
"two_factor_recovery_codes": recoveryCodes,
|
||||
"two_factor_enabled_at": now,
|
||||
"two_factor_last_verified_at": now,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to enable two-factor authentication: %w", err)
|
||||
}
|
||||
|
||||
// Clear user cache
|
||||
if u.cache != nil {
|
||||
u.cache.Del(u.cacheKey(userID))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DisableTwoFactor disables two-factor authentication for user
|
||||
func (u *DefaultUser) DisableTwoFactor(ctx context.Context, userID string, code string) error {
|
||||
// Get current user data
|
||||
m := model.Select(u.model)
|
||||
user, err := m.Find(userID, model.QueryParam{
|
||||
Select: []interface{}{"two_factor_secret", "two_factor_recovery_codes"},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get user: %w", err)
|
||||
}
|
||||
|
||||
if user == nil {
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
// Verify code (either TOTP or recovery code)
|
||||
verified := false
|
||||
if secret, ok := user["two_factor_secret"].(string); ok && secret != "" {
|
||||
verified = u.verifyTOTPWithSecret(secret, code, "SHA1", 6, 30)
|
||||
}
|
||||
|
||||
if !verified {
|
||||
// Try recovery code
|
||||
if recoveryCodes, ok := user["two_factor_recovery_codes"].([]interface{}); ok {
|
||||
for _, rc := range recoveryCodes {
|
||||
if rcStr, ok := rc.(string); ok && rcStr == code {
|
||||
verified = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !verified {
|
||||
return fmt.Errorf("invalid verification code")
|
||||
}
|
||||
|
||||
// Disable two-factor authentication
|
||||
err = m.Update(userID, map[string]interface{}{
|
||||
"two_factor_enabled": false,
|
||||
"two_factor_secret": nil,
|
||||
"two_factor_recovery_codes": nil,
|
||||
"two_factor_enabled_at": nil,
|
||||
"two_factor_last_verified_at": nil,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to disable two-factor authentication: %w", err)
|
||||
}
|
||||
|
||||
// Clear user cache
|
||||
if u.cache != nil {
|
||||
u.cache.Del(u.cacheKey(userID))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// VerifyTOTPCode verifies a TOTP code for user
|
||||
func (u *DefaultUser) VerifyTOTPCode(ctx context.Context, userID string, code string) (bool, error) {
|
||||
// Get user data
|
||||
m := model.Select(u.model)
|
||||
user, err := m.Find(userID, model.QueryParam{
|
||||
Select: []interface{}{"two_factor_enabled", "two_factor_secret", "two_factor_algorithm", "two_factor_digits", "two_factor_period"},
|
||||
})
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to get user: %w", err)
|
||||
}
|
||||
|
||||
if user == nil {
|
||||
return false, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
// Check if two-factor is enabled
|
||||
if enabled, ok := user["two_factor_enabled"].(bool); !ok || !enabled {
|
||||
return false, fmt.Errorf("two-factor authentication is not enabled")
|
||||
}
|
||||
|
||||
// Get TOTP parameters
|
||||
secret, _ := user["two_factor_secret"].(string)
|
||||
algorithm, _ := user["two_factor_algorithm"].(string)
|
||||
digits, _ := user["two_factor_digits"].(int)
|
||||
period, _ := user["two_factor_period"].(int)
|
||||
|
||||
// Set defaults
|
||||
if algorithm == "" {
|
||||
algorithm = "SHA1"
|
||||
}
|
||||
if digits == 0 {
|
||||
digits = 6
|
||||
}
|
||||
if period == 0 {
|
||||
period = 30
|
||||
}
|
||||
|
||||
// Verify code
|
||||
verified := u.verifyTOTPWithSecret(secret, code, algorithm, digits, period)
|
||||
|
||||
if verified {
|
||||
// Update last verified time
|
||||
m.Update(userID, map[string]interface{}{
|
||||
"two_factor_last_verified_at": time.Now(),
|
||||
})
|
||||
|
||||
// Clear user cache
|
||||
if u.cache != nil {
|
||||
u.cache.Del(u.cacheKey(userID))
|
||||
}
|
||||
}
|
||||
|
||||
return verified, nil
|
||||
}
|
||||
|
||||
// GenerateRecoveryCodes generates new recovery codes for user
|
||||
func (u *DefaultUser) GenerateRecoveryCodes(ctx context.Context, userID string) ([]string, error) {
|
||||
// Generate new recovery codes
|
||||
recoveryCodes, err := u.generateRecoveryCodesList()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate recovery codes: %w", err)
|
||||
}
|
||||
|
||||
// Update user record
|
||||
m := model.Select(u.model)
|
||||
err = m.Update(userID, map[string]interface{}{
|
||||
"two_factor_recovery_codes": recoveryCodes,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to update recovery codes: %w", err)
|
||||
}
|
||||
|
||||
// Clear user cache
|
||||
if u.cache != nil {
|
||||
u.cache.Del(u.cacheKey(userID))
|
||||
}
|
||||
|
||||
// Convert to string slice for return
|
||||
result := make([]string, len(recoveryCodes))
|
||||
for i, code := range recoveryCodes {
|
||||
result[i] = code.(string)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// VerifyRecoveryCode verifies and consumes a recovery code
|
||||
func (u *DefaultUser) VerifyRecoveryCode(ctx context.Context, userID string, code string) (bool, error) {
|
||||
// Get user data
|
||||
m := model.Select(u.model)
|
||||
user, err := m.Find(userID, model.QueryParam{
|
||||
Select: []interface{}{"two_factor_enabled", "two_factor_recovery_codes"},
|
||||
})
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to get user: %w", err)
|
||||
}
|
||||
|
||||
if user == nil {
|
||||
return false, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
// Check if two-factor is enabled
|
||||
if enabled, ok := user["two_factor_enabled"].(bool); !ok || !enabled {
|
||||
return false, fmt.Errorf("two-factor authentication is not enabled")
|
||||
}
|
||||
|
||||
// Get recovery codes
|
||||
recoveryCodes, ok := user["two_factor_recovery_codes"].([]interface{})
|
||||
if !ok {
|
||||
return false, fmt.Errorf("no recovery codes found")
|
||||
}
|
||||
|
||||
// Find and remove the used code
|
||||
var newRecoveryCodes []interface{}
|
||||
found := false
|
||||
for _, rc := range recoveryCodes {
|
||||
if rcStr, ok := rc.(string); ok && rcStr == code {
|
||||
found = true
|
||||
// Don't add this code to the new list (consume it)
|
||||
} else {
|
||||
newRecoveryCodes = append(newRecoveryCodes, rc)
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// Update user record with remaining codes
|
||||
err = m.Update(userID, map[string]interface{}{
|
||||
"two_factor_recovery_codes": newRecoveryCodes,
|
||||
"two_factor_last_verified_at": time.Now(),
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to update recovery codes: %w", err)
|
||||
}
|
||||
|
||||
// Clear user cache
|
||||
if u.cache != nil {
|
||||
u.cache.Del(u.cacheKey(userID))
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// Helper methods for TOTP
|
||||
|
||||
// generateQRCodeURL generates a QR code URL for TOTP setup
|
||||
func (u *DefaultUser) generateQRCodeURL(secret, issuer, accountName string) string {
|
||||
// Build the otpauth URL
|
||||
params := url.Values{}
|
||||
params.Set("secret", secret)
|
||||
params.Set("issuer", issuer)
|
||||
params.Set("algorithm", "SHA1")
|
||||
params.Set("digits", "6")
|
||||
params.Set("period", "30")
|
||||
|
||||
label := fmt.Sprintf("%s:%s", issuer, accountName)
|
||||
qrURL := fmt.Sprintf("otpauth://totp/%s?%s", url.QueryEscape(label), params.Encode())
|
||||
|
||||
return qrURL
|
||||
}
|
||||
|
||||
// generateRecoveryCodesList generates a list of recovery codes
|
||||
func (u *DefaultUser) generateRecoveryCodesList() ([]interface{}, error) {
|
||||
codes := make([]interface{}, 10) // Generate 10 recovery codes
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
// Generate 8-character recovery code
|
||||
code := make([]byte, 8)
|
||||
if _, err := rand.Read(code); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Convert to hex string
|
||||
codeStr := fmt.Sprintf("%x", code)
|
||||
codes[i] = codeStr
|
||||
}
|
||||
|
||||
return codes, nil
|
||||
}
|
||||
|
||||
// verifyTOTPWithSecret verifies a TOTP code with given parameters
|
||||
func (u *DefaultUser) verifyTOTPWithSecret(secret, code, algorithm string, digits, period int) bool {
|
||||
// Decode secret
|
||||
secretBytes, err := base32.StdEncoding.DecodeString(secret)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// Get current time
|
||||
now := time.Now().Unix()
|
||||
|
||||
// Check current time window and previous/next windows for clock skew
|
||||
for i := -1; i <= 1; i++ {
|
||||
timeCounter := (now + int64(i*period)) / int64(period)
|
||||
expectedCode := u.generateTOTPCode(secretBytes, timeCounter, algorithm, digits)
|
||||
|
||||
if expectedCode == code {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// generateTOTPCode generates a TOTP code
|
||||
func (u *DefaultUser) generateTOTPCode(secret []byte, timeCounter int64, algorithm string, digits int) string {
|
||||
// Convert time counter to byte array
|
||||
buf := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(buf, uint64(timeCounter))
|
||||
|
||||
// Choose hash algorithm
|
||||
var h hash.Hash
|
||||
switch algorithm {
|
||||
case "SHA256":
|
||||
h = sha256.New()
|
||||
case "SHA512":
|
||||
h = sha512.New()
|
||||
default:
|
||||
h = sha1.New()
|
||||
}
|
||||
|
||||
// HMAC
|
||||
for i := 0; i < len(secret); i++ {
|
||||
h.Write([]byte{secret[i] ^ 0x36})
|
||||
}
|
||||
for i := len(secret); i < h.BlockSize(); i++ {
|
||||
h.Write([]byte{0x36})
|
||||
}
|
||||
h.Write(buf)
|
||||
innerHash := h.Sum(nil)
|
||||
|
||||
h.Reset()
|
||||
for i := 0; i < len(secret); i++ {
|
||||
h.Write([]byte{secret[i] ^ 0x5c})
|
||||
}
|
||||
for i := len(secret); i < h.BlockSize(); i++ {
|
||||
h.Write([]byte{0x5c})
|
||||
}
|
||||
h.Write(innerHash)
|
||||
hmacHash := h.Sum(nil)
|
||||
|
||||
// Dynamic truncation
|
||||
offset := hmacHash[len(hmacHash)-1] & 0x0f
|
||||
binCode := binary.BigEndian.Uint32(hmacHash[offset:offset+4]) & 0x7fffffff
|
||||
|
||||
// Generate digits
|
||||
code := binCode % uint32(math.Pow10(digits))
|
||||
|
||||
return fmt.Sprintf("%0*d", digits, code)
|
||||
}
|
||||
|
||||
// GetUserForAuth retrieves user information for authentication purposes (internal use only)
|
||||
// This method includes sensitive fields like password_hash and should not be exposed to external APIs
|
||||
func (u *DefaultUser) GetUserForAuth(ctx context.Context, identifier string, identifierType string) (interface{}, error) {
|
||||
// Get user from database using the model
|
||||
m := model.Select(u.model)
|
||||
|
||||
var column string
|
||||
switch identifierType {
|
||||
case "username":
|
||||
column = "username"
|
||||
case "email":
|
||||
column = "email"
|
||||
case "subject":
|
||||
column = "subject"
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid identifier type: %s", identifierType)
|
||||
}
|
||||
|
||||
user, err := m.Get(model.QueryParam{
|
||||
Select: AuthUserFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: column, Value: identifier},
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get user for auth: %w", err)
|
||||
}
|
||||
|
||||
if len(user) == 0 {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
return user[0], nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -2,6 +2,7 @@ package user
|
|||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
|
|
@ -11,60 +12,324 @@ import (
|
|||
|
||||
// GetRole retrieves role information by role_id
|
||||
func (u *DefaultUser) GetRole(ctx context.Context, roleID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
m := model.Select(u.roleModel)
|
||||
roles, err := m.Get(model.QueryParam{
|
||||
Select: u.roleFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
if len(roles) == 0 {
|
||||
return nil, fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
return roles[0], nil
|
||||
}
|
||||
|
||||
// CreateRole creates a new user role
|
||||
func (u *DefaultUser) CreateRole(ctx context.Context, roleData maps.MapStrAny) (interface{}, error) {
|
||||
// TODO: implement - role_id should be provided in roleData
|
||||
return nil, nil
|
||||
// Validate required role_id field
|
||||
if _, exists := roleData["role_id"]; !exists {
|
||||
return nil, fmt.Errorf("role_id is required in roleData")
|
||||
}
|
||||
|
||||
// Set default values if not provided
|
||||
if _, exists := roleData["is_active"]; !exists {
|
||||
roleData["is_active"] = true
|
||||
}
|
||||
if _, exists := roleData["is_default"]; !exists {
|
||||
roleData["is_default"] = false
|
||||
}
|
||||
if _, exists := roleData["is_system"]; !exists {
|
||||
roleData["is_system"] = false
|
||||
}
|
||||
if _, exists := roleData["level"]; !exists {
|
||||
roleData["level"] = 0
|
||||
}
|
||||
if _, exists := roleData["sort_order"]; !exists {
|
||||
roleData["sort_order"] = 0
|
||||
}
|
||||
|
||||
m := model.Select(u.roleModel)
|
||||
id, err := m.Create(roleData)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToCreateRole, err)
|
||||
}
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// UpdateRole updates an existing role
|
||||
func (u *DefaultUser) UpdateRole(ctx context.Context, roleID string, roleData maps.MapStrAny) error {
|
||||
// TODO: implement
|
||||
// Remove sensitive fields that should not be updated directly
|
||||
sensitiveFields := []string{"id", "role_id", "created_at"}
|
||||
for _, field := range sensitiveFields {
|
||||
delete(roleData, field)
|
||||
}
|
||||
|
||||
// Skip update if no valid fields remain
|
||||
if len(roleData) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
m := model.Select(u.roleModel)
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, roleData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateRole, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteRole soft deletes a role (if not system role)
|
||||
func (u *DefaultUser) DeleteRole(ctx context.Context, roleID string) error {
|
||||
// TODO: implement
|
||||
// First check if role exists and is not a system role
|
||||
m := model.Select(u.roleModel)
|
||||
roles, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"id", "role_id", "is_system"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
if len(roles) == 0 {
|
||||
return fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
role := roles[0]
|
||||
// Check if this is a system role
|
||||
if isSystem, ok := role["is_system"].(bool); ok && isSystem {
|
||||
return fmt.Errorf("cannot delete system role: %s", roleID)
|
||||
}
|
||||
// Handle different boolean types from database
|
||||
if isSystemInt, ok := role["is_system"].(int64); ok && isSystemInt != 0 {
|
||||
return fmt.Errorf("cannot delete system role: %s", roleID)
|
||||
}
|
||||
|
||||
// Proceed with soft delete
|
||||
affected, err := m.DeleteWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is deleted
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToDeleteRole, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetRoles retrieves roles by query parameters
|
||||
func (u *DefaultUser) GetRoles(ctx context.Context, param model.QueryParam) ([]maps.MapStr, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// Set default select fields if not provided
|
||||
if param.Select == nil {
|
||||
param.Select = u.roleFields
|
||||
}
|
||||
|
||||
m := model.Select(u.roleModel)
|
||||
roles, err := m.Get(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
return roles, nil
|
||||
}
|
||||
|
||||
// PaginateRoles retrieves paginated list of roles
|
||||
func (u *DefaultUser) PaginateRoles(ctx context.Context, param model.QueryParam, page int, pagesize int) (maps.MapStr, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// Set default select fields if not provided
|
||||
if param.Select == nil {
|
||||
param.Select = u.roleFields
|
||||
}
|
||||
|
||||
m := model.Select(u.roleModel)
|
||||
result, err := m.Paginate(param, page, pagesize)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// CountRoles returns total count of roles with optional filters
|
||||
func (u *DefaultUser) CountRoles(ctx context.Context, param model.QueryParam) (int64, error) {
|
||||
// TODO: implement
|
||||
return 0, nil
|
||||
// Use Paginate with a small page size to get the total count
|
||||
// This is more reliable than manual COUNT(*) queries
|
||||
m := model.Select(u.roleModel)
|
||||
result, err := m.Paginate(param, 1, 1) // Get first page with 1 item to get total
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
// Extract total from pagination result
|
||||
if total, ok := result["total"].(int64); ok {
|
||||
return total, nil
|
||||
}
|
||||
|
||||
// Handle different total types returned by Paginate
|
||||
if totalInterface, ok := result["total"]; ok {
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
return int64(v), nil
|
||||
case int32:
|
||||
return int64(v), nil
|
||||
case int64:
|
||||
return v, nil
|
||||
case uint:
|
||||
return int64(v), nil
|
||||
case uint32:
|
||||
return int64(v), nil
|
||||
case uint64:
|
||||
return int64(v), nil
|
||||
default:
|
||||
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
|
||||
}
|
||||
}
|
||||
|
||||
return 0, fmt.Errorf("total not found in pagination result")
|
||||
}
|
||||
|
||||
// GetRolePermissions retrieves permissions for a role
|
||||
func (u *DefaultUser) GetRolePermissions(ctx context.Context, roleID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
m := model.Select(u.roleModel)
|
||||
roles, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"role_id", "permissions", "restricted_permissions"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
if len(roles) == 0 {
|
||||
return nil, fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
role := roles[0]
|
||||
permissions := maps.MapStrAny{
|
||||
"role_id": roleID,
|
||||
"permissions": role["permissions"],
|
||||
"restricted_permissions": role["restricted_permissions"],
|
||||
}
|
||||
|
||||
return permissions, nil
|
||||
}
|
||||
|
||||
// SetRolePermissions sets permissions for a role
|
||||
func (u *DefaultUser) SetRolePermissions(ctx context.Context, roleID string, permissions maps.MapStrAny) error {
|
||||
// TODO: implement
|
||||
// Prepare update data - only allow permission-related fields
|
||||
updateData := maps.MapStrAny{}
|
||||
|
||||
if perms, ok := permissions["permissions"]; ok {
|
||||
updateData["permissions"] = perms
|
||||
}
|
||||
|
||||
if restrictedPerms, ok := permissions["restricted_permissions"]; ok {
|
||||
updateData["restricted_permissions"] = restrictedPerms
|
||||
}
|
||||
|
||||
// Skip update if no permission fields provided
|
||||
if len(updateData) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
m := model.Select(u.roleModel)
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateRole, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateRolePermissions validates if role has specific permissions
|
||||
func (u *DefaultUser) ValidateRolePermissions(ctx context.Context, roleID string, requiredPermissions []string) (bool, error) {
|
||||
// TODO: implement
|
||||
return false, nil
|
||||
if len(requiredPermissions) == 0 {
|
||||
return true, nil // No permissions required
|
||||
}
|
||||
|
||||
// Get role permissions
|
||||
rolePermissions, err := u.GetRolePermissions(ctx, roleID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// Extract permissions and restricted permissions
|
||||
permissions, _ := rolePermissions["permissions"].(map[string]interface{})
|
||||
restrictedPermissions, _ := rolePermissions["restricted_permissions"].([]interface{})
|
||||
|
||||
// Convert restricted permissions to map for faster lookup
|
||||
restrictedMap := make(map[string]bool)
|
||||
for _, perm := range restrictedPermissions {
|
||||
if permStr, ok := perm.(string); ok {
|
||||
restrictedMap[permStr] = true
|
||||
}
|
||||
}
|
||||
|
||||
// Check each required permission
|
||||
for _, requiredPerm := range requiredPermissions {
|
||||
// First check if permission is explicitly restricted
|
||||
if restrictedMap[requiredPerm] {
|
||||
return false, nil // Permission is explicitly denied
|
||||
}
|
||||
|
||||
// Check if permission exists in granted permissions
|
||||
if permissions == nil {
|
||||
return false, nil // No permissions granted
|
||||
}
|
||||
|
||||
// Look for the permission in the permissions object
|
||||
// This is a simple implementation - in practice, you might want more sophisticated permission matching
|
||||
permValue, exists := permissions[requiredPerm]
|
||||
if !exists {
|
||||
return false, nil // Permission not found
|
||||
}
|
||||
|
||||
// Check if permission is enabled (assuming boolean values)
|
||||
if permBool, ok := permValue.(bool); ok && !permBool {
|
||||
return false, nil // Permission exists but is disabled
|
||||
}
|
||||
}
|
||||
|
||||
return true, nil // All required permissions are valid
|
||||
}
|
||||
|
|
|
|||
737
openapi/oauth/providers/user/role_test.go
Normal file
737
openapi/oauth/providers/user/role_test.go
Normal file
|
|
@ -0,0 +1,737 @@
|
|||
package user_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
)
|
||||
|
||||
// TestRoleData represents test role data structure
|
||||
type TestRoleData struct {
|
||||
RoleID string `json:"role_id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsDefault bool `json:"is_default"`
|
||||
IsSystem bool `json:"is_system"`
|
||||
Level int `json:"level"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
Color string `json:"color"`
|
||||
Icon string `json:"icon"`
|
||||
Permissions map[string]interface{} `json:"permissions"`
|
||||
Metadata map[string]interface{} `json:"metadata"`
|
||||
}
|
||||
|
||||
func TestRoleBasicOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] // 8 char UUID
|
||||
|
||||
// Create test role data dynamically
|
||||
testRole := &TestRoleData{
|
||||
RoleID: "testrole_" + testUUID,
|
||||
Name: "Test Role " + testUUID,
|
||||
Description: "Test role for unit testing " + testUUID,
|
||||
IsActive: true,
|
||||
IsDefault: false,
|
||||
IsSystem: false,
|
||||
Level: 10,
|
||||
SortOrder: 100,
|
||||
Color: "#007bff",
|
||||
Icon: "test-icon",
|
||||
Permissions: map[string]interface{}{
|
||||
"read": true,
|
||||
"write": true,
|
||||
"delete": false,
|
||||
},
|
||||
Metadata: map[string]interface{}{
|
||||
"source": "test",
|
||||
"uuid": testUUID,
|
||||
},
|
||||
}
|
||||
|
||||
// Test CreateRole
|
||||
t.Run("CreateRole", func(t *testing.T) {
|
||||
roleData := maps.MapStrAny{
|
||||
"role_id": testRole.RoleID,
|
||||
"name": testRole.Name,
|
||||
"description": testRole.Description,
|
||||
"level": testRole.Level,
|
||||
"sort_order": testRole.SortOrder,
|
||||
"color": testRole.Color,
|
||||
"icon": testRole.Icon,
|
||||
"permissions": testRole.Permissions,
|
||||
"metadata": testRole.Metadata,
|
||||
}
|
||||
|
||||
id, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, id)
|
||||
|
||||
// Verify default values were set
|
||||
assert.Equal(t, true, roleData["is_active"])
|
||||
assert.Equal(t, false, roleData["is_default"])
|
||||
assert.Equal(t, false, roleData["is_system"])
|
||||
// level should remain as provided (10), not be overridden
|
||||
})
|
||||
|
||||
// Test GetRole
|
||||
t.Run("GetRole", func(t *testing.T) {
|
||||
role, err := testProvider.GetRole(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, role)
|
||||
|
||||
// Verify key fields
|
||||
assert.Equal(t, testRole.RoleID, role["role_id"])
|
||||
assert.Equal(t, testRole.Name, role["name"])
|
||||
assert.Equal(t, testRole.Description, role["description"])
|
||||
assert.Equal(t, testRole.Color, role["color"])
|
||||
assert.Equal(t, testRole.Icon, role["icon"])
|
||||
|
||||
// Handle different boolean representations from database
|
||||
isActive := role["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
|
||||
assert.NotNil(t, role["created_at"])
|
||||
})
|
||||
|
||||
// Test UpdateRole
|
||||
t.Run("UpdateRole", func(t *testing.T) {
|
||||
updateData := maps.MapStrAny{
|
||||
"name": "Updated Test Role",
|
||||
"description": "Updated description for testing",
|
||||
"color": "#28a745",
|
||||
"icon": "updated-icon",
|
||||
"level": 20,
|
||||
"metadata": map[string]interface{}{
|
||||
"updated": true,
|
||||
"version": 2,
|
||||
},
|
||||
}
|
||||
|
||||
err := testProvider.UpdateRole(ctx, testRole.RoleID, updateData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify update
|
||||
role, err := testProvider.GetRole(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "Updated Test Role", role["name"])
|
||||
assert.Equal(t, "Updated description for testing", role["description"])
|
||||
assert.Equal(t, "#28a745", role["color"])
|
||||
assert.Equal(t, "updated-icon", role["icon"])
|
||||
|
||||
// Test updating sensitive fields (should be ignored)
|
||||
sensitiveData := maps.MapStrAny{
|
||||
"id": 999,
|
||||
"role_id": "malicious_role_id",
|
||||
"created_at": "2020-01-01T00:00:00Z",
|
||||
}
|
||||
|
||||
err = testProvider.UpdateRole(ctx, testRole.RoleID, sensitiveData)
|
||||
assert.NoError(t, err) // Should not error, just ignore sensitive fields
|
||||
|
||||
// Verify sensitive fields were not changed
|
||||
role, err = testProvider.GetRole(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, testRole.RoleID, role["role_id"]) // Should remain unchanged
|
||||
})
|
||||
|
||||
// Create a system role for delete test
|
||||
t.Run("CreateSystemRole", func(t *testing.T) {
|
||||
systemRoleData := maps.MapStrAny{
|
||||
"role_id": "systemrole_" + testUUID,
|
||||
"name": "System Role " + testUUID,
|
||||
"description": "System role for delete testing",
|
||||
"is_system": true,
|
||||
}
|
||||
|
||||
id, err := testProvider.CreateRole(ctx, systemRoleData)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, id)
|
||||
})
|
||||
|
||||
// Test DeleteRole - System Role Protection
|
||||
t.Run("DeleteRole_SystemRoleProtection", func(t *testing.T) {
|
||||
systemRoleID := "systemrole_" + testUUID
|
||||
err := testProvider.DeleteRole(ctx, systemRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "cannot delete system role")
|
||||
|
||||
// Verify system role still exists
|
||||
role, err := testProvider.GetRole(ctx, systemRoleID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, role)
|
||||
})
|
||||
|
||||
// Test DeleteRole - Normal Role (at the end)
|
||||
t.Run("DeleteRole", func(t *testing.T) {
|
||||
err := testProvider.DeleteRole(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify role was deleted
|
||||
_, err = testProvider.GetRole(ctx, testRole.RoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestRolePermissionOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Create a role for permission testing
|
||||
testRole := &TestRoleData{
|
||||
RoleID: "permrole_" + testUUID,
|
||||
Name: "Permission Test Role " + testUUID,
|
||||
Description: "Role for testing permissions",
|
||||
IsActive: true,
|
||||
Permissions: map[string]interface{}{
|
||||
"users.read": true,
|
||||
"users.write": true,
|
||||
"users.delete": false,
|
||||
"admin.access": true,
|
||||
},
|
||||
}
|
||||
|
||||
// Create role
|
||||
roleData := maps.MapStrAny{
|
||||
"role_id": testRole.RoleID,
|
||||
"name": testRole.Name,
|
||||
"description": testRole.Description,
|
||||
"permissions": testRole.Permissions,
|
||||
"restricted_permissions": []string{
|
||||
"system.config",
|
||||
"root.access",
|
||||
},
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test GetRolePermissions
|
||||
t.Run("GetRolePermissions", func(t *testing.T) {
|
||||
permissions, err := testProvider.GetRolePermissions(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, permissions)
|
||||
|
||||
assert.Equal(t, testRole.RoleID, permissions["role_id"])
|
||||
assert.NotNil(t, permissions["permissions"])
|
||||
assert.NotNil(t, permissions["restricted_permissions"])
|
||||
|
||||
// Verify permissions structure
|
||||
permsMap, ok := permissions["permissions"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, true, permsMap["users.read"])
|
||||
assert.Equal(t, true, permsMap["users.write"])
|
||||
assert.Equal(t, false, permsMap["users.delete"])
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetRolePermissions
|
||||
t.Run("SetRolePermissions", func(t *testing.T) {
|
||||
newPermissions := maps.MapStrAny{
|
||||
"permissions": map[string]interface{}{
|
||||
"users.read": true,
|
||||
"users.write": false, // Changed
|
||||
"users.delete": true, // Changed
|
||||
"posts.read": true, // New
|
||||
},
|
||||
"restricted_permissions": []string{
|
||||
"system.config",
|
||||
"dangerous.operation", // New restriction
|
||||
},
|
||||
}
|
||||
|
||||
err := testProvider.SetRolePermissions(ctx, testRole.RoleID, newPermissions)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify permissions were updated
|
||||
permissions, err := testProvider.GetRolePermissions(ctx, testRole.RoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
permsMap, ok := permissions["permissions"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, true, permsMap["users.read"])
|
||||
assert.Equal(t, false, permsMap["users.write"]) // Should be updated
|
||||
assert.Equal(t, true, permsMap["users.delete"]) // Should be updated
|
||||
assert.Equal(t, true, permsMap["posts.read"]) // Should be new
|
||||
}
|
||||
})
|
||||
|
||||
// Test ValidateRolePermissions
|
||||
t.Run("ValidateRolePermissions_ValidPermissions", func(t *testing.T) {
|
||||
requiredPermissions := []string{"users.read", "posts.read"}
|
||||
valid, err := testProvider.ValidateRolePermissions(ctx, testRole.RoleID, requiredPermissions)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
})
|
||||
|
||||
t.Run("ValidateRolePermissions_InvalidPermissions", func(t *testing.T) {
|
||||
requiredPermissions := []string{"users.write"} // This was set to false
|
||||
valid, err := testProvider.ValidateRolePermissions(ctx, testRole.RoleID, requiredPermissions)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should be false because users.write is disabled
|
||||
})
|
||||
|
||||
t.Run("ValidateRolePermissions_RestrictedPermissions", func(t *testing.T) {
|
||||
requiredPermissions := []string{"system.config"} // This is in restricted list
|
||||
valid, err := testProvider.ValidateRolePermissions(ctx, testRole.RoleID, requiredPermissions)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should be false because it's restricted
|
||||
})
|
||||
|
||||
t.Run("ValidateRolePermissions_EmptyRequirements", func(t *testing.T) {
|
||||
requiredPermissions := []string{}
|
||||
valid, err := testProvider.ValidateRolePermissions(ctx, testRole.RoleID, requiredPermissions)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Should be true when no permissions required
|
||||
})
|
||||
|
||||
t.Run("ValidateRolePermissions_NonExistentPermission", func(t *testing.T) {
|
||||
requiredPermissions := []string{"nonexistent.permission"}
|
||||
valid, err := testProvider.ValidateRolePermissions(ctx, testRole.RoleID, requiredPermissions)
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should be false for nonexistent permissions
|
||||
})
|
||||
}
|
||||
|
||||
func TestRoleListOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create multiple test roles for list operations
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
testRoles := []TestRoleData{
|
||||
{
|
||||
RoleID: "listrole_" + testUUID + "_1",
|
||||
Name: "List Role 1",
|
||||
Description: "First role for list testing",
|
||||
IsActive: true,
|
||||
Level: 10,
|
||||
},
|
||||
{
|
||||
RoleID: "listrole_" + testUUID + "_2",
|
||||
Name: "List Role 2",
|
||||
Description: "Second role for list testing",
|
||||
IsActive: true,
|
||||
Level: 20,
|
||||
},
|
||||
{
|
||||
RoleID: "listrole_" + testUUID + "_3",
|
||||
Name: "List Role 3",
|
||||
Description: "Third role for list testing",
|
||||
IsActive: false, // Different status for filtering
|
||||
Level: 30,
|
||||
},
|
||||
{
|
||||
RoleID: "listrole_" + testUUID + "_4",
|
||||
Name: "List Role 4",
|
||||
Description: "Fourth role for list testing",
|
||||
IsActive: true,
|
||||
Level: 40,
|
||||
},
|
||||
{
|
||||
RoleID: "listrole_" + testUUID + "_5",
|
||||
Name: "List Role 5",
|
||||
Description: "Fifth role for list testing",
|
||||
IsActive: true,
|
||||
Level: 50,
|
||||
},
|
||||
}
|
||||
|
||||
// Create roles in database
|
||||
for _, roleData := range testRoles {
|
||||
roleMap := maps.MapStrAny{
|
||||
"role_id": roleData.RoleID,
|
||||
"name": roleData.Name,
|
||||
"description": roleData.Description,
|
||||
"is_active": roleData.IsActive,
|
||||
"level": roleData.Level,
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateRole(ctx, roleMap)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Test GetRoles
|
||||
t.Run("GetRoles_All", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
},
|
||||
}
|
||||
roles, err := testProvider.GetRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, len(roles), 5) // At least our 5 test roles
|
||||
|
||||
// Check that basic fields are returned by default
|
||||
if len(roles) > 0 {
|
||||
role := roles[0]
|
||||
assert.Contains(t, role, "role_id")
|
||||
assert.Contains(t, role, "name")
|
||||
assert.Contains(t, role, "description")
|
||||
assert.Contains(t, role, "is_active")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetRoles_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
roles, err := testProvider.GetRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, len(roles), 4) // At least 4 active roles
|
||||
|
||||
// All returned roles should be active
|
||||
for _, role := range roles {
|
||||
if strings.Contains(role["role_id"].(string), "listrole_"+testUUID+"_") {
|
||||
// Handle different boolean representations from database
|
||||
isActive := role["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetRoles_WithCustomFields", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Select: []interface{}{"role_id", "name", "is_active", "level"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
},
|
||||
Limit: 3,
|
||||
}
|
||||
roles, err := testProvider.GetRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.LessOrEqual(t, len(roles), 3) // Respects limit
|
||||
|
||||
if len(roles) > 0 {
|
||||
role := roles[0]
|
||||
assert.Contains(t, role, "role_id")
|
||||
assert.Contains(t, role, "name")
|
||||
assert.Contains(t, role, "is_active")
|
||||
assert.Contains(t, role, "level")
|
||||
}
|
||||
})
|
||||
|
||||
// Test PaginateRoles
|
||||
t.Run("PaginateRoles_FirstPage", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
},
|
||||
Orders: []model.QueryOrder{
|
||||
{Column: "level", Option: "asc"},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateRoles(ctx, param, 1, 3)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
// Check pagination structure
|
||||
assert.Contains(t, result, "data")
|
||||
assert.Contains(t, result, "total")
|
||||
assert.Contains(t, result, "page")
|
||||
assert.Contains(t, result, "pagesize")
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.LessOrEqual(t, len(data), 3) // Page size limit
|
||||
|
||||
// Handle different total types
|
||||
totalInterface, exists := result["total"]
|
||||
assert.True(t, exists)
|
||||
|
||||
var total int64
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
total = int64(v)
|
||||
case int32:
|
||||
total = int64(v)
|
||||
case int64:
|
||||
total = v
|
||||
case uint:
|
||||
total = int64(v)
|
||||
case uint32:
|
||||
total = int64(v)
|
||||
case uint64:
|
||||
total = int64(v)
|
||||
default:
|
||||
t.Errorf("unexpected total type: %T, value: %v", totalInterface, totalInterface)
|
||||
}
|
||||
assert.GreaterOrEqual(t, total, int64(5)) // At least 5 roles
|
||||
|
||||
assert.Equal(t, 1, result["page"])
|
||||
assert.Equal(t, 3, result["pagesize"])
|
||||
})
|
||||
|
||||
t.Run("PaginateRoles_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateRoles(ctx, param, 1, 10)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.GreaterOrEqual(t, len(data), 4) // At least 4 active roles
|
||||
|
||||
// Verify is_active filter works
|
||||
for _, role := range data {
|
||||
if strings.Contains(role["role_id"].(string), "listrole_"+testUUID+"_") {
|
||||
// Handle different boolean representations from database
|
||||
isActive := role["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Test CountRoles
|
||||
t.Run("CountRoles_All", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(5)) // At least 5 roles
|
||||
})
|
||||
|
||||
t.Run("CountRoles_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(4)) // At least 4 active roles
|
||||
})
|
||||
|
||||
t.Run("CountRoles_SpecificLevel", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: "listrole_" + testUUID + "_%"},
|
||||
{Column: "level", OP: ">=", Value: 30},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
// We created 3 roles with level >= 30 (30, 40, 50), but be flexible with database state
|
||||
assert.GreaterOrEqual(t, count, int64(1)) // At least 1 role with level >= 30
|
||||
assert.LessOrEqual(t, count, int64(5)) // But not more than 5 (our total test roles)
|
||||
})
|
||||
|
||||
t.Run("CountRoles_NoResults", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: "nonexistent_role_id"},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, int64(0), count)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRoleErrorHandling(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
nonExistentRoleID := "nonexistent_role_" + testUUID
|
||||
|
||||
t.Run("GetRole_NotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetRole(ctx, nonExistentRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("CreateRole_MissingRoleID", func(t *testing.T) {
|
||||
roleData := maps.MapStrAny{
|
||||
"name": "Test Role",
|
||||
"description": "Role without role_id",
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role_id is required")
|
||||
})
|
||||
|
||||
t.Run("UpdateRole_NotFound", func(t *testing.T) {
|
||||
updateData := maps.MapStrAny{"name": "Test"}
|
||||
err := testProvider.UpdateRole(ctx, nonExistentRoleID, updateData)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("DeleteRole_NotFound", func(t *testing.T) {
|
||||
err := testProvider.DeleteRole(ctx, nonExistentRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("GetRolePermissions_NotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetRolePermissions(ctx, nonExistentRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("SetRolePermissions_NotFound", func(t *testing.T) {
|
||||
permissions := maps.MapStrAny{
|
||||
"permissions": map[string]interface{}{"test": true},
|
||||
}
|
||||
err := testProvider.SetRolePermissions(ctx, nonExistentRoleID, permissions)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("ValidateRolePermissions_NotFound", func(t *testing.T) {
|
||||
requiredPermissions := []string{"test.permission"}
|
||||
_, err := testProvider.ValidateRolePermissions(ctx, nonExistentRoleID, requiredPermissions)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("GetRoles_EmptyResult", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: nonExistentRoleID},
|
||||
},
|
||||
}
|
||||
roles, err := testProvider.GetRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 0, len(roles)) // Empty slice, not nil
|
||||
})
|
||||
|
||||
t.Run("PaginateRoles_EmptyResult", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: nonExistentRoleID},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateRoles(ctx, param, 1, 10)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, 0, len(data))
|
||||
|
||||
// Handle different total types
|
||||
totalInterface, exists := result["total"]
|
||||
assert.True(t, exists)
|
||||
|
||||
var total int64
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
total = int64(v)
|
||||
case int32:
|
||||
total = int64(v)
|
||||
case int64:
|
||||
total = v
|
||||
case uint:
|
||||
total = int64(v)
|
||||
case uint32:
|
||||
total = int64(v)
|
||||
case uint64:
|
||||
total = int64(v)
|
||||
default:
|
||||
t.Errorf("unexpected total type: %T, value: %v", totalInterface, totalInterface)
|
||||
}
|
||||
assert.Equal(t, int64(0), total)
|
||||
})
|
||||
|
||||
t.Run("UpdateRole_EmptyData", func(t *testing.T) {
|
||||
// First create a role for this test
|
||||
testRoleID := "emptyupdate_" + testUUID
|
||||
roleData := maps.MapStrAny{
|
||||
"role_id": testRoleID,
|
||||
"name": "Test Role for Empty Update",
|
||||
}
|
||||
_, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test with empty update data (should not error, just do nothing)
|
||||
emptyData := maps.MapStrAny{}
|
||||
err = testProvider.UpdateRole(ctx, testRoleID, emptyData)
|
||||
assert.NoError(t, err) // Should not error, just skip update
|
||||
})
|
||||
|
||||
t.Run("SetRolePermissions_EmptyData", func(t *testing.T) {
|
||||
// First create a role for this test
|
||||
testRoleID := "emptyperm_" + testUUID
|
||||
roleData := maps.MapStrAny{
|
||||
"role_id": testRoleID,
|
||||
"name": "Test Role for Empty Permissions",
|
||||
}
|
||||
_, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test with empty permission data (should not error, just do nothing)
|
||||
emptyData := maps.MapStrAny{}
|
||||
err = testProvider.SetRolePermissions(ctx, testRoleID, emptyData)
|
||||
assert.NoError(t, err) // Should not error, just skip update
|
||||
})
|
||||
|
||||
t.Run("CountRoles_ComplexFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "is_active", Value: true},
|
||||
{Column: "level", OP: ">=", Value: 10},
|
||||
{Column: "is_system", Value: false},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountRoles(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(0)) // Should handle complex filters without error
|
||||
})
|
||||
}
|
||||
|
|
@ -2,6 +2,7 @@ package user
|
|||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
|
|
@ -11,54 +12,277 @@ import (
|
|||
|
||||
// GetType retrieves type information by type_id
|
||||
func (u *DefaultUser) GetType(ctx context.Context, typeID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
m := model.Select(u.typeModel)
|
||||
types, err := m.Get(model.QueryParam{
|
||||
Select: u.typeFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
if len(types) == 0 {
|
||||
return nil, fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
return types[0], nil
|
||||
}
|
||||
|
||||
// CreateType creates a new user type
|
||||
func (u *DefaultUser) CreateType(ctx context.Context, typeData maps.MapStrAny) (interface{}, error) {
|
||||
// TODO: implement - type_id should be provided in typeData
|
||||
return nil, nil
|
||||
// Validate required type_id field
|
||||
if _, exists := typeData["type_id"]; !exists {
|
||||
return nil, fmt.Errorf("type_id is required in typeData")
|
||||
}
|
||||
|
||||
// Set default values if not provided
|
||||
if _, exists := typeData["is_active"]; !exists {
|
||||
typeData["is_active"] = true
|
||||
}
|
||||
if _, exists := typeData["is_default"]; !exists {
|
||||
typeData["is_default"] = false
|
||||
}
|
||||
if _, exists := typeData["sort_order"]; !exists {
|
||||
typeData["sort_order"] = 0
|
||||
}
|
||||
if _, exists := typeData["max_sessions"]; !exists {
|
||||
typeData["max_sessions"] = nil // Allow unlimited sessions by default
|
||||
}
|
||||
if _, exists := typeData["session_timeout"]; !exists {
|
||||
typeData["session_timeout"] = 0 // No timeout by default
|
||||
}
|
||||
|
||||
m := model.Select(u.typeModel)
|
||||
id, err := m.Create(typeData)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToCreateType, err)
|
||||
}
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// UpdateType updates an existing type
|
||||
func (u *DefaultUser) UpdateType(ctx context.Context, typeID string, typeData maps.MapStrAny) error {
|
||||
// TODO: implement
|
||||
// Remove sensitive fields that should not be updated directly
|
||||
sensitiveFields := []string{"id", "type_id", "created_at"}
|
||||
for _, field := range sensitiveFields {
|
||||
delete(typeData, field)
|
||||
}
|
||||
|
||||
// Skip update if no valid fields remain
|
||||
if len(typeData) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
m := model.Select(u.typeModel)
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, typeData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateType, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteType soft deletes a type
|
||||
func (u *DefaultUser) DeleteType(ctx context.Context, typeID string) error {
|
||||
// TODO: implement
|
||||
// First check if type exists
|
||||
m := model.Select(u.typeModel)
|
||||
types, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"id", "type_id"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
if len(types) == 0 {
|
||||
return fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
// Proceed with soft delete
|
||||
affected, err := m.DeleteWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is deleted
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToDeleteType, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTypes retrieves types by query parameters
|
||||
func (u *DefaultUser) GetTypes(ctx context.Context, param model.QueryParam) ([]maps.MapStr, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// Set default select fields if not provided
|
||||
if param.Select == nil {
|
||||
param.Select = u.typeFields
|
||||
}
|
||||
|
||||
m := model.Select(u.typeModel)
|
||||
types, err := m.Get(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
return types, nil
|
||||
}
|
||||
|
||||
// PaginateTypes retrieves paginated list of types
|
||||
func (u *DefaultUser) PaginateTypes(ctx context.Context, param model.QueryParam, page int, pagesize int) (maps.MapStr, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// Set default select fields if not provided
|
||||
if param.Select == nil {
|
||||
param.Select = u.typeFields
|
||||
}
|
||||
|
||||
m := model.Select(u.typeModel)
|
||||
result, err := m.Paginate(param, page, pagesize)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// CountTypes returns total count of types with optional filters
|
||||
func (u *DefaultUser) CountTypes(ctx context.Context, param model.QueryParam) (int64, error) {
|
||||
// TODO: implement
|
||||
return 0, nil
|
||||
// Use Paginate with a small page size to get the total count
|
||||
// This is more reliable than manual COUNT(*) queries
|
||||
m := model.Select(u.typeModel)
|
||||
result, err := m.Paginate(param, 1, 1) // Get first page with 1 item to get total
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
// Extract total from pagination result
|
||||
if total, ok := result["total"].(int64); ok {
|
||||
return total, nil
|
||||
}
|
||||
|
||||
// Handle different total types returned by Paginate
|
||||
if totalInterface, ok := result["total"]; ok {
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
return int64(v), nil
|
||||
case int32:
|
||||
return int64(v), nil
|
||||
case int64:
|
||||
return v, nil
|
||||
case uint:
|
||||
return int64(v), nil
|
||||
case uint32:
|
||||
return int64(v), nil
|
||||
case uint64:
|
||||
return int64(v), nil
|
||||
default:
|
||||
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
|
||||
}
|
||||
}
|
||||
|
||||
return 0, fmt.Errorf("total not found in pagination result")
|
||||
}
|
||||
|
||||
// GetTypeConfiguration retrieves configuration for a type (schema, features, limits, etc.)
|
||||
func (u *DefaultUser) GetTypeConfiguration(ctx context.Context, typeID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
m := model.Select(u.typeModel)
|
||||
types, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"type_id", "schema", "features", "limits", "password_policy", "metadata"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
if len(types) == 0 {
|
||||
return nil, fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
typeRecord := types[0]
|
||||
config := maps.MapStrAny{
|
||||
"type_id": typeID,
|
||||
"schema": typeRecord["schema"],
|
||||
"features": typeRecord["features"],
|
||||
"limits": typeRecord["limits"],
|
||||
"password_policy": typeRecord["password_policy"],
|
||||
"metadata": typeRecord["metadata"],
|
||||
}
|
||||
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// SetTypeConfiguration sets configuration for a type
|
||||
func (u *DefaultUser) SetTypeConfiguration(ctx context.Context, typeID string, config maps.MapStrAny) error {
|
||||
// TODO: implement
|
||||
// Prepare update data - only allow configuration-related fields
|
||||
updateData := maps.MapStrAny{}
|
||||
|
||||
if schema, ok := config["schema"]; ok {
|
||||
updateData["schema"] = schema
|
||||
}
|
||||
|
||||
if features, ok := config["features"]; ok {
|
||||
updateData["features"] = features
|
||||
}
|
||||
|
||||
if limits, ok := config["limits"]; ok {
|
||||
updateData["limits"] = limits
|
||||
}
|
||||
|
||||
if passwordPolicy, ok := config["password_policy"]; ok {
|
||||
updateData["password_policy"] = passwordPolicy
|
||||
}
|
||||
|
||||
if metadata, ok := config["metadata"]; ok {
|
||||
updateData["metadata"] = metadata
|
||||
}
|
||||
|
||||
// Skip update if no configuration fields provided
|
||||
if len(updateData) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
m := model.Select(u.typeModel)
|
||||
affected, err := m.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateType, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
812
openapi/oauth/providers/user/type_test.go
Normal file
812
openapi/oauth/providers/user/type_test.go
Normal file
|
|
@ -0,0 +1,812 @@
|
|||
package user_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
)
|
||||
|
||||
// TestTypeData represents test type data structure
|
||||
type TestTypeData struct {
|
||||
TypeID string `json:"type_id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsDefault bool `json:"is_default"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
DefaultRoleID string `json:"default_role_id"`
|
||||
MaxSessions *int `json:"max_sessions"`
|
||||
SessionTimeout int `json:"session_timeout"`
|
||||
Schema map[string]interface{} `json:"schema"`
|
||||
Features map[string]interface{} `json:"features"`
|
||||
Limits map[string]interface{} `json:"limits"`
|
||||
PasswordPolicy map[string]interface{} `json:"password_policy"`
|
||||
Metadata map[string]interface{} `json:"metadata"`
|
||||
}
|
||||
|
||||
func TestTypeBasicOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] // 8 char UUID
|
||||
|
||||
// Create test type data dynamically
|
||||
maxSessions := 5
|
||||
testType := &TestTypeData{
|
||||
TypeID: "testtype_" + testUUID,
|
||||
Name: "Test Type " + testUUID,
|
||||
Description: "Test type for unit testing " + testUUID,
|
||||
IsActive: true,
|
||||
IsDefault: false,
|
||||
SortOrder: 100,
|
||||
DefaultRoleID: "user",
|
||||
MaxSessions: &maxSessions,
|
||||
SessionTimeout: 3600,
|
||||
Schema: map[string]interface{}{
|
||||
"version": "1.0",
|
||||
"fields": map[string]interface{}{
|
||||
"profile": map[string]interface{}{
|
||||
"required": true,
|
||||
"type": "object",
|
||||
},
|
||||
},
|
||||
},
|
||||
Features: map[string]interface{}{
|
||||
"mfa_enabled": true,
|
||||
"api_access": true,
|
||||
"export_data": false,
|
||||
"custom_branding": true,
|
||||
},
|
||||
Limits: map[string]interface{}{
|
||||
"storage_mb": 1024,
|
||||
"api_calls_day": 10000,
|
||||
"team_members": 50,
|
||||
"projects": 10,
|
||||
},
|
||||
PasswordPolicy: map[string]interface{}{
|
||||
"min_length": 8,
|
||||
"require_uppercase": true,
|
||||
"require_lowercase": true,
|
||||
"require_numbers": true,
|
||||
"require_symbols": false,
|
||||
"max_age_days": 90,
|
||||
},
|
||||
Metadata: map[string]interface{}{
|
||||
"source": "test",
|
||||
"uuid": testUUID,
|
||||
"version": "1.0",
|
||||
},
|
||||
}
|
||||
|
||||
// Test CreateType
|
||||
t.Run("CreateType", func(t *testing.T) {
|
||||
typeData := maps.MapStrAny{
|
||||
"type_id": testType.TypeID,
|
||||
"name": testType.Name,
|
||||
"description": testType.Description,
|
||||
"sort_order": testType.SortOrder,
|
||||
"default_role_id": testType.DefaultRoleID,
|
||||
"max_sessions": testType.MaxSessions,
|
||||
"session_timeout": testType.SessionTimeout,
|
||||
"schema": testType.Schema,
|
||||
"features": testType.Features,
|
||||
"limits": testType.Limits,
|
||||
"password_policy": testType.PasswordPolicy,
|
||||
"metadata": testType.Metadata,
|
||||
}
|
||||
|
||||
id, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, id)
|
||||
|
||||
// Verify default values were set
|
||||
assert.Equal(t, true, typeData["is_active"])
|
||||
assert.Equal(t, false, typeData["is_default"])
|
||||
// sort_order, max_sessions, session_timeout should remain as provided
|
||||
})
|
||||
|
||||
// Test GetType
|
||||
t.Run("GetType", func(t *testing.T) {
|
||||
typeRecord, err := testProvider.GetType(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, typeRecord)
|
||||
|
||||
// Verify key fields
|
||||
assert.Equal(t, testType.TypeID, typeRecord["type_id"])
|
||||
assert.Equal(t, testType.Name, typeRecord["name"])
|
||||
assert.Equal(t, testType.Description, typeRecord["description"])
|
||||
assert.Equal(t, testType.DefaultRoleID, typeRecord["default_role_id"])
|
||||
|
||||
// Handle different boolean representations from database
|
||||
isActive := typeRecord["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
|
||||
assert.NotNil(t, typeRecord["created_at"])
|
||||
})
|
||||
|
||||
// Test UpdateType
|
||||
t.Run("UpdateType", func(t *testing.T) {
|
||||
newMaxSessions := 10
|
||||
updateData := maps.MapStrAny{
|
||||
"name": "Updated Test Type",
|
||||
"description": "Updated description for testing",
|
||||
"sort_order": 200,
|
||||
"default_role_id": "admin",
|
||||
"max_sessions": &newMaxSessions,
|
||||
"session_timeout": 7200,
|
||||
"metadata": map[string]interface{}{
|
||||
"updated": true,
|
||||
"version": "2.0",
|
||||
},
|
||||
}
|
||||
|
||||
err := testProvider.UpdateType(ctx, testType.TypeID, updateData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify update
|
||||
typeRecord, err := testProvider.GetType(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "Updated Test Type", typeRecord["name"])
|
||||
assert.Equal(t, "Updated description for testing", typeRecord["description"])
|
||||
assert.Equal(t, "admin", typeRecord["default_role_id"])
|
||||
|
||||
// Test updating sensitive fields (should be ignored)
|
||||
sensitiveData := maps.MapStrAny{
|
||||
"id": 999,
|
||||
"type_id": "malicious_type_id",
|
||||
"created_at": "2020-01-01T00:00:00Z",
|
||||
}
|
||||
|
||||
err = testProvider.UpdateType(ctx, testType.TypeID, sensitiveData)
|
||||
assert.NoError(t, err) // Should not error, just ignore sensitive fields
|
||||
|
||||
// Verify sensitive fields were not changed
|
||||
typeRecord, err = testProvider.GetType(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, testType.TypeID, typeRecord["type_id"]) // Should remain unchanged
|
||||
})
|
||||
|
||||
// Test DeleteType (at the end)
|
||||
t.Run("DeleteType", func(t *testing.T) {
|
||||
err := testProvider.DeleteType(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify type was deleted
|
||||
_, err = testProvider.GetType(ctx, testType.TypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestTypeConfigurationOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Create a type for configuration testing
|
||||
testType := &TestTypeData{
|
||||
TypeID: "configtype_" + testUUID,
|
||||
Name: "Config Test Type " + testUUID,
|
||||
Description: "Type for testing configuration",
|
||||
IsActive: true,
|
||||
Schema: map[string]interface{}{
|
||||
"version": "1.0",
|
||||
"type": "premium",
|
||||
},
|
||||
Features: map[string]interface{}{
|
||||
"api_access": true,
|
||||
"advanced_reports": true,
|
||||
"custom_integrations": false,
|
||||
"scope_limits": []interface{}{
|
||||
"read", "write", "admin.read",
|
||||
},
|
||||
},
|
||||
Limits: map[string]interface{}{
|
||||
"storage_gb": 10,
|
||||
"users": 100,
|
||||
"api_calls": 50000,
|
||||
},
|
||||
PasswordPolicy: map[string]interface{}{
|
||||
"min_length": 12,
|
||||
"require_symbols": true,
|
||||
"history_count": 5,
|
||||
},
|
||||
Metadata: map[string]interface{}{
|
||||
"plan": "premium",
|
||||
"tier": 2,
|
||||
"features": "advanced",
|
||||
},
|
||||
}
|
||||
|
||||
// Create type
|
||||
typeData := maps.MapStrAny{
|
||||
"type_id": testType.TypeID,
|
||||
"name": testType.Name,
|
||||
"description": testType.Description,
|
||||
"schema": testType.Schema,
|
||||
"features": testType.Features,
|
||||
"limits": testType.Limits,
|
||||
"password_policy": testType.PasswordPolicy,
|
||||
"metadata": testType.Metadata,
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test GetTypeConfiguration
|
||||
t.Run("GetTypeConfiguration", func(t *testing.T) {
|
||||
config, err := testProvider.GetTypeConfiguration(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, config)
|
||||
|
||||
assert.Equal(t, testType.TypeID, config["type_id"])
|
||||
assert.NotNil(t, config["schema"])
|
||||
assert.NotNil(t, config["features"])
|
||||
assert.NotNil(t, config["limits"])
|
||||
assert.NotNil(t, config["password_policy"])
|
||||
assert.NotNil(t, config["metadata"])
|
||||
|
||||
// Verify schema structure
|
||||
schemaMap, ok := config["schema"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, "1.0", schemaMap["version"])
|
||||
assert.Equal(t, "premium", schemaMap["type"])
|
||||
}
|
||||
|
||||
// Verify features structure
|
||||
featuresMap, ok := config["features"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, true, featuresMap["api_access"])
|
||||
assert.Equal(t, true, featuresMap["advanced_reports"])
|
||||
assert.Equal(t, false, featuresMap["custom_integrations"])
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetTypeConfiguration
|
||||
t.Run("SetTypeConfiguration", func(t *testing.T) {
|
||||
newConfig := maps.MapStrAny{
|
||||
"schema": map[string]interface{}{
|
||||
"version": "2.0",
|
||||
"type": "enterprise", // Changed
|
||||
},
|
||||
"features": map[string]interface{}{
|
||||
"api_access": true,
|
||||
"advanced_reports": true,
|
||||
"custom_integrations": true, // Changed
|
||||
"white_label": true, // New
|
||||
"scope_limits": []interface{}{
|
||||
"read", "write", "admin.read", "admin.write", // Extended
|
||||
},
|
||||
},
|
||||
"limits": map[string]interface{}{
|
||||
"storage_gb": 50, // Increased
|
||||
"users": 500, // Increased
|
||||
"api_calls": 100000, // Increased
|
||||
},
|
||||
"password_policy": map[string]interface{}{
|
||||
"min_length": 16, // Increased
|
||||
"require_symbols": true,
|
||||
"history_count": 10, // Increased
|
||||
"complexity_score": 8, // New
|
||||
},
|
||||
"metadata": map[string]interface{}{
|
||||
"plan": "enterprise", // Changed
|
||||
"tier": 3, // Changed
|
||||
"features": "premium",
|
||||
"updated_by": "test", // New
|
||||
},
|
||||
}
|
||||
|
||||
err := testProvider.SetTypeConfiguration(ctx, testType.TypeID, newConfig)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify configuration was updated
|
||||
config, err := testProvider.GetTypeConfiguration(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify schema update
|
||||
schemaMap, ok := config["schema"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, "2.0", schemaMap["version"])
|
||||
assert.Equal(t, "enterprise", schemaMap["type"]) // Should be updated
|
||||
}
|
||||
|
||||
// Verify features update
|
||||
featuresMap, ok := config["features"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, true, featuresMap["custom_integrations"]) // Should be updated
|
||||
assert.Equal(t, true, featuresMap["white_label"]) // Should be new
|
||||
}
|
||||
|
||||
// Verify limits update
|
||||
limitsMap, ok := config["limits"].(map[string]interface{})
|
||||
if ok {
|
||||
// Handle different numeric types from database
|
||||
storageInterface := limitsMap["storage_gb"]
|
||||
switch v := storageInterface.(type) {
|
||||
case int:
|
||||
assert.Equal(t, 50, v)
|
||||
case int32:
|
||||
assert.Equal(t, int32(50), v)
|
||||
case int64:
|
||||
assert.Equal(t, int64(50), v)
|
||||
case float64:
|
||||
assert.Equal(t, float64(50), v)
|
||||
default:
|
||||
t.Errorf("unexpected storage_gb type: %T, value: %v", storageInterface, storageInterface)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetTypeConfiguration with partial data
|
||||
t.Run("SetTypeConfiguration_PartialUpdate", func(t *testing.T) {
|
||||
partialConfig := maps.MapStrAny{
|
||||
"metadata": map[string]interface{}{
|
||||
"plan": "enterprise",
|
||||
"tier": 3,
|
||||
"features": "premium",
|
||||
"updated": true, // New field
|
||||
"timestamp": "2024-01-01", // New field
|
||||
},
|
||||
}
|
||||
|
||||
err := testProvider.SetTypeConfiguration(ctx, testType.TypeID, partialConfig)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify only metadata was updated, other configs remain
|
||||
config, err := testProvider.GetTypeConfiguration(ctx, testType.TypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Schema should remain from previous update
|
||||
schemaMap, ok := config["schema"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, "2.0", schemaMap["version"])
|
||||
}
|
||||
|
||||
// Metadata should be updated
|
||||
metadataMap, ok := config["metadata"].(map[string]interface{})
|
||||
if ok {
|
||||
assert.Equal(t, true, metadataMap["updated"])
|
||||
assert.Equal(t, "2024-01-01", metadataMap["timestamp"])
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetTypeConfiguration with empty data (should not error)
|
||||
t.Run("SetTypeConfiguration_EmptyData", func(t *testing.T) {
|
||||
emptyConfig := maps.MapStrAny{}
|
||||
err := testProvider.SetTypeConfiguration(ctx, testType.TypeID, emptyConfig)
|
||||
assert.NoError(t, err) // Should not error, just skip update
|
||||
})
|
||||
}
|
||||
|
||||
func TestTypeListOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create multiple test types for list operations
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
testTypes := []TestTypeData{
|
||||
{
|
||||
TypeID: "listtype_" + testUUID + "_1",
|
||||
Name: "List Type 1",
|
||||
Description: "First type for list testing",
|
||||
IsActive: true,
|
||||
SortOrder: 10,
|
||||
},
|
||||
{
|
||||
TypeID: "listtype_" + testUUID + "_2",
|
||||
Name: "List Type 2",
|
||||
Description: "Second type for list testing",
|
||||
IsActive: true,
|
||||
SortOrder: 20,
|
||||
},
|
||||
{
|
||||
TypeID: "listtype_" + testUUID + "_3",
|
||||
Name: "List Type 3",
|
||||
Description: "Third type for list testing",
|
||||
IsActive: false, // Different status for filtering
|
||||
SortOrder: 30,
|
||||
},
|
||||
{
|
||||
TypeID: "listtype_" + testUUID + "_4",
|
||||
Name: "List Type 4",
|
||||
Description: "Fourth type for list testing",
|
||||
IsActive: true,
|
||||
SortOrder: 40,
|
||||
},
|
||||
{
|
||||
TypeID: "listtype_" + testUUID + "_5",
|
||||
Name: "List Type 5",
|
||||
Description: "Fifth type for list testing",
|
||||
IsActive: true,
|
||||
SortOrder: 50,
|
||||
},
|
||||
}
|
||||
|
||||
// Create types in database
|
||||
for _, typeData := range testTypes {
|
||||
typeMap := maps.MapStrAny{
|
||||
"type_id": typeData.TypeID,
|
||||
"name": typeData.Name,
|
||||
"description": typeData.Description,
|
||||
"is_active": typeData.IsActive,
|
||||
"sort_order": typeData.SortOrder,
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateType(ctx, typeMap)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Test GetTypes
|
||||
t.Run("GetTypes_All", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
},
|
||||
}
|
||||
types, err := testProvider.GetTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, len(types), 5) // At least our 5 test types
|
||||
|
||||
// Check that basic fields are returned by default
|
||||
if len(types) > 0 {
|
||||
typeRecord := types[0]
|
||||
assert.Contains(t, typeRecord, "type_id")
|
||||
assert.Contains(t, typeRecord, "name")
|
||||
assert.Contains(t, typeRecord, "description")
|
||||
assert.Contains(t, typeRecord, "is_active")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetTypes_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
types, err := testProvider.GetTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, len(types), 4) // At least 4 active types
|
||||
|
||||
// All returned types should be active
|
||||
for _, typeRecord := range types {
|
||||
if strings.Contains(typeRecord["type_id"].(string), "listtype_"+testUUID+"_") {
|
||||
// Handle different boolean representations from database
|
||||
isActive := typeRecord["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetTypes_WithCustomFields", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Select: []interface{}{"type_id", "name", "is_active", "sort_order"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
},
|
||||
Limit: 3,
|
||||
}
|
||||
types, err := testProvider.GetTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.LessOrEqual(t, len(types), 3) // Respects limit
|
||||
|
||||
if len(types) > 0 {
|
||||
typeRecord := types[0]
|
||||
assert.Contains(t, typeRecord, "type_id")
|
||||
assert.Contains(t, typeRecord, "name")
|
||||
assert.Contains(t, typeRecord, "is_active")
|
||||
assert.Contains(t, typeRecord, "sort_order")
|
||||
}
|
||||
})
|
||||
|
||||
// Test PaginateTypes
|
||||
t.Run("PaginateTypes_FirstPage", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
},
|
||||
Orders: []model.QueryOrder{
|
||||
{Column: "sort_order", Option: "asc"},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateTypes(ctx, param, 1, 3)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
// Check pagination structure
|
||||
assert.Contains(t, result, "data")
|
||||
assert.Contains(t, result, "total")
|
||||
assert.Contains(t, result, "page")
|
||||
assert.Contains(t, result, "pagesize")
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.LessOrEqual(t, len(data), 3) // Page size limit
|
||||
|
||||
// Handle different total types
|
||||
totalInterface, exists := result["total"]
|
||||
assert.True(t, exists)
|
||||
|
||||
var total int64
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
total = int64(v)
|
||||
case int32:
|
||||
total = int64(v)
|
||||
case int64:
|
||||
total = v
|
||||
case uint:
|
||||
total = int64(v)
|
||||
case uint32:
|
||||
total = int64(v)
|
||||
case uint64:
|
||||
total = int64(v)
|
||||
default:
|
||||
t.Errorf("unexpected total type: %T, value: %v", totalInterface, totalInterface)
|
||||
}
|
||||
assert.GreaterOrEqual(t, total, int64(5)) // At least 5 types
|
||||
|
||||
assert.Equal(t, 1, result["page"])
|
||||
assert.Equal(t, 3, result["pagesize"])
|
||||
})
|
||||
|
||||
t.Run("PaginateTypes_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateTypes(ctx, param, 1, 10)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.GreaterOrEqual(t, len(data), 4) // At least 4 active types
|
||||
|
||||
// Verify is_active filter works
|
||||
for _, typeRecord := range data {
|
||||
if strings.Contains(typeRecord["type_id"].(string), "listtype_"+testUUID+"_") {
|
||||
// Handle different boolean representations from database
|
||||
isActive := typeRecord["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Test CountTypes
|
||||
t.Run("CountTypes_All", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(5)) // At least 5 types
|
||||
})
|
||||
|
||||
t.Run("CountTypes_WithFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
{Column: "is_active", Value: true},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(4)) // At least 4 active types
|
||||
})
|
||||
|
||||
t.Run("CountTypes_SpecificSortOrder", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: "listtype_" + testUUID + "_%"},
|
||||
{Column: "sort_order", OP: ">=", Value: 30},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
// We created 3 types with sort_order >= 30 (30, 40, 50), but be flexible with database state
|
||||
assert.GreaterOrEqual(t, count, int64(1)) // At least 1 type with sort_order >= 30
|
||||
assert.LessOrEqual(t, count, int64(5)) // But not more than 5 (our total test types)
|
||||
})
|
||||
|
||||
t.Run("CountTypes_NoResults", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: "nonexistent_type_id"},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, int64(0), count)
|
||||
})
|
||||
}
|
||||
|
||||
func TestTypeErrorHandling(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
nonExistentTypeID := "nonexistent_type_" + testUUID
|
||||
|
||||
t.Run("GetType_NotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetType(ctx, nonExistentTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("CreateType_MissingTypeID", func(t *testing.T) {
|
||||
typeData := maps.MapStrAny{
|
||||
"name": "Test Type",
|
||||
"description": "Type without type_id",
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type_id is required")
|
||||
})
|
||||
|
||||
t.Run("UpdateType_NotFound", func(t *testing.T) {
|
||||
updateData := maps.MapStrAny{"name": "Test"}
|
||||
err := testProvider.UpdateType(ctx, nonExistentTypeID, updateData)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("DeleteType_NotFound", func(t *testing.T) {
|
||||
err := testProvider.DeleteType(ctx, nonExistentTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("GetTypeConfiguration_NotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetTypeConfiguration(ctx, nonExistentTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("SetTypeConfiguration_NotFound", func(t *testing.T) {
|
||||
config := maps.MapStrAny{
|
||||
"schema": map[string]interface{}{"test": true},
|
||||
}
|
||||
err := testProvider.SetTypeConfiguration(ctx, nonExistentTypeID, config)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("GetTypes_EmptyResult", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: nonExistentTypeID},
|
||||
},
|
||||
}
|
||||
types, err := testProvider.GetTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 0, len(types)) // Empty slice, not nil
|
||||
})
|
||||
|
||||
t.Run("PaginateTypes_EmptyResult", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: nonExistentTypeID},
|
||||
},
|
||||
}
|
||||
result, err := testProvider.PaginateTypes(ctx, param, 1, 10)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
|
||||
data, ok := result["data"].([]maps.MapStr)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, 0, len(data))
|
||||
|
||||
// Handle different total types
|
||||
totalInterface, exists := result["total"]
|
||||
assert.True(t, exists)
|
||||
|
||||
var total int64
|
||||
switch v := totalInterface.(type) {
|
||||
case int:
|
||||
total = int64(v)
|
||||
case int32:
|
||||
total = int64(v)
|
||||
case int64:
|
||||
total = v
|
||||
case uint:
|
||||
total = int64(v)
|
||||
case uint32:
|
||||
total = int64(v)
|
||||
case uint64:
|
||||
total = int64(v)
|
||||
default:
|
||||
t.Errorf("unexpected total type: %T, value: %v", totalInterface, totalInterface)
|
||||
}
|
||||
assert.Equal(t, int64(0), total)
|
||||
})
|
||||
|
||||
t.Run("UpdateType_EmptyData", func(t *testing.T) {
|
||||
// First create a type for this test
|
||||
testTypeID := "emptyupdate_" + testUUID
|
||||
typeData := maps.MapStrAny{
|
||||
"type_id": testTypeID,
|
||||
"name": "Test Type for Empty Update",
|
||||
}
|
||||
_, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test with empty update data (should not error, just do nothing)
|
||||
emptyData := maps.MapStrAny{}
|
||||
err = testProvider.UpdateType(ctx, testTypeID, emptyData)
|
||||
assert.NoError(t, err) // Should not error, just skip update
|
||||
})
|
||||
|
||||
t.Run("SetTypeConfiguration_EmptyData", func(t *testing.T) {
|
||||
// First create a type for this test
|
||||
testTypeID := "emptyconfig_" + testUUID
|
||||
typeData := maps.MapStrAny{
|
||||
"type_id": testTypeID,
|
||||
"name": "Test Type for Empty Configuration",
|
||||
}
|
||||
_, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test with empty configuration data (should not error, just do nothing)
|
||||
emptyData := maps.MapStrAny{}
|
||||
err = testProvider.SetTypeConfiguration(ctx, testTypeID, emptyData)
|
||||
assert.NoError(t, err) // Should not error, just skip update
|
||||
})
|
||||
|
||||
t.Run("CountTypes_ComplexFilters", func(t *testing.T) {
|
||||
param := model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "is_active", Value: true},
|
||||
{Column: "sort_order", OP: ">=", Value: 10},
|
||||
{Column: "is_default", Value: false},
|
||||
},
|
||||
}
|
||||
count, err := testProvider.CountTypes(ctx, param)
|
||||
assert.NoError(t, err)
|
||||
assert.GreaterOrEqual(t, count, int64(0)) // Should handle complex filters without error
|
||||
})
|
||||
}
|
||||
|
|
@ -6,6 +6,7 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/log"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
|
@ -255,9 +256,48 @@ func (u *DefaultUser) UpdateUser(ctx context.Context, userID string, userData ma
|
|||
return nil
|
||||
}
|
||||
|
||||
// DeleteUser soft deletes a user account
|
||||
// DeleteUser soft deletes a user account and all associated data
|
||||
func (u *DefaultUser) DeleteUser(ctx context.Context, userID string) error {
|
||||
// First verify the user exists
|
||||
m := model.Select(u.model)
|
||||
users, err := m.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id"},
|
||||
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)
|
||||
}
|
||||
|
||||
// Clean up associated data before deleting the user
|
||||
// Note: We log warnings for cleanup failures but don't fail the user deletion
|
||||
|
||||
// 1. Delete all OAuth accounts for this user
|
||||
err = u.DeleteUserOAuthAccounts(ctx, userID)
|
||||
if err != nil {
|
||||
log.Warn("Failed to delete OAuth accounts for user %s: %v", userID, err)
|
||||
}
|
||||
|
||||
// 2. Clear user role assignment (set role_id to null)
|
||||
err = u.ClearUserRole(ctx, userID)
|
||||
if err != nil {
|
||||
log.Warn("Failed to clear role assignment for user %s: %v", userID, err)
|
||||
}
|
||||
|
||||
// 3. Clear user type assignment (set type_id to null)
|
||||
err = u.ClearUserType(ctx, userID)
|
||||
if err != nil {
|
||||
log.Warn("Failed to clear type assignment for user %s: %v", userID, err)
|
||||
}
|
||||
|
||||
// 4. Finally, delete the user account
|
||||
affected, err := m.DeleteWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ package user
|
|||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
)
|
||||
|
||||
|
|
@ -10,30 +12,360 @@ import (
|
|||
|
||||
// GetUserRole retrieves user's role information
|
||||
func (u *DefaultUser) GetUserRole(ctx context.Context, userID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// First get the user's role_id
|
||||
userModel := model.Select(u.model)
|
||||
users, err := userModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "role_id"},
|
||||
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]
|
||||
roleID, ok := user["role_id"].(string)
|
||||
if !ok || roleID == "" {
|
||||
return nil, fmt.Errorf("user %s has no role assigned", userID)
|
||||
}
|
||||
|
||||
// Now get the full role information
|
||||
roleModel := model.Select(u.roleModel)
|
||||
roles, err := roleModel.Get(model.QueryParam{
|
||||
Select: u.roleFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
if len(roles) == 0 {
|
||||
return nil, fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
return roles[0], nil
|
||||
}
|
||||
|
||||
// SetUserRole assigns a role to a user
|
||||
func (u *DefaultUser) SetUserRole(ctx context.Context, userID string, roleID string) error {
|
||||
// TODO: implement
|
||||
// First validate that the role exists
|
||||
roleModel := model.Select(u.roleModel)
|
||||
roles, err := roleModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"role_id", "is_active"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", Value: roleID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetRole, err)
|
||||
}
|
||||
|
||||
if len(roles) == 0 {
|
||||
return fmt.Errorf(ErrRoleNotFound)
|
||||
}
|
||||
|
||||
// Check if role is active
|
||||
role := roles[0]
|
||||
if isActive, ok := role["is_active"].(bool); ok && !isActive {
|
||||
return fmt.Errorf("cannot assign inactive role: %s", roleID)
|
||||
}
|
||||
// Handle different boolean types from database
|
||||
if isActiveInt, ok := role["is_active"].(int64); ok && isActiveInt == 0 {
|
||||
return fmt.Errorf("cannot assign inactive role: %s", roleID)
|
||||
}
|
||||
|
||||
// Update user's role_id
|
||||
updateData := maps.MapStrAny{
|
||||
"role_id": roleID,
|
||||
}
|
||||
|
||||
userModel := model.Select(u.model)
|
||||
affected, err := userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateUser, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearUserRole removes role assignment from a user (sets role_id to null)
|
||||
func (u *DefaultUser) ClearUserRole(ctx context.Context, userID string) error {
|
||||
// First check if user exists
|
||||
userModel := model.Select(u.model)
|
||||
users, err := userModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id"},
|
||||
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)
|
||||
}
|
||||
|
||||
// Update role_id to null (even if it's already null, this should succeed)
|
||||
updateData := maps.MapStrAny{
|
||||
"role_id": nil, // Set role_id to null to clear role assignment
|
||||
}
|
||||
|
||||
_, err = userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateUser, err)
|
||||
}
|
||||
|
||||
// Don't check affected rows - setting null to null is still a successful operation
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUserType retrieves user's type information
|
||||
func (u *DefaultUser) GetUserType(ctx context.Context, userID string) (maps.MapStrAny, error) {
|
||||
// TODO: implement
|
||||
return nil, nil
|
||||
// First get the user's type_id
|
||||
userModel := model.Select(u.model)
|
||||
users, err := userModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id", "type_id"},
|
||||
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]
|
||||
typeID, ok := user["type_id"].(string)
|
||||
if !ok || typeID == "" {
|
||||
return nil, fmt.Errorf("user %s has no type assigned", userID)
|
||||
}
|
||||
|
||||
// Now get the full type information
|
||||
typeModel := model.Select(u.typeModel)
|
||||
types, err := typeModel.Get(model.QueryParam{
|
||||
Select: u.typeFields,
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
if len(types) == 0 {
|
||||
return nil, fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
return types[0], nil
|
||||
}
|
||||
|
||||
// SetUserType assigns a type to a user
|
||||
func (u *DefaultUser) SetUserType(ctx context.Context, userID string, typeID string) error {
|
||||
// TODO: implement
|
||||
// First validate that the type exists
|
||||
typeModel := model.Select(u.typeModel)
|
||||
types, err := typeModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"type_id", "is_active"},
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", Value: typeID},
|
||||
},
|
||||
Limit: 1,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToGetType, err)
|
||||
}
|
||||
|
||||
if len(types) == 0 {
|
||||
return fmt.Errorf(ErrTypeNotFound)
|
||||
}
|
||||
|
||||
// Check if type is active
|
||||
typeRecord := types[0]
|
||||
if isActive, ok := typeRecord["is_active"].(bool); ok && !isActive {
|
||||
return fmt.Errorf("cannot assign inactive type: %s", typeID)
|
||||
}
|
||||
// Handle different boolean types from database
|
||||
if isActiveInt, ok := typeRecord["is_active"].(int64); ok && isActiveInt == 0 {
|
||||
return fmt.Errorf("cannot assign inactive type: %s", typeID)
|
||||
}
|
||||
|
||||
// Update user's type_id
|
||||
updateData := maps.MapStrAny{
|
||||
"type_id": typeID,
|
||||
}
|
||||
|
||||
userModel := model.Select(u.model)
|
||||
affected, err := userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateUser, err)
|
||||
}
|
||||
|
||||
if affected == 0 {
|
||||
return fmt.Errorf(ErrUserNotFound)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearUserType removes type assignment from a user (sets type_id to null)
|
||||
func (u *DefaultUser) ClearUserType(ctx context.Context, userID string) error {
|
||||
// First check if user exists
|
||||
userModel := model.Select(u.model)
|
||||
users, err := userModel.Get(model.QueryParam{
|
||||
Select: []interface{}{"user_id"},
|
||||
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)
|
||||
}
|
||||
|
||||
// Update type_id to null (even if it's already null, this should succeed)
|
||||
updateData := maps.MapStrAny{
|
||||
"type_id": nil, // Set type_id to null to clear type assignment
|
||||
}
|
||||
|
||||
_, err = userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userID},
|
||||
},
|
||||
Limit: 1, // Safety: ensure only one record is updated
|
||||
}, updateData)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf(ErrFailedToUpdateUser, err)
|
||||
}
|
||||
|
||||
// Don't check affected rows - setting null to null is still a successful operation
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateUserScope validates if a user has access to requested scopes based on role and type
|
||||
func (u *DefaultUser) ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error) {
|
||||
// TODO: implement
|
||||
return false, nil
|
||||
if len(scopes) == 0 {
|
||||
return true, nil // No scopes required
|
||||
}
|
||||
|
||||
// Get user's role
|
||||
userRole, err := u.GetUserRole(ctx, userID)
|
||||
if err != nil {
|
||||
// If user has no role, check if scopes are required
|
||||
if err.Error() == fmt.Sprintf("user %s has no role assigned", userID) {
|
||||
// Users without roles have minimal access (empty scopes only)
|
||||
return len(scopes) == 0, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
|
||||
// Extract role_id for permission validation
|
||||
roleID, ok := userRole["role_id"].(string)
|
||||
if !ok {
|
||||
return false, fmt.Errorf("invalid role_id format")
|
||||
}
|
||||
|
||||
// Use role-based permission validation
|
||||
valid, err := u.ValidateRolePermissions(ctx, roleID, scopes)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// If role validation passes, check type-specific restrictions if applicable
|
||||
if valid {
|
||||
// Get user's type for additional validation
|
||||
userType, err := u.GetUserType(ctx, userID)
|
||||
if err != nil {
|
||||
// If user has no type, role validation is sufficient
|
||||
if err.Error() == fmt.Sprintf("user %s has no type assigned", userID) {
|
||||
return valid, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
|
||||
// Get type configuration to check for additional scope restrictions
|
||||
typeID, ok := userType["type_id"].(string)
|
||||
if !ok {
|
||||
return false, fmt.Errorf("invalid type_id format")
|
||||
}
|
||||
|
||||
typeConfig, err := u.GetTypeConfiguration(ctx, typeID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// Check if type has specific scope limitations
|
||||
if features, ok := typeConfig["features"].(map[string]interface{}); ok {
|
||||
if scopeLimits, exists := features["scope_limits"]; exists {
|
||||
if limitList, ok := scopeLimits.([]interface{}); ok {
|
||||
// If type has scope limits, ensure all requested scopes are allowed
|
||||
allowedScopes := make(map[string]bool)
|
||||
for _, scope := range limitList {
|
||||
if scopeStr, ok := scope.(string); ok {
|
||||
allowedScopes[scopeStr] = true
|
||||
}
|
||||
}
|
||||
|
||||
// Check each requested scope against type limits
|
||||
for _, scope := range scopes {
|
||||
if !allowedScopes[scope] {
|
||||
return false, nil // Scope not allowed by type
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return valid, nil
|
||||
}
|
||||
|
|
|
|||
815
openapi/oauth/providers/user/user_role_type_test.go
Normal file
815
openapi/oauth/providers/user/user_role_type_test.go
Normal file
|
|
@ -0,0 +1,815 @@
|
|||
package user_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/yaoapp/gou/model"
|
||||
"github.com/yaoapp/kun/maps"
|
||||
)
|
||||
|
||||
func TestUserRoleOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] // 8 char UUID
|
||||
|
||||
// Step 1: Create a test user first
|
||||
testUser := createTestUserData("roleuser" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
// Step 2: Create test roles for assignment
|
||||
testRoles := []maps.MapStrAny{
|
||||
{
|
||||
"role_id": "adminrole_" + testUUID,
|
||||
"name": "Admin Role " + testUUID,
|
||||
"description": "Administrator role for testing",
|
||||
"is_active": true,
|
||||
"level": 100,
|
||||
},
|
||||
{
|
||||
"role_id": "userrole_" + testUUID,
|
||||
"name": "User Role " + testUUID,
|
||||
"description": "Regular user role for testing",
|
||||
"is_active": true,
|
||||
"level": 10,
|
||||
},
|
||||
{
|
||||
"role_id": "inactiverole_" + testUUID,
|
||||
"name": "Inactive Role " + testUUID,
|
||||
"description": "Inactive role for testing",
|
||||
"is_active": false,
|
||||
"level": 0,
|
||||
},
|
||||
}
|
||||
|
||||
// Create roles in database
|
||||
for _, roleData := range testRoles {
|
||||
_, err := testProvider.CreateRole(ctx, roleData)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
adminRoleID := "adminrole_" + testUUID
|
||||
userRoleID := "userrole_" + testUUID
|
||||
inactiveRoleID := "inactiverole_" + testUUID
|
||||
|
||||
// Test SetUserRole
|
||||
t.Run("SetUserRole", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, testUserID, adminRoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify role was assigned by getting user info
|
||||
user, err := testProvider.GetUser(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, adminRoleID, user["role_id"])
|
||||
})
|
||||
|
||||
// Test GetUserRole
|
||||
t.Run("GetUserRole", func(t *testing.T) {
|
||||
role, err := testProvider.GetUserRole(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, role)
|
||||
|
||||
// Verify we got the correct role information
|
||||
assert.Equal(t, adminRoleID, role["role_id"])
|
||||
assert.Equal(t, "Admin Role "+testUUID, role["name"])
|
||||
assert.Equal(t, "Administrator role for testing", role["description"])
|
||||
|
||||
// Handle different boolean representations from database
|
||||
isActive := role["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetUserRole - Change to different role
|
||||
t.Run("SetUserRole_ChangeRole", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, testUserID, userRoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify role was changed
|
||||
role, err := testProvider.GetUserRole(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, userRoleID, role["role_id"])
|
||||
assert.Equal(t, "User Role "+testUUID, role["name"])
|
||||
})
|
||||
|
||||
// Test SetUserRole - Inactive Role (should fail)
|
||||
t.Run("SetUserRole_InactiveRole", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, testUserID, inactiveRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "cannot assign inactive role")
|
||||
|
||||
// Verify role was not changed
|
||||
role, err := testProvider.GetUserRole(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, userRoleID, role["role_id"]) // Should still be the previous role
|
||||
})
|
||||
|
||||
// Test ClearUserRole
|
||||
t.Run("ClearUserRole", func(t *testing.T) {
|
||||
err := testProvider.ClearUserRole(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify role was cleared
|
||||
user, err := testProvider.GetUser(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Nil(t, user["role_id"]) // Should be null/nil
|
||||
|
||||
// GetUserRole should now fail
|
||||
_, err = testProvider.GetUserRole(ctx, testUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no role assigned")
|
||||
})
|
||||
|
||||
// Test SetUserRole again after clearing
|
||||
t.Run("SetUserRole_AfterClear", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, testUserID, adminRoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify role was assigned again
|
||||
role, err := testProvider.GetUserRole(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, adminRoleID, role["role_id"])
|
||||
})
|
||||
}
|
||||
|
||||
func TestUserTypeOperations(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Step 1: Create a test user first
|
||||
testUser := createTestUserData("typeuser" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
// Step 2: Create test types for assignment
|
||||
testTypes := []maps.MapStrAny{
|
||||
{
|
||||
"type_id": "basictype_" + testUUID,
|
||||
"name": "Basic Type " + testUUID,
|
||||
"description": "Basic user type for testing",
|
||||
"is_active": true,
|
||||
"sort_order": 10,
|
||||
},
|
||||
{
|
||||
"type_id": "premiumtype_" + testUUID,
|
||||
"name": "Premium Type " + testUUID,
|
||||
"description": "Premium user type for testing",
|
||||
"is_active": true,
|
||||
"sort_order": 20,
|
||||
},
|
||||
{
|
||||
"type_id": "inactivetype_" + testUUID,
|
||||
"name": "Inactive Type " + testUUID,
|
||||
"description": "Inactive type for testing",
|
||||
"is_active": false,
|
||||
"sort_order": 0,
|
||||
},
|
||||
}
|
||||
|
||||
// Create types in database
|
||||
for _, typeData := range testTypes {
|
||||
_, err := testProvider.CreateType(ctx, typeData)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
basicTypeID := "basictype_" + testUUID
|
||||
premiumTypeID := "premiumtype_" + testUUID
|
||||
inactiveTypeID := "inactivetype_" + testUUID
|
||||
|
||||
// Test SetUserType
|
||||
t.Run("SetUserType", func(t *testing.T) {
|
||||
err := testProvider.SetUserType(ctx, testUserID, basicTypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify type was assigned by getting user info
|
||||
user, err := testProvider.GetUser(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, basicTypeID, user["type_id"])
|
||||
})
|
||||
|
||||
// Test GetUserType
|
||||
t.Run("GetUserType", func(t *testing.T) {
|
||||
userType, err := testProvider.GetUserType(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, userType)
|
||||
|
||||
// Verify we got the correct type information
|
||||
assert.Equal(t, basicTypeID, userType["type_id"])
|
||||
assert.Equal(t, "Basic Type "+testUUID, userType["name"])
|
||||
assert.Equal(t, "Basic user type for testing", userType["description"])
|
||||
|
||||
// Handle different boolean representations from database
|
||||
isActive := userType["is_active"]
|
||||
switch v := isActive.(type) {
|
||||
case bool:
|
||||
assert.True(t, v)
|
||||
case int, int32, int64:
|
||||
assert.NotEqual(t, 0, v) // Any non-zero value is true
|
||||
default:
|
||||
t.Errorf("unexpected is_active type: %T, value: %v", isActive, isActive)
|
||||
}
|
||||
})
|
||||
|
||||
// Test SetUserType - Change to different type
|
||||
t.Run("SetUserType_ChangeType", func(t *testing.T) {
|
||||
err := testProvider.SetUserType(ctx, testUserID, premiumTypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify type was changed
|
||||
userType, err := testProvider.GetUserType(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, premiumTypeID, userType["type_id"])
|
||||
assert.Equal(t, "Premium Type "+testUUID, userType["name"])
|
||||
})
|
||||
|
||||
// Test SetUserType - Inactive Type (should fail)
|
||||
t.Run("SetUserType_InactiveType", func(t *testing.T) {
|
||||
err := testProvider.SetUserType(ctx, testUserID, inactiveTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "cannot assign inactive type")
|
||||
|
||||
// Verify type was not changed
|
||||
userType, err := testProvider.GetUserType(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, premiumTypeID, userType["type_id"]) // Should still be the previous type
|
||||
})
|
||||
|
||||
// Test ClearUserType
|
||||
t.Run("ClearUserType", func(t *testing.T) {
|
||||
err := testProvider.ClearUserType(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify type was cleared
|
||||
_, err = testProvider.GetUserType(ctx, testUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no type assigned")
|
||||
|
||||
// Verify user still exists
|
||||
user, err := testProvider.GetUser(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, testUserID, user["user_id"])
|
||||
assert.Nil(t, user["type_id"]) // type_id should be null
|
||||
})
|
||||
|
||||
// Test SetUserType - After Clear
|
||||
t.Run("SetUserType_AfterClear", func(t *testing.T) {
|
||||
// Re-assign a type after clearing
|
||||
err := testProvider.SetUserType(ctx, testUserID, basicTypeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify type was assigned
|
||||
userType, err := testProvider.GetUserType(ctx, testUserID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, basicTypeID, userType["type_id"])
|
||||
assert.Equal(t, "Basic Type "+testUUID, userType["name"])
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidateUserScope(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Step 1: Create test role with specific permissions
|
||||
testRole := maps.MapStrAny{
|
||||
"role_id": "scoperole_" + testUUID,
|
||||
"name": "Scope Test Role " + testUUID,
|
||||
"description": "Role for testing scope validation",
|
||||
"is_active": true,
|
||||
"permissions": map[string]interface{}{
|
||||
"read": true,
|
||||
"write": true,
|
||||
"admin.read": true,
|
||||
"admin.write": false,
|
||||
"delete": false,
|
||||
},
|
||||
"restricted_permissions": []string{
|
||||
"system.config",
|
||||
"root.access",
|
||||
},
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateRole(ctx, testRole)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Step 2: Create test type with scope limitations
|
||||
testType := maps.MapStrAny{
|
||||
"type_id": "scopetype_" + testUUID,
|
||||
"name": "Scope Test Type " + testUUID,
|
||||
"description": "Type for testing scope validation",
|
||||
"is_active": true,
|
||||
"features": map[string]interface{}{
|
||||
"api_access": true,
|
||||
"scope_limits": []interface{}{
|
||||
"read", "write", "admin.read", // Allowed scopes
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = testProvider.CreateType(ctx, testType)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Step 3: Create test user and assign role and type
|
||||
testUser := createTestUserData("scopeuser" + testUUID)
|
||||
_, testUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
roleID := "scoperole_" + testUUID
|
||||
typeID := "scopetype_" + testUUID
|
||||
|
||||
// Assign role and type to user
|
||||
err = testProvider.SetUserRole(ctx, testUserID, roleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
err = testProvider.SetUserType(ctx, testUserID, typeID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test various scope validation scenarios
|
||||
t.Run("ValidateUserScope_EmptyScopes", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Empty scopes should always be valid
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_ValidSingleScope", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"read"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // "read" is allowed by both role and type
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_ValidMultipleScopes", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"read", "write"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Both "read" and "write" are allowed
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_ValidAdminReadScope", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"admin.read"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // "admin.read" is allowed by both role and type
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_InvalidRolePermission", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"admin.write"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // "admin.write" is denied by role permissions
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_RestrictedPermission", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"system.config"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // "system.config" is in restricted permissions
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_TypeScopeLimitation", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"delete"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // "delete" is not in type's scope_limits
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_MixedValidInvalid", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"read", "delete"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should fail because "delete" is not allowed
|
||||
})
|
||||
|
||||
t.Run("ValidateUserScope_NonExistentScope", func(t *testing.T) {
|
||||
valid, err := testProvider.ValidateUserScope(ctx, testUserID, []string{"nonexistent.permission"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Non-existent permissions should be denied
|
||||
})
|
||||
|
||||
// Test user without role
|
||||
t.Run("ValidateUserScope_UserWithoutRole", func(t *testing.T) {
|
||||
// Create a user without role assignment
|
||||
userWithoutRole := createTestUserData("noroleuser" + testUUID)
|
||||
_, userWithoutRoleID := setupTestUser(t, ctx, userWithoutRole)
|
||||
|
||||
// Clear any default role that might have been set
|
||||
err := testProvider.ClearUserRole(ctx, userWithoutRoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// User without role should only have access to empty scopes
|
||||
valid, err := testProvider.ValidateUserScope(ctx, userWithoutRoleID, []string{})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Empty scopes should be valid
|
||||
|
||||
// Users without roles have minimal access (empty scopes only)
|
||||
valid, err = testProvider.ValidateUserScope(ctx, userWithoutRoleID, []string{"read"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Should return false - users without roles can only access empty scopes
|
||||
})
|
||||
|
||||
// Test user without type (type restrictions should not apply)
|
||||
t.Run("ValidateUserScope_UserWithoutType", func(t *testing.T) {
|
||||
// Create a user with role but without type
|
||||
userWithoutType := createTestUserData("notypeuser" + testUUID)
|
||||
userWithoutType.TypeID = "" // Explicitly clear type_id
|
||||
_, userWithoutTypeID := setupTestUser(t, ctx, userWithoutType)
|
||||
|
||||
// Assign role but no type
|
||||
err := testProvider.SetUserRole(ctx, userWithoutTypeID, roleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Manually clear type_id to ensure user has no type
|
||||
userModel := model.Select("__yao.user")
|
||||
_, err = userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userWithoutTypeID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, maps.MapStrAny{
|
||||
"type_id": nil, // Set type_id to null
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify user has no type assigned
|
||||
_, err = testProvider.GetUserType(ctx, userWithoutTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no type assigned")
|
||||
|
||||
// Should be able to access permissions allowed by role (no type restrictions)
|
||||
valid, err := testProvider.ValidateUserScope(ctx, userWithoutTypeID, []string{"read", "write"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Role allows these, no type restrictions
|
||||
|
||||
// Should still be restricted by role permissions
|
||||
valid, err = testProvider.ValidateUserScope(ctx, userWithoutTypeID, []string{"admin.write"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Role denies this
|
||||
})
|
||||
|
||||
// Test type without scope limits
|
||||
t.Run("ValidateUserScope_TypeWithoutScopeLimits", func(t *testing.T) {
|
||||
// Create a type without scope limits
|
||||
openType := maps.MapStrAny{
|
||||
"type_id": "opentype_" + testUUID,
|
||||
"name": "Open Type " + testUUID,
|
||||
"description": "Type without scope limitations",
|
||||
"is_active": true,
|
||||
"features": map[string]interface{}{
|
||||
"api_access": true,
|
||||
// No scope_limits - should allow anything the role permits
|
||||
},
|
||||
}
|
||||
|
||||
_, err := testProvider.CreateType(ctx, openType)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Create user with role and open type
|
||||
openUser := createTestUserData("openuser" + testUUID)
|
||||
_, openUserID := setupTestUser(t, ctx, openUser)
|
||||
|
||||
err = testProvider.SetUserRole(ctx, openUserID, roleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
err = testProvider.SetUserType(ctx, openUserID, "opentype_"+testUUID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Should be able to access any permission allowed by role
|
||||
valid, err := testProvider.ValidateUserScope(ctx, openUserID, []string{"read", "write", "admin.read"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid) // Type has no limitations, role allows these
|
||||
|
||||
// Should still be restricted by role permissions
|
||||
valid, err = testProvider.ValidateUserScope(ctx, openUserID, []string{"admin.write"})
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, valid) // Role denies this
|
||||
})
|
||||
}
|
||||
|
||||
func TestUserRoleErrorHandling(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to avoid conflicts
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
nonExistentUserID := "nonexistent_user_" + testUUID
|
||||
nonExistentRoleID := "nonexistent_role_" + testUUID
|
||||
|
||||
// Create a valid user for some tests
|
||||
testUser := createTestUserData("erroruser" + testUUID)
|
||||
_, validUserID := setupTestUser(t, ctx, testUser)
|
||||
|
||||
// Create a valid role for some tests
|
||||
validRoleData := maps.MapStrAny{
|
||||
"role_id": "validrole_" + testUUID,
|
||||
"name": "Valid Role " + testUUID,
|
||||
"description": "Valid role for error testing",
|
||||
"is_active": true,
|
||||
}
|
||||
_, err := testProvider.CreateRole(ctx, validRoleData)
|
||||
assert.NoError(t, err)
|
||||
validRoleID := "validrole_" + testUUID
|
||||
|
||||
t.Run("GetUserRole_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetUserRole(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("GetUserRole_NoRoleAssigned", func(t *testing.T) {
|
||||
// Create a user without a role assignment
|
||||
userWithoutRole := createTestUserData("noroleuser" + testUUID)
|
||||
_, userWithoutRoleID := setupTestUser(t, ctx, userWithoutRole)
|
||||
|
||||
// Clear any default role that might have been set
|
||||
err := testProvider.ClearUserRole(ctx, userWithoutRoleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, err = testProvider.GetUserRole(ctx, userWithoutRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no role assigned")
|
||||
})
|
||||
|
||||
t.Run("SetUserRole_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, nonExistentUserID, validRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("SetUserRole_RoleNotFound", func(t *testing.T) {
|
||||
err := testProvider.SetUserRole(ctx, validUserID, nonExistentRoleID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "role not found")
|
||||
})
|
||||
|
||||
t.Run("ClearUserRole_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.ClearUserRole(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("ClearUserRole_NoRoleTooClear", func(t *testing.T) {
|
||||
// Create a user without a role assignment
|
||||
userWithoutRole := createTestUserData("clearnouser" + testUUID)
|
||||
_, userWithoutRoleID := setupTestUser(t, ctx, userWithoutRole)
|
||||
|
||||
// Clear any default role that might have been set
|
||||
err := testProvider.ClearUserRole(ctx, userWithoutRoleID)
|
||||
assert.NoError(t, err) // Should succeed even if no role was assigned
|
||||
|
||||
// Try to clear again (should still succeed)
|
||||
err = testProvider.ClearUserRole(ctx, userWithoutRoleID)
|
||||
assert.NoError(t, err) // Should not error even if no role exists
|
||||
})
|
||||
|
||||
// Create a valid type for some tests
|
||||
validTypeData := maps.MapStrAny{
|
||||
"type_id": "validtype_" + testUUID,
|
||||
"name": "Valid Type " + testUUID,
|
||||
"description": "Valid type for error testing",
|
||||
"is_active": true,
|
||||
}
|
||||
_, err = testProvider.CreateType(ctx, validTypeData)
|
||||
assert.NoError(t, err)
|
||||
validTypeID := "validtype_" + testUUID
|
||||
|
||||
// Test user type error handling
|
||||
t.Run("GetUserType_UserNotFound", func(t *testing.T) {
|
||||
_, err := testProvider.GetUserType(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("GetUserType_NoTypeAssigned", func(t *testing.T) {
|
||||
// Create a user without a type assignment
|
||||
userWithoutType := createTestUserData("notypeuser" + testUUID)
|
||||
userWithoutType.TypeID = "" // Explicitly clear type_id
|
||||
_, userWithoutTypeID := setupTestUser(t, ctx, userWithoutType)
|
||||
|
||||
// Manually clear type_id to ensure user has no type
|
||||
userModel := model.Select("__yao.user")
|
||||
_, err = userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userWithoutTypeID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, maps.MapStrAny{
|
||||
"type_id": nil, // Set type_id to null
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, err = testProvider.GetUserType(ctx, userWithoutTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no type assigned")
|
||||
})
|
||||
|
||||
t.Run("SetUserType_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.SetUserType(ctx, nonExistentUserID, validTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("SetUserType_TypeNotFound", func(t *testing.T) {
|
||||
nonExistentTypeID := "nonexistent_type_" + testUUID
|
||||
err := testProvider.SetUserType(ctx, validUserID, nonExistentTypeID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "type not found")
|
||||
})
|
||||
|
||||
t.Run("ClearUserType_UserNotFound", func(t *testing.T) {
|
||||
err := testProvider.ClearUserType(ctx, nonExistentUserID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("ClearUserType_NoTypeTooClear", func(t *testing.T) {
|
||||
// Create a user without type assignment
|
||||
userWithoutType := createTestUserData("clearnotypeuser" + testUUID)
|
||||
userWithoutType.TypeID = "" // Explicitly clear type_id
|
||||
_, userWithoutTypeID := setupTestUser(t, ctx, userWithoutType)
|
||||
|
||||
// Manually clear type_id to ensure user has no type
|
||||
userModel := model.Select("__yao.user")
|
||||
_, err = userModel.UpdateWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "user_id", Value: userWithoutTypeID},
|
||||
},
|
||||
Limit: 1,
|
||||
}, maps.MapStrAny{
|
||||
"type_id": nil, // Set type_id to null
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Try to clear again (should still succeed)
|
||||
err = testProvider.ClearUserType(ctx, userWithoutTypeID)
|
||||
assert.NoError(t, err) // Should not error even if no type exists
|
||||
})
|
||||
|
||||
// Test scope validation error handling
|
||||
t.Run("ValidateUserScope_UserNotFound", func(t *testing.T) {
|
||||
scopes := []string{"read", "write"}
|
||||
valid, err := testProvider.ValidateUserScope(ctx, nonExistentUserID, scopes)
|
||||
assert.Error(t, err)
|
||||
assert.False(t, valid)
|
||||
assert.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestUserRoleIntegration(t *testing.T) {
|
||||
prepare(t)
|
||||
defer clean()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Use UUID to ensure unique identifiers
|
||||
testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8]
|
||||
|
||||
// Create multiple users and roles for integration testing
|
||||
users := make([]string, 3)
|
||||
for i := 0; i < 3; i++ {
|
||||
userData := createTestUserData("integuser" + testUUID + string('0'+rune(i)))
|
||||
_, userID := setupTestUser(t, ctx, userData)
|
||||
users[i] = userID
|
||||
}
|
||||
|
||||
roles := []string{
|
||||
"adminrole_" + testUUID,
|
||||
"userrole_" + testUUID,
|
||||
"guestrole_" + testUUID,
|
||||
}
|
||||
|
||||
roleData := []maps.MapStrAny{
|
||||
{
|
||||
"role_id": roles[0],
|
||||
"name": "Admin Role " + testUUID,
|
||||
"description": "Administrator role",
|
||||
"is_active": true,
|
||||
"level": 100,
|
||||
},
|
||||
{
|
||||
"role_id": roles[1],
|
||||
"name": "User Role " + testUUID,
|
||||
"description": "Regular user role",
|
||||
"is_active": true,
|
||||
"level": 10,
|
||||
},
|
||||
{
|
||||
"role_id": roles[2],
|
||||
"name": "Guest Role " + testUUID,
|
||||
"description": "Guest user role",
|
||||
"is_active": true,
|
||||
"level": 1,
|
||||
},
|
||||
}
|
||||
|
||||
// Create roles
|
||||
for _, role := range roleData {
|
||||
_, err := testProvider.CreateRole(ctx, role)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
t.Run("CompleteUserRoleFlow", func(t *testing.T) {
|
||||
userID := users[0]
|
||||
|
||||
// Step 1: Assign admin role
|
||||
err := testProvider.SetUserRole(ctx, userID, roles[0])
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Step 2: Verify role assignment
|
||||
role, err := testProvider.GetUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, roles[0], role["role_id"])
|
||||
assert.Equal(t, "Admin Role "+testUUID, role["name"])
|
||||
|
||||
// Step 3: Change to user role
|
||||
err = testProvider.SetUserRole(ctx, userID, roles[1])
|
||||
assert.NoError(t, err)
|
||||
|
||||
role, err = testProvider.GetUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, roles[1], role["role_id"])
|
||||
|
||||
// Step 4: Clear role
|
||||
err = testProvider.ClearUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Step 5: Verify role was cleared
|
||||
_, err = testProvider.GetUserRole(ctx, userID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no role assigned")
|
||||
|
||||
// Step 6: Reassign role
|
||||
err = testProvider.SetUserRole(ctx, userID, roles[2])
|
||||
assert.NoError(t, err)
|
||||
|
||||
role, err = testProvider.GetUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, roles[2], role["role_id"])
|
||||
})
|
||||
|
||||
t.Run("MultipleUsersRoleAssignment", func(t *testing.T) {
|
||||
// Assign different roles to different users
|
||||
for i, userID := range users {
|
||||
err := testProvider.SetUserRole(ctx, userID, roles[i])
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Verify each user has the correct role
|
||||
for i, userID := range users {
|
||||
role, err := testProvider.GetUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, roles[i], role["role_id"])
|
||||
}
|
||||
|
||||
// Clear all roles
|
||||
for _, userID := range users {
|
||||
err := testProvider.ClearUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Verify all roles were cleared
|
||||
for _, userID := range users {
|
||||
_, err := testProvider.GetUserRole(ctx, userID)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has no role assigned")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("RoleConsistency", func(t *testing.T) {
|
||||
userID := users[0]
|
||||
roleID := roles[0]
|
||||
|
||||
// Assign role
|
||||
err := testProvider.SetUserRole(ctx, userID, roleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Get role through user role method
|
||||
userRole, err := testProvider.GetUserRole(ctx, userID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Get role directly through role method
|
||||
directRole, err := testProvider.GetRole(ctx, roleID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Both should return the same role information
|
||||
assert.Equal(t, directRole["role_id"], userRole["role_id"])
|
||||
assert.Equal(t, directRole["name"], userRole["name"])
|
||||
assert.Equal(t, directRole["description"], userRole["description"])
|
||||
assert.Equal(t, directRole["is_active"], userRole["is_active"])
|
||||
assert.Equal(t, directRole["level"], userRole["level"])
|
||||
})
|
||||
}
|
||||
|
|
@ -112,13 +112,45 @@ func cleanupTestData() {
|
|||
})
|
||||
}
|
||||
|
||||
// Clean roles (should be done before users due to potential role_id references)
|
||||
roleModel := model.Select("__yao.user_role")
|
||||
rolePatterns := []string{
|
||||
"test%", "%testrole%", "%listrole%", "%permrole%", "%adminrole%", "%userrole%",
|
||||
"%inactiverole%", "%systemrole%", "%validrole%", "%emptyupdate%", "%emptyperm%",
|
||||
"%guestrole%", "%scoperole%",
|
||||
}
|
||||
for _, pattern := range rolePatterns {
|
||||
roleModel.DestroyWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "role_id", OP: "like", Value: pattern},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Clean types (should be done before users due to potential type_id references)
|
||||
typeModel := model.Select("__yao.user_type")
|
||||
typePatterns := []string{
|
||||
"test%", "%testtype%", "%listtype%", "%configtype%", "%basictype%", "%premiumtype%",
|
||||
"%inactivetype%", "%validtype%", "%emptyupdate%", "%emptyconfig%", "%scopetype%",
|
||||
"%opentype%",
|
||||
}
|
||||
for _, pattern := range typePatterns {
|
||||
typeModel.DestroyWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
{Column: "type_id", OP: "like", Value: pattern},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Clean users
|
||||
userModel := model.Select("__yao.user")
|
||||
|
||||
// Delete test users by pattern (using hard delete)
|
||||
userPatterns := []string{
|
||||
"test-%", "test_%", "%testuser%", "%oauthtest%", "%oauthlist%",
|
||||
"%oautherror%", "%deletetest%",
|
||||
"%oautherror%", "%deletetest%", "%roleuser%", "%typeuser%", "%scopeuser%",
|
||||
"%erroruser%", "%noroleuser%", "%clearnouser%", "%integuser%", "%notypeuser%",
|
||||
"%openuser%", "%clearnotypeuser%",
|
||||
}
|
||||
for _, pattern := range userPatterns {
|
||||
userModel.DestroyWhere(model.QueryParam{
|
||||
|
|
@ -129,7 +161,11 @@ func cleanupTestData() {
|
|||
}
|
||||
|
||||
// Also clean by username pattern
|
||||
usernamePatterns := []string{"testuser%", "%oauth_%", "%deletetest%"}
|
||||
usernamePatterns := []string{
|
||||
"testuser%", "%oauth_%", "%deletetest%", "%roleuser%", "%typeuser%",
|
||||
"%scopeuser%", "%erroruser%", "%noroleuser%", "%clearnouser%", "%integuser%",
|
||||
"%notypeuser%", "%openuser%", "%clearnotypeuser%",
|
||||
}
|
||||
for _, pattern := range usernamePatterns {
|
||||
userModel.DestroyWhere(model.QueryParam{
|
||||
Wheres: []model.QueryWhere{
|
||||
|
|
|
|||
|
|
@ -174,8 +174,10 @@ type UserProvider interface {
|
|||
// User Role and Type Management
|
||||
GetUserRole(ctx context.Context, userID string) (maps.MapStrAny, error)
|
||||
SetUserRole(ctx context.Context, userID string, roleID string) error
|
||||
ClearUserRole(ctx context.Context, userID string) error
|
||||
GetUserType(ctx context.Context, userID string) (maps.MapStrAny, error)
|
||||
SetUserType(ctx context.Context, userID string, typeID string) error
|
||||
ClearUserType(ctx context.Context, userID string) error
|
||||
ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error)
|
||||
|
||||
// User MFA Management
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue