// 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" "io" "strings" "time" ) const ( // Header is the default HTTP Authorization header name Header = "Authorization" // DefaultExpiry is the default token expiration time DefaultExpiry = time.Hour * 12 // KeySize is the secret key size KeySize = 64 typ = "JWT" ) // DefaultSecretReader is the default secret key generator var DefaultSecretReader = rand.Reader 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") ErrInvalidKeySize = errors.New("invalid secret key size") ) // JWT represents the JSON Web Token signing and blacklisting infrastructure type JWT struct { key []byte expiry time.Duration blacklist Blacklist done 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 Blacklist, secret io.Reader) (*JWT, error) { if expiry <= 0 { expiry = DefaultExpiry } if secret == nil { secret = DefaultSecretReader } secret = io.LimitReader(secret, KeySize) key := make([]byte, KeySize) i, err := secret.Read(key) if err != nil { return nil, errors.New("secret reader error: " + err.Error()) } if i < KeySize { return nil, ErrInvalidKeySize } jwt := &JWT{ key: key, expiry: expiry, blacklist: blacklist, } if blacklist != nil { jwt.done = make(chan struct{}) go jwt.clean() } return jwt, nil } 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 == nil { 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.blacklist.Add(t.Sig(), exp) return nil } // blacklisted checks if a token is blacklisted func (jwt *JWT) blacklisted(sig string) error { if sig == "" { return ErrEmptySignature } if jwt.blacklist.Check(sig) { return ErrBlacklisted } return nil } // clean looks for expired blacklisted tokens and removes them func (jwt *JWT) clean() { for { t := time.NewTimer(time.Hour) select { case <-t.C: case <-jwt.done: t.Stop() return } now := time.Now().UTC().Unix() for k, v := range jwt.blacklist.Map() { if now > v { jwt.blacklist.Remove(k) } } } } // 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 jwt.blacklist != nil { 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 == nil { return ErrBlacklistNotEnabled } 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 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, } }