From b4a92fe472c89fa93c028d4d3762b985094f1bd1 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 13:27:32 +0200 Subject: [PATCH] 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 }