Browse Source

fix: skip non-string elements in JWT group claim arrays

Align JWT group parsing with OAuth2 getGroupField: type-assert array
elements to string and skip invalid entries instead of fmt.Sprintf
formatting them into malformed group names like %!s(int=42).

Addresses CodeRabbit review on #1104.

Co-authored-by: Cursor <cursoragent@cursor.com>
Bereket Abraham 1 ngày trước cách đây
mục cha
commit
a850ff9dcb
2 tập tin đã thay đổi với 63 bổ sung14 xóa
  1. 27 14
      service/internal/auth/otjwt/jwt.go
  2. 36 0
      service/internal/auth/otjwt/jwt_test.go

+ 27 - 14
service/internal/auth/otjwt/jwt.go

@@ -235,20 +235,33 @@ func parseJwt(cfg *config.Config, token string) *authTypes.AuthenticatedUser {
 }
 
 func parseGroupClaim(groupClaim string, claims jwt.MapClaims, sep string) string {
-	usergroup := ""
-	if val, ok := claims[groupClaim]; ok {
-		if array, ok := val.([]any); ok {
-			groups := make([]string, len(array))
-			for i, v := range array {
-				groups[i] = fmt.Sprintf("%s", v)
-			}
-			if sep == "" {
-				sep = " "
-			}
-			usergroup = strings.Join(groups, sep)
-		} else {
-			usergroup = fmt.Sprintf("%s", val)
+	val, ok := claims[groupClaim]
+	if !ok {
+		return ""
+	}
+
+	array, isArray := val.([]any)
+	if isArray {
+		return joinJWTGroupArray(array, groupClaim, sep)
+	}
+
+	return fmt.Sprintf("%s", val)
+}
+
+func joinJWTGroupArray(arrayVal []any, groupClaim string, sep string) string {
+	if sep == "" {
+		sep = " "
+	}
+
+	groups := make([]string, 0, len(arrayVal))
+	for _, element := range arrayVal {
+		groupName, isString := element.(string)
+		if !isString {
+			log.Warnf("Skipping non-string group entry in JWT claim %v: %v", groupClaim, element)
+			continue
 		}
+		groups = append(groups, groupName)
 	}
-	return usergroup
+
+	return strings.Join(groups, sep)
 }

+ 36 - 0
service/internal/auth/otjwt/jwt_test.go

@@ -202,6 +202,42 @@ func makeJWTRequest(t *testing.T, srv *httptest.Server, tokenStr string) *http.R
 	return res
 }
 
+func TestJWTHeaderSkipsNonStringGroupArrayElements(t *testing.T) {
+	privateKey, publicKeyPath := createKeys(t)
+	defer func() { _ = os.Remove(publicKeyPath) }()
+
+	cfg := config.DefaultConfig()
+	cfg.AuthJwtPubKeyPath = publicKeyPath
+	cfg.AuthJwtClaimUsername = "sub"
+	cfg.AuthJwtClaimUserGroup = "olivetinGroup"
+	cfg.AuthJwtHeader = "Authorization"
+
+	tokenStr := createJWTTokenWithGroups(t, privateKey, []any{"admins", 42, "ops"})
+
+	mux := newMux()
+	mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
+		context := &authpublic.AuthCheckingContext{
+			Request: r,
+			Config:  cfg,
+		}
+		user := CheckUserFromJwtHeader(context)
+
+		if user == nil {
+			w.WriteHeader(http.StatusForbidden)
+			return
+		}
+
+		assert.Equal(t, "test", user.Username)
+		assert.Equal(t, "admins ops", user.UsergroupLine)
+	})
+
+	srv := httptest.NewServer(mux)
+	defer srv.Close()
+
+	res := makeJWTRequest(t, srv, tokenStr) //nolint:bodyclose // closed by verifyJWTResponse
+	verifyJWTResponse(t, res, http.StatusOK)
+}
+
 func TestJWTHeaderWithCustomGroupSeparator(t *testing.T) {
 	privateKey, publicKeyPath := createKeys(t)
 	defer func() { _ = os.Remove(publicKeyPath) }()