jwt/jwt.go
2018-09-13 16:02:17 +02:00

262 lines
6.4 KiB
Go

// 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"
"sync"
"time"
)
const (
// HTTPHeader is the default HTTP Authorization header name
HTTPHeader = "Authorization"
// DefaultExpiry is the default token expiration time
DefaultExpiry = time.Hour * 12
// KeySize is the secret key size
KeySize = 64
// Typ is the JWT type
Typ = "JWT"
// TypClaim is the typ claim name
TypClaim = "typ"
// AlgClaim is the alg claim name
AlgClaim = "alg"
// ExpClaim is the exp claim name
ExpClaim = "exp"
// NbfClaim is the nbf claim name
NbfClaim = "nbf"
// TokenSeparator is the tokens separator char
TokenSeparator = "."
)
// 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")
ErrMissingNbf = errors.New("missing nbf claim")
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
stopOnce sync.Once
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.
// The secret key size needs to be at least 64 bytes.
// If secret is nil the DefaultSecretReader 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
}
// Expiry returns the configured expiry
func (jwt *JWT) Expiry() time.Duration {
return jwt.expiry
}
// sum calculates the HMAC hash sum of the token
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
}
if err := jwt.blacklisted(t.Sig()); err != nil {
return err
}
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}
exp := t.Claims.GetInt64(ExpClaim)
switch {
case exp == 0:
return ErrMissingExp
case Now() > exp:
return ErrExp
}
if !hmac.Equal(jwt.sum(t.Data(), h), t.RawSig()) {
return ErrInvalid
}
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 := Now()
for k, v := range jwt.blacklist.Map() {
if now > v {
jwt.blacklist.Remove(k)
}
}
}
}
// Sign will sign the provided token using the secret key
// If the 'exp' and 'nbf' claims do not exist, they will be written to the default values in UNIX format:
// 'exp': Now + expiry specified at New() or DefaultExpiry (UTC)
// 'nbf': Now (UTC)
func (jwt *JWT) Sign(t *Token) (err error) {
if t == nil {
return ErrTokenIsNil
}
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}
now := time.Now()
if _, ok := t.Claims.Get(ExpClaim); !ok {
t.Claims.Set(ExpClaim, newExp(now, jwt.expiry))
}
if _, ok := t.Claims.Get(NbfClaim); !ok {
t.Claims.Set(NbfClaim, NewNbf(now))
}
head, err := json.Marshal(t.Header)
if err != nil {
return
}
claims, err := json.Marshal(t.Claims)
if err != nil {
return
}
t.data = base64.URLEncoding.EncodeToString(head) + TokenSeparator + base64.URLEncoding.EncodeToString(claims)
t.rawSignature = jwt.sum(t.data, h)
t.signature = base64.URLEncoding.EncodeToString(t.rawSignature)
t.raw = t.data + TokenSeparator + t.signature
return
}
// Verify will verify the provided token using the secret key
func (jwt *JWT) Verify(t *Token) error {
if t == nil {
return ErrTokenIsNil
}
if jwt.blacklist != nil {
if err := jwt.blacklisted(t.Sig()); err != nil {
return err
}
}
now := Now()
if t.Header.GetString(TypClaim) != Typ {
return ErrNoJWT
}
h := ParseHash(t.Header.GetString(AlgClaim))
if h == nil {
return ErrUnsupportedAlg
}
exp := t.Claims.GetInt64(ExpClaim)
switch {
case exp == 0:
return ErrMissingExp
case now > exp:
return ErrExp
}
nbf := t.Claims.GetInt64(NbfClaim)
switch {
case nbf == 0:
return ErrMissingNbf
case 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
func (jwt *JWT) Stop() error {
if jwt.blacklist == nil {
return ErrBlacklistNotEnabled
}
jwt.stopOnce.Do(func() {
close(jwt.done)
})
return nil
}