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
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
133
jwt.go
133
jwt.go
|
|
@ -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
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return ErrNoJWT
|
return ErrNoJWT
|
||||||
}
|
}
|
||||||
var h Hash
|
h := ParseHash(header.GetString("alg"))
|
||||||
switch a := header["alg"].(type) {
|
if h == nil {
|
||||||
case string:
|
|
||||||
h = ParseHash(a)
|
|
||||||
if h == nil {
|
|
||||||
return ErrUnsupportedAlg
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return ErrUnsupportedAlg
|
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()) {
|
||||||
|
|
|
||||||
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"))
|
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"))
|
||||||
|
|
|
||||||
6
token.go
6
token.go
|
|
@ -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,
|
||||||
},
|
},
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue