From f2d57c6e1868d1dc5ab021876511f0c5c46568b8 Mon Sep 17 00:00:00 2001 From: ston1th Date: Thu, 13 Sep 2018 11:39:49 +0200 Subject: [PATCH] major refactoring --- claims.go | 28 +++++++++-- jwt.go | 133 +++++++++++++++++++--------------------------------- jwt_test.go | 15 ++++++ token.go | 6 +-- 4 files changed, 90 insertions(+), 92 deletions(-) diff --git a/claims.go b/claims.go index 74e26c9..0071f30 100644 --- a/claims.go +++ b/claims.go @@ -29,17 +29,37 @@ func (c Claims) GetBool(key string) (b bool) { // GetInt returns an int from the claims map 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 { i, _ = v.(int) } return } -// GetFloat returns a float from the claims map -func (c Claims) GetFloat(key string) (f float64) { - if v, ok := c.Get(key); ok && v != nil { - f, _ = v.(float64) +// GetInt64 returns an int64 from the claims map +func (c Claims) GetInt64(key string) (i int64) { + if f, ok := c.getFloat64(key); ok { + 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 } diff --git a/jwt.go b/jwt.go index 32a9012..eff51c5 100644 --- a/jwt.go +++ b/jwt.go @@ -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()) { diff --git a/jwt_test.go b/jwt_test.go index 103eb32..ee6dea7 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -42,6 +42,20 @@ func TestValidate(t *testing.T) { 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) if err != nil { t.Error(err) @@ -50,6 +64,7 @@ func TestValidate(t *testing.T) { if err == nil { t.Error(errors.New("double invalidate")) } + err = jwt.Verify(nt) if err == nil { t.Error(errors.New("token should be blacklisted")) diff --git a/token.go b/token.go index 3a08da2..4157f90 100644 --- a/token.go +++ b/token.go @@ -14,7 +14,7 @@ type Token struct { data string signature string rawSignature []byte - header map[string]interface{} + header Claims Claims Claims } @@ -34,7 +34,7 @@ func (t *Token) RawSig() []byte { } // Header returns the Tokens header -func (t *Token) Header() map[string]interface{} { +func (t *Token) Header() Claims { return t.header } @@ -54,7 +54,7 @@ func NewToken(claims Claims, hash Hash) *Token { hash = NewHS256() } return &Token{ - header: map[string]interface{}{ + header: Claims{ "alg": hash.Alg(), "typ": typ, },