added config options and nonce

This commit is contained in:
ston1th 2025-10-15 21:49:13 +02:00
commit d2cf954016
10 changed files with 146 additions and 85 deletions

View file

@ -1,4 +1,4 @@
Copyright (C) 2018 Marius Schellenberger Copyright (C) 2025 Marius Schellenberger
All rights reserved. All rights reserved.
Redistribution and use in source and binary forms, with or without Redistribution and use in source and binary forms, with or without

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
@ -30,37 +30,28 @@ func (c Claims) GetBool(key string) (b bool) {
// GetInt returns an int from the claims map // GetInt returns an int from the claims map
func (c Claims) GetInt(key string) (i int) { func (c Claims) GetInt(key string) (i int) {
if f, ok := c.getFloat64(key); ok { i = int(c.GetInt64(key))
return int(f)
}
if v, ok := c.Get(key); ok && v != nil {
i, _ = v.(int)
}
return return
} }
// GetInt64 returns an int64 from the claims map // GetInt64 returns an int64 from the claims map
func (c Claims) GetInt64(key string) (i int64) { func (c Claims) GetInt64(key string) (i int64) {
if f, ok := c.getFloat64(key); ok {
return int64(f)
}
if v, ok := c.Get(key); ok && v != nil { if v, ok := c.Get(key); ok && v != nil {
i, _ = v.(int64) switch val := v.(type) {
case int64:
i = val
case float64:
i = int64(val)
}
} }
return return
} }
// getFloat64 returns a float64 and ok from the claims map // getFloat64 returns a float64 and ok from the claims map
func (c Claims) getFloat64(key string) (f float64, fok bool) {
if v, ok := c.Get(key); ok && v != nil {
f, fok = v.(float64)
}
return
}
// GetFloat64 returns a float64 from the claims map
func (c Claims) GetFloat64(key string) (f float64) { func (c Claims) GetFloat64(key string) (f float64) {
f, _ = c.getFloat64(key) if v, ok := c.Get(key); ok && v != nil {
f, _ = v.(float64)
}
return return
} }

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt

33
hash.go
View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
@ -14,19 +14,25 @@ const (
HS512Name = "HS512" HS512Name = "HS512"
) )
// ParseHash returns the Hash type equal to the input string // NewHash returns the Hash type equal to the input string
func ParseHash(alg string) Hash { func NewHash(alg string) Hash {
switch alg { switch alg {
case HS256Name: case HS256Name:
return NewHS256() return hs256
case HS384Name: case HS384Name:
return NewHS384() return hs384
case HS512Name: case HS512Name:
return NewHS512() return hs512
} }
return nil return nil
} }
var (
hs256 = HS256{}
hs384 = HS384{}
hs512 = HS512{}
)
// Hash is the hashsum interface for signing the jwt // Hash is the hashsum interface for signing the jwt
type Hash interface { type Hash interface {
Hash() hash.Hash Hash() hash.Hash
@ -36,11 +42,6 @@ type Hash interface {
// HS256 implements the Hash interface with SHA256 // HS256 implements the Hash interface with SHA256
type HS256 struct{} type HS256 struct{}
// NewHS256 returns a new HS256 instance
func NewHS256() Hash {
return HS256{}
}
// Alg returns the algorithm name "HS256" // Alg returns the algorithm name "HS256"
func (HS256) Alg() string { func (HS256) Alg() string {
return HS256Name return HS256Name
@ -54,11 +55,6 @@ func (HS256) Hash() hash.Hash {
// HS384 implements the Hash interface with SHA384 // HS384 implements the Hash interface with SHA384
type HS384 struct{} type HS384 struct{}
// NewHS384 returns a new HS384 instance
func NewHS384() Hash {
return HS384{}
}
// Alg returns the algorithm name "HS384" // Alg returns the algorithm name "HS384"
func (HS384) Alg() string { func (HS384) Alg() string {
return HS384Name return HS384Name
@ -72,11 +68,6 @@ func (HS384) Hash() hash.Hash {
// HS512 implements the Hash interface with SHA512 // HS512 implements the Hash interface with SHA512
type HS512 struct{} type HS512 struct{}
// NewHS512 returns a new HS512 instance
func NewHS512() Hash {
return HS512{}
}
// Alg returns the algorithm name "HS512" // Alg returns the algorithm name "HS512"
func (HS512) Alg() string { func (HS512) Alg() string {
return HS512Name return HS512Name

74
jwt.go
View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
// Package jwt provides a easy to use JSON Web Token and blacklisting library // Package jwt provides a easy to use JSON Web Token and blacklisting library
package jwt package jwt
@ -31,6 +31,8 @@ const (
ExpClaim = "exp" ExpClaim = "exp"
// NbfClaim is the nbf claim name // NbfClaim is the nbf claim name
NbfClaim = "nbf" NbfClaim = "nbf"
// NonceClaim
NonceClaim = "nonce"
// TokenSeparator is the tokens separator char // TokenSeparator is the tokens separator char
TokenSeparator = "." TokenSeparator = "."
) )
@ -57,14 +59,39 @@ var (
ErrInvalidKeySize = errors.New(jwtErr + "invalid secret key size") ErrInvalidKeySize = errors.New(jwtErr + "invalid secret key size")
) )
type JWTOption func(*JWT)
func WithExpiry(expiry time.Duration) JWTOption {
return func(jwt *JWT) {
jwt.expiry = expiry
}
}
func WithBlacklist(blacklist Blacklist) JWTOption {
return func(jwt *JWT) {
jwt.blacklist = blacklist
}
}
func WithSecret(secret io.Reader) JWTOption {
return func(jwt *JWT) {
jwt.secretReader = secret
}
}
func WithNonce() JWTOption {
return func(jwt *JWT) {
jwt.nonce = true
}
}
// JWT represents the JSON Web Token signing and blacklisting infrastructure // JWT represents the JSON Web Token signing and blacklisting infrastructure
type JWT struct { type JWT struct {
key []byte secretReader io.Reader
expiry time.Duration key []byte
expiry time.Duration
blacklist Blacklist blacklist Blacklist
stopOnce sync.Once nonce bool
done chan struct{} stopOnce sync.Once
done chan struct{}
} }
// New returns a new JWT object with the given expiry timeout. // New returns a new JWT object with the given expiry timeout.
@ -73,14 +100,18 @@ type JWT struct {
// If secret is nil the DefaultSecretReader is used. // If secret is nil the DefaultSecretReader is used.
// If blacklisting is enabled, the JWT object leaks a goroutine to garbage-collect expired blacklisted tokens. // If blacklisting is enabled, the JWT object leaks a goroutine to garbage-collect expired blacklisted tokens.
// Call the Stop() method to exit the goroutine. // Call the Stop() method to exit the goroutine.
func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, error) { func New(options ...JWTOption) (*JWT, error) {
if expiry <= 0 { jwt := &JWT{}
expiry = DefaultExpiry for _, option := range options {
option(jwt)
} }
if secret == nil { if jwt.expiry <= 0 {
secret = DefaultSecretReader jwt.expiry = DefaultExpiry
} }
secret = io.LimitReader(secret, KeySize) if jwt.secretReader == nil {
jwt.secretReader = DefaultSecretReader
}
secret := io.LimitReader(jwt.secretReader, KeySize)
key := make([]byte, KeySize) key := make([]byte, KeySize)
i, err := secret.Read(key) i, err := secret.Read(key)
if err != nil { if err != nil {
@ -89,12 +120,8 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err
if i < KeySize { if i < KeySize {
return nil, ErrInvalidKeySize return nil, ErrInvalidKeySize
} }
jwt := &JWT{ jwt.key = key
key: key, if jwt.blacklist != nil {
expiry: expiry,
blacklist: blacklist,
}
if blacklist != nil {
jwt.done = make(chan struct{}) jwt.done = make(chan struct{})
go jwt.clean() go jwt.clean()
} }
@ -128,7 +155,7 @@ func (jwt *JWT) Invalidate(t *Token) error {
if t.Header.GetString(TypClaim) != Typ { if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT return ErrNoJWT
} }
h := ParseHash(t.Header.GetString(AlgClaim)) h := NewHash(t.Header.GetString(AlgClaim))
if h == nil { if h == nil {
return ErrUnsupportedAlg return ErrUnsupportedAlg
} }
@ -190,7 +217,7 @@ func (jwt *JWT) Sign(t *Token) (err error) {
if t.Header.GetString(TypClaim) != Typ { if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT return ErrNoJWT
} }
h := ParseHash(t.Header.GetString(AlgClaim)) h := NewHash(t.Header.GetString(AlgClaim))
if h == nil { if h == nil {
return ErrUnsupportedAlg return ErrUnsupportedAlg
} }
@ -201,6 +228,9 @@ func (jwt *JWT) Sign(t *Token) (err error) {
if _, ok := t.Claims.Get(NbfClaim); !ok { if _, ok := t.Claims.Get(NbfClaim); !ok {
t.Claims.Set(NbfClaim, NewNbf(now)) t.Claims.Set(NbfClaim, NewNbf(now))
} }
if jwt.nonce {
t.Claims.Set(NonceClaim, enc.EncodeToString(key16()))
}
head, err := json.Marshal(t.Header) head, err := json.Marshal(t.Header)
if err != nil { if err != nil {
return return
@ -230,7 +260,7 @@ func (jwt *JWT) Verify(t *Token) error {
if t.Header.GetString(TypClaim) != Typ { if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT return ErrNoJWT
} }
h := ParseHash(t.Header.GetString(AlgClaim)) h := NewHash(t.Header.GetString(AlgClaim))
if h == nil { if h == nil {
return ErrUnsupportedAlg return ErrUnsupportedAlg
} }

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
@ -17,7 +17,10 @@ var claims = Claims{
} }
func TestValidate(t *testing.T) { func TestValidate(t *testing.T) {
jwt, err := New(time.Second, NewMemBlacklist(), nil) jwt, err := New(
WithExpiry(time.Second),
WithBlacklist(NewMemBlacklist()),
)
if err != nil { if err != nil {
t.Error(err) t.Error(err)
} }
@ -83,8 +86,35 @@ func TestValidate(t *testing.T) {
} }
} }
func TestNonce(t *testing.T) {
jwt, err := New(
WithExpiry(time.Second),
WithNonce(),
)
if err != nil {
t.Error(err)
}
token := NewToken(claims.Copy(), nil)
err = jwt.Sign(token)
if err != nil {
t.Error(err)
}
nt, err := DecodeToken(token.String())
if err != nil {
t.Error(err)
}
err = jwt.Verify(nt)
if err != nil {
t.Error(err)
}
nonce := nt.Claims.GetString(NonceClaim)
if nonce == "" {
t.Error(errors.New("nonce is empty"))
}
}
func TestNoBlacklist(t *testing.T) { func TestNoBlacklist(t *testing.T) {
jwt, err := New(time.Second, nil, nil) jwt, err := New(WithExpiry(time.Second))
if err != nil { if err != nil {
t.Error(err) t.Error(err)
} }
@ -118,21 +148,27 @@ func TestNoBlacklist(t *testing.T) {
} }
func TestEmptySecretReader(t *testing.T) { func TestEmptySecretReader(t *testing.T) {
_, err := New(time.Second, nil, new(bytes.Buffer)) _, err := New(
WithExpiry(time.Second),
WithSecret(new(bytes.Buffer)),
)
if err == nil { if err == nil {
t.Error(errors.New("error should be secret reader error")) t.Error(errors.New("error should be secret reader error"))
} }
} }
func TestInvalidSecretReader(t *testing.T) { func TestInvalidSecretReader(t *testing.T) {
_, err := New(time.Second, nil, bytes.NewBufferString("123")) _, err := New(
WithExpiry(time.Second),
WithSecret(bytes.NewBufferString("123")),
)
if err != ErrInvalidKeySize { if err != ErrInvalidKeySize {
t.Error(errors.New("error should be invalid key size")) t.Error(errors.New("error should be invalid key size"))
} }
} }
func BenchmarkDecodeToken(b *testing.B) { func BenchmarkDecodeToken(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), nil) t := NewToken(claims.Copy(), nil)
jwt.Sign(t) jwt.Sign(t)
token := t.String() token := t.String()
@ -145,19 +181,19 @@ func BenchmarkDecodeToken(b *testing.B) {
func BenchmarkNew(b *testing.B) { func BenchmarkNew(b *testing.B) {
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
New(0, nil, nil) New()
} }
} }
func BenchmarkNewWithBlacklist(b *testing.B) { func BenchmarkNewWithBlacklist(b *testing.B) {
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
New(0, NewMemBlacklist(), nil) New(WithBlacklist(NewMemBlacklist()))
} }
} }
func BenchmarkSignHS256(b *testing.B) { func BenchmarkSignHS256(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), nil) t := NewToken(claims.Copy(), nil)
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
@ -166,8 +202,8 @@ func BenchmarkSignHS256(b *testing.B) {
} }
func BenchmarkSignHS384(b *testing.B) { func BenchmarkSignHS384(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), NewHS384()) t := NewToken(claims.Copy(), NewHash(HS384Name))
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
jwt.Sign(t) jwt.Sign(t)
@ -175,8 +211,8 @@ func BenchmarkSignHS384(b *testing.B) {
} }
func BenchmarkSignHS512(b *testing.B) { func BenchmarkSignHS512(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), NewHS512()) t := NewToken(claims.Copy(), NewHash(HS512Name))
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
jwt.Sign(t) jwt.Sign(t)
@ -184,7 +220,7 @@ func BenchmarkSignHS512(b *testing.B) {
} }
func BenchmarkVerifyHS256(b *testing.B) { func BenchmarkVerifyHS256(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), nil) t := NewToken(claims.Copy(), nil)
jwt.Sign(t) jwt.Sign(t)
b.ReportAllocs() b.ReportAllocs()
@ -194,8 +230,8 @@ func BenchmarkVerifyHS256(b *testing.B) {
} }
func BenchmarkVerifyHS384(b *testing.B) { func BenchmarkVerifyHS384(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), NewHS384()) t := NewToken(claims.Copy(), NewHash(HS384Name))
jwt.Sign(t) jwt.Sign(t)
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
@ -204,8 +240,8 @@ func BenchmarkVerifyHS384(b *testing.B) {
} }
func BenchmarkVerifyHS512(b *testing.B) { func BenchmarkVerifyHS512(b *testing.B) {
jwt, _ := New(0, nil, nil) jwt, _ := New()
t := NewToken(claims.Copy(), NewHS512()) t := NewToken(claims.Copy(), NewHash(HS512Name))
jwt.Sign(t) jwt.Sign(t)
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {

11
rand.go Normal file
View file

@ -0,0 +1,11 @@
// Copyright (C) 2025 Marius Schellenberger
package jwt
import "crypto/rand"
func key16() []byte {
b := make([]byte, 16)
rand.Read(b)
return b
}

View file

@ -1,3 +1,5 @@
// Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
import "time" import "time"

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
@ -45,7 +45,7 @@ func NewToken(claims Claims, hash Hash) *Token {
claims = make(Claims) claims = make(Claims)
} }
if hash == nil { if hash == nil {
hash = NewHS256() hash = NewHash(HS256Name)
} }
return &Token{ return &Token{
Header: Claims{ Header: Claims{