Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 101 additions & 10 deletions api/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"log/slog"
"net/http"
"net/url"
"slices"
"strings"
"time"

Expand Down Expand Up @@ -62,6 +63,9 @@ func NewOIDC(conf *config.Configuration, db *database.GormDatabase, userChangeNo
Provider: provider,
UserChangeNotifier: userChangeNotifier,
UsernameClaim: conf.OIDC.UsernameClaim,
GroupsClaim: conf.OIDC.GroupsClaim,
GroupsUser: conf.OIDC.GroupsUser,
GroupsAdmin: conf.OIDC.GroupsAdmin,
PasswordStrength: conf.PassStrength,
SecureCookie: conf.Server.SecureCookie,
AutoRegister: conf.OIDC.AutoRegister,
Expand Down Expand Up @@ -91,6 +95,9 @@ type OIDCAPI struct {
Provider rp.RelyingParty
UserChangeNotifier *UserChangeNotifier
UsernameClaim string
GroupsClaim string
GroupsUser []string
GroupsAdmin []string
PasswordStrength int
SecureCookie bool
AutoRegister bool
Expand Down Expand Up @@ -216,7 +223,7 @@ func (a *OIDCAPI) promptURLParams() []rp.URLParamOpt {
// $ref: "#/definitions/Error"
func (a *OIDCAPI) CallbackHandler() gin.HandlerFunc {
callback := func(w http.ResponseWriter, r *http.Request, tokens *oidc.Tokens[*oidc.IDTokenClaims], state string, provider rp.RelyingParty, info *oidc.UserInfo) {
user, status, err := a.resolveUser(tokens.IDTokenClaims.GetIssuer(), info)
user, status, err := a.resolveUser(tokens.IDTokenClaims, info)
if err != nil {
http.Error(w, err.Error(), status)
return
Expand Down Expand Up @@ -385,7 +392,7 @@ func (a *OIDCAPI) ExternalTokenHandler(ctx *gin.Context) {
ctx.AbortWithError(http.StatusInternalServerError, fmt.Errorf("failed to get user info: %w", err))
return
}
user, status, resolveErr := a.resolveUser(tokens.IDTokenClaims.GetIssuer(), info)
user, status, resolveErr := a.resolveUser(tokens.IDTokenClaims, info)
if resolveErr != nil {
ctx.AbortWithError(status, resolveErr)
return
Expand Down Expand Up @@ -416,7 +423,8 @@ func (a *OIDCAPI) generateState() (string, error) {
// this OIDC identity, which requires GOTIFY_OIDC_LINK_BY_USERNAME and
// that the user is not already bound to a different identity.
// 3. Otherwise auto-register a new user, which requires GOTIFY_OIDC_AUTOREGISTER.
func (a *OIDCAPI) resolveUser(issuer string, info *oidc.UserInfo) (*model.User, int, error) {
func (a *OIDCAPI) resolveUser(idToken *oidc.IDTokenClaims, info *oidc.UserInfo) (*model.User, int, error) {
issuer := idToken.GetIssuer()
if issuer == "" {
return nil, http.StatusInternalServerError, errors.New("issuer claim was empty")
}
Expand All @@ -436,11 +444,25 @@ func (a *OIDCAPI) resolveUser(issuer string, info *oidc.UserInfo) (*model.User,
if err != nil {
return nil, http.StatusInternalServerError, fmt.Errorf("database error: %w", err)
}

hasAdminGroup, status, err := a.resolvePermission(idToken.Claims, info.Claims)
if err != nil {
log.Err(err).Str("oidc_id", oidcID).Interface("idTokenClaims", idToken.Claims).Interface("userinfoClaims", info.Claims).Msg("OIDC: resolve permission")
return nil, status, err
}

if user != nil {
if len(a.GroupsAdmin) > 0 && user.Admin != hasAdminGroup {
user.Admin = hasAdminGroup
if err := a.DB.UpdateUser(user); err != nil {
return nil, http.StatusInternalServerError, fmt.Errorf("database error: %w", err)
}
log.Warn().Str("oidc_id", oidcID).Str("username", user.Name).Bool("admin", user.Admin).Msg("OIDC change permission")
}
return user, 0, nil
}

usernameRaw, ok := info.Claims[a.UsernameClaim]
usernameRaw, ok := lookupClaim(a.UsernameClaim, idToken.Claims, info.Claims)
if !ok {
return nil, http.StatusInternalServerError, fmt.Errorf("username claim %q is missing", a.UsernameClaim)
}
Expand All @@ -454,12 +476,12 @@ func (a *OIDCAPI) resolveUser(issuer string, info *oidc.UserInfo) (*model.User,
return nil, http.StatusInternalServerError, fmt.Errorf("database error: %w", err)
}
if byUsername != nil {
return a.linkExistingUser(byUsername, oidcID)
return a.linkExistingUser(byUsername, oidcID, hasAdminGroup)
}
return a.registerUser(username, oidcID)
return a.registerUser(username, oidcID, hasAdminGroup)
}

func (a *OIDCAPI) linkExistingUser(user *model.User, oidcID string) (*model.User, int, error) {
func (a *OIDCAPI) linkExistingUser(user *model.User, oidcID string, hasAdminGroup bool) (*model.User, int, error) {
if !a.LinkByUsername {
log.Warn().Str("oidc_id", oidcID).Str("username", user.Name).Msgf("OIDC login rejected: a local user with the username already exists and %s is disabled", config.EnvOIDCLinkByUsername)
return nil, http.StatusForbidden, fmt.Errorf("a local user with the username %s already exists and linking by username is disabled", user.Name)
Expand All @@ -469,21 +491,34 @@ func (a *OIDCAPI) linkExistingUser(user *model.User, oidcID string) (*model.User
return nil, http.StatusForbidden, fmt.Errorf("the user %s is already bound to a different OIDC identity", user.Name)
}
user.OIDCID = &oidcID
if len(a.GroupsAdmin) > 0 {
user.Admin = hasAdminGroup
}
if err := a.DB.UpdateUser(user); err != nil {
return nil, http.StatusInternalServerError, fmt.Errorf("failed to bind user to OIDC identity: %w", err)
}
log.Warn().Str("oidc_id", oidcID).Str("username", user.Name).Bool("admin", user.Admin).Msg("OIDC link by username")
return user, 0, nil
}

func (a *OIDCAPI) registerUser(username, oidcID string) (*model.User, int, error) {
func (a *OIDCAPI) registerUser(username, oidcID string, hasAdminGroup bool) (*model.User, int, error) {
if !a.AutoRegister {
return nil, http.StatusForbidden, errors.New("user does not exist and auto-registration is disabled")
}
user := &model.User{Name: username, Admin: false, Pass: nil, OIDCID: &oidcID}
user := &model.User{
Name: username,
Pass: nil,
OIDCID: &oidcID,
}

if len(a.GroupsAdmin) > 0 {
user.Admin = hasAdminGroup
}

if err := a.DB.CreateUser(user); err != nil {
return nil, http.StatusInternalServerError, fmt.Errorf("failed to create user: %w", err)
}
log.Info().Str("oidc_id", oidcID).Str("username", user.Name).Msg("OIDC auto registration")
log.Info().Str("oidc_id", oidcID).Str("username", user.Name).Bool("admin", user.Admin).Msg("OIDC auto registration")
if err := a.UserChangeNotifier.fireUserAdded(user.ID); err != nil {
log.Error().Err(err).Uint("user_id", user.ID).Msg("Could not notify user change")
}
Expand Down Expand Up @@ -514,3 +549,59 @@ func (a *OIDCAPI) popPendingSession(key string) (*pendingOIDCSession, bool) {
}
return nil, false
}

func (a *OIDCAPI) resolvePermission(idTokenClaims, userInfoClaims map[string]any) (bool, int, error) {
if a.GroupsClaim == "" {
return false, 0, nil
}

groupsRaw, ok := lookupClaim(a.GroupsClaim, idTokenClaims, userInfoClaims)
if !ok {
return false, http.StatusInternalServerError, fmt.Errorf("groups claim %q is missing", a.GroupsClaim)
}

var groups []string
switch groupsRaw := groupsRaw.(type) {
case []string:
groups = groupsRaw
case []any:
for _, groupRaw := range groupsRaw {
group, ok := groupRaw.(string)
if !ok {
return false, http.StatusInternalServerError, fmt.Errorf("groups claim %q contains a non-string element: %#v", a.GroupsClaim, groupRaw)
}
groups = append(groups, group)
}
case string:
groups = append(groups, groupsRaw)
default:
return false, http.StatusInternalServerError, fmt.Errorf("groups claim %q is not a string or string array: %#v", a.GroupsClaim, groupsRaw)
}

switch {
case containsAny(a.GroupsAdmin, groups):
return true, 0, nil
case len(a.GroupsUser) == 0 || containsAny(a.GroupsUser, groups):
return false, 0, nil
default:
return false, http.StatusForbidden, errors.New("user is not in any allowed group")
}
}

func lookupClaim(name string, idTokenClaims, userInfoClaims map[string]any) (any, bool) {
if value, ok := idTokenClaims[name]; ok {
return value, true
}
value, ok := userInfoClaims[name]
return value, ok
}

func containsAny(configured, actual []string) bool {
for _, value := range actual {
if slices.Contains(configured, value) {
return true
}
}

return false
}
Loading
Loading