major refactoring
This commit is contained in:
parent
826d035d3e
commit
f2d57c6e18
4 changed files with 87 additions and 89 deletions
28
claims.go
28
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
|
||||
}
|
||||
|
||||
|
|
|
|||
131
jwt.go
131
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 {
|
||||
if header.GetString("typ") != typ {
|
||||
return ErrNoJWT
|
||||
}
|
||||
default:
|
||||
return ErrNoJWT
|
||||
}
|
||||
var h Hash
|
||||
switch a := header["alg"].(type) {
|
||||
case string:
|
||||
h = ParseHash(a)
|
||||
h := ParseHash(header.GetString("alg"))
|
||||
if h == nil {
|
||||
return ErrUnsupportedAlg
|
||||
}
|
||||
default:
|
||||
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()) {
|
||||
|
|
|
|||
15
jwt_test.go
15
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"))
|
||||
|
|
|
|||
6
token.go
6
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,
|
||||
},
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue