char/topaz-flake
git clone https://git.t4t.associates/char/topaz-flake
744b7fb
main
1diff --git a/internal/common/env_config.go b/internal/common/env_config.go 2index82a421f ..fc4824f 100644 3--- a/internal/common/env_config.go 4+++ b/internal/common/env_config.go 5@@ -54,6 +54,8 @@ type EnvConfigSchema struct { 6 VersionCheckDisabled bool `env:"VERSION_CHECK_DISABLED"` 7 StaticApiKey string `env:"STATIC_API_KEY" options:"file"` 8 9+ OidcEmailDomainRewrites string `env:"OIDC_EMAIL_DOMAIN_REWRITES"` 10+ 11 FileBackend string `env:"FILE_BACKEND" options:"toLower"` 12 UploadPath string `env:"UPLOAD_PATH"` 13 S3Bucket string `env:"S3_BUCKET"` 14diff --git a/internal/service/oidc_service.go b/internal/service/oidc_service.go 15indexe493a5b ..9cd4116 100644 16--- a/internal/service/oidc_service.go 17+++ b/internal/service/oidc_service.go 18@@ -587,7 +587,7 @@ func (s *OidcService) createTokenFromRefreshToken(ctx context.Context, input dto 19 } 20 21 // Load the profile, which we need for the ID token 22- userClaims, err := s.getUserClaims(ctx, &storedRefreshToken.User, storedRefreshToken.Scopes(), tx) 23+ userClaims, err := s.getUserClaims(ctx, &storedRefreshToken.User, storedRefreshToken.Scopes(), input.ClientID, tx) 24 if err != nil { 25 return CreatedTokens{}, err 26 } 27@@ -1923,7 +1923,7 @@ func (s *OidcService) GetClientPreview(ctx context.Context, clientID string, use 28 return nil, &common.OidcAccessDeniedError{} 29 } 30 31- userClaims, err := s.getUserClaims(ctx, &user, scopes, tx) 32+ userClaims, err := s.getUserClaims(ctx, &user, scopes, clientID, tx) 33 if err != nil { 34 return nil, err 35 } 36@@ -1976,15 +1976,15 @@ func (s *OidcService) getUserClaimsForClientInternal(ctx context.Context, userID 37 return nil, err 38 } 39 40- return s.getUserClaims(ctx, &authorizedOidcClient.User, authorizedOidcClient.Scopes(), tx) 41+ return s.getUserClaims(ctx, &authorizedOidcClient.User, authorizedOidcClient.Scopes(), clientID, tx) 42 } 43 44-func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scopes []string, tx *gorm.DB) (map[string]any, error) { 45+func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scopes []string, clientID string, tx *gorm.DB) (map[string]any, error) { 46 claims := make(map[string]any, 10) 47 48 claims["sub"] = user.ID 49 if slices.Contains(scopes, "email") { 50- claims["email"] = user.Email 51+ claims["email"] = rewriteEmailDomainForClient(user.Email, clientID) 52 claims["email_verified"] = user.EmailVerified 53 } 54 55@@ -2027,12 +2027,42 @@ func (s *OidcService) getUserClaims(ctx context.Context, user *model.User, scope 56 } 57 58 if slices.Contains(scopes, "email") { 59- claims["email"] = user.Email 60+ claims["email"] = rewriteEmailDomainForClient(user.Email, clientID) 61 } 62 63 return claims, nil 64 } 65 66+func rewriteEmailDomainForClient(email *string, clientID string) *string { 67+ if email == nil { 68+ return nil 69+ } 70+ 71+ domain := oidcEmailRewriteDomain(clientID) 72+ if domain == "" { 73+ return email 74+ } 75+ 76+ localPart, _, ok := strings.Cut(*email, "@") 77+ if !ok || localPart == "" { 78+ return email 79+ } 80+ 81+ rewritten := localPart + "@" + domain 82+ return &rewritten 83+} 84+ 85+func oidcEmailRewriteDomain(clientID string) string { 86+ for entry := range strings.SplitSeq(common.EnvConfig.OidcEmailDomainRewrites, ",") { 87+ client, domain, ok := strings.Cut(strings.TrimSpace(entry), "=") 88+ if ok && strings.TrimSpace(client) == clientID { 89+ return strings.TrimSpace(domain) 90+ } 91+ } 92+ 93+ return "" 94+} 95+ 96 func (s *OidcService) IsClientAccessibleToUser(ctx context.Context, clientID string, userID string) (bool, error) { 97 var user model.User 98 err := s.db.WithContext(ctx).Preload("UserGroups").First(&user, "id = ?", userID).Error