diff --git a/openapi/oauth/providers/user/exists_test.go b/openapi/oauth/providers/user/exists_test.go new file mode 100644 index 00000000..c007f341 --- /dev/null +++ b/openapi/oauth/providers/user/exists_test.go @@ -0,0 +1,243 @@ +package user_test + +import ( + "context" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/yaoapp/kun/maps" +) + +// TestExistsMethods tests all resource existence check methods +func TestExistsMethods(t *testing.T) { + prepare(t) + defer clean() + + ctx := context.Background() + + // Use UUID to ensure unique identifiers across test runs + testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] // 8 char UUID + + testUserID := "test-user-for-exists-" + testUUID + testUsername := "testexistsuser" + testUUID + testEmail := "testexists" + testUUID + "@example.com" + + t.Run("UserExists", func(t *testing.T) { + // Test non-existent user + exists, err := testProvider.UserExists(ctx, "nonexistent-user-id") + assert.NoError(t, err) + assert.False(t, exists) + + // Create a test user + userData := maps.MapStrAny{ + "user_id": testUserID, + "preferred_username": testUsername, + "email": testEmail, + "password": "password123", + "status": "active", + } + + _, err = testProvider.CreateUser(ctx, userData) + assert.NoError(t, err) + + // Test existing user + exists, err = testProvider.UserExists(ctx, testUserID) + assert.NoError(t, err) + assert.True(t, exists) + }) + + t.Run("UserExistsByEmail", func(t *testing.T) { + // Test non-existent email + exists, err := testProvider.UserExistsByEmail(ctx, "nonexistent@example.com") + assert.NoError(t, err) + assert.False(t, exists) + + // Test existing email (using unique email from test setup) + exists, err = testProvider.UserExistsByEmail(ctx, testEmail) + assert.NoError(t, err) + assert.True(t, exists) + }) + + t.Run("UserExistsByPreferredUsername", func(t *testing.T) { + // Test non-existent username + exists, err := testProvider.UserExistsByPreferredUsername(ctx, "nonexistentuser") + assert.NoError(t, err) + assert.False(t, exists) + + // Test existing username (using unique username from test setup) + exists, err = testProvider.UserExistsByPreferredUsername(ctx, testUsername) + assert.NoError(t, err) + assert.True(t, exists) + }) + + // Define unique IDs for roles and types + testRoleID := "test-role-for-exists-" + testUUID + testTypeID := "test-type-for-exists-" + testUUID + + t.Run("RoleExists", func(t *testing.T) { + // Test non-existent role + exists, err := testProvider.RoleExists(ctx, "nonexistent-role") + assert.NoError(t, err) + assert.False(t, exists) + + // Create a test role + roleData := maps.MapStrAny{ + "role_id": testRoleID, + "name": "Test Exists Role " + testUUID, + "description": "Role for testing exists method", + "is_active": true, + } + + _, err = testProvider.CreateRole(ctx, roleData) + assert.NoError(t, err) + + // Test existing role + exists, err = testProvider.RoleExists(ctx, testRoleID) + assert.NoError(t, err) + assert.True(t, exists) + }) + + t.Run("TypeExists", func(t *testing.T) { + // Test non-existent type + exists, err := testProvider.TypeExists(ctx, "nonexistent-type") + assert.NoError(t, err) + assert.False(t, exists) + + // Create a test type + typeData := maps.MapStrAny{ + "type_id": testTypeID, + "name": "Test Exists Type " + testUUID, + "description": "Type for testing exists method", + "is_active": true, + } + + _, err = testProvider.CreateType(ctx, typeData) + assert.NoError(t, err) + + // Test existing type + exists, err = testProvider.TypeExists(ctx, testTypeID) + assert.NoError(t, err) + assert.True(t, exists) + }) + + t.Run("OAuthAccountExists", func(t *testing.T) { + // Test non-existent OAuth account + exists, err := testProvider.OAuthAccountExists(ctx, "nonexistent-provider", "nonexistent-subject") + assert.NoError(t, err) + assert.False(t, exists) + + // Create a test OAuth account (using unique identifiers) + testOAuthProvider := "test-provider-" + testUUID + testSubject := "test-subject-for-exists-" + testUUID + oauthData := maps.MapStrAny{ + "provider": testOAuthProvider, + "sub": testSubject, + "name": "Test OAuth User " + testUUID, + "email": "testoauth" + testUUID + "@example.com", + "is_active": true, + } + + _, err = testProvider.CreateOAuthAccount(ctx, testUserID, oauthData) + assert.NoError(t, err) + + // Test existing OAuth account + exists, err = testProvider.OAuthAccountExists(ctx, testOAuthProvider, testSubject) + assert.NoError(t, err) + assert.True(t, exists) + }) + + t.Run("UserHasRole", func(t *testing.T) { + // Test user without role + hasRole, err := testProvider.UserHasRole(ctx, testUserID) + assert.NoError(t, err) + assert.False(t, hasRole) + + // Assign role to user + err = testProvider.SetUserRole(ctx, testUserID, testRoleID) + assert.NoError(t, err) + + // Test user with role + hasRole, err = testProvider.UserHasRole(ctx, testUserID) + assert.NoError(t, err) + assert.True(t, hasRole) + + // Test non-existent user + _, err = testProvider.UserHasRole(ctx, "nonexistent-user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "user not found") + }) + + t.Run("UserHasType", func(t *testing.T) { + // Test user without type + hasType, err := testProvider.UserHasType(ctx, testUserID) + assert.NoError(t, err) + assert.False(t, hasType) + + // Assign type to user + err = testProvider.SetUserType(ctx, testUserID, testTypeID) + assert.NoError(t, err) + + // Test user with type + hasType, err = testProvider.UserHasType(ctx, testUserID) + assert.NoError(t, err) + assert.True(t, hasType) + + // Test non-existent user + _, err = testProvider.UserHasType(ctx, "nonexistent-user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "user not found") + }) +} + +// TestExistsPerformance tests the performance benefit of Exists methods vs full Get methods +func TestExistsPerformance(t *testing.T) { + prepare(t) + defer clean() + + ctx := context.Background() + + // Use UUID to ensure unique identifiers + perfUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] // 8 char UUID + perfUserID := "perf-test-user-" + perfUUID + perfUsername := "perfuser" + perfUUID + perfEmail := "perf" + perfUUID + "@example.com" + + // Create a test user for performance comparison + userData := maps.MapStrAny{ + "user_id": perfUserID, + "preferred_username": perfUsername, + "email": perfEmail, + "password": "password123", + "status": "active", + } + + _, err := testProvider.CreateUser(ctx, userData) + assert.NoError(t, err) + + t.Run("UserExists_vs_GetUser", func(t *testing.T) { + // Both should work, but UserExists should be more efficient + // (we can't easily measure performance in unit tests, but we verify functionality) + + // Test UserExists + exists, err := testProvider.UserExists(ctx, perfUserID) + assert.NoError(t, err) + assert.True(t, exists) + + // Test GetUser (more expensive) + user, err := testProvider.GetUser(ctx, perfUserID) + assert.NoError(t, err) + assert.NotNil(t, user) + assert.Equal(t, perfUserID, user["user_id"]) + + // Both methods should give consistent results for existence + exists, err = testProvider.UserExists(ctx, "nonexistent-user") + assert.NoError(t, err) + assert.False(t, exists) + + _, err = testProvider.GetUser(ctx, "nonexistent-user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "user not found") + }) +} diff --git a/openapi/oauth/providers/user/oauth_account.go b/openapi/oauth/providers/user/oauth_account.go index 80330ec0..56c82aaa 100644 --- a/openapi/oauth/providers/user/oauth_account.go +++ b/openapi/oauth/providers/user/oauth_account.go @@ -58,6 +58,25 @@ func (u *DefaultUser) GetOAuthAccount(ctx context.Context, provider string, subj return accounts[0], nil } +// OAuthAccountExists checks if an OAuth account exists by provider and subject (lightweight query) +func (u *DefaultUser) OAuthAccountExists(ctx context.Context, provider string, subject string) (bool, error) { + m := model.Select(u.oauthAccountModel) + accounts, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "provider", Value: provider}, + {Column: "sub", Value: subject}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetOAuthAccount, err) + } + + return len(accounts) > 0, nil +} + // GetUserOAuthAccounts retrieves all OAuth accounts for a user func (u *DefaultUser) GetUserOAuthAccounts(ctx context.Context, userID string) ([]maps.MapStrAny, error) { m := model.Select(u.oauthAccountModel) diff --git a/openapi/oauth/providers/user/role.go b/openapi/oauth/providers/user/role.go index 05a48368..db198018 100644 --- a/openapi/oauth/providers/user/role.go +++ b/openapi/oauth/providers/user/role.go @@ -32,6 +32,24 @@ func (u *DefaultUser) GetRole(ctx context.Context, roleID string) (maps.MapStrAn return roles[0], nil } +// RoleExists checks if a role exists by role_id (lightweight query) +func (u *DefaultUser) RoleExists(ctx context.Context, roleID string) (bool, error) { + m := model.Select(u.roleModel) + roles, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "role_id", Value: roleID}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetRole, err) + } + + return len(roles) > 0, nil +} + // CreateRole creates a new user role func (u *DefaultUser) CreateRole(ctx context.Context, roleData maps.MapStrAny) (interface{}, error) { // Validate required role_id field diff --git a/openapi/oauth/providers/user/type.go b/openapi/oauth/providers/user/type.go index 60638d12..13fe46bf 100644 --- a/openapi/oauth/providers/user/type.go +++ b/openapi/oauth/providers/user/type.go @@ -32,6 +32,24 @@ func (u *DefaultUser) GetType(ctx context.Context, typeID string) (maps.MapStrAn return types[0], nil } +// TypeExists checks if a type exists by type_id (lightweight query) +func (u *DefaultUser) TypeExists(ctx context.Context, typeID string) (bool, error) { + m := model.Select(u.typeModel) + types, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "type_id", Value: typeID}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetType, err) + } + + return len(types) > 0, nil +} + // CreateType creates a new user type func (u *DefaultUser) CreateType(ctx context.Context, typeData maps.MapStrAny) (interface{}, error) { // Validate required type_id field diff --git a/openapi/oauth/providers/user/user_basic.go b/openapi/oauth/providers/user/user_basic.go index 306c62d2..7ab1087c 100644 --- a/openapi/oauth/providers/user/user_basic.go +++ b/openapi/oauth/providers/user/user_basic.go @@ -35,6 +35,60 @@ func (u *DefaultUser) GetUser(ctx context.Context, userID string) (maps.MapStrAn return users[0], nil } +// UserExists checks if a user exists by user_id (lightweight query) +func (u *DefaultUser) UserExists(ctx context.Context, userID string) (bool, error) { + m := model.Select(u.model) + users, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "user_id", Value: userID}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetUser, err) + } + + return len(users) > 0, nil +} + +// UserExistsByEmail checks if a user exists by email (lightweight query) +func (u *DefaultUser) UserExistsByEmail(ctx context.Context, email string) (bool, error) { + m := model.Select(u.model) + users, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "email", Value: email}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetUser, err) + } + + return len(users) > 0, nil +} + +// UserExistsByPreferredUsername checks if a user exists by preferred_username (lightweight query) +func (u *DefaultUser) UserExistsByPreferredUsername(ctx context.Context, preferredUsername string) (bool, error) { + m := model.Select(u.model) + users, err := m.Get(model.QueryParam{ + Select: []interface{}{"id"}, // Only select ID for existence check + Wheres: []model.QueryWhere{ + {Column: "preferred_username", Value: preferredUsername}, + }, + Limit: 1, // Only need to know if at least one exists + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetUser, err) + } + + return len(users) > 0, nil +} + // GetUserByPreferredUsername retrieves user by preferred_username (OIDC standard) func (u *DefaultUser) GetUserByPreferredUsername(ctx context.Context, preferredUsername string) (maps.MapStrAny, error) { m := model.Select(u.model) diff --git a/openapi/oauth/providers/user/user_role_type.go b/openapi/oauth/providers/user/user_role_type.go index 52e1816b..70f463a3 100644 --- a/openapi/oauth/providers/user/user_role_type.go +++ b/openapi/oauth/providers/user/user_role_type.go @@ -151,6 +151,30 @@ func (u *DefaultUser) ClearUserRole(ctx context.Context, userID string) error { return nil } +// UserHasRole checks if a user has a role assigned (lightweight query) +func (u *DefaultUser) UserHasRole(ctx context.Context, userID string) (bool, error) { + userModel := model.Select(u.model) + users, err := userModel.Get(model.QueryParam{ + Select: []interface{}{"role_id"}, // Only select role_id field + Wheres: []model.QueryWhere{ + {Column: "user_id", Value: userID}, + }, + Limit: 1, + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetUser, err) + } + + if len(users) == 0 { + return false, fmt.Errorf(ErrUserNotFound) + } + + user := users[0] + roleID, ok := user["role_id"].(string) + return ok && roleID != "", nil +} + // GetUserType retrieves user's type information func (u *DefaultUser) GetUserType(ctx context.Context, userID string) (maps.MapStrAny, error) { // First get the user's type_id @@ -292,6 +316,30 @@ func (u *DefaultUser) ClearUserType(ctx context.Context, userID string) error { return nil } +// UserHasType checks if a user has a type assigned (lightweight query) +func (u *DefaultUser) UserHasType(ctx context.Context, userID string) (bool, error) { + userModel := model.Select(u.model) + users, err := userModel.Get(model.QueryParam{ + Select: []interface{}{"type_id"}, // Only select type_id field + Wheres: []model.QueryWhere{ + {Column: "user_id", Value: userID}, + }, + Limit: 1, + }) + + if err != nil { + return false, fmt.Errorf(ErrFailedToGetUser, err) + } + + if len(users) == 0 { + return false, fmt.Errorf(ErrUserNotFound) + } + + user := users[0] + typeID, ok := user["type_id"].(string) + return ok && typeID != "", 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) { if len(scopes) == 0 { diff --git a/openapi/oauth/providers/user/user_test.go b/openapi/oauth/providers/user/user_test.go index c46d443f..55d6e98f 100644 --- a/openapi/oauth/providers/user/user_test.go +++ b/openapi/oauth/providers/user/user_test.go @@ -117,7 +117,7 @@ func cleanupTestData() { rolePatterns := []string{ "test%", "%testrole%", "%listrole%", "%permrole%", "%adminrole%", "%userrole%", "%inactiverole%", "%systemrole%", "%validrole%", "%emptyupdate%", "%emptyperm%", - "%guestrole%", "%scoperole%", + "%guestrole%", "%scoperole%", "%test-role-for-exists%", } for _, pattern := range rolePatterns { roleModel.DestroyWhere(model.QueryParam{ @@ -132,7 +132,7 @@ func cleanupTestData() { typePatterns := []string{ "test%", "%testtype%", "%listtype%", "%configtype%", "%basictype%", "%premiumtype%", "%inactivetype%", "%validtype%", "%emptyupdate%", "%emptyconfig%", "%scopetype%", - "%opentype%", + "%opentype%", "%test-type-for-exists%", } for _, pattern := range typePatterns { typeModel.DestroyWhere(model.QueryParam{ @@ -150,7 +150,7 @@ func cleanupTestData() { "test-%", "test_%", "%testuser%", "%oauthtest%", "%oauthlist%", "%oautherror%", "%deletetest%", "%roleuser%", "%typeuser%", "%scopeuser%", "%erroruser%", "%noroleuser%", "%clearnouser%", "%integuser%", "%notypeuser%", - "%openuser%", "%clearnotypeuser%", + "%openuser%", "%clearnotypeuser%", "%test-user-for-exists%", "%perf-test-user%", } for _, pattern := range userPatterns { userModel.DestroyWhere(model.QueryParam{ @@ -164,7 +164,7 @@ func cleanupTestData() { usernamePatterns := []string{ "testuser%", "%oauth_%", "%deletetest%", "%roleuser%", "%typeuser%", "%scopeuser%", "%erroruser%", "%noroleuser%", "%clearnouser%", "%integuser%", - "%notypeuser%", "%openuser%", "%clearnotypeuser%", + "%notypeuser%", "%openuser%", "%clearnotypeuser%", "%testexistsuser%", "%perfuser%", } for _, pattern := range usernamePatterns { userModel.DestroyWhere(model.QueryParam{ diff --git a/openapi/oauth/types/interfaces.go b/openapi/oauth/types/interfaces.go index dd01f75d..ae2accc8 100644 --- a/openapi/oauth/types/interfaces.go +++ b/openapi/oauth/types/interfaces.go @@ -152,6 +152,9 @@ type UserProvider interface { // User Basic Operations GetUser(ctx context.Context, userID string) (maps.MapStrAny, error) + UserExists(ctx context.Context, userID string) (bool, error) + UserExistsByEmail(ctx context.Context, email string) (bool, error) + UserExistsByPreferredUsername(ctx context.Context, preferredUsername string) (bool, error) GetUserByPreferredUsername(ctx context.Context, preferredUsername string) (maps.MapStrAny, error) GetUserByEmail(ctx context.Context, email string) (maps.MapStrAny, error) @@ -175,9 +178,11 @@ type UserProvider interface { 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 + UserHasRole(ctx context.Context, userID string) (bool, 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 + UserHasType(ctx context.Context, userID string) (bool, error) ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error) // User MFA Management @@ -196,6 +201,7 @@ type UserProvider interface { CreateOAuthAccount(ctx context.Context, userID string, oauthData maps.MapStrAny) (interface{}, error) GetOAuthAccount(ctx context.Context, provider string, subject string) (maps.MapStrAny, error) + OAuthAccountExists(ctx context.Context, provider string, subject string) (bool, error) GetUserOAuthAccounts(ctx context.Context, userID string) ([]maps.MapStrAny, error) UpdateOAuthAccount(ctx context.Context, provider string, subject string, oauthData maps.MapStrAny) error DeleteOAuthAccount(ctx context.Context, provider string, subject string) error @@ -210,6 +216,7 @@ type UserProvider interface { // ============================================================================ GetRole(ctx context.Context, roleID string) (maps.MapStrAny, error) + RoleExists(ctx context.Context, roleID string) (bool, error) CreateRole(ctx context.Context, roleData maps.MapStrAny) (interface{}, error) UpdateRole(ctx context.Context, roleID string, roleData maps.MapStrAny) error DeleteRole(ctx context.Context, roleID string) error @@ -227,6 +234,7 @@ type UserProvider interface { // ============================================================================ GetType(ctx context.Context, typeID string) (maps.MapStrAny, error) + TypeExists(ctx context.Context, typeID string) (bool, error) CreateType(ctx context.Context, typeData maps.MapStrAny) (interface{}, error) UpdateType(ctx context.Context, typeID string, typeData maps.MapStrAny) error DeleteType(ctx context.Context, typeID string) error