major refactoring

This commit is contained in:
ston1th 2018-09-13 11:39:49 +02:00
commit f2d57c6e18
4 changed files with 87 additions and 89 deletions

133
jwt.go
View file

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