char/topaz-flake

git clone https://git.t4t.associates/char/topaz-flake

Charlotte Sompatch pocket id to support email domain rewrites744b7fb

main
3.6 KiB98 linesraw
1diff --git a/internal/common/env_config.go b/internal/common/env_config.go
2index 82a421f..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
15index e493a5b..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