Compare commits

...

3 commits

Author SHA1 Message Date
473d54569c added config options and nonce 2025-10-15 21:49:13 +02:00
1613055dfd use raw base64 as defined by rfc7515 2018-12-28 19:10:22 +01:00
a00a154398 added errors to blacklist 2018-09-22 16:34:10 +02:00
10 changed files with 171 additions and 100 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
@ -6,10 +6,10 @@ import "sync"
// Blacklist is the blacklisting storage interface // Blacklist is the blacklisting storage interface
type Blacklist interface { type Blacklist interface {
Add(string, int64) Add(string, int64) error
Remove(string) Remove(string) error
Check(string) bool Check(string) bool
Map() BlacklistMap Map() (BlacklistMap, error)
} }
// BlacklistMap is the blacklist map structure // BlacklistMap is the blacklist map structure
@ -28,17 +28,19 @@ func NewMemBlacklist() *MemBlacklist {
} }
// Add adds a new token signature with expiration time to the blacklist // Add adds a new token signature with expiration time to the blacklist
func (mb *MemBlacklist) Add(sig string, exp int64) { func (mb *MemBlacklist) Add(sig string, exp int64) error {
mb.Lock() mb.Lock()
mb.list[sig] = exp mb.list[sig] = exp
mb.Unlock() mb.Unlock()
return nil
} }
// Remove deletes a token signature from the blacklist // Remove deletes a token signature from the blacklist
func (mb *MemBlacklist) Remove(sig string) { func (mb *MemBlacklist) Remove(sig string) error {
mb.Lock() mb.Lock()
delete(mb.list, sig) delete(mb.list, sig)
mb.Unlock() mb.Unlock()
return nil
} }
// Check returns true if a token signature is blacklisted and false otherwise // Check returns true if a token signature is blacklisted and false otherwise
@ -50,7 +52,7 @@ func (mb *MemBlacklist) Check(sig string) (ok bool) {
} }
// Map returns the blacklist in the form of a iterable map structure for cleanup // Map returns the blacklist in the form of a iterable map structure for cleanup
func (mb *MemBlacklist) Map() (list BlacklistMap) { func (mb *MemBlacklist) Map() (list BlacklistMap, err error) {
list = make(BlacklistMap) list = make(BlacklistMap)
mb.RLock() mb.RLock()
for k, v := range mb.list { for k, v := range mb.list {

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
} }

7
encoding.go Normal file
View file

@ -0,0 +1,7 @@
// Copyright (C) 2025 Marius Schellenberger
package jwt
import "encoding/base64"
var enc = base64.RawURLEncoding

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

78
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
@ -6,7 +6,6 @@ package jwt
import ( import (
"crypto/hmac" "crypto/hmac"
"crypto/rand" "crypto/rand"
"encoding/base64"
"encoding/json" "encoding/json"
"errors" "errors"
"io" "io"
@ -32,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 = "."
) )
@ -58,12 +59,37 @@ 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 {
secretReader io.Reader
key []byte key []byte
expiry time.Duration expiry time.Duration
blacklist Blacklist blacklist Blacklist
nonce bool
stopOnce sync.Once stopOnce sync.Once
done chan struct{} done chan struct{}
} }
@ -74,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 {
@ -90,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()
} }
@ -129,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
} }
@ -143,8 +169,7 @@ func (jwt *JWT) Invalidate(t *Token) error {
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
return ErrInvalid return ErrInvalid
} }
jwt.blacklist.Add(t.Sig(), exp) return jwt.blacklist.Add(t.Sig(), exp)
return nil
} }
// blacklisted checks if a token is blacklisted // blacklisted checks if a token is blacklisted
@ -169,7 +194,11 @@ func (jwt *JWT) clean() {
return return
} }
now := Now() now := Now()
for k, v := range jwt.blacklist.Map() { m, err := jwt.blacklist.Map()
if err != nil {
continue
}
for k, v := range m {
if now > v { if now > v {
jwt.blacklist.Remove(k) jwt.blacklist.Remove(k)
} }
@ -188,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
} }
@ -199,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
@ -207,9 +239,9 @@ func (jwt *JWT) Sign(t *Token) (err error) {
if err != nil { if err != nil {
return return
} }
t.data = base64.URLEncoding.EncodeToString(head) + TokenSeparator + base64.URLEncoding.EncodeToString(claims) t.data = enc.EncodeToString(head) + TokenSeparator + enc.EncodeToString(claims)
t.rawSignature = jwt.sum(t.data, h) t.rawSignature = jwt.sum(t.data, h)
t.signature = base64.URLEncoding.EncodeToString(t.rawSignature) t.signature = enc.EncodeToString(t.rawSignature)
t.raw = t.data + TokenSeparator + t.signature t.raw = t.data + TokenSeparator + t.signature
return return
} }
@ -228,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,9 +1,8 @@
// Copyright (C) 2018 Marius Schellenberger // Copyright (C) 2025 Marius Schellenberger
package jwt package jwt
import ( import (
"encoding/base64"
"encoding/json" "encoding/json"
"strings" "strings"
) )
@ -46,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{
@ -75,7 +74,7 @@ func DecodeToken(token string) (t *Token, err error) {
} }
} }
t = new(Token) t = new(Token)
header, err := base64.URLEncoding.DecodeString(parts[0]) header, err := enc.DecodeString(parts[0])
if err != nil { if err != nil {
return return
} }
@ -83,7 +82,7 @@ func DecodeToken(token string) (t *Token, err error) {
if err != nil { if err != nil {
return return
} }
claims, err := base64.URLEncoding.DecodeString(parts[1]) claims, err := enc.DecodeString(parts[1])
if err != nil { if err != nil {
return return
} }
@ -91,7 +90,7 @@ func DecodeToken(token string) (t *Token, err error) {
if err != nil { if err != nil {
return return
} }
t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) t.rawSignature, err = enc.DecodeString(parts[2])
t.signature = parts[2] t.signature = parts[2]
t.data = parts[0] + TokenSeparator + parts[1] t.data = parts[0] + TokenSeparator + parts[1]
t.raw = token t.raw = token