jwt/jwt.go
2016-09-14 11:45:12 +02:00

346 lines
7.3 KiB
Go

// 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()
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) {
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,
}
}