mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-11 20:29:04 +02:00
fix: keep integer custom claim values exact in ID tokens and userinfo (#1829)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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{}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user