// Copyright (C) 2018 Marius Schellenberger // Package jwt provides a easy to use JSON Web Token and blacklisting library package jwt import ( "crypto/hmac" "crypto/rand" "encoding/base64" "encoding/json" "errors" "strings" "sync" "time" ) const ( // Header is the default HTTP Authorization header name Header = "Authorization" // DefaultExpiry is the default token expiration time DefaultExpiry = time.Hour * 12 typ = "JWT" keySize = 64 ) var ( ErrNoJWT = errors.New("not a json web token") ErrEmptyToken = errors.New("token is empty") ErrUnsupportedAlg = errors.New("unsupported algorithm") ErrInvalid = errors.New("token validation failed") ErrBlacklisted = errors.New("token blacklisted") ErrBlacklistNotEnabled = errors.New("blacklisting is not enabled") ErrNBF = errors.New("token not valid yet") ErrEXP = errors.New("token expired") ErrMissingEXP = errors.New("missing exp claim") ErrMissingTokenParts = errors.New("missing token parts") ErrEmptySignature = errors.New("token signature is empty") ErrTokenIsNil = errors.New("token is nil") ) // JWT represents the JSON Web Token signing and blacklisting infrastructure type JWT struct { key []byte expiry time.Duration // protects list sync.RWMutex blacklist bool list map[string]int64 stop chan struct{} } // New returns a new JWT object with the given expiry timeout. // If the timeout is less or equal to zero the default expiry (12 hours) is used. // If blacklisting is enabled, the JWT object leaks a goroutine to garbage-collect expired blacklisted tokens. // Call the Stop() method to exit the goroutine. func New(expiry time.Duration, blacklist bool) *JWT { if expiry <= 0 { expiry = DefaultExpiry } key := make([]byte, keySize) rand.Read(key) jwt := &JWT{ key: key, expiry: expiry, blacklist: blacklist, } if blacklist { jwt.list = make(map[string]int64) jwt.stop = make(chan struct{}) go jwt.clean() } return jwt } func (jwt *JWT) sum(token string, h Hash) []byte { mac := hmac.New(h.Hash, jwt.key) mac.Write([]byte(token)) return mac.Sum(nil) } // Invalidate checks if a token is already blacklisted // If the token is not blacklisted, it will get blacklisted func (jwt *JWT) Invalidate(t *Token) error { if !jwt.blacklist { return ErrBlacklistNotEnabled } 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 } jwt.Lock() defer jwt.Unlock() jwt.list[t.Sig()] = exp return nil } // blacklisted checks if a token is blacklisted func (jwt *JWT) blacklisted(sig string) error { if sig == "" { return ErrEmptySignature } jwt.RLock() defer jwt.RUnlock() _, ok := jwt.list[sig] if ok { return ErrBlacklisted } return nil } // clean will look for expired blacklisted tokens and removes them func (jwt *JWT) clean() { for { t := time.NewTimer(time.Hour) select { case <-t.C: case <-jwt.stop: t.Stop() return } now := time.Now().UTC().Unix() jwt.Lock() for k, v := range jwt.list { if now > v { delete(jwt.list, k) } } jwt.Unlock() } } // Sign will sign the provided token using the secret key // This will overwrite existing 'exp' and 'nbf' claims func (jwt *JWT) Sign(t *Token) (err error) { if t == nil { return ErrTokenIsNil } header := t.Header() switch t := header["typ"].(type) { case string: if t != typ { return ErrNoJWT } default: return ErrNoJWT } var h Hash switch a := header["alg"].(type) { case string: h = ParseHash(a) 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() head, err := json.Marshal(header) if err != nil { return } claims, err := json.Marshal(t.Claims) if err != nil { 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))}, ".") return } // Verify will verify the provided token using the secret key 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 expOK && jwt.blacklist { err := jwt.blacklisted(t.Sig()) if err != nil { return err } } if expOK && now > exp { return ErrEXP } if nbfOK && now < nbf { return ErrNBF } if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) { return ErrInvalid } return nil } // Stop will end the cleaner goroutinge execution // Calling Stop twice or more will panic func (jwt *JWT) Stop() error { if !jwt.blacklist { return ErrBlacklistNotEnabled } select { case jwt.stop <- struct{}{}: close(jwt.stop) default: } 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 map[string]interface{} } // 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 h is nil HS256 will be used func NewToken(claims map[string]interface{}, h Hash) *Token { if h == nil { h = NewHS256() } return &Token{ header: map[string]interface{}{ "alg": h.Alg(), "typ": typ, }, Claims: claims, } }