custom exp and nbf

This commit is contained in:
ston1th 2018-09-13 13:27:32 +02:00
commit b4a92fe472
4 changed files with 60 additions and 29 deletions

View file

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

74
jwt.go
View file

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

View file

@ -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)

View file

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