diff --git a/LICENSE b/LICENSE index 74cee2b..115e96e 100644 --- a/LICENSE +++ b/LICENSE @@ -1,4 +1,4 @@ -Copyright (C) 2025 Marius Schellenberger +Copyright (C) 2018 Marius Schellenberger All rights reserved. Redistribution and use in source and binary forms, with or without diff --git a/blacklist.go b/blacklist.go index 180c5fd..b60a928 100644 --- a/blacklist.go +++ b/blacklist.go @@ -1,4 +1,4 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger package jwt diff --git a/claims.go b/claims.go index a5f4364..10cbcfe 100644 --- a/claims.go +++ b/claims.go @@ -1,4 +1,4 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger package jwt @@ -30,31 +30,40 @@ func (c Claims) GetBool(key string) (b bool) { // GetInt returns an int from the claims map func (c Claims) GetInt(key string) (i int) { - i = int(c.GetInt64(key)) + if f, ok := c.getFloat64(key); ok { + return int(f) + } + if v, ok := c.Get(key); ok && v != nil { + i, _ = v.(int) + } 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 { - switch val := v.(type) { - case int64: - i = val - case float64: - i = int64(val) - } + i, _ = v.(int64) } return } // getFloat64 returns a float64 and ok from the claims map -func (c Claims) GetFloat64(key string) (f float64) { +func (c Claims) getFloat64(key string) (f float64, fok bool) { if v, ok := c.Get(key); ok && v != nil { - f, _ = v.(float64) + 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) + return +} + // Set sets the value of key in the claims map, if not nil func (c Claims) Set(key string, v interface{}) { if c == nil { diff --git a/encoding.go b/encoding.go deleted file mode 100644 index 3693f5f..0000000 --- a/encoding.go +++ /dev/null @@ -1,7 +0,0 @@ -// Copyright (C) 2025 Marius Schellenberger - -package jwt - -import "encoding/base64" - -var enc = base64.RawURLEncoding diff --git a/hash.go b/hash.go index e47c0cc..2c224ab 100644 --- a/hash.go +++ b/hash.go @@ -1,4 +1,4 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger package jwt @@ -14,25 +14,19 @@ const ( HS512Name = "HS512" ) -// NewHash returns the Hash type equal to the input string -func NewHash(alg string) Hash { +// ParseHash returns the Hash type equal to the input string +func ParseHash(alg string) Hash { switch alg { case HS256Name: - return hs256 + return NewHS256() case HS384Name: - return hs384 + return NewHS384() case HS512Name: - return hs512 + return NewHS512() } return nil } -var ( - hs256 = HS256{} - hs384 = HS384{} - hs512 = HS512{} -) - // Hash is the hashsum interface for signing the jwt type Hash interface { Hash() hash.Hash @@ -42,6 +36,11 @@ 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 @@ -55,6 +54,11 @@ 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 @@ -68,6 +72,11 @@ 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 diff --git a/jwt.go b/jwt.go index f9ce6af..d77c7dd 100644 --- a/jwt.go +++ b/jwt.go @@ -1,4 +1,4 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger // Package jwt provides a easy to use JSON Web Token and blacklisting library package jwt @@ -6,6 +6,7 @@ package jwt import ( "crypto/hmac" "crypto/rand" + "encoding/base64" "encoding/json" "errors" "io" @@ -31,8 +32,6 @@ const ( ExpClaim = "exp" // NbfClaim is the nbf claim name NbfClaim = "nbf" - // NonceClaim - NonceClaim = "nonce" // TokenSeparator is the tokens separator char TokenSeparator = "." ) @@ -59,39 +58,14 @@ 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 { - secretReader io.Reader - key []byte - expiry time.Duration - blacklist Blacklist - nonce bool - stopOnce sync.Once - done chan struct{} + key []byte + expiry time.Duration + + blacklist Blacklist + stopOnce sync.Once + done chan struct{} } // New returns a new JWT object with the given expiry timeout. @@ -100,18 +74,14 @@ 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(options ...JWTOption) (*JWT, error) { - jwt := &JWT{} - for _, option := range options { - option(jwt) +func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, error) { + if expiry <= 0 { + expiry = DefaultExpiry } - if jwt.expiry <= 0 { - jwt.expiry = DefaultExpiry + if secret == nil { + secret = DefaultSecretReader } - if jwt.secretReader == nil { - jwt.secretReader = DefaultSecretReader - } - secret := io.LimitReader(jwt.secretReader, KeySize) + secret = io.LimitReader(secret, KeySize) key := make([]byte, KeySize) i, err := secret.Read(key) if err != nil { @@ -120,8 +90,12 @@ func New(options ...JWTOption) (*JWT, error) { if i < KeySize { return nil, ErrInvalidKeySize } - jwt.key = key - if jwt.blacklist != nil { + jwt := &JWT{ + key: key, + expiry: expiry, + blacklist: blacklist, + } + if blacklist != nil { jwt.done = make(chan struct{}) go jwt.clean() } @@ -155,7 +129,7 @@ func (jwt *JWT) Invalidate(t *Token) error { if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := NewHash(t.Header.GetString(AlgClaim)) + h := ParseHash(t.Header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } @@ -217,7 +191,7 @@ func (jwt *JWT) Sign(t *Token) (err error) { if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := NewHash(t.Header.GetString(AlgClaim)) + h := ParseHash(t.Header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } @@ -228,9 +202,6 @@ 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 @@ -239,9 +210,9 @@ func (jwt *JWT) Sign(t *Token) (err error) { if err != nil { return } - t.data = enc.EncodeToString(head) + TokenSeparator + enc.EncodeToString(claims) + t.data = base64.URLEncoding.EncodeToString(head) + TokenSeparator + base64.URLEncoding.EncodeToString(claims) t.rawSignature = jwt.sum(t.data, h) - t.signature = enc.EncodeToString(t.rawSignature) + t.signature = base64.URLEncoding.EncodeToString(t.rawSignature) t.raw = t.data + TokenSeparator + t.signature return } @@ -260,7 +231,7 @@ func (jwt *JWT) Verify(t *Token) error { if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := NewHash(t.Header.GetString(AlgClaim)) + h := ParseHash(t.Header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } diff --git a/jwt_test.go b/jwt_test.go index 664cf8a..59789d3 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -1,4 +1,4 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger package jwt @@ -17,10 +17,7 @@ var claims = Claims{ } func TestValidate(t *testing.T) { - jwt, err := New( - WithExpiry(time.Second), - WithBlacklist(NewMemBlacklist()), - ) + jwt, err := New(time.Second, NewMemBlacklist(), nil) if err != nil { t.Error(err) } @@ -86,35 +83,8 @@ 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(WithExpiry(time.Second)) + jwt, err := New(time.Second, nil, nil) if err != nil { t.Error(err) } @@ -148,27 +118,21 @@ func TestNoBlacklist(t *testing.T) { } func TestEmptySecretReader(t *testing.T) { - _, err := New( - WithExpiry(time.Second), - WithSecret(new(bytes.Buffer)), - ) + _, err := New(time.Second, nil, new(bytes.Buffer)) if err == nil { t.Error(errors.New("error should be secret reader error")) } } func TestInvalidSecretReader(t *testing.T) { - _, err := New( - WithExpiry(time.Second), - WithSecret(bytes.NewBufferString("123")), - ) + _, err := New(time.Second, nil, bytes.NewBufferString("123")) if err != ErrInvalidKeySize { t.Error(errors.New("error should be invalid key size")) } } func BenchmarkDecodeToken(b *testing.B) { - jwt, _ := New() + jwt, _ := New(0, nil, nil) t := NewToken(claims.Copy(), nil) jwt.Sign(t) token := t.String() @@ -181,19 +145,19 @@ func BenchmarkDecodeToken(b *testing.B) { func BenchmarkNew(b *testing.B) { b.ReportAllocs() for i := 0; i < b.N; i++ { - New() + New(0, nil, nil) } } func BenchmarkNewWithBlacklist(b *testing.B) { b.ReportAllocs() for i := 0; i < b.N; i++ { - New(WithBlacklist(NewMemBlacklist())) + New(0, NewMemBlacklist(), nil) } } func BenchmarkSignHS256(b *testing.B) { - jwt, _ := New() + jwt, _ := New(0, nil, nil) t := NewToken(claims.Copy(), nil) b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -202,8 +166,8 @@ func BenchmarkSignHS256(b *testing.B) { } func BenchmarkSignHS384(b *testing.B) { - jwt, _ := New() - t := NewToken(claims.Copy(), NewHash(HS384Name)) + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS384()) b.ReportAllocs() for i := 0; i < b.N; i++ { jwt.Sign(t) @@ -211,8 +175,8 @@ func BenchmarkSignHS384(b *testing.B) { } func BenchmarkSignHS512(b *testing.B) { - jwt, _ := New() - t := NewToken(claims.Copy(), NewHash(HS512Name)) + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS512()) b.ReportAllocs() for i := 0; i < b.N; i++ { jwt.Sign(t) @@ -220,7 +184,7 @@ func BenchmarkSignHS512(b *testing.B) { } func BenchmarkVerifyHS256(b *testing.B) { - jwt, _ := New() + jwt, _ := New(0, nil, nil) t := NewToken(claims.Copy(), nil) jwt.Sign(t) b.ReportAllocs() @@ -230,8 +194,8 @@ func BenchmarkVerifyHS256(b *testing.B) { } func BenchmarkVerifyHS384(b *testing.B) { - jwt, _ := New() - t := NewToken(claims.Copy(), NewHash(HS384Name)) + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS384()) jwt.Sign(t) b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -240,8 +204,8 @@ func BenchmarkVerifyHS384(b *testing.B) { } func BenchmarkVerifyHS512(b *testing.B) { - jwt, _ := New() - t := NewToken(claims.Copy(), NewHash(HS512Name)) + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS512()) jwt.Sign(t) b.ReportAllocs() for i := 0; i < b.N; i++ { diff --git a/rand.go b/rand.go deleted file mode 100644 index bd1b191..0000000 --- a/rand.go +++ /dev/null @@ -1,11 +0,0 @@ -// Copyright (C) 2025 Marius Schellenberger - -package jwt - -import "crypto/rand" - -func key16() []byte { - b := make([]byte, 16) - rand.Read(b) - return b -} diff --git a/time.go b/time.go index 77e151f..71b10ca 100644 --- a/time.go +++ b/time.go @@ -1,5 +1,3 @@ -// Copyright (C) 2025 Marius Schellenberger - package jwt import "time" diff --git a/token.go b/token.go index 772415c..83aee75 100644 --- a/token.go +++ b/token.go @@ -1,8 +1,9 @@ -// Copyright (C) 2025 Marius Schellenberger +// Copyright (C) 2018 Marius Schellenberger package jwt import ( + "encoding/base64" "encoding/json" "strings" ) @@ -45,7 +46,7 @@ func NewToken(claims Claims, hash Hash) *Token { claims = make(Claims) } if hash == nil { - hash = NewHash(HS256Name) + hash = NewHS256() } return &Token{ Header: Claims{ @@ -74,7 +75,7 @@ func DecodeToken(token string) (t *Token, err error) { } } t = new(Token) - header, err := enc.DecodeString(parts[0]) + header, err := base64.URLEncoding.DecodeString(parts[0]) if err != nil { return } @@ -82,7 +83,7 @@ func DecodeToken(token string) (t *Token, err error) { if err != nil { return } - claims, err := enc.DecodeString(parts[1]) + claims, err := base64.URLEncoding.DecodeString(parts[1]) if err != nil { return } @@ -90,7 +91,7 @@ func DecodeToken(token string) (t *Token, err error) { if err != nil { return } - t.rawSignature, err = enc.DecodeString(parts[2]) + t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) t.signature = parts[2] t.data = parts[0] + TokenSeparator + parts[1] t.raw = token