367 lines
7.9 KiB
Go
367 lines
7.9 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"
|
|
"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,
|
|
}
|
|
}
|