diff --git a/internal/common/env_config.go b/internal/common/env_config.go index 82a421f..fc4824f 100644 --- a/internal/common/env_config.go +++ b/internal/common/env_config.go @@ -54,6 +54,8 @@ type EnvConfigSchema struct { VersionCheckDisabled bool `env:"VERSION_CHECK_DISABLED"` StaticApiKey string `env:"STATIC_API_KEY" options:"file"` + OidcEmailDomainRewrites string `env:"OIDC_EMAIL_DOMAIN_REWRITES"` + FileBackend string `env:"FILE_BACKEND" options:"toLower"` UploadPath string `env:"UPLOAD_PATH"` S3Bucket string `env:"S3_BUCKET"` diff --git a/internal/service/oidc_service.go b/internal/service/oidc_service.go index e493a5b..9cd4116 100644 --- a/internal/service/oidc_service.go +++ b/internal/service/oidc_service.go @@ -587,7 +587,7 @@ func (s *OidcService) createTokenFromRefreshToken(ctx context.Context, input dto } // Load the profile, which we need for the ID token - userClaims, err := s.getUserClaims(ctx, &storedRefreshToken.User, storedRefreshToken.Scopes(), tx) + userClaims, err := s.getUserClaims(ctx, &storedRefreshToken.User, storedRefreshToken.Scopes(), input.ClientID, tx) if err != nil { return CreatedTokens{}, err } @@ -1923,7 +1923,7 @@ func (s *OidcService) GetClientPreview(ctx context.Context, clientID string, use return nil, &common.OidcAccessDeniedError{} } - userClaims, err := s.getUserClaims(ctx, &user, scopes, tx) + userClaims, err := s.getUserClaims(ctx, &user, scopes, clientID, tx) if err != nil { return nil, err } @@ -1976,15 +1976,15 @@ func (s *OidcService) getUserClaimsForClientInternal(ctx context.Context, userID return nil, err } - return s.getUserClaims(ctx, &authorizedOidcClient.User, authorizedOidcClient.Scopes(), tx) + return s.getUserClaims(ctx, &authorizedOidcClient.User, authorizedOidcClient.Scopes(), clientID, tx) } -func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scopes []string, tx *gorm.DB) (map[string]any, error) { +func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scopes []string, clientID string, tx *gorm.DB) (map[string]any, error) { claims := make(map[string]any, 10) claims["sub"] = user.ID if slices.Contains(scopes, "email") { - claims["email"] = user.Email + claims["email"] = rewriteEmailDomainForClient(user.Email, clientID) claims["email_verified"] = user.EmailVerified } @@ -2027,12 +2027,42 @@ func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scope } if slices.Contains(scopes, "email") { - claims["email"] = user.Email + claims["email"] = rewriteEmailDomainForClient(user.Email, clientID) } return claims, nil } +func rewriteEmailDomainForClient(email *string, clientID string) *string { + if email == nil { + return nil + } + + domain := oidcEmailRewriteDomain(clientID) + if domain == "" { + return email + } + + localPart, _, ok := strings.Cut(*email, "@") + if !ok || localPart == "" { + return email + } + + rewritten := localPart + "@" + domain + return &rewritten +} + +func oidcEmailRewriteDomain(clientID string) string { + for entry := range strings.SplitSeq(common.EnvConfig.OidcEmailDomainRewrites, ",") { + client, domain, ok := strings.Cut(strings.TrimSpace(entry), "=") + if ok && strings.TrimSpace(client) == clientID { + return strings.TrimSpace(domain) + } + } + + return "" +} + func (s *OidcService) IsClientAccessibleToUser(ctx context.Context, clientID string, userID string) (bool, error) { var user model.User err := s.db.WithContext(ctx).Preload("UserGroups").First(&user, "id = ?", userID).Error