Merge pull request #1075 from trheyi/main

Refactor user provider methods to return string IDs and update test d…
This commit is contained in:
Max 2025-08-03 10:26:17 +08:00 committed by GitHub
commit 7fbfb58b78
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
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 // GetUserProvider returns the user provider for the service
func (s *Service) GetUserProvider() types.UserProvider { func (s *Service) GetUserProvider() (types.UserProvider, error) {
return s.userProvider return s.userProvider, nil
} }
// GetClientProvider returns the client provider for the service // 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 // Create test users using the updated user provider interface
userProvider := service.GetUserProvider() userProvider, _ := service.GetUserProvider()
for i, testUser := range testUsers { for i, testUser := range testUsers {
// Convert TestUser to the format expected by CreateUser // Convert TestUser to the format expected by CreateUser
userData := map[string]interface{}{ userData := map[string]interface{}{
@ -505,16 +505,13 @@ func setupTestData(t *testing.T, service *Service) {
createdUserID, err := userProvider.CreateUser(ctx, userData) createdUserID, err := userProvider.CreateUser(ctx, userData)
require.NoError(t, err, "Failed to create test user %d: %s", i, testUser.Description) 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 // The CreateUser method now returns the user_id as string directly
if userID, ok := createdUserID.(int64); ok { testUser.UserID = createdUserID
testUser.ID = userID
} else if userID, ok := createdUserID.(int); ok { // For backward compatibility, also store as string in userData if needed
testUser.ID = int64(userID) userData["user_id"] = createdUserID
} else {
testUser.ID = int64(0) // Fallback for interface{} types
}
// Extract the auto-generated user_id from userData (CreateUser sets it) // Extract the auto-generated user_id from userData (CreateUser sets it)
if generatedUserID, ok := userData["user_id"].(string); ok { 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) { t.Run("get user provider", func(t *testing.T) {
userProvider := service.GetUserProvider() userProvider, _ := service.GetUserProvider()
assert.NotNil(t, userProvider) assert.NotNil(t, userProvider)
assert.Implements(t, (*types.UserProvider)(nil), userProvider) assert.Implements(t, (*types.UserProvider)(nil), userProvider)
}) })
@ -864,7 +861,7 @@ func TestProviderInitialization(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Create custom providers (for this test, we'll use the default ones) // Create custom providers (for this test, we'll use the default ones)
customUserProvider := tempService.GetUserProvider() customUserProvider, _ := tempService.GetUserProvider()
customClientProvider := tempService.GetClientProvider() customClientProvider := tempService.GetClientProvider()
config := &Config{ 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 // 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 // Validate required role_id field
if _, exists := roleData["role_id"]; !exists { 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 // 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) m := model.Select(u.roleModel)
id, err := m.Create(roleData) id, err := m.Create(roleData)
if err != nil { 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 // 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 // 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 // Validate required type_id field
if _, exists := typeData["type_id"]; !exists { 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 // 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) m := model.Select(u.typeModel)
id, err := m.Create(typeData) id, err := m.Create(typeData)
if err != nil { 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 // 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 // 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 // Auto-generate user_id if not provided
if _, exists := userData["user_id"]; !exists { if _, exists := userData["user_id"]; !exists {
userID, err := u.GenerateUserID(ctx, true) // Force safe mode to ensure uniqueness userID, err := u.GenerateUserID(ctx, true) // Force safe mode to ensure uniqueness
if err != nil { if err != nil {
return nil, fmt.Errorf(ErrFailedToGenerateUserID, err) return "", fmt.Errorf(ErrFailedToGenerateUserID, err)
} }
userData["user_id"] = userID userData["user_id"] = userID
} }
@ -268,10 +268,16 @@ func (u *DefaultUser) CreateUser(ctx context.Context, userData maps.MapStrAny) (
m := model.Select(u.model) m := model.Select(u.model)
id, err := m.Create(userData) id, err := m.Create(userData)
if err != nil { 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) // 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 UpdatePassword(ctx context.Context, userID string, newPassword string) error
ResetPassword(ctx context.Context, userID string) (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 UpdateUser(ctx context.Context, userID string, userData maps.MapStrAny) error
DeleteUser(ctx context.Context, userID string) error DeleteUser(ctx context.Context, userID string) error
UpdateUserLastLogin(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) GetRole(ctx context.Context, roleID string) (maps.MapStrAny, error)
RoleExists(ctx context.Context, roleID string) (bool, 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 UpdateRole(ctx context.Context, roleID string, roleData maps.MapStrAny) error
DeleteRole(ctx context.Context, roleID string) error DeleteRole(ctx context.Context, roleID string) error
@ -235,7 +235,7 @@ type UserProvider interface {
GetType(ctx context.Context, typeID string) (maps.MapStrAny, error) GetType(ctx context.Context, typeID string) (maps.MapStrAny, error)
TypeExists(ctx context.Context, typeID string) (bool, 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 UpdateType(ctx context.Context, typeID string, typeData maps.MapStrAny) error
DeleteType(ctx context.Context, typeID string) 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) userInfo, err = provider.GetUserInfo(tokenResponse.AccessToken, tokenResponse.TokenType)
} }
// Create / Update / User then login (Generate access_token and id_token)
if err != nil { if err != nil {
errorResp := &response.ErrorResponse{ errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code, Code: response.ErrInvalidRequest.Code,
@ -204,12 +202,18 @@ func authback(c *gin.Context) {
return return
} }
// Respond with success // LoginThirdParty(providerID, userInfo)
response.RespondWithSuccess(c, response.StatusOK, map[string]interface{}{ loginResponse, err := LoginThirdParty(providerID, userInfo)
"params": params, if err != nil {
"token": tokenResponse, errorResp := &response.ErrorResponse{
"user": userInfo, 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 // 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 { if fullConfig.ThirdParty != nil {
publicConfig.ThirdParty = &ThirdParty{} 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 // Deep copy Providers with sensitive data removal
if fullConfig.ThirdParty.Providers != nil { if fullConfig.ThirdParty.Providers != nil {
publicProviders := make([]*Provider, len(fullConfig.ThirdParty.Providers)) publicProviders := make([]*Provider, len(fullConfig.ThirdParty.Providers))
@ -412,7 +404,7 @@ func createPublicConfig(fullConfig *Config) Config {
Color: provider.Color, Color: provider.Color,
TextColor: provider.TextColor, TextColor: provider.TextColor,
// Only expose display fields for frontend // 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 publicProviders[i] = &publicProvider

View file

@ -55,7 +55,6 @@ type TokenConfig struct {
// ThirdParty represents the third party login configuration // ThirdParty represents the third party login configuration
type ThirdParty struct { type ThirdParty struct {
Register *RegisterConfig `json:"register,omitempty"`
Providers []*Provider `json:"providers,omitempty"` Providers []*Provider `json:"providers,omitempty"`
} }
@ -80,6 +79,7 @@ type Provider struct {
UserInfoSource string `json:"user_info_source,omitempty"` // "endpoint" (default) | "id_token" | "access_token" UserInfoSource string `json:"user_info_source,omitempty"` // "endpoint" (default) | "id_token" | "access_token"
Endpoints *Endpoints `json:"endpoints,omitempty"` Endpoints *Endpoints `json:"endpoints,omitempty"`
Mapping interface{} `json:"mapping,omitempty"` // string (preset) | map[string]string (custom) | nil (generic) 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 // 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 // OIDCAddress is an alias for OIDC standard address claim type
type OIDCAddress = oauthtypes.OIDCAddress 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 // Built-in preset mapping types
const ( const (
MappingGoogle = "google" 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.Empty(t, publicProvider.Scopes, "Scopes should be empty in ThirdParty providers")
assert.Nil(t, publicProvider.Endpoints, "Endpoints should be nil 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.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.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.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.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.Token, "Token endpoint should be string")
assert.IsType(t, "", provider.Endpoints.UserInfo, "UserInfo 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)
}
} }
}) })
} }