diff --git a/blacklist.go b/blacklist.go index 024796d..9ecf662 100644 --- a/blacklist.go +++ b/blacklist.go @@ -22,7 +22,7 @@ type MemBlacklist struct { list BlacklistMap } -// NewMemBlacklist +// NewMemBlacklist implements the Blacklist interface using an in-memory map func NewMemBlacklist() *MemBlacklist { return &MemBlacklist{list: make(BlacklistMap)} } diff --git a/jwt.go b/jwt.go index 1268cb2..32a9012 100644 --- a/jwt.go +++ b/jwt.go @@ -11,6 +11,7 @@ import ( "errors" "io" "strings" + "sync" "time" ) @@ -50,6 +51,7 @@ type JWT struct { expiry time.Duration blacklist Blacklist + stopOnce sync.Once done chan struct{} } @@ -85,6 +87,7 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err return jwt, nil } +// sum calculates the HMAC hash sum of the token func (jwt *JWT) sum(token string, h Hash) []byte { mac := hmac.New(h.Hash, jwt.key) mac.Write([]byte(token)) @@ -248,8 +251,7 @@ func (jwt *JWT) Verify(t *Token) error { nbfOK = true } if jwt.blacklist != nil { - err := jwt.blacklisted(t.Sig()) - if err != nil { + if err := jwt.blacklisted(t.Sig()); err != nil { return err } } @@ -266,106 +268,12 @@ func (jwt *JWT) Verify(t *Token) error { } // Stop will end the cleaner goroutinge execution -// Calling Stop twice or more will panic func (jwt *JWT) Stop() error { if jwt.blacklist == nil { return ErrBlacklistNotEnabled } - close(jwt.done) + jwt.stopOnce.Do(func() { + close(jwt.done) + }) return nil } - -// DecodeToken decodes a raw string token into a *Token object -func DecodeToken(token string) (t *Token, err error) { - if token == "" { - return nil, ErrEmptyToken - } - parts := strings.Split(token, ".") - if len(parts) < 3 { - return nil, ErrMissingTokenParts - } - if len(parts) > 3 { - return nil, ErrNoJWT - } - for _, v := range parts { - if v == "" { - return nil, ErrNoJWT - } - } - t = new(Token) - header, err := base64.URLEncoding.DecodeString(parts[0]) - if err != nil { - return - } - err = json.Unmarshal(header, &t.header) - if err != nil { - return - } - claims, err := base64.URLEncoding.DecodeString(parts[1]) - if err != nil { - return - } - err = json.Unmarshal(claims, &t.Claims) - if err != nil { - return - } - t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) - t.signature = parts[2] - t.data = strings.Join(parts[:2], ".") - t.raw = token - return -} - -// Token is the JWT token representation -type Token struct { - raw string - data string - signature string - rawSignature []byte - header map[string]interface{} - Claims Claims -} - -// String returns the tokens encoded string -func (t *Token) String() string { - return t.raw -} - -// Sig returns the Tokens URLEncoded signature -func (t *Token) Sig() string { - return t.signature -} - -// RawSig returns the Tokens raw signature -func (t *Token) RawSig() []byte { - return t.rawSignature -} - -// Header returns the Tokens header -func (t *Token) Header() map[string]interface{} { - return t.header -} - -// Data returns the first two token fields -func (t *Token) Data() string { - return t.data -} - -// NewToken returns a new *Token using the provided hash algorithm and claims -// If claims is nil, an empty map is used -// If hash is nil, HS256 is used -func NewToken(claims Claims, hash Hash) *Token { - if claims == nil { - claims = make(Claims) - } - if hash == nil { - hash = NewHS256() - } - return &Token{ - header: map[string]interface{}{ - "alg": hash.Alg(), - "typ": typ, - }, - Claims: claims, - } -} diff --git a/jwt_test.go b/jwt_test.go index 17a519b..103eb32 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -58,6 +58,11 @@ func TestValidate(t *testing.T) { if err != nil { t.Error(err) } + // should not panic + err = jwt.Stop() + if err != nil { + t.Error(err) + } } func TestNoBlacklist(t *testing.T) { diff --git a/token.go b/token.go new file mode 100644 index 0000000..3a08da2 --- /dev/null +++ b/token.go @@ -0,0 +1,104 @@ +// Copyright (C) 2018 Marius Schellenberger + +package jwt + +import ( + "encoding/base64" + "encoding/json" + "strings" +) + +// Token is the JWT token representation +type Token struct { + raw string + data string + signature string + rawSignature []byte + header map[string]interface{} + Claims Claims +} + +// String returns the tokens encoded string +func (t *Token) String() string { + return t.raw +} + +// Sig returns the Tokens URLEncoded signature +func (t *Token) Sig() string { + return t.signature +} + +// RawSig returns the Tokens raw signature +func (t *Token) RawSig() []byte { + return t.rawSignature +} + +// Header returns the Tokens header +func (t *Token) Header() map[string]interface{} { + return t.header +} + +// Data returns the first two token fields +func (t *Token) Data() string { + return t.data +} + +// NewToken returns a new *Token using the provided hash algorithm and claims +// If claims is nil, an empty map is used +// If hash is nil, HS256 is used +func NewToken(claims Claims, hash Hash) *Token { + if claims == nil { + claims = make(Claims) + } + if hash == nil { + hash = NewHS256() + } + return &Token{ + header: map[string]interface{}{ + "alg": hash.Alg(), + "typ": typ, + }, + Claims: claims, + } +} + +// DecodeToken decodes a raw string token into a *Token object +func DecodeToken(token string) (t *Token, err error) { + if token == "" { + return nil, ErrEmptyToken + } + parts := strings.Split(token, ".") + if len(parts) < 3 { + return nil, ErrMissingTokenParts + } + if len(parts) > 3 { + return nil, ErrNoJWT + } + for _, v := range parts { + if v == "" { + return nil, ErrNoJWT + } + } + t = new(Token) + header, err := base64.URLEncoding.DecodeString(parts[0]) + if err != nil { + return + } + err = json.Unmarshal(header, &t.header) + if err != nil { + return + } + claims, err := base64.URLEncoding.DecodeString(parts[1]) + if err != nil { + return + } + err = json.Unmarshal(claims, &t.Claims) + if err != nil { + return + } + t.rawSignature, err = base64.URLEncoding.DecodeString(parts[2]) + t.signature = parts[2] + t.data = strings.Join(parts[:2], ".") + t.raw = token + return +}