diff --git a/controller/custom_oauth.go b/controller/custom_oauth.go index c21ec7910bce..37a3170db87c 100644 --- a/controller/custom_oauth.go +++ b/controller/custom_oauth.go @@ -32,6 +32,7 @@ type CustomOAuthProviderResponse struct { UsernameField string `json:"username_field"` DisplayNameField string `json:"display_name_field"` EmailField string `json:"email_field"` + GroupField string `json:"group_field"` WellKnown string `json:"well_known"` AuthStyle int `json:"auth_style"` AccessPolicy string `json:"access_policy"` @@ -62,6 +63,7 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro UsernameField: p.UsernameField, DisplayNameField: p.DisplayNameField, EmailField: p.EmailField, + GroupField: p.GroupField, WellKnown: p.WellKnown, AuthStyle: p.AuthStyle, AccessPolicy: p.AccessPolicy, @@ -127,6 +129,7 @@ type CreateCustomOAuthProviderRequest struct { UsernameField string `json:"username_field"` DisplayNameField string `json:"display_name_field"` EmailField string `json:"email_field"` + GroupField string `json:"group_field"` WellKnown string `json:"well_known"` AuthStyle int `json:"auth_style"` AccessPolicy string `json:"access_policy"` @@ -245,6 +248,7 @@ func CreateCustomOAuthProvider(c *gin.Context) { UsernameField: req.UsernameField, DisplayNameField: req.DisplayNameField, EmailField: req.EmailField, + GroupField: req.GroupField, WellKnown: req.WellKnown, AuthStyle: req.AuthStyle, AccessPolicy: req.AccessPolicy, @@ -282,6 +286,7 @@ type UpdateCustomOAuthProviderRequest struct { UsernameField string `json:"username_field"` DisplayNameField string `json:"display_name_field"` EmailField string `json:"email_field"` + GroupField string `json:"group_field"` WellKnown *string `json:"well_known"` // Optional: if nil, keep existing AuthStyle *int `json:"auth_style"` // Optional: if nil, keep existing AccessPolicy *string `json:"access_policy"` // Optional: if nil, keep existing @@ -368,6 +373,9 @@ func UpdateCustomOAuthProvider(c *gin.Context) { if req.EmailField != "" { provider.EmailField = req.EmailField } + if req.GroupField != "" { + provider.GroupField = req.GroupField + } if req.WellKnown != nil { provider.WellKnown = *req.WellKnown } diff --git a/controller/oauth.go b/controller/oauth.go index 9951f22b035f..0083c6901e4b 100644 --- a/controller/oauth.go +++ b/controller/oauth.go @@ -9,6 +9,7 @@ import ( "github.com/QuantumNous/new-api/i18n" "github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/oauth" + "github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -209,6 +210,29 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o if user.Id == 0 { return nil, &OAuthUserDeletedError{} } + + // Update user's group if OAuth provides a different group and it's in available groups + if oauthUser.Group != "" && oauthUser.Group != user.Group { + common.SysLog(fmt.Sprintf("[OAuth] User %d current group: '%s', OAuth group: '%s'", user.Id, user.Group, oauthUser.Group)) + // Check if group exists in group ratio settings + if ratio_setting.ContainsGroupRatio(oauthUser.Group) { + user.Group = oauthUser.Group + if err := user.Update(false); err != nil { + common.SysError(fmt.Sprintf("[OAuth] Failed to update user %d group to '%s': %s", user.Id, oauthUser.Group, err.Error())) + } else { + common.SysLog(fmt.Sprintf("[OAuth] Updated user %d group to '%s' from OAuth provider", user.Id, oauthUser.Group)) + } + } else { + common.SysLog(fmt.Sprintf("[OAuth] OAuth group '%s' not in group ratio settings for user %d, keeping current group '%s'", oauthUser.Group, user.Id, user.Group)) + } + } else { + if oauthUser.Group == "" { + common.SysLog(fmt.Sprintf("[OAuth] User %d OAuth group is empty, skipping group update", user.Id)) + } else if oauthUser.Group == user.Group { + common.SysLog(fmt.Sprintf("[OAuth] User %d group '%s' already matches OAuth group, no update needed", user.Id, user.Group)) + } + } + return user, nil } @@ -262,6 +286,17 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o user.Role = common.RoleCommonUser user.Status = common.UserStatusEnabled + // Auto-assign group from OAuth provider if configured + if oauthUser.Group != "" { + // Check if the group from OAuth is in the platform's group ratio settings + if ratio_setting.ContainsGroupRatio(oauthUser.Group) { + user.Group = oauthUser.Group + common.SysLog(fmt.Sprintf("[OAuth] Auto-assigned group '%s' to new user from OAuth provider (matched group ratio settings)", oauthUser.Group)) + } else { + common.SysLog(fmt.Sprintf("[OAuth] Group '%s' from OAuth provider not found in group ratio settings, using default 'default'", oauthUser.Group)) + } + } + // Handle affiliate code affCode := session.Get("aff") inviterId := 0 diff --git a/go.mod b/go.mod index 29f797e0b7c1..bc400524ae0a 100644 --- a/go.mod +++ b/go.mod @@ -96,7 +96,7 @@ require ( github.com/icza/bitio v1.1.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect - github.com/jackc/pgx/v5 v5.7.1 // indirect + github.com/jackc/pgx/v5 v5.9.1 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/jfreymuth/vorbis v1.0.2 // indirect github.com/jinzhu/inflection v1.0.0 // indirect diff --git a/go.sum b/go.sum index 221e14672254..b90a7e1a26f0 100644 --- a/go.sum +++ b/go.sum @@ -152,8 +152,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= -github.com/jackc/pgx/v5 v5.7.1 h1:x7SYsPBYDkHDksogeSmZZ5xzThcTgRz++I5E+ePFUcs= -github.com/jackc/pgx/v5 v5.7.1/go.mod h1:e7O26IywZZ+naJtWWos6i6fvWK+29etgITqrqHLfoZA= +github.com/jackc/pgx/v5 v5.9.1 h1:uwrxJXBnx76nyISkhr33kQLlUqjv7et7b9FjCen/tdc= +github.com/jackc/pgx/v5 v5.9.1/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jfreymuth/oggvorbis v1.0.5 h1:u+Ck+R0eLSRhgq8WTmffYnrVtSztJcYrl588DM4e3kQ= diff --git a/model/custom_oauth_provider.go b/model/custom_oauth_provider.go index 12b4d11113d0..e6daf809d5a0 100644 --- a/model/custom_oauth_provider.go +++ b/model/custom_oauth_provider.go @@ -55,6 +55,7 @@ type CustomOAuthProvider struct { UsernameField string `json:"username_field" gorm:"type:varchar(128);default:'preferred_username'"` // Username field path DisplayNameField string `json:"display_name_field" gorm:"type:varchar(128);default:'name'"` // Display name field path EmailField string `json:"email_field" gorm:"type:varchar(128);default:'email'"` // Email field path + GroupField string `json:"group_field" gorm:"type:varchar(128);default:''"` // Group field path for auto-assigning user group, e.g., "groups", "roles", "data.group" // Advanced options WellKnown string `json:"well_known" gorm:"type:varchar(512)"` // OIDC discovery endpoint (optional) diff --git a/oauth/generic.go b/oauth/generic.go index 11bbb9b625f6..8835836ab242 100644 --- a/oauth/generic.go +++ b/oauth/generic.go @@ -243,6 +243,16 @@ func (p *GenericOAuthProvider) GetUserInfo(ctx context.Context, token *OAuthToke username := gjson.Get(bodyStr, p.config.UsernameField).String() displayName := gjson.Get(bodyStr, p.config.DisplayNameField).String() email := gjson.Get(bodyStr, p.config.EmailField).String() + group := "" + if p.config.GroupField != "" { + groupResult := gjson.Get(bodyStr, p.config.GroupField) + // If result is an array, take the first element + if groupResult.IsArray() && len(groupResult.Array()) > 0 { + group = groupResult.Array()[0].String() + } else { + group = groupResult.String() + } + } // If user ID field returns a number, convert it if userId == "" { @@ -260,8 +270,8 @@ func (p *GenericOAuthProvider) GetUserInfo(ctx context.Context, token *OAuthToke return nil, NewOAuthError(i18n.MsgOAuthUserInfoEmpty, map[string]any{"Provider": p.config.Name}) } - logger.LogDebug(ctx, "[OAuth-Generic-%s] GetUserInfo success: id=%s, username=%s, name=%s, email=%s", - p.config.Slug, userId, username, displayName, email) + common.SysLog(fmt.Sprintf("[OAuth-Generic-%s] GetUserInfo success: id=%s, username=%s, name=%s, email=%s, group=%s", + p.config.Slug, userId, username, displayName, email, group)) policyRaw := strings.TrimSpace(p.config.AccessPolicy) if policyRaw != "" { @@ -284,6 +294,7 @@ func (p *GenericOAuthProvider) GetUserInfo(ctx context.Context, token *OAuthToke Username: username, DisplayName: displayName, Email: email, + Group: group, Extra: map[string]any{ "provider": p.config.Slug, }, diff --git a/oauth/types.go b/oauth/types.go index 383e6f351302..ef050f39c6e2 100644 --- a/oauth/types.go +++ b/oauth/types.go @@ -20,6 +20,8 @@ type OAuthUser struct { DisplayName string // Email is the email from the OAuth provider Email string + // Group is the group/role from the OAuth provider (optional, used for auto-assigning user group) + Group string // Extra contains any additional provider-specific data Extra map[string]any } diff --git a/web/src/components/settings/CustomOAuthSetting.jsx b/web/src/components/settings/CustomOAuthSetting.jsx index 0912160bee5d..a52b88b178ff 100644 --- a/web/src/components/settings/CustomOAuthSetting.jsx +++ b/web/src/components/settings/CustomOAuthSetting.jsx @@ -962,6 +962,17 @@ const CustomOAuthSetting = ({ serverAddress }) => { + + + + + +