Refactor user provider methods to return string IDs and update test data handling

- Updated GetUserProvider to return an error alongside the user provider for better error handling.
- Modified CreateUser, CreateRole, and CreateType methods to return user and role IDs as strings instead of interfaces, enhancing type safety.
- Adjusted test data setup in oauth_test.go to reflect changes in user ID handling and ensure compatibility with the updated user provider interface.
- Removed unnecessary nil checks and improved assertions in tests for clarity and reliability.
This commit is contained in:
Max 2025-08-03 10:25:09 +08:00
parent 97f053e41d
commit 5301023cf2
12 changed files with 205 additions and 48 deletions

View file

@ -160,8 +160,8 @@ func (s *Service) GetConfig() *Config {
}
// GetUserProvider returns the user provider for the service
func (s *Service) GetUserProvider() types.UserProvider {
return s.userProvider
func (s *Service) GetUserProvider() (types.UserProvider, error) {
return s.userProvider, nil
}
// GetClientProvider returns the client provider for the service

View file

@ -484,7 +484,7 @@ func setupTestData(t *testing.T, service *Service) {
}
// Create test users using the updated user provider interface
userProvider := service.GetUserProvider()
userProvider, _ := service.GetUserProvider()
for i, testUser := range testUsers {
// Convert TestUser to the format expected by CreateUser
userData := map[string]interface{}{
@ -505,16 +505,13 @@ func setupTestData(t *testing.T, service *Service) {
createdUserID, err := userProvider.CreateUser(ctx, userData)
require.NoError(t, err, "Failed to create test user %d: %s", i, testUser.Description)
require.NotNil(t, createdUserID, "Created user ID should not be nil")
require.NotEmpty(t, createdUserID, "Created user ID should not be empty")
// Update the test user with the created database ID and auto-generated user_id
if userID, ok := createdUserID.(int64); ok {
testUser.ID = userID
} else if userID, ok := createdUserID.(int); ok {
testUser.ID = int64(userID)
} else {
testUser.ID = int64(0) // Fallback for interface{} types
}
// The CreateUser method now returns the user_id as string directly
testUser.UserID = createdUserID
// For backward compatibility, also store as string in userData if needed
userData["user_id"] = createdUserID
// Extract the auto-generated user_id from userData (CreateUser sets it)
if generatedUserID, ok := userData["user_id"].(string); ok {
@ -727,7 +724,7 @@ func TestServiceGetters(t *testing.T) {
})
t.Run("get user provider", func(t *testing.T) {
userProvider := service.GetUserProvider()
userProvider, _ := service.GetUserProvider()
assert.NotNil(t, userProvider)
assert.Implements(t, (*types.UserProvider)(nil), userProvider)
})
@ -864,7 +861,7 @@ func TestProviderInitialization(t *testing.T) {
require.NoError(t, err)
// Create custom providers (for this test, we'll use the default ones)
customUserProvider := tempService.GetUserProvider()
customUserProvider, _ := tempService.GetUserProvider()
customClientProvider := tempService.GetClientProvider()
config := &Config{

View file

@ -51,10 +51,10 @@ func (u *DefaultUser) RoleExists(ctx context.Context, roleID string) (bool, erro
}
// CreateRole creates a new user role
func (u *DefaultUser) CreateRole(ctx context.Context, roleData maps.MapStrAny) (interface{}, error) {
func (u *DefaultUser) CreateRole(ctx context.Context, roleData maps.MapStrAny) (string, error) {
// Validate required role_id field
if _, exists := roleData["role_id"]; !exists {
return nil, fmt.Errorf("role_id is required in roleData")
return "", fmt.Errorf("role_id is required in roleData")
}
// Set default values if not provided
@ -77,10 +77,16 @@ func (u *DefaultUser) CreateRole(ctx context.Context, roleData maps.MapStrAny) (
m := model.Select(u.roleModel)
id, err := m.Create(roleData)
if err != nil {
return nil, fmt.Errorf(ErrFailedToCreateRole, err)
return "", fmt.Errorf(ErrFailedToCreateRole, err)
}
return id, nil
// Return the role_id as string (preferred approach)
if roleID, ok := roleData["role_id"].(string); ok {
return roleID, nil
}
// Fallback: convert the returned int id to string
return fmt.Sprintf("%d", id), nil
}
// UpdateRole updates an existing role

View file

@ -51,10 +51,10 @@ func (u *DefaultUser) TypeExists(ctx context.Context, typeID string) (bool, erro
}
// CreateType creates a new user type
func (u *DefaultUser) CreateType(ctx context.Context, typeData maps.MapStrAny) (interface{}, error) {
func (u *DefaultUser) CreateType(ctx context.Context, typeData maps.MapStrAny) (string, error) {
// Validate required type_id field
if _, exists := typeData["type_id"]; !exists {
return nil, fmt.Errorf("type_id is required in typeData")
return "", fmt.Errorf("type_id is required in typeData")
}
// Set default values if not provided
@ -77,10 +77,16 @@ func (u *DefaultUser) CreateType(ctx context.Context, typeData maps.MapStrAny) (
m := model.Select(u.typeModel)
id, err := m.Create(typeData)
if err != nil {
return nil, fmt.Errorf(ErrFailedToCreateType, err)
return "", fmt.Errorf(ErrFailedToCreateType, err)
}
return id, nil
// Return the type_id as string (preferred approach)
if typeID, ok := typeData["type_id"].(string); ok {
return typeID, nil
}
// Fallback: convert the returned int id to string
return fmt.Sprintf("%d", id), nil
}
// UpdateType updates an existing type

View file

@ -244,12 +244,12 @@ func (u *DefaultUser) ResetPassword(ctx context.Context, userID string) (string,
}
// CreateUser creates a new user with OIDC standard fields
func (u *DefaultUser) CreateUser(ctx context.Context, userData maps.MapStrAny) (interface{}, error) {
func (u *DefaultUser) CreateUser(ctx context.Context, userData maps.MapStrAny) (string, error) {
// Auto-generate user_id if not provided
if _, exists := userData["user_id"]; !exists {
userID, err := u.GenerateUserID(ctx, true) // Force safe mode to ensure uniqueness
if err != nil {
return nil, fmt.Errorf(ErrFailedToGenerateUserID, err)
return "", fmt.Errorf(ErrFailedToGenerateUserID, err)
}
userData["user_id"] = userID
}
@ -268,10 +268,16 @@ func (u *DefaultUser) CreateUser(ctx context.Context, userData maps.MapStrAny) (
m := model.Select(u.model)
id, err := m.Create(userData)
if err != nil {
return nil, fmt.Errorf(ErrFailedToCreateUser, err)
return "", fmt.Errorf(ErrFailedToCreateUser, err)
}
return id, nil
// Return the user_id as string (preferred approach)
if userID, ok := userData["user_id"].(string); ok {
return userID, nil
}
// Fallback: convert the returned int id to string
return fmt.Sprintf("%d", id), nil
}
// UpdateUser updates user information (excludes sensitive fields like password, MFA)

View file

@ -163,7 +163,7 @@ type UserProvider interface {
UpdatePassword(ctx context.Context, userID string, newPassword string) error
ResetPassword(ctx context.Context, userID string) (string, error)
CreateUser(ctx context.Context, userData maps.MapStrAny) (interface{}, error)
CreateUser(ctx context.Context, userData maps.MapStrAny) (string, error)
UpdateUser(ctx context.Context, userID string, userData maps.MapStrAny) error
DeleteUser(ctx context.Context, userID string) error
UpdateUserLastLogin(ctx context.Context, userID string) error
@ -217,7 +217,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)
CreateRole(ctx context.Context, roleData maps.MapStrAny) (string, error)
UpdateRole(ctx context.Context, roleID string, roleData maps.MapStrAny) error
DeleteRole(ctx context.Context, roleID string) error
@ -235,7 +235,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)
CreateType(ctx context.Context, typeData maps.MapStrAny) (string, error)
UpdateType(ctx context.Context, typeID string, typeData maps.MapStrAny) error
DeleteType(ctx context.Context, typeID string) error

View file

@ -0,0 +1,20 @@
package types
// Map converts the OIDCUserInfo to a map[string]interface{}
func (user OIDCUserInfo) Map() map[string]interface{} {
return map[string]interface{}{
"sub": user.Sub,
"name": user.Name,
"given_name": user.GivenName,
"family_name": user.FamilyName,
"middle_name": user.MiddleName,
"nickname": user.Nickname,
"preferred_username": user.PreferredUsername,
"profile": user.Profile,
"picture": user.Picture,
"website": user.Website,
"email": user.Email,
"email_verified": user.EmailVerified,
"gender": user.Gender,
}
}

View file

@ -193,8 +193,6 @@ func authback(c *gin.Context) {
userInfo, err = provider.GetUserInfo(tokenResponse.AccessToken, tokenResponse.TokenType)
}
// Create / Update / User then login (Generate access_token and id_token)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
@ -204,12 +202,18 @@ func authback(c *gin.Context) {
return
}
// Respond with success
response.RespondWithSuccess(c, response.StatusOK, map[string]interface{}{
"params": params,
"token": tokenResponse,
"user": userInfo,
})
// LoginThirdParty(providerID, userInfo)
loginResponse, err := LoginThirdParty(providerID, userInfo)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Failed to login: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
response.RespondWithSuccess(c, response.StatusOK, loginResponse)
}
// getOAuthAuthorizationURL generates OAuth authorization URL for a provider

105
openapi/signin/login.go Normal file
View file

@ -0,0 +1,105 @@
package signin
import (
"context"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/openapi/oauth"
"github.com/yaoapp/yao/openapi/oauth/providers/user"
oauthtypes "github.com/yaoapp/yao/openapi/oauth/types"
)
// LoginThirdParty is the handler for third party login
func LoginThirdParty(providerID string, userinfo *oauthtypes.OIDCUserInfo) (*LoginResponse, error) {
// Get provider
provider, err := GetProvider(providerID)
if err != nil {
return nil, err
}
// Check if user exists
userProvider, err := oauth.OAuth.GetUserProvider()
if err != nil {
return nil, err
}
// Auto register user if not exists
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var userID string
// Auto register user if not exists
if provider.Register != nil && provider.Register.Auto {
userID, err = userProvider.GetOAuthUserID(ctx, providerID, userinfo.Sub)
if err != nil && err.Error() == user.ErrOAuthAccountNotFound {
userData := map[string]interface{}{
"name": userinfo.Name,
"given_name": userinfo.GivenName,
"family_name": userinfo.FamilyName,
"picture": userinfo.Picture,
"role_id": provider.Register.Role,
"status": "active",
}
// Auto register user
userID, err = userProvider.CreateUser(ctx, userData)
if err != nil {
return nil, err
}
// Create OAuth account
userData = userinfo.Map()
userData["provider"] = providerID
_, err = userProvider.CreateOAuthAccount(ctx, userID, userData)
if err != nil {
return nil, err
}
}
}
// Get User ID from OAuth account
userID, err = userProvider.GetOAuthUserID(ctx, providerID, userinfo.Sub)
if err != nil {
return nil, err
}
return LoginByUserID(userID)
}
// LoginByUserID is the handler for login
func LoginByUserID(userid string) (*LoginResponse, error) {
// Get User
userProvider, err := oauth.OAuth.GetUserProvider()
if err != nil {
return nil, err
}
// Get User
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
user, err := userProvider.GetUser(ctx, userid)
if err != nil {
return nil, err
}
// Update Last Login
err = userProvider.UpdateUserLastLogin(ctx, userid)
if err != nil {
log.Warn("Failed to update last login: %s", err.Error())
}
return &LoginResponse{
AccessToken: "mock_access_token",
IDToken: "mock_id_token",
RefreshToken: "mock_refresh_token",
ExpiresIn: 3600,
TokenType: "Bearer",
Scope: "openid profile email",
User: user,
}, nil
}

View file

@ -393,14 +393,6 @@ func createPublicConfig(fullConfig *Config) Config {
if fullConfig.ThirdParty != nil {
publicConfig.ThirdParty = &ThirdParty{}
// Deep copy Register configuration
if fullConfig.ThirdParty.Register != nil {
publicConfig.ThirdParty.Register = &RegisterConfig{
Auto: fullConfig.ThirdParty.Register.Auto,
Role: fullConfig.ThirdParty.Register.Role,
}
}
// Deep copy Providers with sensitive data removal
if fullConfig.ThirdParty.Providers != nil {
publicProviders := make([]*Provider, len(fullConfig.ThirdParty.Providers))
@ -412,7 +404,7 @@ func createPublicConfig(fullConfig *Config) Config {
Color: provider.Color,
TextColor: provider.TextColor,
// Only expose display fields for frontend
// Remove sensitive fields: ClientID, ClientSecret, ClientSecretGenerator, Scopes, Endpoints, Mapping
// Remove sensitive fields: ClientID, ClientSecret, ClientSecretGenerator, Scopes, Endpoints, Mapping, Register
}
publicProviders[i] = &publicProvider

View file

@ -55,8 +55,7 @@ type TokenConfig struct {
// ThirdParty represents the third party login configuration
type ThirdParty struct {
Register *RegisterConfig `json:"register,omitempty"`
Providers []*Provider `json:"providers,omitempty"`
Providers []*Provider `json:"providers,omitempty"`
}
// RegisterConfig represents the auto register configuration
@ -80,6 +79,7 @@ type Provider struct {
UserInfoSource string `json:"user_info_source,omitempty"` // "endpoint" (default) | "id_token" | "access_token"
Endpoints *Endpoints `json:"endpoints,omitempty"`
Mapping interface{} `json:"mapping,omitempty"` // string (preset) | map[string]string (custom) | nil (generic)
Register *RegisterConfig `json:"register,omitempty"`
}
// SecretGenerator represents the client secret generator configuration
@ -151,6 +151,17 @@ type OAuthUserInfoResponse = oauthtypes.OIDCUserInfo
// OIDCAddress is an alias for OIDC standard address claim type
type OIDCAddress = oauthtypes.OIDCAddress
// LoginResponse represents the response for login
type LoginResponse struct {
AccessToken string `json:"access_token"`
IDToken string `json:"id_token,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
ExpiresIn int `json:"expires_in,omitempty"`
TokenType string `json:"token_type,omitempty"`
Scope string `json:"scope,omitempty"`
User map[string]interface{} `json:"user,omitempty"`
}
// Built-in preset mapping types
const (
MappingGoogle = "google"

View file

@ -78,6 +78,7 @@ func TestSigninGetConfigs(t *testing.T) {
assert.Empty(t, publicProvider.Scopes, "Scopes should be empty in ThirdParty providers")
assert.Nil(t, publicProvider.Endpoints, "Endpoints should be nil in ThirdParty providers")
assert.Empty(t, publicProvider.Mapping, "Mapping should be empty in ThirdParty providers")
assert.Nil(t, publicProvider.Register, "Register config should be nil in public config (sensitive data)")
}
}
}
@ -178,6 +179,7 @@ func TestSigninConfigStructure(t *testing.T) {
assert.Empty(t, provider.Scopes, "Provider scopes should be empty in ThirdParty providers")
assert.Nil(t, provider.Mapping, "Provider mapping should be nil in ThirdParty providers")
assert.Nil(t, provider.Endpoints, "Provider endpoints should be nil in ThirdParty providers")
assert.Nil(t, provider.Register, "Provider register should be nil in ThirdParty providers (it's in global map now)")
}
}
}
@ -225,6 +227,14 @@ func TestSigninGlobalProvidersMap(t *testing.T) {
assert.IsType(t, "", provider.Endpoints.Token, "Token endpoint should be string")
assert.IsType(t, "", provider.Endpoints.UserInfo, "UserInfo endpoint should be string")
}
// Test register configuration (should be present in global providers)
if provider.Register != nil {
assert.IsType(t, false, provider.Register.Auto, "Register auto should be boolean")
assert.IsType(t, "", provider.Register.Role, "Register role should be string")
t.Logf("Provider '%s' has register config: auto=%t, role=%s",
providerID, provider.Register.Auto, provider.Register.Role)
}
}
})
}