From 6e61945c2c80f613e36c514f5d84a1b2ae88a7ea Mon Sep 17 00:00:00 2001 From: ston1th Date: Wed, 12 Sep 2018 23:14:09 +0200 Subject: [PATCH 01/12] added claims type --- claims.go | 52 ++++++++++++++++++++++++++++++++++++++++++++++++++++ jwt.go | 16 ++++++++++------ jwt_test.go | 4 ++-- 3 files changed, 64 insertions(+), 8 deletions(-) create mode 100644 claims.go diff --git a/claims.go b/claims.go new file mode 100644 index 0000000..74e26c9 --- /dev/null +++ b/claims.go @@ -0,0 +1,52 @@ +// Copyright (C) 2018 Marius Schellenberger + +package jwt + +// Claims is the claim type of the token +type Claims map[string]interface{} + +// Get returns a value from the claims map +func (c Claims) Get(key string) (v interface{}, ok bool) { + v, ok = c[key] + return +} + +// GetString returns a string from the claims map +func (c Claims) GetString(key string) (s string) { + if v, ok := c.Get(key); ok && v != nil { + s, _ = v.(string) + } + return +} + +// GetBool returns a bool from the claims map +func (c Claims) GetBool(key string) (b bool) { + if v, ok := c.Get(key); ok && v != nil { + b, _ = v.(bool) + } + return +} + +// GetInt returns an int from the claims map +func (c Claims) GetInt(key string) (i int) { + if v, ok := c.Get(key); ok && v != nil { + i, _ = v.(int) + } + return +} + +// GetFloat returns a float from the claims map +func (c Claims) GetFloat(key string) (f float64) { + if v, ok := c.Get(key); ok && v != nil { + f, _ = v.(float64) + } + 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 { + return + } + c[key] = v +} diff --git a/jwt.go b/jwt.go index 2106f13..1268cb2 100644 --- a/jwt.go +++ b/jwt.go @@ -323,7 +323,7 @@ type Token struct { signature string rawSignature []byte header map[string]interface{} - Claims map[string]interface{} + Claims Claims } // String returns the tokens encoded string @@ -352,14 +352,18 @@ func (t *Token) Data() string { } // NewToken returns a new *Token using the provided hash algorithm and claims -// If h is nil HS256 will be used -func NewToken(claims map[string]interface{}, h Hash) *Token { - if h == nil { - h = NewHS256() +// If claims is nil, an empty map is used +// If hash is nil, HS256 is used +func NewToken(claims Claims, hash Hash) *Token { + if claims == nil { + claims = make(Claims) + } + if hash == nil { + hash = NewHS256() } return &Token{ header: map[string]interface{}{ - "alg": h.Alg(), + "alg": hash.Alg(), "typ": typ, }, Claims: claims, diff --git a/jwt_test.go b/jwt_test.go index 30cfbfe..17a519b 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -14,7 +14,7 @@ func TestValidate(t *testing.T) { if err != nil { t.Error(err) } - token := NewToken(map[string]interface{}{ + token := NewToken(Claims{ "sub": "1234567890", "name": "John Doe", "admin": true, @@ -65,7 +65,7 @@ func TestNoBlacklist(t *testing.T) { if err != nil { t.Error(err) } - token := NewToken(map[string]interface{}{ + token := NewToken(Claims{ "sub": "1234567890", "name": "John Doe", "admin": true, From 826d035d3e3b03cf2058182977e897b584cb4086 Mon Sep 17 00:00:00 2001 From: ston1th Date: Wed, 12 Sep 2018 23:35:03 +0200 Subject: [PATCH 02/12] split up library and removed Stop() panic when called twice --- blacklist.go | 2 +- jwt.go | 106 ++++----------------------------------------------- jwt_test.go | 5 +++ token.go | 104 ++++++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 117 insertions(+), 100 deletions(-) create mode 100644 token.go diff --git a/blacklist.go b/blacklist.go index 024796d..9ecf662 100644 --- a/blacklist.go +++ b/blacklist.go @@ -22,7 +22,7 @@ type MemBlacklist struct { list BlacklistMap } -// NewMemBlacklist +// NewMemBlacklist implements the Blacklist interface using an in-memory map func NewMemBlacklist() *MemBlacklist { return &MemBlacklist{list: make(BlacklistMap)} } diff --git a/jwt.go b/jwt.go index 1268cb2..32a9012 100644 --- a/jwt.go +++ b/jwt.go @@ -11,6 +11,7 @@ import ( "errors" "io" "strings" + "sync" "time" ) @@ -50,6 +51,7 @@ type JWT struct { expiry time.Duration blacklist Blacklist + stopOnce sync.Once done chan struct{} } @@ -85,6 +87,7 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err return jwt, nil } +// sum calculates the HMAC hash sum of the token func (jwt *JWT) sum(token string, h Hash) []byte { mac := hmac.New(h.Hash, jwt.key) mac.Write([]byte(token)) @@ -248,8 +251,7 @@ func (jwt *JWT) Verify(t *Token) error { nbfOK = true } if jwt.blacklist != nil { - err := jwt.blacklisted(t.Sig()) - if err != nil { + if err := jwt.blacklisted(t.Sig()); err != nil { return err } } @@ -266,106 +268,12 @@ func (jwt *JWT) Verify(t *Token) error { } // Stop will end the cleaner goroutinge execution -// Calling Stop twice or more will panic func (jwt *JWT) Stop() error { if jwt.blacklist == nil { return ErrBlacklistNotEnabled } - close(jwt.done) + jwt.stopOnce.Do(func() { + close(jwt.done) + }) return nil } - -// DecodeToken decodes a raw string token into a *Token object -func DecodeToken(token string) (t *Token, err error) { - if token == "" { - return nil, ErrEmptyToken - } - parts := strings.Split(token, ".") - if len(parts) < 3 { - return nil, ErrMissingTokenParts - } - if len(parts) > 3 { - return nil, ErrNoJWT - } - for _, v := range parts { - if v == "" { - return nil, ErrNoJWT - } - } - t = new(Token) - header, err := base64.URLEncoding.DecodeString(parts[0]) - if err != nil { - return - } - err = json.Unmarshal(header, &t.header) - if err != nil { - return - } - claims, err := base64.URLEncoding.DecodeString(parts[1]) - if err != nil { - return - } - err = json.Unmarshal(claims, &t.Claims) - if err != nil { - return - } - t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) - t.signature = parts[2] - t.data = strings.Join(parts[:2], ".") - t.raw = token - return -} - -// Token is the JWT token representation -type Token struct { - raw string - data string - signature string - rawSignature []byte - header map[string]interface{} - Claims Claims -} - -// String returns the tokens encoded string -func (t *Token) String() string { - return t.raw -} - -// Sig returns the Tokens URLEncoded signature -func (t *Token) Sig() string { - return t.signature -} - -// RawSig returns the Tokens raw signature -func (t *Token) RawSig() []byte { - return t.rawSignature -} - -// Header returns the Tokens header -func (t *Token) Header() map[string]interface{} { - return t.header -} - -// Data returns the first two token fields -func (t *Token) Data() string { - return t.data -} - -// NewToken returns a new *Token using the provided hash algorithm and claims -// If claims is nil, an empty map is used -// If hash is nil, HS256 is used -func NewToken(claims Claims, hash Hash) *Token { - if claims == nil { - claims = make(Claims) - } - if hash == nil { - hash = NewHS256() - } - return &Token{ - header: map[string]interface{}{ - "alg": hash.Alg(), - "typ": typ, - }, - Claims: claims, - } -} diff --git a/jwt_test.go b/jwt_test.go index 17a519b..103eb32 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -58,6 +58,11 @@ func TestValidate(t *testing.T) { if err != nil { t.Error(err) } + // should not panic + err = jwt.Stop() + if err != nil { + t.Error(err) + } } func TestNoBlacklist(t *testing.T) { diff --git a/token.go b/token.go new file mode 100644 index 0000000..3a08da2 --- /dev/null +++ b/token.go @@ -0,0 +1,104 @@ +// Copyright (C) 2018 Marius Schellenberger + +package jwt + +import ( + "encoding/base64" + "encoding/json" + "strings" +) + +// Token is the JWT token representation +type Token struct { + raw string + data string + signature string + rawSignature []byte + header map[string]interface{} + Claims Claims +} + +// String returns the tokens encoded string +func (t *Token) String() string { + return t.raw +} + +// Sig returns the Tokens URLEncoded signature +func (t *Token) Sig() string { + return t.signature +} + +// RawSig returns the Tokens raw signature +func (t *Token) RawSig() []byte { + return t.rawSignature +} + +// Header returns the Tokens header +func (t *Token) Header() map[string]interface{} { + return t.header +} + +// Data returns the first two token fields +func (t *Token) Data() string { + return t.data +} + +// NewToken returns a new *Token using the provided hash algorithm and claims +// If claims is nil, an empty map is used +// If hash is nil, HS256 is used +func NewToken(claims Claims, hash Hash) *Token { + if claims == nil { + claims = make(Claims) + } + if hash == nil { + hash = NewHS256() + } + return &Token{ + header: map[string]interface{}{ + "alg": hash.Alg(), + "typ": typ, + }, + Claims: claims, + } +} + +// DecodeToken decodes a raw string token into a *Token object +func DecodeToken(token string) (t *Token, err error) { + if token == "" { + return nil, ErrEmptyToken + } + parts := strings.Split(token, ".") + if len(parts) < 3 { + return nil, ErrMissingTokenParts + } + if len(parts) > 3 { + return nil, ErrNoJWT + } + for _, v := range parts { + if v == "" { + return nil, ErrNoJWT + } + } + t = new(Token) + header, err := base64.URLEncoding.DecodeString(parts[0]) + if err != nil { + return + } + err = json.Unmarshal(header, &t.header) + if err != nil { + return + } + claims, err := base64.URLEncoding.DecodeString(parts[1]) + if err != nil { + return + } + err = json.Unmarshal(claims, &t.Claims) + if err != nil { + return + } + t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) + t.signature = parts[2] + t.data = strings.Join(parts[:2], ".") + t.raw = token + return +} From f2d57c6e1868d1dc5ab021876511f0c5c46568b8 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 11:39:49 +0200 Subject: [PATCH 03/12] major refactoring --- claims.go | 28 +++++++++-- jwt.go | 133 +++++++++++++++++++--------------------------------- jwt_test.go | 15 ++++++ token.go | 6 +-- 4 files changed, 90 insertions(+), 92 deletions(-) diff --git a/claims.go b/claims.go index 74e26c9..0071f30 100644 --- a/claims.go +++ b/claims.go @@ -29,17 +29,37 @@ 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) } return } -// GetFloat returns a float from the claims map -func (c Claims) GetFloat(key string) (f float64) { - if v, ok := c.Get(key); ok && v != nil { - f, _ = v.(float64) +// 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) + } + 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) return } diff --git a/jwt.go b/jwt.go index 32a9012..eff51c5 100644 --- a/jwt.go +++ b/jwt.go @@ -16,8 +16,8 @@ import ( ) const ( - // Header is the default HTTP Authorization header name - Header = "Authorization" + // HTTPHeader is the default HTTP Authorization header name + HTTPHeader = "Authorization" // DefaultExpiry is the default token expiration time DefaultExpiry = time.Hour * 12 // KeySize is the secret key size @@ -38,6 +38,7 @@ var ( ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled") ErrNBF = errors.New("token not valid yet") ErrEXP = errors.New("token expired") + ErrMissingNBF = errors.New("missing nbf claim") ErrMissingEXP = errors.New("missing exp claim") ErrMissingTokenParts = errors.New("missing token parts") ErrEmptySignature = errors.New("token signature is empty") @@ -103,26 +104,27 @@ func (jwt *JWT) Invalidate(t *Token) error { if t == nil { return ErrTokenIsNil } - var ( - expOK bool - exp int64 - ) - switch n := t.Claims["exp"].(type) { - case json.Number: - if i, err := n.Int64(); err == nil { - exp = i - expOK = true - } - case float64: - exp = int64(n) - expOK = true - } - if !expOK { - return ErrMissingEXP - } if err := jwt.blacklisted(t.Sig()); err != nil { return err } + header := t.Header() + if header.GetString("typ") != typ { + return ErrNoJWT + } + h := ParseHash(header.GetString("alg")) + if h == nil { + return ErrUnsupportedAlg + } + exp := t.Claims.GetInt64("exp") + switch { + case exp == 0: + return ErrMissingEXP + case time.Now().UTC().Unix() > exp: + return ErrEXP + } + if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { + return ErrInvalid + } jwt.blacklist.Add(t.Sig(), exp) return nil } @@ -164,27 +166,16 @@ func (jwt *JWT) Sign(t *Token) (err error) { return ErrTokenIsNil } header := t.Header() - switch t := header["typ"].(type) { - case string: - if t != typ { - return ErrNoJWT - } - default: + if header.GetString("typ") != typ { return ErrNoJWT } - var h Hash - switch a := header["alg"].(type) { - case string: - h = ParseHash(a) - if h == nil { - return ErrUnsupportedAlg - } - default: + h := ParseHash(header.GetString("alg")) + if h == nil { return ErrUnsupportedAlg } now := time.Now().UTC() - t.Claims["exp"] = now.Add(jwt.expiry).Unix() - t.Claims["nbf"] = now.Unix() + t.Claims.Set("exp", now.Add(jwt.expiry).Unix()) + t.Claims.Set("nbf", now.Unix()) head, err := json.Marshal(header) if err != nil { return @@ -194,7 +185,9 @@ func (jwt *JWT) Sign(t *Token) (err error) { return } t.data = strings.Join([]string{base64.URLEncoding.EncodeToString(head), base64.URLEncoding.EncodeToString(claims)}, ".") - t.raw = strings.Join([]string{t.data, base64.URLEncoding.EncodeToString(jwt.sum(t.data, h))}, ".") + t.rawSignature = jwt.sum(t.data, h) + t.signature = base64.URLEncoding.EncodeToString(t.rawSignature) + t.raw = strings.Join([]string{t.data, t.signature}, ".") return } @@ -203,62 +196,32 @@ func (jwt *JWT) Verify(t *Token) error { if t == nil { return ErrTokenIsNil } - now := time.Now().UTC().Unix() - var ( - exp int64 - nbf int64 - expOK bool - nbfOK bool - h Hash - ) - header := t.Header() - switch t := header["typ"].(type) { - case string: - if t != typ { - return ErrNoJWT - } - default: - return ErrNoJWT - } - switch a := header["alg"].(type) { - case string: - h = ParseHash(a) - if h == nil { - return ErrUnsupportedAlg - } - default: - return ErrUnsupportedAlg - } - - switch n := t.Claims["exp"].(type) { - case json.Number: - if i, err := n.Int64(); err == nil { - exp = i - expOK = true - } - case float64: - exp = int64(n) - expOK = true - } - switch n := t.Claims["nbf"].(type) { - case json.Number: - if i, err := n.Int64(); err == nil { - nbf = i - nbfOK = true - } - case float64: - nbf = int64(n) - nbfOK = true - } if jwt.blacklist != nil { if err := jwt.blacklisted(t.Sig()); err != nil { return err } } - if expOK && now > exp { + now := time.Now().UTC().Unix() + header := t.Header() + if header.GetString("typ") != typ { + return ErrNoJWT + } + h := ParseHash(header.GetString("alg")) + if h == nil { + return ErrUnsupportedAlg + } + exp := t.Claims.GetInt64("exp") + switch { + case exp == 0: + return ErrMissingEXP + case now > exp: return ErrEXP } - if nbfOK && now < nbf { + nbf := t.Claims.GetInt64("nbf") + switch { + case nbf == 0: + return ErrMissingNBF + case now < nbf: return ErrNBF } if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { diff --git a/jwt_test.go b/jwt_test.go index 103eb32..ee6dea7 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -42,6 +42,20 @@ func TestValidate(t *testing.T) { t.Error(errors.New("token should be expired")) } + err = jwt.Invalidate(nt) + if err == nil { + t.Error(errors.New("token is expired")) + } + + err = jwt.Sign(nt) + if err != nil { + t.Error(err) + } + err = jwt.Verify(nt) + if err != nil { + t.Error(err) + } + err = jwt.Invalidate(nt) if err != nil { t.Error(err) @@ -50,6 +64,7 @@ func TestValidate(t *testing.T) { if err == nil { t.Error(errors.New("double invalidate")) } + err = jwt.Verify(nt) if err == nil { t.Error(errors.New("token should be blacklisted")) diff --git a/token.go b/token.go index 3a08da2..4157f90 100644 --- a/token.go +++ b/token.go @@ -14,7 +14,7 @@ type Token struct { data string signature string rawSignature []byte - header map[string]interface{} + header Claims Claims Claims } @@ -34,7 +34,7 @@ func (t *Token) RawSig() []byte { } // Header returns the Tokens header -func (t *Token) Header() map[string]interface{} { +func (t *Token) Header() Claims { return t.header } @@ -54,7 +54,7 @@ func NewToken(claims Claims, hash Hash) *Token { hash = NewHS256() } return &Token{ - header: map[string]interface{}{ + header: Claims{ "alg": hash.Alg(), "typ": typ, }, From b4a92fe472c89fa93c028d4d3762b985094f1bd1 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 13:27:32 +0200 Subject: [PATCH 04/12] custom exp and nbf --- claims.go | 6 +++++ jwt.go | 74 +++++++++++++++++++++++++++++++++++------------------ jwt_test.go | 1 + token.go | 8 +++--- 4 files changed, 60 insertions(+), 29 deletions(-) diff --git a/claims.go b/claims.go index 0071f30..b0c99c9 100644 --- a/claims.go +++ b/claims.go @@ -3,6 +3,7 @@ package jwt // Claims is the claim type of the token +// The Claims map is not goroutine safe type Claims map[string]interface{} // Get returns a value from the claims map @@ -70,3 +71,8 @@ func (c Claims) Set(key string, v interface{}) { } c[key] = v } + +// Delete removes the key from the claims map +func (c Claims) Delete(key string) { + delete(c, key) +} diff --git a/jwt.go b/jwt.go index eff51c5..4c52a10 100644 --- a/jwt.go +++ b/jwt.go @@ -23,7 +23,18 @@ const ( // KeySize is the secret key size KeySize = 64 - typ = "JWT" + // Typ is the JWT type + Typ = "JWT" + // TypClaim is the typ claim name + TypClaim = "typ" + // AlgClaim is the alg claim name + AlgClaim = "alg" + // ExpClaim is the exp claim name + ExpClaim = "exp" + // NbfClaim is the nbf claim name + NbfClaim = "nbf" + // TokenSeparator is the tokens separator char + TokenSeparator = "." ) // DefaultSecretReader is the default secret key generator @@ -36,10 +47,10 @@ var ( ErrInvalid = errors.New("token validation failed") ErrBlacklisted = errors.New("token blacklisted") ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled") - ErrNBF = errors.New("token not valid yet") - ErrEXP = errors.New("token expired") - ErrMissingNBF = errors.New("missing nbf claim") - ErrMissingEXP = errors.New("missing exp claim") + ErrNbf = errors.New("token not valid yet") + ErrExp = errors.New("token expired") + ErrMissingNbf = errors.New("missing nbf claim") + ErrMissingExp = errors.New("missing exp claim") ErrMissingTokenParts = errors.New("missing token parts") ErrEmptySignature = errors.New("token signature is empty") ErrTokenIsNil = errors.New("token is nil") @@ -58,6 +69,8 @@ type JWT struct { // New returns a new JWT object with the given expiry timeout. // If the timeout is less or equal to zero the default expiry (12 hours) is used. +// The secret key size needs to be at least 64 bytes. +// 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) { @@ -88,6 +101,11 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err return jwt, nil } +// Expiry returns the configured expiry +func (jwt *JWT) Expiry() time.Duration { + return jwt.expiry +} + // sum calculates the HMAC hash sum of the token func (jwt *JWT) sum(token string, h Hash) []byte { mac := hmac.New(h.Hash, jwt.key) @@ -108,19 +126,19 @@ func (jwt *JWT) Invalidate(t *Token) error { return err } header := t.Header() - if header.GetString("typ") != typ { + if header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(header.GetString("alg")) + h := ParseHash(header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } - exp := t.Claims.GetInt64("exp") + exp := t.Claims.GetInt64(ExpClaim) switch { case exp == 0: - return ErrMissingEXP + return ErrMissingExp case time.Now().UTC().Unix() > exp: - return ErrEXP + return ErrExp } if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { return ErrInvalid @@ -160,22 +178,28 @@ func (jwt *JWT) clean() { } // Sign will sign the provided token using the secret key -// This will overwrite existing 'exp' and 'nbf' claims +// If the 'exp' and 'nbf' claims do not exist, they will be written to the default values in UNIX format: +// 'exp': Now + expiry specified at New() or DefaultExpiry (UTC) +// 'nbf': Now (UTC) func (jwt *JWT) Sign(t *Token) (err error) { if t == nil { return ErrTokenIsNil } header := t.Header() - if header.GetString("typ") != typ { + if header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(header.GetString("alg")) + h := ParseHash(header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } now := time.Now().UTC() - t.Claims.Set("exp", now.Add(jwt.expiry).Unix()) - t.Claims.Set("nbf", now.Unix()) + if _, ok := t.Claims.Get(ExpClaim); !ok { + t.Claims.Set(ExpClaim, now.Add(jwt.expiry).Unix()) + } + if _, ok := t.Claims.Get(NbfClaim); !ok { + t.Claims.Set(NbfClaim, now.Unix()) + } head, err := json.Marshal(header) if err != nil { return @@ -184,10 +208,10 @@ func (jwt *JWT) Sign(t *Token) (err error) { if err != nil { return } - t.data = strings.Join([]string{base64.URLEncoding.EncodeToString(head), base64.URLEncoding.EncodeToString(claims)}, ".") + t.data = strings.Join([]string{base64.URLEncoding.EncodeToString(head), base64.URLEncoding.EncodeToString(claims)}, TokenSeparator) t.rawSignature = jwt.sum(t.data, h) t.signature = base64.URLEncoding.EncodeToString(t.rawSignature) - t.raw = strings.Join([]string{t.data, t.signature}, ".") + t.raw = strings.Join([]string{t.data, t.signature}, TokenSeparator) return } @@ -203,26 +227,26 @@ func (jwt *JWT) Verify(t *Token) error { } now := time.Now().UTC().Unix() header := t.Header() - if header.GetString("typ") != typ { + if header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(header.GetString("alg")) + h := ParseHash(header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } - exp := t.Claims.GetInt64("exp") + exp := t.Claims.GetInt64(ExpClaim) switch { case exp == 0: - return ErrMissingEXP + return ErrMissingExp case now > exp: - return ErrEXP + return ErrExp } - nbf := t.Claims.GetInt64("nbf") + nbf := t.Claims.GetInt64(NbfClaim) switch { case nbf == 0: - return ErrMissingNBF + return ErrMissingNbf case now < nbf: - return ErrNBF + return ErrNbf } if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { return ErrInvalid diff --git a/jwt_test.go b/jwt_test.go index ee6dea7..02c7eeb 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -47,6 +47,7 @@ func TestValidate(t *testing.T) { t.Error(errors.New("token is expired")) } + nt.Claims.Delete(ExpClaim) err = jwt.Sign(nt) if err != nil { t.Error(err) diff --git a/token.go b/token.go index 4157f90..b7e0b89 100644 --- a/token.go +++ b/token.go @@ -55,8 +55,8 @@ func NewToken(claims Claims, hash Hash) *Token { } return &Token{ header: Claims{ - "alg": hash.Alg(), - "typ": typ, + AlgClaim: hash.Alg(), + TypClaim: Typ, }, Claims: claims, } @@ -67,7 +67,7 @@ func DecodeToken(token string) (t *Token, err error) { if token == "" { return nil, ErrEmptyToken } - parts := strings.Split(token, ".") + parts := strings.Split(token, TokenSeparator) if len(parts) < 3 { return nil, ErrMissingTokenParts } @@ -98,7 +98,7 @@ func DecodeToken(token string) (t *Token, err error) { } t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) t.signature = parts[2] - t.data = strings.Join(parts[:2], ".") + t.data = strings.Join(parts[:2], TokenSeparator) t.raw = token return } From fb52ae4a3009a42fb126dbbffaccc1bec803c2d8 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 14:18:11 +0200 Subject: [PATCH 05/12] added benchmarks --- claims.go | 9 +++++ jwt.go | 22 +++++------ jwt_test.go | 103 ++++++++++++++++++++++++++++++++++++++++++++++------ token.go | 13 ++----- 4 files changed, 113 insertions(+), 34 deletions(-) diff --git a/claims.go b/claims.go index b0c99c9..10cbcfe 100644 --- a/claims.go +++ b/claims.go @@ -76,3 +76,12 @@ func (c Claims) Set(key string, v interface{}) { func (c Claims) Delete(key string) { delete(c, key) } + +// Copy returns a new Claims map +func (c Claims) Copy() (n Claims) { + n = make(Claims) + for k, v := range c { + n[k] = v + } + return +} diff --git a/jwt.go b/jwt.go index 4c52a10..e586aa3 100644 --- a/jwt.go +++ b/jwt.go @@ -10,7 +10,6 @@ import ( "encoding/json" "errors" "io" - "strings" "sync" "time" ) @@ -125,11 +124,10 @@ func (jwt *JWT) Invalidate(t *Token) error { if err := jwt.blacklisted(t.Sig()); err != nil { return err } - header := t.Header() - if header.GetString(TypClaim) != Typ { + if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(header.GetString(AlgClaim)) + h := ParseHash(t.Header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } @@ -185,11 +183,10 @@ func (jwt *JWT) Sign(t *Token) (err error) { if t == nil { return ErrTokenIsNil } - header := t.Header() - if header.GetString(TypClaim) != Typ { + if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(header.GetString(AlgClaim)) + h := ParseHash(t.Header.GetString(AlgClaim)) if h == nil { return ErrUnsupportedAlg } @@ -200,7 +197,7 @@ func (jwt *JWT) Sign(t *Token) (err error) { if _, ok := t.Claims.Get(NbfClaim); !ok { t.Claims.Set(NbfClaim, now.Unix()) } - head, err := json.Marshal(header) + head, err := json.Marshal(t.Header) if err != nil { return } @@ -208,10 +205,10 @@ func (jwt *JWT) Sign(t *Token) (err error) { if err != nil { return } - t.data = strings.Join([]string{base64.URLEncoding.EncodeToString(head), base64.URLEncoding.EncodeToString(claims)}, TokenSeparator) + t.data = base64.URLEncoding.EncodeToString(head) + TokenSeparator + base64.URLEncoding.EncodeToString(claims) t.rawSignature = jwt.sum(t.data, h) t.signature = base64.URLEncoding.EncodeToString(t.rawSignature) - t.raw = strings.Join([]string{t.data, t.signature}, TokenSeparator) + t.raw = t.data + TokenSeparator + t.signature return } @@ -226,11 +223,10 @@ func (jwt *JWT) Verify(t *Token) error { } } now := time.Now().UTC().Unix() - header := t.Header() - if header.GetString(TypClaim) != Typ { + if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } - h := ParseHash(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 02c7eeb..59789d3 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -9,17 +9,19 @@ import ( "time" ) +var claims = Claims{ + "sub": "1234567890", + "name": "John Doe", + "admin": true, + "fizz": "buzz", +} + func TestValidate(t *testing.T) { jwt, err := New(time.Second, NewMemBlacklist(), nil) if err != nil { t.Error(err) } - token := NewToken(Claims{ - "sub": "1234567890", - "name": "John Doe", - "admin": true, - "fizz": "buzz", - }, nil) + token := NewToken(claims.Copy(), nil) err = jwt.Sign(token) if err != nil { t.Error(err) @@ -86,12 +88,7 @@ func TestNoBlacklist(t *testing.T) { if err != nil { t.Error(err) } - token := NewToken(Claims{ - "sub": "1234567890", - "name": "John Doe", - "admin": true, - "fizz": "buzz", - }, nil) + token := NewToken(claims.Copy(), nil) err = jwt.Sign(token) if err != nil { t.Error(err) @@ -133,3 +130,85 @@ func TestInvalidSecretReader(t *testing.T) { t.Error(errors.New("error should be invalid key size")) } } + +func BenchmarkDecodeToken(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), nil) + jwt.Sign(t) + token := t.String() + b.ReportAllocs() + for i := 0; i < b.N; i++ { + DecodeToken(token) + } +} + +func BenchmarkNew(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + New(0, nil, nil) + } +} + +func BenchmarkNewWithBlacklist(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + New(0, NewMemBlacklist(), nil) + } +} + +func BenchmarkSignHS256(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), nil) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkSignHS384(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS384()) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkSignHS512(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS512()) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkVerifyHS256(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), nil) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } +} + +func BenchmarkVerifyHS384(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS384()) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } +} + +func BenchmarkVerifyHS512(b *testing.B) { + jwt, _ := New(0, nil, nil) + t := NewToken(claims.Copy(), NewHS512()) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } +} diff --git a/token.go b/token.go index b7e0b89..34760a3 100644 --- a/token.go +++ b/token.go @@ -14,7 +14,7 @@ type Token struct { data string signature string rawSignature []byte - header Claims + Header Claims Claims Claims } @@ -33,11 +33,6 @@ func (t *Token) RawSig() []byte { return t.rawSignature } -// Header returns the Tokens header -func (t *Token) Header() Claims { - return t.header -} - // Data returns the first two token fields func (t *Token) Data() string { return t.data @@ -54,7 +49,7 @@ func NewToken(claims Claims, hash Hash) *Token { hash = NewHS256() } return &Token{ - header: Claims{ + Header: Claims{ AlgClaim: hash.Alg(), TypClaim: Typ, }, @@ -84,7 +79,7 @@ func DecodeToken(token string) (t *Token, err error) { if err != nil { return } - err = json.Unmarshal(header, &t.header) + err = json.Unmarshal(header, &t.Header) if err != nil { return } @@ -98,7 +93,7 @@ func DecodeToken(token string) (t *Token, err error) { } t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) t.signature = parts[2] - t.data = strings.Join(parts[:2], TokenSeparator) + t.data = parts[0] + TokenSeparator + parts[1] t.raw = token return } From 1bdb6009cf9586edb1b8bcf3618f32436a8b6929 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 15:58:24 +0200 Subject: [PATCH 06/12] added utc time functions --- jwt.go | 10 +++++----- time.go | 22 ++++++++++++++++++++++ 2 files changed, 27 insertions(+), 5 deletions(-) create mode 100644 time.go diff --git a/jwt.go b/jwt.go index e586aa3..0165de2 100644 --- a/jwt.go +++ b/jwt.go @@ -135,7 +135,7 @@ func (jwt *JWT) Invalidate(t *Token) error { switch { case exp == 0: return ErrMissingExp - case time.Now().UTC().Unix() > exp: + case Now() > exp: return ErrExp } if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { @@ -190,12 +190,12 @@ func (jwt *JWT) Sign(t *Token) (err error) { if h == nil { return ErrUnsupportedAlg } - now := time.Now().UTC() + now := time.Now() if _, ok := t.Claims.Get(ExpClaim); !ok { - t.Claims.Set(ExpClaim, now.Add(jwt.expiry).Unix()) + t.Claims.Set(ExpClaim, newExp(now, jwt.expiry)) } if _, ok := t.Claims.Get(NbfClaim); !ok { - t.Claims.Set(NbfClaim, now.Unix()) + t.Claims.Set(NbfClaim, NewNbf(now)) } head, err := json.Marshal(t.Header) if err != nil { @@ -222,7 +222,7 @@ func (jwt *JWT) Verify(t *Token) error { return err } } - now := time.Now().UTC().Unix() + now := Now() if t.Header.GetString(TypClaim) != Typ { return ErrNoJWT } diff --git a/time.go b/time.go new file mode 100644 index 0000000..71b10ca --- /dev/null +++ b/time.go @@ -0,0 +1,22 @@ +package jwt + +import "time" + +// Now returns the current time in UTC Unix format +func Now() int64 { + return NewNbf(time.Now()) +} + +// NewNbf returns a new 'not before' date +func NewNbf(t time.Time) int64 { + return t.UTC().Unix() +} + +// NewExp returns a new expiration date +func NewExp(d time.Duration) int64 { + return newExp(time.Now(), d) +} + +func newExp(t time.Time, d time.Duration) int64 { + return t.UTC().Add(d).Unix() +} From b071c473f0696d1812b1097c9fa344e38598fcba Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 16:02:17 +0200 Subject: [PATCH 07/12] cleanup --- jwt.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/jwt.go b/jwt.go index 0165de2..05fe60e 100644 --- a/jwt.go +++ b/jwt.go @@ -166,7 +166,7 @@ func (jwt *JWT) clean() { t.Stop() return } - now := time.Now().UTC().Unix() + now := Now() for k, v := range jwt.blacklist.Map() { if now > v { jwt.blacklist.Remove(k) From 4076e5cf1740ea8a33bd124016655040ff34d131 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 16:29:45 +0200 Subject: [PATCH 08/12] fixed typos, added error prefix and added usage to README.md --- README.md | 62 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- jwt.go | 34 ++++++++++++++++-------------- token.go | 8 +++---- 3 files changed, 83 insertions(+), 21 deletions(-) diff --git a/README.md b/README.md index ad89f76..74aa8a3 100644 --- a/README.md +++ b/README.md @@ -1 +1,61 @@ -# jwt - a simple JSON Web Token library in Go +# jwt - a simple JSON Web Token library written in Go + +## Usage + +``` +package main + +import ( + "fmt" + "git.giftfish.de/ston1th/jwt" + "log" + "time" +) + +func main() { + // create a new signing instance + JWT, err := jwt.New(time.Hour, jwt.NewMemBlacklist(), nil) + if err != nil { + log.Fatal(err) + } + + // create a new token + t := jwt.NewToken(jwt.Claims{ + "username": "admin", + }, nil) + + // sign the token + err = JWT.Sign(t) + if err != nil { + log.Fatal(err) + } + + // print the token string + token := t.String() + fmt.Println(token) + + // decode the token string + newToken, err := jwt.DecodeToken(token) + if err != nil { + log.Fatal(err) + } + + // verify the decoded token + err = JWT.Verify(newToken) + if err != nil { + log.Fatal(err) + } + + // add the token to the blacklist + err = JWT.Invalidate(newToken) + if err != nil { + log.Fatal(err) + } + + // try to verify the blacklisted token, it should fail + err = JWT.Verify(newToken) + if err != nil { + log.Fatal(err) + } +} +``` diff --git a/jwt.go b/jwt.go index 05fe60e..76217f4 100644 --- a/jwt.go +++ b/jwt.go @@ -39,21 +39,23 @@ const ( // DefaultSecretReader is the default secret key generator var DefaultSecretReader = rand.Reader +const jwtErr = "jwt: " + var ( - ErrNoJWT = errors.New("not a json web token") - ErrEmptyToken = errors.New("token is empty") - ErrUnsupportedAlg = errors.New("unsupported algorithm") - ErrInvalid = errors.New("token validation failed") - ErrBlacklisted = errors.New("token blacklisted") - ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled") - ErrNbf = errors.New("token not valid yet") - ErrExp = errors.New("token expired") - ErrMissingNbf = errors.New("missing nbf claim") - ErrMissingExp = errors.New("missing exp claim") - ErrMissingTokenParts = errors.New("missing token parts") - ErrEmptySignature = errors.New("token signature is empty") - ErrTokenIsNil = errors.New("token is nil") - ErrInvalidKeySize = errors.New("invalid secret key size") + ErrNoJWT = errors.New(jwtErr + "not a json web token") + ErrEmptyToken = errors.New(jwtErr + "token is empty") + ErrUnsupportedAlg = errors.New(jwtErr + "unsupported algorithm") + ErrInvalid = errors.New(jwtErr + "token validation failed") + ErrBlacklisted = errors.New(jwtErr + "token blacklisted") + ErrBlacklistNotEnabled = errors.New(jwtErr + "blacklisting is not enabled") + ErrNbf = errors.New(jwtErr + "token not valid yet") + ErrExp = errors.New(jwtErr + "token expired") + ErrMissingNbf = errors.New(jwtErr + "missing nbf claim") + ErrMissingExp = errors.New(jwtErr + "missing exp claim") + ErrMissingTokenParts = errors.New(jwtErr + "missing token parts") + ErrEmptySignature = errors.New(jwtErr + "token signature is empty") + ErrTokenIsNil = errors.New(jwtErr + "token is nil") + ErrInvalidKeySize = errors.New(jwtErr + "invalid secret key size") ) // JWT represents the JSON Web Token signing and blacklisting infrastructure @@ -83,7 +85,7 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err key := make([]byte, KeySize) i, err := secret.Read(key) if err != nil { - return nil, errors.New("secret reader error: " + err.Error()) + return nil, errors.New(jwtErr + "secret reader error: " + err.Error()) } if i < KeySize { return nil, ErrInvalidKeySize @@ -250,7 +252,7 @@ func (jwt *JWT) Verify(t *Token) error { return nil } -// Stop will end the cleaner goroutinge execution +// Stop will end the cleaner goroutine execution func (jwt *JWT) Stop() error { if jwt.blacklist == nil { return ErrBlacklistNotEnabled diff --git a/token.go b/token.go index 34760a3..83aee75 100644 --- a/token.go +++ b/token.go @@ -8,7 +8,7 @@ import ( "strings" ) -// Token is the JWT token representation +// Token represents a JWT token type Token struct { raw string data string @@ -23,17 +23,17 @@ func (t *Token) String() string { return t.raw } -// Sig returns the Tokens URLEncoded signature +// Sig returns the tokens URLEncoded signature func (t *Token) Sig() string { return t.signature } -// RawSig returns the Tokens raw signature +// RawSig returns the tokens raw signature func (t *Token) RawSig() []byte { return t.rawSignature } -// Data returns the first two token fields +// Data returns the first two URLEncoded token fields func (t *Token) Data() string { return t.data } From 45825d685eb36328c3dec9fcf1affa3f34dea227 Mon Sep 17 00:00:00 2001 From: ston1th Date: Wed, 19 Sep 2018 19:42:25 +0200 Subject: [PATCH 09/12] added go.mod --- go.mod | 1 + 1 file changed, 1 insertion(+) create mode 100644 go.mod diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..6752b42 --- /dev/null +++ b/go.mod @@ -0,0 +1 @@ +module git.giftfish.de/ston1th/jwt/v3 From a00a154398922890bf52dbc09dfb60f6b801b1aa Mon Sep 17 00:00:00 2001 From: ston1th Date: Sat, 22 Sep 2018 16:34:10 +0200 Subject: [PATCH 10/12] added errors to blacklist --- blacklist.go | 14 ++++++++------ jwt.go | 9 ++++++--- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/blacklist.go b/blacklist.go index 9ecf662..b60a928 100644 --- a/blacklist.go +++ b/blacklist.go @@ -6,10 +6,10 @@ import "sync" // Blacklist is the blacklisting storage interface type Blacklist interface { - Add(string, int64) - Remove(string) + Add(string, int64) error + Remove(string) error Check(string) bool - Map() BlacklistMap + Map() (BlacklistMap, error) } // 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 -func (mb *MemBlacklist) Add(sig string, exp int64) { +func (mb *MemBlacklist) Add(sig string, exp int64) error { mb.Lock() mb.list[sig] = exp mb.Unlock() + return nil } // Remove deletes a token signature from the blacklist -func (mb *MemBlacklist) Remove(sig string) { +func (mb *MemBlacklist) Remove(sig string) error { mb.Lock() delete(mb.list, sig) mb.Unlock() + return nil } // 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 -func (mb *MemBlacklist) Map() (list BlacklistMap) { +func (mb *MemBlacklist) Map() (list BlacklistMap, err error) { list = make(BlacklistMap) mb.RLock() for k, v := range mb.list { diff --git a/jwt.go b/jwt.go index 76217f4..d77c7dd 100644 --- a/jwt.go +++ b/jwt.go @@ -143,8 +143,7 @@ func (jwt *JWT) Invalidate(t *Token) error { if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { return ErrInvalid } - jwt.blacklist.Add(t.Sig(), exp) - return nil + return jwt.blacklist.Add(t.Sig(), exp) } // blacklisted checks if a token is blacklisted @@ -169,7 +168,11 @@ func (jwt *JWT) clean() { return } 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 { jwt.blacklist.Remove(k) } From 1613055dfdfcb7c4835c3f2bd30e76f1c86b29e6 Mon Sep 17 00:00:00 2001 From: ston1th Date: Fri, 28 Dec 2018 19:10:22 +0100 Subject: [PATCH 11/12] use raw base64 as defined by rfc7515 --- encoding.go | 7 +++++++ jwt.go | 5 ++--- token.go | 7 +++---- 3 files changed, 12 insertions(+), 7 deletions(-) create mode 100644 encoding.go diff --git a/encoding.go b/encoding.go new file mode 100644 index 0000000..73b81cf --- /dev/null +++ b/encoding.go @@ -0,0 +1,7 @@ +// Copyright (C) 2018 Marius Schellenberger + +package jwt + +import "encoding/base64" + +var enc = base64.RawURLEncoding diff --git a/jwt.go b/jwt.go index d77c7dd..44f82ec 100644 --- a/jwt.go +++ b/jwt.go @@ -6,7 +6,6 @@ package jwt import ( "crypto/hmac" "crypto/rand" - "encoding/base64" "encoding/json" "errors" "io" @@ -210,9 +209,9 @@ func (jwt *JWT) Sign(t *Token) (err error) { if err != nil { 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.signature = base64.URLEncoding.EncodeToString(t.rawSignature) + t.signature = enc.EncodeToString(t.rawSignature) t.raw = t.data + TokenSeparator + t.signature return } diff --git a/token.go b/token.go index 83aee75..a1f1dd1 100644 --- a/token.go +++ b/token.go @@ -3,7 +3,6 @@ package jwt import ( - "encoding/base64" "encoding/json" "strings" ) @@ -75,7 +74,7 @@ func DecodeToken(token string) (t *Token, err error) { } } t = new(Token) - header, err := base64.URLEncoding.DecodeString(parts[0]) + header, err := enc.DecodeString(parts[0]) if err != nil { return } @@ -83,7 +82,7 @@ func DecodeToken(token string) (t *Token, err error) { if err != nil { return } - claims, err := base64.URLEncoding.DecodeString(parts[1]) + claims, err := enc.DecodeString(parts[1]) if err != nil { return } @@ -91,7 +90,7 @@ func DecodeToken(token string) (t *Token, err error) { if err != nil { return } - t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) + t.rawSignature, err = enc.DecodeString(parts[2]) t.signature = parts[2] t.data = parts[0] + TokenSeparator + parts[1] t.raw = token From 473d54569c2f10f5073af9ad1b504745a6cdcb92 Mon Sep 17 00:00:00 2001 From: ston1th Date: Wed, 15 Oct 2025 21:49:13 +0200 Subject: [PATCH 12/12] added config options and nonce --- LICENSE | 2 +- blacklist.go | 2 +- claims.go | 31 ++++++++-------------- encoding.go | 2 +- hash.go | 33 +++++++++-------------- jwt.go | 74 ++++++++++++++++++++++++++++++++++++---------------- jwt_test.go | 72 +++++++++++++++++++++++++++++++++++++------------- rand.go | 11 ++++++++ time.go | 2 ++ token.go | 4 +-- 10 files changed, 147 insertions(+), 86 deletions(-) create mode 100644 rand.go diff --git a/LICENSE b/LICENSE index 115e96e..74cee2b 100644 --- a/LICENSE +++ b/LICENSE @@ -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 diff --git a/blacklist.go b/blacklist.go index b60a928..180c5fd 100644 --- a/blacklist.go +++ b/blacklist.go @@ -1,4 +1,4 @@ -// Copyright (C) 2018 Marius Schellenberger +// Copyright (C) 2025 Marius Schellenberger package jwt diff --git a/claims.go b/claims.go index 10cbcfe..a5f4364 100644 --- a/claims.go +++ b/claims.go @@ -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 } diff --git a/encoding.go b/encoding.go index 73b81cf..3693f5f 100644 --- a/encoding.go +++ b/encoding.go @@ -1,4 +1,4 @@ -// Copyright (C) 2018 Marius Schellenberger +// Copyright (C) 2025 Marius Schellenberger package jwt diff --git a/hash.go b/hash.go index 2c224ab..e47c0cc 100644 --- a/hash.go +++ b/hash.go @@ -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 diff --git a/jwt.go b/jwt.go index 44f82ec..f9ce6af 100644 --- a/jwt.go +++ b/jwt.go @@ -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 } diff --git a/jwt_test.go b/jwt_test.go index 59789d3..664cf8a 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -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++ { diff --git a/rand.go b/rand.go new file mode 100644 index 0000000..bd1b191 --- /dev/null +++ b/rand.go @@ -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 +} diff --git a/time.go b/time.go index 71b10ca..77e151f 100644 --- a/time.go +++ b/time.go @@ -1,3 +1,5 @@ +// Copyright (C) 2025 Marius Schellenberger + package jwt import "time" diff --git a/token.go b/token.go index a1f1dd1..772415c 100644 --- a/token.go +++ b/token.go @@ -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{