split up library and removed Stop() panic when called twice
This commit is contained in:
parent
6e61945c2c
commit
826d035d3e
4 changed files with 117 additions and 100 deletions
|
|
@ -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
104
jwt.go
|
|
@ -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,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -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
104
token.go
Normal 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
|
||||||
|
}
|
||||||
Loading…
Add table
Add a link
Reference in a new issue