From 89827f076f53531351669b1bca772079e7b83a32 Mon Sep 17 00:00:00 2001 From: Jay Hemnani <193022578+jayhemnani9910@users.noreply.github.com> Date: Sun, 11 Oct 2026 04:34:40 +0530 Subject: [PATCH] fix: keep integer custom claim values exact in ID tokens and userinfo (#1829) --- backend/internal/oidc/claims_service.go | 56 +++++++++++++++++--- backend/internal/oidc/claims_service_test.go | 23 ++++++++ backend/internal/oidc/session.go | 31 +++++++++++ backend/internal/oidc/session_test.go | 16 ++++++ 4 files changed, 119 insertions(+), 7 deletions(-) diff --git a/backend/internal/oidc/claims_service.go b/backend/internal/oidc/claims_service.go index f284fa64..b7b0073e 100644 --- a/backend/internal/oidc/claims_service.go +++ b/backend/internal/oidc/claims_service.go @@ -4,7 +4,9 @@ import ( "context" "encoding/json" "errors" + "io" "slices" + "strings" "github.com/pocket-id/fosite" "github.com/pocket-id/pocket-id/backend/internal/common" @@ -127,13 +129,7 @@ func (s *ClaimsService) GetUserClaims(ctx context.Context, userID string, scopes } for _, customClaim := range customClaims { - // A custom claim value can be a JSON document or a plain string - var jsonValue any - if err := json.Unmarshal([]byte(customClaim.Value), &jsonValue); err == nil { - claims[customClaim.Key] = jsonValue - } else { - claims[customClaim.Key] = customClaim.Value - } + claims[customClaim.Key] = parseCustomClaimValue(customClaim.Value) } claims["given_name"] = user.FirstName @@ -164,3 +160,49 @@ func (s *ClaimsService) GetUserClaims(ctx context.Context, userID string, scopes return claims, nil } + +// parseCustomClaimValue decodes a custom claim value that is a JSON document and returns any other value as a plain string +// Integers are kept as int64 because float64 drops digits past 2^53 and is written in exponent form in signed tokens +func parseCustomClaimValue(value string) any { + decoder := json.NewDecoder(strings.NewReader(value)) + decoder.UseNumber() + + var parsed any + if err := decoder.Decode(&parsed); err != nil || decoder.Decode(new(any)) != io.EOF { + return value + } + + normalized, err := normalizeJSONNumbers(parsed) + if err != nil { + return value + } + return normalized +} + +// normalizeJSONNumbers replaces every json.Number with an int64 when it is an integer and with a float64 otherwise +func normalizeJSONNumbers(value any) (any, error) { + switch v := value.(type) { + case json.Number: + if i, err := v.Int64(); err == nil { + return i, nil + } + return v.Float64() + case map[string]any: + for key, item := range v { + normalized, err := normalizeJSONNumbers(item) + if err != nil { + return nil, err + } + v[key] = normalized + } + case []any: + for i, item := range v { + normalized, err := normalizeJSONNumbers(item) + if err != nil { + return nil, err + } + v[i] = normalized + } + } + return value, nil +} diff --git a/backend/internal/oidc/claims_service_test.go b/backend/internal/oidc/claims_service_test.go index 194b8cfc..ea44492e 100644 --- a/backend/internal/oidc/claims_service_test.go +++ b/backend/internal/oidc/claims_service_test.go @@ -152,6 +152,29 @@ func TestClaimsServiceGetUserClaims(t *testing.T) { }) } +// TestClaimsServiceGetUserClaimsKeepsIntegerCustomClaims checks that integer custom claims are not turned into float64 +func TestClaimsServiceGetUserClaimsKeepsIntegerCustomClaims(t *testing.T) { + db := testutils.NewDatabaseForTest(t) + require.NoError(t, db.Create(&model.User{Base: model.Base{ID: "user-1"}, Username: "tim"}).Error) + + customClaims := fakeCustomClaimSource{claims: []model.CustomClaim{ + {Key: "discord_id", Value: "1234567890123456789"}, + {Key: "uid_number", Value: "1000000"}, + {Key: "ratio", Value: "1.5"}, + {Key: "nested", Value: `{"ids":[1234567890123456789]}`}, + {Key: "not_json", Value: "1 2"}, + }} + service := newClaimsService(db, customClaims, "", nil) + + claims, err := service.GetUserClaims(t.Context(), "user-1", []string{"profile"}) + require.NoError(t, err) + require.Equal(t, int64(1234567890123456789), claims["discord_id"]) + require.Equal(t, int64(1000000), claims["uid_number"]) + require.InDelta(t, 1.5, claims["ratio"], 0) + require.Equal(t, map[string]any{"ids": []any{int64(1234567890123456789)}}, claims["nested"]) + require.Equal(t, "1 2", claims["not_json"]) +} + // TestClaimsServiceAppliesSigningAlgToIDTokenHeader verifies the ID token header carries the // signing algorithm so fosite derives the at_hash/c_hash digest from it (e.g. RS384 -> // SHA-384, ES512 -> SHA-512) instead of always defaulting to SHA-256. diff --git a/backend/internal/oidc/session.go b/backend/internal/oidc/session.go index aec212ed..ad28ce6b 100644 --- a/backend/internal/oidc/session.go +++ b/backend/internal/oidc/session.go @@ -1,6 +1,7 @@ package oidc import ( + "bytes" "encoding/json" "time" "uuid" @@ -104,6 +105,36 @@ func (s *Session) Clone() fosite.Session { return &clone } +// UnmarshalJSON decodes a stored session and keeps integer ID token claims as int64 +// A plain decode turns them into float64, which drops digits past 2^53 and is written in exponent form in the ID token +func (s *Session) UnmarshalJSON(data []byte) error { + type plainSession Session + err := json.Unmarshal(data, (*plainSession)(s)) + if err != nil || s.Claims == nil { + return err + } + + // Decode the extra ID token claims a second time with exact numbers + var exact struct { + Claims struct { + Extra map[string]any `json:"ext"` + } `json:"id_token_claims"` + } + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.UseNumber() + err = decoder.Decode(&exact) + if err != nil { + return err + } + + extra, err := normalizeJSONNumbers(exact.Claims.Extra) + if err != nil { + return err + } + s.Claims.Extra, _ = extra.(map[string]any) + return nil +} + func (s *Session) IDTokenClaims() *fositejwt.IDTokenClaims { if s.Claims == nil { s.Claims = &fositejwt.IDTokenClaims{} diff --git a/backend/internal/oidc/session_test.go b/backend/internal/oidc/session_test.go index c641de1a..de88fc59 100644 --- a/backend/internal/oidc/session_test.go +++ b/backend/internal/oidc/session_test.go @@ -32,3 +32,19 @@ func TestNewAuthenticatedSessionDefaultsTimes(t *testing.T) { require.False(t, session.Claims.AuthTime.After(after)) require.Equal(t, session.Claims.AuthTime, session.Claims.RequestedAt) } + +func TestSessionCloneKeepsIntegerIDTokenClaims(t *testing.T) { + session := NewAuthenticatedSession("user-id", "passkey", time.Time{}, time.Time{}) + session.Claims.Extra["discord_id"] = int64(1234567890123456789) + session.Claims.Extra["uid_number"] = int64(1000000) + session.Claims.Extra["nested"] = map[string]any{"ids": []any{int64(1234567890123456789)}} + + // The stored authorize session goes through the same JSON round trip before the ID token is signed + cloned, ok := session.Clone().(*Session) + require.True(t, ok) + require.Equal(t, int64(1234567890123456789), cloned.Claims.Extra["discord_id"]) + require.Equal(t, int64(1000000), cloned.Claims.Extra["uid_number"]) + require.Equal(t, map[string]any{"ids": []any{int64(1234567890123456789)}}, cloned.Claims.Extra["nested"]) + require.Equal(t, "user-id", cloned.Claims.Subject) + require.Equal(t, session.Claims.AuthTime, cloned.Claims.AuthTime) +}