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

View file

@ -29,17 +29,37 @@ func (c Claims) GetBool(key string) (b bool) {
// GetInt returns an int from the claims map // GetInt returns an int from the claims map
func (c Claims) GetInt(key string) (i int) { 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 { if v, ok := c.Get(key); ok && v != nil {
i, _ = v.(int) i, _ = v.(int)
} }
return return
} }
// GetFloat returns a float from the claims map // GetInt64 returns an int64 from the claims map
func (c Claims) GetFloat(key string) (f float64) { func (c Claims) GetInt64(key string) (i int64) {
if v, ok := c.Get(key); ok && v != nil { if f, ok := c.getFloat64(key); ok {
f, _ = v.(float64) 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 return
} }

131
jwt.go
View file

@ -16,8 +16,8 @@ import (
) )
const ( const (
// Header is the default HTTP Authorization header name // HTTPHeader is the default HTTP Authorization header name
Header = "Authorization" HTTPHeader = "Authorization"
// DefaultExpiry is the default token expiration time // DefaultExpiry is the default token expiration time
DefaultExpiry = time.Hour * 12 DefaultExpiry = time.Hour * 12
// KeySize is the secret key size // KeySize is the secret key size
@ -38,6 +38,7 @@ var (
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")
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")
@ -103,26 +104,27 @@ func (jwt *JWT) Invalidate(t *Token) error {
if t == nil { if t == nil {
return ErrTokenIsNil 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 { if err := jwt.blacklisted(t.Sig()); err != nil {
return err 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) jwt.blacklist.Add(t.Sig(), exp)
return nil return nil
} }
@ -164,27 +166,16 @@ func (jwt *JWT) Sign(t *Token) (err error) {
return ErrTokenIsNil return ErrTokenIsNil
} }
header := t.Header() header := t.Header()
switch t := header["typ"].(type) { if header.GetString("typ") != typ {
case string:
if t != typ {
return ErrNoJWT return ErrNoJWT
} }
default: h := ParseHash(header.GetString("alg"))
return ErrNoJWT
}
var h Hash
switch a := header["alg"].(type) {
case string:
h = ParseHash(a)
if h == nil { if h == nil {
return ErrUnsupportedAlg return ErrUnsupportedAlg
} }
default:
return ErrUnsupportedAlg
}
now := time.Now().UTC() now := time.Now().UTC()
t.Claims["exp"] = now.Add(jwt.expiry).Unix() t.Claims.Set("exp", now.Add(jwt.expiry).Unix())
t.Claims["nbf"] = now.Unix() t.Claims.Set("nbf", now.Unix())
head, err := json.Marshal(header) head, err := json.Marshal(header)
if err != nil { if err != nil {
return return
@ -194,7 +185,9 @@ func (jwt *JWT) Sign(t *Token) (err error) {
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)}, ".")
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 return
} }
@ -203,62 +196,32 @@ func (jwt *JWT) Verify(t *Token) error {
if t == nil { if t == nil {
return ErrTokenIsNil 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 jwt.blacklist != nil {
if err := jwt.blacklisted(t.Sig()); err != nil { if err := jwt.blacklisted(t.Sig()); err != nil {
return err 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 return ErrEXP
} }
if nbfOK && now < nbf { nbf := t.Claims.GetInt64("nbf")
switch {
case nbf == 0:
return ErrMissingNBF
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()) {

View file

@ -42,6 +42,20 @@ func TestValidate(t *testing.T) {
t.Error(errors.New("token should be expired")) 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) err = jwt.Invalidate(nt)
if err != nil { if err != nil {
t.Error(err) t.Error(err)
@ -50,6 +64,7 @@ func TestValidate(t *testing.T) {
if err == nil { if err == nil {
t.Error(errors.New("double invalidate")) t.Error(errors.New("double invalidate"))
} }
err = jwt.Verify(nt) err = jwt.Verify(nt)
if err == nil { if err == nil {
t.Error(errors.New("token should be blacklisted")) t.Error(errors.New("token should be blacklisted"))

View file

@ -14,7 +14,7 @@ type Token struct {
data string data string
signature string signature string
rawSignature []byte rawSignature []byte
header map[string]interface{} header Claims
Claims Claims Claims Claims
} }
@ -34,7 +34,7 @@ func (t *Token) RawSig() []byte {
} }
// Header returns the Tokens header // Header returns the Tokens header
func (t *Token) Header() map[string]interface{} { func (t *Token) Header() Claims {
return t.header return t.header
} }
@ -54,7 +54,7 @@ func NewToken(claims Claims, hash Hash) *Token {
hash = NewHS256() hash = NewHS256()
} }
return &Token{ return &Token{
header: map[string]interface{}{ header: Claims{
"alg": hash.Alg(), "alg": hash.Alg(),
"typ": typ, "typ": typ,
}, },