Add lightweight existence check methods for OAuth accounts, roles, types, and users

- Implemented methods to check the existence of OAuth accounts, roles, types, and users by their respective identifiers, enhancing the user provider's functionality.
- Added error handling for these methods to ensure robust feedback in case of failures.
- Updated the user provider interface to include the new existence check methods, improving overall usability and maintainability.
This commit is contained in:
Max 2025-08-03 07:14:09 +08:00
parent e3e4b1bb77
commit 21c7caed5d
8 changed files with 412 additions and 4 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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