split up library and removed Stop() panic when called twice

This commit is contained in:
ston1th 2018-09-12 23:35:03 +02:00
commit 826d035d3e
4 changed files with 117 additions and 100 deletions

View file

@ -22,7 +22,7 @@ type MemBlacklist struct {
list BlacklistMap list BlacklistMap
} }
// NewMemBlacklist // NewMemBlacklist implements the Blacklist interface using an in-memory map
func NewMemBlacklist() *MemBlacklist { func NewMemBlacklist() *MemBlacklist {
return &MemBlacklist{list: make(BlacklistMap)} return &MemBlacklist{list: make(BlacklistMap)}
} }

104
jwt.go
View file

@ -11,6 +11,7 @@ import (
"errors" "errors"
"io" "io"
"strings" "strings"
"sync"
"time" "time"
) )
@ -50,6 +51,7 @@ type JWT struct {
expiry time.Duration expiry time.Duration
blacklist Blacklist blacklist Blacklist
stopOnce sync.Once
done chan struct{} done chan struct{}
} }
@ -85,6 +87,7 @@ func New(expiry time.Duration, blacklist Blacklist, secret io.Reader) (*JWT, err
return jwt, nil return jwt, nil
} }
// sum calculates the HMAC hash sum of the token
func (jwt *JWT) sum(token string, h Hash) []byte { func (jwt *JWT) sum(token string, h Hash) []byte {
mac := hmac.New(h.Hash, jwt.key) mac := hmac.New(h.Hash, jwt.key)
mac.Write([]byte(token)) mac.Write([]byte(token))
@ -248,8 +251,7 @@ func (jwt *JWT) Verify(t *Token) error {
nbfOK = true nbfOK = true
} }
if jwt.blacklist != nil { if jwt.blacklist != nil {
err := jwt.blacklisted(t.Sig()) if err := jwt.blacklisted(t.Sig()); err != nil {
if err != nil {
return err return err
} }
} }
@ -266,106 +268,12 @@ func (jwt *JWT) Verify(t *Token) error {
} }
// Stop will end the cleaner goroutinge execution // Stop will end the cleaner goroutinge execution
// Calling Stop twice or more will panic
func (jwt *JWT) Stop() error { func (jwt *JWT) Stop() error {
if jwt.blacklist == nil { if jwt.blacklist == nil {
return ErrBlacklistNotEnabled return ErrBlacklistNotEnabled
} }
jwt.stopOnce.Do(func() {
close(jwt.done) close(jwt.done)
})
return nil 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,
}
}

View file

@ -58,6 +58,11 @@ func TestValidate(t *testing.T) {
if err != nil { if err != nil {
t.Error(err) t.Error(err)
} }
// should not panic
err = jwt.Stop()
if err != nil {
t.Error(err)
}
} }
func TestNoBlacklist(t *testing.T) { func TestNoBlacklist(t *testing.T) {

104
token.go Normal file
View file

@ -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
}