custom exp and nbf
This commit is contained in:
parent
f2d57c6e18
commit
b4a92fe472
4 changed files with 60 additions and 29 deletions
|
|
@ -3,6 +3,7 @@
|
||||||
package jwt
|
package jwt
|
||||||
|
|
||||||
// Claims is the claim type of the token
|
// Claims is the claim type of the token
|
||||||
|
// The Claims map is not goroutine safe
|
||||||
type Claims map[string]interface{}
|
type Claims map[string]interface{}
|
||||||
|
|
||||||
// Get returns a value from the claims map
|
// Get returns a value from the claims map
|
||||||
|
|
@ -70,3 +71,8 @@ func (c Claims) Set(key string, v interface{}) {
|
||||||
}
|
}
|
||||||
c[key] = v
|
c[key] = v
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Delete removes the key from the claims map
|
||||||
|
func (c Claims) Delete(key string) {
|
||||||
|
delete(c, key)
|
||||||
|
}
|
||||||
|
|
|
||||||
74
jwt.go
74
jwt.go
|
|
@ -23,7 +23,18 @@ const (
|
||||||
// KeySize is the secret key size
|
// KeySize is the secret key size
|
||||||
KeySize = 64
|
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
|
// DefaultSecretReader is the default secret key generator
|
||||||
|
|
@ -36,10 +47,10 @@ var (
|
||||||
ErrInvalid = errors.New("token validation failed")
|
ErrInvalid = errors.New("token validation failed")
|
||||||
ErrBlacklisted = errors.New("token blacklisted")
|
ErrBlacklisted = errors.New("token blacklisted")
|
||||||
ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled")
|
ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled")
|
||||||
ErrNBF = errors.New("token not valid yet")
|
ErrNbf = errors.New("token not valid yet")
|
||||||
ErrEXP = errors.New("token expired")
|
ErrExp = errors.New("token expired")
|
||||||
ErrMissingNBF = errors.New("missing nbf claim")
|
ErrMissingNbf = errors.New("missing nbf claim")
|
||||||
ErrMissingEXP = errors.New("missing exp claim")
|
ErrMissingExp = errors.New("missing exp claim")
|
||||||
ErrMissingTokenParts = errors.New("missing token parts")
|
ErrMissingTokenParts = errors.New("missing token parts")
|
||||||
ErrEmptySignature = errors.New("token signature is empty")
|
ErrEmptySignature = errors.New("token signature is empty")
|
||||||
ErrTokenIsNil = errors.New("token is nil")
|
ErrTokenIsNil = errors.New("token is nil")
|
||||||
|
|
@ -58,6 +69,8 @@ type JWT struct {
|
||||||
|
|
||||||
// New returns a new JWT object with the given expiry timeout.
|
// 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.
|
// 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.
|
// 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(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
|
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
|
// sum calculates the HMAC hash sum of the token
|
||||||
func (jwt *JWT) sum(token string, h Hash) []byte {
|
func (jwt *JWT) sum(token string, h Hash) []byte {
|
||||||
mac := hmac.New(h.Hash, jwt.key)
|
mac := hmac.New(h.Hash, jwt.key)
|
||||||
|
|
@ -108,19 +126,19 @@ func (jwt *JWT) Invalidate(t *Token) error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
header := t.Header()
|
header := t.Header()
|
||||||
if header.GetString("typ") != typ {
|
if header.GetString(TypClaim) != Typ {
|
||||||
return ErrNoJWT
|
return ErrNoJWT
|
||||||
}
|
}
|
||||||
h := ParseHash(header.GetString("alg"))
|
h := ParseHash(header.GetString(AlgClaim))
|
||||||
if h == nil {
|
if h == nil {
|
||||||
return ErrUnsupportedAlg
|
return ErrUnsupportedAlg
|
||||||
}
|
}
|
||||||
exp := t.Claims.GetInt64("exp")
|
exp := t.Claims.GetInt64(ExpClaim)
|
||||||
switch {
|
switch {
|
||||||
case exp == 0:
|
case exp == 0:
|
||||||
return ErrMissingEXP
|
return ErrMissingExp
|
||||||
case time.Now().UTC().Unix() > exp:
|
case time.Now().UTC().Unix() > exp:
|
||||||
return ErrEXP
|
return ErrExp
|
||||||
}
|
}
|
||||||
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
|
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
|
||||||
return ErrInvalid
|
return ErrInvalid
|
||||||
|
|
@ -160,22 +178,28 @@ func (jwt *JWT) clean() {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sign will sign the provided token using the secret key
|
// 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) {
|
func (jwt *JWT) Sign(t *Token) (err error) {
|
||||||
if t == nil {
|
if t == nil {
|
||||||
return ErrTokenIsNil
|
return ErrTokenIsNil
|
||||||
}
|
}
|
||||||
header := t.Header()
|
header := t.Header()
|
||||||
if header.GetString("typ") != typ {
|
if header.GetString(TypClaim) != Typ {
|
||||||
return ErrNoJWT
|
return ErrNoJWT
|
||||||
}
|
}
|
||||||
h := ParseHash(header.GetString("alg"))
|
h := ParseHash(header.GetString(AlgClaim))
|
||||||
if h == nil {
|
if h == nil {
|
||||||
return ErrUnsupportedAlg
|
return ErrUnsupportedAlg
|
||||||
}
|
}
|
||||||
now := time.Now().UTC()
|
now := time.Now().UTC()
|
||||||
t.Claims.Set("exp", now.Add(jwt.expiry).Unix())
|
if _, ok := t.Claims.Get(ExpClaim); !ok {
|
||||||
t.Claims.Set("nbf", now.Unix())
|
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)
|
head, err := json.Marshal(header)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
|
|
@ -184,10 +208,10 @@ func (jwt *JWT) Sign(t *Token) (err error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
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.rawSignature = jwt.sum(t.data, h)
|
||||||
t.signature = base64.URLEncoding.EncodeToString(t.rawSignature)
|
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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -203,26 +227,26 @@ func (jwt *JWT) Verify(t *Token) error {
|
||||||
}
|
}
|
||||||
now := time.Now().UTC().Unix()
|
now := time.Now().UTC().Unix()
|
||||||
header := t.Header()
|
header := t.Header()
|
||||||
if header.GetString("typ") != typ {
|
if header.GetString(TypClaim) != Typ {
|
||||||
return ErrNoJWT
|
return ErrNoJWT
|
||||||
}
|
}
|
||||||
h := ParseHash(header.GetString("alg"))
|
h := ParseHash(header.GetString(AlgClaim))
|
||||||
if h == nil {
|
if h == nil {
|
||||||
return ErrUnsupportedAlg
|
return ErrUnsupportedAlg
|
||||||
}
|
}
|
||||||
exp := t.Claims.GetInt64("exp")
|
exp := t.Claims.GetInt64(ExpClaim)
|
||||||
switch {
|
switch {
|
||||||
case exp == 0:
|
case exp == 0:
|
||||||
return ErrMissingEXP
|
return ErrMissingExp
|
||||||
case now > exp:
|
case now > exp:
|
||||||
return ErrEXP
|
return ErrExp
|
||||||
}
|
}
|
||||||
nbf := t.Claims.GetInt64("nbf")
|
nbf := t.Claims.GetInt64(NbfClaim)
|
||||||
switch {
|
switch {
|
||||||
case nbf == 0:
|
case nbf == 0:
|
||||||
return ErrMissingNBF
|
return ErrMissingNbf
|
||||||
case now < nbf:
|
case now < nbf:
|
||||||
return ErrNBF
|
return ErrNbf
|
||||||
}
|
}
|
||||||
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
|
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
|
||||||
return ErrInvalid
|
return ErrInvalid
|
||||||
|
|
|
||||||
|
|
@ -47,6 +47,7 @@ func TestValidate(t *testing.T) {
|
||||||
t.Error(errors.New("token is expired"))
|
t.Error(errors.New("token is expired"))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
nt.Claims.Delete(ExpClaim)
|
||||||
err = jwt.Sign(nt)
|
err = jwt.Sign(nt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Error(err)
|
t.Error(err)
|
||||||
|
|
|
||||||
8
token.go
8
token.go
|
|
@ -55,8 +55,8 @@ func NewToken(claims Claims, hash Hash) *Token {
|
||||||
}
|
}
|
||||||
return &Token{
|
return &Token{
|
||||||
header: Claims{
|
header: Claims{
|
||||||
"alg": hash.Alg(),
|
AlgClaim: hash.Alg(),
|
||||||
"typ": typ,
|
TypClaim: Typ,
|
||||||
},
|
},
|
||||||
Claims: claims,
|
Claims: claims,
|
||||||
}
|
}
|
||||||
|
|
@ -67,7 +67,7 @@ func DecodeToken(token string) (t *Token, err error) {
|
||||||
if token == "" {
|
if token == "" {
|
||||||
return nil, ErrEmptyToken
|
return nil, ErrEmptyToken
|
||||||
}
|
}
|
||||||
parts := strings.Split(token, ".")
|
parts := strings.Split(token, TokenSeparator)
|
||||||
if len(parts) < 3 {
|
if len(parts) < 3 {
|
||||||
return nil, ErrMissingTokenParts
|
return nil, ErrMissingTokenParts
|
||||||
}
|
}
|
||||||
|
|
@ -98,7 +98,7 @@ func DecodeToken(token string) (t *Token, err error) {
|
||||||
}
|
}
|
||||||
t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2])
|
t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2])
|
||||||
t.signature = parts[2]
|
t.signature = parts[2]
|
||||||
t.data = strings.Join(parts[:2], ".")
|
t.data = strings.Join(parts[:2], TokenSeparator)
|
||||||
t.raw = token
|
t.raw = token
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue