// Package jwt provides a easy to use JSON Web Token and blacklisting library package jwt import ( "crypto/hmac" "crypto/rand" "crypto/sha256" "crypto/sha512" "encoding/base64" "encoding/json" "errors" "hash" "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 ) // Hash represents the three diffrernt hash types type Hash int const ( HS256 Hash = iota //SHA256 HS384 //SHA384 HS512 //SHA512 unsupported ) // String returns the string representation of Hash func (h Hash) String() string { switch h { case HS256: return "HS256" case HS384: return "HS384" case HS512: return "HS512" } return "" } var ( ErrNoJWT = errors.New("no json web token") ErrUnsupportedAlg = errors.New("unsupported algoritm") 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("missíng token parts") ErrEmptySignature = errors.New("token signature is empty") ) // parseHash returns the hash.Hash type equal to the input string func parseHash(alg string) (h func() hash.Hash) { switch alg { case "HS256": h = sha256.New case "HS384": h = sha512.New384 case "HS512": h = sha512.New } return } // 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, error) { if expiry <= 0 { expiry = DefaultExpiry } key := make([]byte, keySize) _, err := rand.Read(key) if err != nil { return nil, err } 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, nil } func (jwt *JWT) sum(token string, h func() hash.Hash) []byte { mac := hmac.New(h, 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 } 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.Signature); err != nil { return err } jwt.Lock() defer jwt.Unlock() jwt.list[t.Signature] = 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() defer jwt.Unlock() for k, v := range jwt.list { if now > v { delete(jwt.list, 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) { now := time.Now().UTC() t.Claims["exp"] = now.Add(jwt.expiry).Unix() t.Claims["nbf"] = now.Unix() h, err := json.Marshal(t.Header) if err != nil { return } c, err := json.Marshal(t.Claims) if err != nil { return } var hf func() hash.Hash switch a := t.Header["alg"].(type) { case string: hf = parseHash(a) if hf == nil { return ErrUnsupportedAlg } default: return ErrUnsupportedAlg } t.Data = strings.Join([]string{base64.URLEncoding.EncodeToString(h), base64.URLEncoding.EncodeToString(c)}, ".") t.Raw = strings.Join([]string{t.Data, base64.URLEncoding.EncodeToString(jwt.sum(t.Data, hf))}, ".") return } // Verify will verify the provided token using the secret key func (jwt *JWT) Verify(t *Token) error { now := time.Now().UTC().Unix() var ( exp int64 nbf int64 expOK bool nbfOK bool hf func() hash.Hash ) switch t := t.Header["typ"].(type) { case string: if t != typ { return ErrNoJWT } default: return ErrNoJWT } switch a := t.Header["alg"].(type) { case string: hf = parseHash(a) if hf == 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.Signature) if err != nil { return err } } if expOK && now > exp { return ErrEXP } if nbfOK && now < nbf { return ErrNBF } if !hmac.Equal(jwt.sum(t.Data, hf), t.RawSignature) { 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) { parts := strings.Split(token, ".") if len(parts) < 3 { return nil, ErrMissingTokenParts } 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 Exp int64 Header map[string]interface{} Claims map[string]interface{} } // NewToken returns a new *Token using the provided hash algorithm and claims func NewToken(hash Hash, claims map[string]interface{}) *Token { return &Token{ Header: map[string]interface{}{ "alg": hash.String(), "typ": typ, }, Claims: claims, } }