Compare commits

...

1 commit

Author SHA1 Message Date
473d54569c added config options and nonce 2025-10-15 21:49:13 +02:00
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.
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

View file

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

View file

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

33
hash.go
View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger
// Copyright (C) 2025 Marius Schellenberger
package jwt
@ -14,19 +14,25 @@ const (
HS512Name = "HS512"
)
// ParseHash returns the Hash type equal to the input string
func ParseHash(alg string) Hash {
// NewHash returns the Hash type equal to the input string
func NewHash(alg string) Hash {
switch alg {
case HS256Name:
return NewHS256()
return hs256
case HS384Name:
return NewHS384()
return hs384
case HS512Name:
return NewHS512()
return hs512
}
return nil
}
var (
hs256 = HS256{}
hs384 = HS384{}
hs512 = HS512{}
)
// Hash is the hashsum interface for signing the jwt
type Hash interface {
Hash() hash.Hash
@ -36,11 +42,6 @@ type Hash interface {
// HS256 implements the Hash interface with SHA256
type HS256 struct{}
// NewHS256 returns a new HS256 instance
func NewHS256() Hash {
return HS256{}
}
// Alg returns the algorithm name "HS256"
func (HS256) Alg() string {
return HS256Name
@ -54,11 +55,6 @@ func (HS256) Hash() hash.Hash {
// HS384 implements the Hash interface with SHA384
type HS384 struct{}
// NewHS384 returns a new HS384 instance
func NewHS384() Hash {
return HS384{}
}
// Alg returns the algorithm name "HS384"
func (HS384) Alg() string {
return HS384Name
@ -72,11 +68,6 @@ func (HS384) Hash() hash.Hash {
// HS512 implements the Hash interface with SHA512
type HS512 struct{}
// NewHS512 returns a new HS512 instance
func NewHS512() Hash {
return HS512{}
}
// Alg returns the algorithm name "HS512"
func (HS512) Alg() string {
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
@ -31,6 +31,8 @@ const (
ExpClaim = "exp"
// NbfClaim is the nbf claim name
NbfClaim = "nbf"
// NonceClaim
NonceClaim = "nonce"
// TokenSeparator is the tokens separator char
TokenSeparator = "."
)
@ -57,14 +59,39 @@ var (
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
type JWT struct {
key []byte
expiry time.Duration
blacklist Blacklist
stopOnce sync.Once
done chan struct{}
secretReader io.Reader
key []byte
expiry time.Duration
blacklist Blacklist
nonce bool
stopOnce sync.Once
done chan struct{}
}
// 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 blacklisting is enabled, the JWT object leaks a goroutine to garbage-collect expired blacklisted tokens.
// Call the Stop() method to exit the goroutine.
func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, error) {
if expiry <= 0 {
expiry = DefaultExpiry
func New(options ...JWTOption) (*JWT, error) {
jwt := &JWT{}
for _, option := range options {
option(jwt)
}
if secret == nil {
secret = DefaultSecretReader
if jwt.expiry <= 0 {
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)
i, err := secret.Read(key)
if err != nil {
@ -89,12 +120,8 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err
if i < KeySize {
return nil, ErrInvalidKeySize
}
jwt := &JWT{
key: key,
expiry: expiry,
blacklist: blacklist,
}
if blacklist != nil {
jwt.key = key
if jwt.blacklist != nil {
jwt.done = make(chan struct{})
go jwt.clean()
}
@ -128,7 +155,7 @@ func (jwt *JWT) Invalidate(t *Token) error {
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
h := NewHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}
@ -190,7 +217,7 @@ func (jwt *JWT) Sign(t *Token) (err error) {
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
h := NewHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}
@ -201,6 +228,9 @@ func (jwt *JWT) Sign(t *Token) (err error) {
if _, ok := t.Claims.Get(NbfClaim); !ok {
t.Claims.Set(NbfClaim, NewNbf(now))
}
if jwt.nonce {
t.Claims.Set(NonceClaim, enc.EncodeToString(key16()))
}
head, err := json.Marshal(t.Header)
if err != nil {
return
@ -230,7 +260,7 @@ func (jwt *JWT) Verify(t *Token) error {
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
h := NewHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}

View file

@ -1,4 +1,4 @@
// Copyright (C) 2018 Marius Schellenberger
// Copyright (C) 2025 Marius Schellenberger
package jwt
@ -17,7 +17,10 @@ var claims = Claims{
}
func TestValidate(t *testing.T) {
jwt, err := New(time.Second, NewMemBlacklist(), nil)
jwt, err := New(
WithExpiry(time.Second),
WithBlacklist(NewMemBlacklist()),
)
if err != nil {
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) {
jwt, err := New(time.Second, nil, nil)
jwt, err := New(WithExpiry(time.Second))
if err != nil {
t.Error(err)
}
@ -118,21 +148,27 @@ func TestNoBlacklist(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 {
t.Error(errors.New("error should be secret reader error"))
}
}
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 {
t.Error(errors.New("error should be invalid key size"))
}
}
func BenchmarkDecodeToken(b *testing.B) {
jwt, _ := New(0, nil, nil)
jwt, _ := New()
t := NewToken(claims.Copy(), nil)
jwt.Sign(t)
token := t.String()
@ -145,19 +181,19 @@ func BenchmarkDecodeToken(b *testing.B) {
func BenchmarkNew(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
New(0, nil, nil)
New()
}
}
func BenchmarkNewWithBlacklist(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
New(0, NewMemBlacklist(), nil)
New(WithBlacklist(NewMemBlacklist()))
}
}
func BenchmarkSignHS256(b *testing.B) {
jwt, _ := New(0, nil, nil)
jwt, _ := New()
t := NewToken(claims.Copy(), nil)
b.ReportAllocs()
for i := 0; i < b.N; i++ {
@ -166,8 +202,8 @@ func BenchmarkSignHS256(b *testing.B) {
}
func BenchmarkSignHS384(b *testing.B) {
jwt, _ := New(0, nil, nil)
t := NewToken(claims.Copy(), NewHS384())
jwt, _ := New()
t := NewToken(claims.Copy(), NewHash(HS384Name))
b.ReportAllocs()
for i := 0; i < b.N; i++ {
jwt.Sign(t)
@ -175,8 +211,8 @@ func BenchmarkSignHS384(b *testing.B) {
}
func BenchmarkSignHS512(b *testing.B) {
jwt, _ := New(0, nil, nil)
t := NewToken(claims.Copy(), NewHS512())
jwt, _ := New()
t := NewToken(claims.Copy(), NewHash(HS512Name))
b.ReportAllocs()
for i := 0; i < b.N; i++ {
jwt.Sign(t)
@ -184,7 +220,7 @@ func BenchmarkSignHS512(b *testing.B) {
}
func BenchmarkVerifyHS256(b *testing.B) {
jwt, _ := New(0, nil, nil)
jwt, _ := New()
t := NewToken(claims.Copy(), nil)
jwt.Sign(t)
b.ReportAllocs()
@ -194,8 +230,8 @@ func BenchmarkVerifyHS256(b *testing.B) {
}
func BenchmarkVerifyHS384(b *testing.B) {
jwt, _ := New(0, nil, nil)
t := NewToken(claims.Copy(), NewHS384())
jwt, _ := New()
t := NewToken(claims.Copy(), NewHash(HS384Name))
jwt.Sign(t)
b.ReportAllocs()
for i := 0; i < b.N; i++ {
@ -204,8 +240,8 @@ func BenchmarkVerifyHS384(b *testing.B) {
}
func BenchmarkVerifyHS512(b *testing.B) {
jwt, _ := New(0, nil, nil)
t := NewToken(claims.Copy(), NewHS512())
jwt, _ := New()
t := NewToken(claims.Copy(), NewHash(HS512Name))
jwt.Sign(t)
b.ReportAllocs()
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
import "time"

View file

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