diff --git a/LICENSE b/LICENSE index 9e52d9b..74cee2b 100644 --- a/LICENSE +++ b/LICENSE @@ -1,4 +1,4 @@ -Copyright (C) 2016 Marius Schellenberger +Copyright (C) 2025 Marius Schellenberger All rights reserved. Redistribution and use in source and binary forms, with or without diff --git a/README.md b/README.md index ad89f76..74aa8a3 100644 --- a/README.md +++ b/README.md @@ -1 +1,61 @@ -# jwt - a simple JSON Web Token library in Go +# jwt - a simple JSON Web Token library written in Go + +## Usage + +``` +package main + +import ( + "fmt" + "git.giftfish.de/ston1th/jwt" + "log" + "time" +) + +func main() { + // create a new signing instance + JWT, err := jwt.New(time.Hour, jwt.NewMemBlacklist(), nil) + if err != nil { + log.Fatal(err) + } + + // create a new token + t := jwt.NewToken(jwt.Claims{ + "username": "admin", + }, nil) + + // sign the token + err = JWT.Sign(t) + if err != nil { + log.Fatal(err) + } + + // print the token string + token := t.String() + fmt.Println(token) + + // decode the token string + newToken, err := jwt.DecodeToken(token) + if err != nil { + log.Fatal(err) + } + + // verify the decoded token + err = JWT.Verify(newToken) + if err != nil { + log.Fatal(err) + } + + // add the token to the blacklist + err = JWT.Invalidate(newToken) + if err != nil { + log.Fatal(err) + } + + // try to verify the blacklisted token, it should fail + err = JWT.Verify(newToken) + if err != nil { + log.Fatal(err) + } +} +``` diff --git a/blacklist.go b/blacklist.go new file mode 100644 index 0000000..180c5fd --- /dev/null +++ b/blacklist.go @@ -0,0 +1,63 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import "sync" + +// Blacklist is the blacklisting storage interface +type Blacklist interface { + Add(string, int64) error + Remove(string) error + Check(string) bool + Map() (BlacklistMap, error) +} + +// BlacklistMap is the blacklist map structure +type BlacklistMap map[string]int64 + +// MemBlacklist implements the Blacklist interface +type MemBlacklist struct { + // protects list + sync.RWMutex + list BlacklistMap +} + +// NewMemBlacklist implements the Blacklist interface using an in-memory map +func NewMemBlacklist() *MemBlacklist { + return &MemBlacklist{list: make(BlacklistMap)} +} + +// Add adds a new token signature with expiration time to the blacklist +func (mb *MemBlacklist) Add(sig string, exp int64) error { + mb.Lock() + mb.list[sig] = exp + mb.Unlock() + return nil +} + +// Remove deletes a token signature from the blacklist +func (mb *MemBlacklist) Remove(sig string) error { + mb.Lock() + delete(mb.list, sig) + mb.Unlock() + return nil +} + +// Check returns true if a token signature is blacklisted and false otherwise +func (mb *MemBlacklist) Check(sig string) (ok bool) { + mb.RLock() + _, ok = mb.list[sig] + mb.RUnlock() + return +} + +// Map returns the blacklist in the form of a iterable map structure for cleanup +func (mb *MemBlacklist) Map() (list BlacklistMap, err error) { + list = make(BlacklistMap) + mb.RLock() + for k, v := range mb.list { + list[k] = v + } + mb.RUnlock() + return +} diff --git a/claims.go b/claims.go new file mode 100644 index 0000000..a5f4364 --- /dev/null +++ b/claims.go @@ -0,0 +1,78 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +// Claims is the claim type of the token +// The Claims map is not goroutine safe +type Claims map[string]interface{} + +// Get returns a value from the claims map +func (c Claims) Get(key string) (v interface{}, ok bool) { + v, ok = c[key] + return +} + +// GetString returns a string from the claims map +func (c Claims) GetString(key string) (s string) { + if v, ok := c.Get(key); ok && v != nil { + s, _ = v.(string) + } + return +} + +// GetBool returns a bool from the claims map +func (c Claims) GetBool(key string) (b bool) { + if v, ok := c.Get(key); ok && v != nil { + b, _ = v.(bool) + } + return +} + +// GetInt returns an int from the claims map +func (c Claims) GetInt(key string) (i int) { + i = int(c.GetInt64(key)) + return +} + +// GetInt64 returns an int64 from the claims map +func (c Claims) GetInt64(key string) (i int64) { + if v, ok := c.Get(key); ok && v != nil { + switch val := v.(type) { + case int64: + i = val + case float64: + i = int64(val) + } + } + return +} + +// getFloat64 returns a float64 and ok from the claims map +func (c Claims) GetFloat64(key string) (f float64) { + if v, ok := c.Get(key); ok && v != nil { + f, _ = v.(float64) + } + return +} + +// Set sets the value of key in the claims map, if not nil +func (c Claims) Set(key string, v interface{}) { + if c == nil { + return + } + c[key] = v +} + +// Delete removes the key from the claims map +func (c Claims) Delete(key string) { + delete(c, key) +} + +// Copy returns a new Claims map +func (c Claims) Copy() (n Claims) { + n = make(Claims) + for k, v := range c { + n[k] = v + } + return +} diff --git a/encoding.go b/encoding.go new file mode 100644 index 0000000..3693f5f --- /dev/null +++ b/encoding.go @@ -0,0 +1,7 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import "encoding/base64" + +var enc = base64.RawURLEncoding diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..6752b42 --- /dev/null +++ b/go.mod @@ -0,0 +1 @@ +module git.giftfish.de/ston1th/jwt/v3 diff --git a/hash.go b/hash.go new file mode 100644 index 0000000..e47c0cc --- /dev/null +++ b/hash.go @@ -0,0 +1,79 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import ( + "crypto/sha256" + "crypto/sha512" + "hash" +) + +const ( + HS256Name = "HS256" + HS384Name = "HS384" + HS512Name = "HS512" +) + +// NewHash returns the Hash type equal to the input string +func NewHash(alg string) Hash { + switch alg { + case HS256Name: + return hs256 + case HS384Name: + return hs384 + case HS512Name: + return hs512 + } + return nil +} + +var ( + hs256 = HS256{} + hs384 = HS384{} + hs512 = HS512{} +) + +// Hash is the hashsum interface for signing the jwt +type Hash interface { + Hash() hash.Hash + Alg() string +} + +// HS256 implements the Hash interface with SHA256 +type HS256 struct{} + +// Alg returns the algorithm name "HS256" +func (HS256) Alg() string { + return HS256Name +} + +// Hash returns a new hash.Hash instance of the SHA256 algorithm +func (HS256) Hash() hash.Hash { + return sha256.New() +} + +// HS384 implements the Hash interface with SHA384 +type HS384 struct{} + +// Alg returns the algorithm name "HS384" +func (HS384) Alg() string { + return HS384Name +} + +// Hash returns a new hash.Hash instance of the SHA384 algorithm +func (HS384) Hash() hash.Hash { + return sha512.New384() +} + +// HS512 implements the Hash interface with SHA512 +type HS512 struct{} + +// Alg returns the algorithm name "HS512" +func (HS512) Alg() string { + return HS512Name +} + +// Hash returns a new hash.Hash instance of the SHA512 algorithm +func (HS512) Hash() hash.Hash { + return sha512.New() +} diff --git a/jwt.go b/jwt.go index 03c545c..f9ce6af 100644 --- a/jwt.go +++ b/jwt.go @@ -1,118 +1,141 @@ +// Copyright (C) 2025 Marius Schellenberger + // 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" + "io" "sync" "time" ) const ( - // Header is the default HTTP Authorization header name - Header = "Authorization" + // HTTPHeader is the default HTTP Authorization header name + HTTPHeader = "Authorization" // DefaultExpiry is the default token expiration time DefaultExpiry = time.Hour * 12 - typ = "JWT" - keySize = 64 + // 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" + // NonceClaim + NonceClaim = "nonce" + // TokenSeparator is the tokens separator char + TokenSeparator = "." ) -// Hash represents the three diffrernt hash types -type Hash int +// DefaultSecretReader is the default secret key generator +var DefaultSecretReader = rand.Reader -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 "" -} +const jwtErr = "jwt: " 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") + ErrNoJWT = errors.New(jwtErr + "not a json web token") + ErrEmptyToken = errors.New(jwtErr + "token is empty") + ErrUnsupportedAlg = errors.New(jwtErr + "unsupported algorithm") + ErrInvalid = errors.New(jwtErr + "token validation failed") + ErrBlacklisted = errors.New(jwtErr + "token blacklisted") + ErrBlacklistNotEnabled = errors.New(jwtErr + "blacklisting is not enabled") + ErrNbf = errors.New(jwtErr + "token not valid yet") + ErrExp = errors.New(jwtErr + "token expired") + ErrMissingNbf = errors.New(jwtErr + "missing nbf claim") + ErrMissingExp = errors.New(jwtErr + "missing exp claim") + ErrMissingTokenParts = errors.New(jwtErr + "missing token parts") + ErrEmptySignature = errors.New(jwtErr + "token signature is empty") + ErrTokenIsNil = errors.New(jwtErr + "token is nil") + ErrInvalidKeySize = errors.New(jwtErr + "invalid secret key size") ) -// 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 +type JWTOption func(*JWT) + +func WithExpiry(expiry time.Duration) JWTOption { + return func(jwt *JWT) { + jwt.expiry = expiry + } +} + +func WithBlacklist(blacklist Blacklist) JWTOption { + return func(jwt *JWT) { + jwt.blacklist = blacklist + } +} +func WithSecret(secret io.Reader) JWTOption { + return func(jwt *JWT) { + jwt.secretReader = secret + } +} +func WithNonce() JWTOption { + return func(jwt *JWT) { + jwt.nonce = true } - 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{} + secretReader io.Reader + key []byte + expiry time.Duration + blacklist Blacklist + nonce bool + 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 bool) (*JWT, error) { - if expiry <= 0 { - expiry = DefaultExpiry +func New(options ...JWTOption) (*JWT, error) { + jwt := &JWT{} + for _, option := range options { + option(jwt) } - key := make([]byte, keySize) - _, err := rand.Read(key) + if jwt.expiry <= 0 { + jwt.expiry = DefaultExpiry + } + if jwt.secretReader == nil { + jwt.secretReader = DefaultSecretReader + } + secret := io.LimitReader(jwt.secretReader, KeySize) + key := make([]byte, KeySize) + i, err := secret.Read(key) if err != nil { - return nil, err + return nil, errors.New(jwtErr + "secret reader error: " + err.Error()) } - jwt := &JWT{ - key: key, - expiry: expiry, - blacklist: blacklist, + if i < KeySize { + return nil, ErrInvalidKeySize } - if blacklist { - jwt.list = make(map[string]int64) - jwt.stop = make(chan struct{}) + jwt.key = key + if jwt.blacklist != nil { + jwt.done = 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) +// 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) } @@ -120,33 +143,33 @@ func (jwt *JWT) sum(token string, h func() hash.Hash) []byte { // 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 { + if jwt.blacklist == nil { 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 t == nil { + return ErrTokenIsNil } - if !expOK { - return ErrMissingEXP - } - if err := jwt.blacklisted(t.Signature); err != nil { + if err := jwt.blacklisted(t.Sig()); err != nil { return err } - jwt.Lock() - defer jwt.Unlock() - jwt.list[t.Signature] = exp - return nil + if t.Header.GetString(TypClaim) != Typ { + return ErrNoJWT + } + h := NewHash(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 + } + return jwt.blacklist.Add(t.Sig(), exp) } // blacklisted checks if a token is blacklisted @@ -154,193 +177,120 @@ func (jwt *JWT) blacklisted(sig string) error { if sig == "" { return ErrEmptySignature } - jwt.RLock() - defer jwt.RUnlock() - _, ok := jwt.list[sig] - if ok { + if jwt.blacklist.Check(sig) { return ErrBlacklisted } return nil } -// clean will look for expired blacklisted tokens and removes them +// 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.stop: + case <-jwt.done: t.Stop() return } - now := time.Now().UTC().Unix() - jwt.Lock() - defer jwt.Unlock() - for k, v := range jwt.list { + now := Now() + m, err := jwt.blacklist.Map() + if err != nil { + continue + } + for k, v := range m { if now > v { - delete(jwt.list, k) + jwt.blacklist.Remove(k) } } } } // Sign will sign the provided token using the secret key -// This will overwrite existing 'exp' and 'nbf' claims +// 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) { - 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 + if t == nil { + return ErrTokenIsNil } - c, err := json.Marshal(t.Claims) - if err != nil { - return + if t.Header.GetString(TypClaim) != Typ { + return ErrNoJWT } - var hf func() hash.Hash - switch a := t.Header["alg"].(type) { - case string: - hf = parseHash(a) - if hf == nil { - return ErrUnsupportedAlg - } - default: + h := NewHash(t.Header.GetString(AlgClaim)) + if h == nil { 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))}, ".") + 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)) + } + if jwt.nonce { + t.Claims.Set(NonceClaim, enc.EncodeToString(key16())) + } + head, err := json.Marshal(t.Header) + if err != nil { + return + } + claims, err := json.Marshal(t.Claims) + if err != nil { + return + } + t.data = enc.EncodeToString(head) + TokenSeparator + enc.EncodeToString(claims) + t.rawSignature = jwt.sum(t.data, h) + t.signature = enc.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 { - 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 + if t == nil { + return ErrTokenIsNil } - 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 { + if jwt.blacklist != nil { + if err := jwt.blacklisted(t.Sig()); err != nil { return err } } - if expOK && now > exp { - return ErrEXP + now := Now() + if t.Header.GetString(TypClaim) != Typ { + return ErrNoJWT } - if nbfOK && now < nbf { - return ErrNBF + h := NewHash(t.Header.GetString(AlgClaim)) + if h == nil { + return ErrUnsupportedAlg } - if !hmac.Equal(jwt.sum(t.Data, hf), t.RawSignature) { + 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 -// Calling Stop twice or more will panic +// Stop will end the cleaner goroutine execution func (jwt *JWT) Stop() error { - if !jwt.blacklist { + if jwt.blacklist == nil { return ErrBlacklistNotEnabled } - select { - case jwt.stop <- struct{}{}: - close(jwt.stop) - default: - } + jwt.stopOnce.Do(func() { + close(jwt.done) + }) 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, - } -} diff --git a/jwt_test.go b/jwt_test.go index 0d04ad5..664cf8a 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -1,25 +1,250 @@ +// Copyright (C) 2025 Marius Schellenberger + package jwt import ( + "bytes" + "errors" "testing" "time" ) -func TestValidate(t *testing.T) { - jwt, _ := New(0, true) - token := NewToken(HS256, map[string]interface{}{ - "sub": "1234567890", - "name": "John Doe", - "admin": true, - "fizz": "buzz", - }) - err := jwt.Sign(token) - t.Log(err, token) - nt, err := DecodeToken(token.Raw) - t.Log(nt, err) - time.Sleep(time.Second) - t.Log(jwt.Verify(nt)) - jwt.Invalidate(nt) - t.Log(jwt.Verify(nt)) - jwt.Stop() +var claims = Claims{ + "sub": "1234567890", + "name": "John Doe", + "admin": true, + "fizz": "buzz", +} + +func TestValidate(t *testing.T) { + jwt, err := New( + WithExpiry(time.Second), + WithBlacklist(NewMemBlacklist()), + ) + if err != nil { + t.Error(err) + } + token := NewToken(claims.Copy(), nil) + err = jwt.Sign(token) + if err != nil { + t.Error(err) + } + _, err = DecodeToken("") + if err == nil { + t.Error(errors.New("token is empty")) + } + nt, err := DecodeToken(token.String()) + if err != nil { + t.Error(err) + } + err = jwt.Verify(nt) + if err != nil { + t.Error(err) + } + time.Sleep(time.Second * 2) + err = jwt.Verify(nt) + if err == nil { + t.Error(errors.New("token should be expired")) + } + + err = jwt.Invalidate(nt) + if err == nil { + t.Error(errors.New("token is expired")) + } + + nt.Claims.Delete(ExpClaim) + err = jwt.Sign(nt) + if err != nil { + t.Error(err) + } + err = jwt.Verify(nt) + if err != nil { + t.Error(err) + } + + err = jwt.Invalidate(nt) + if err != nil { + t.Error(err) + } + err = jwt.Invalidate(nt) + if err == nil { + t.Error(errors.New("double invalidate")) + } + + err = jwt.Verify(nt) + if err == nil { + t.Error(errors.New("token should be blacklisted")) + } + err = jwt.Stop() + if err != nil { + t.Error(err) + } + // should not panic + err = jwt.Stop() + if err != nil { + t.Error(err) + } +} + +func TestNonce(t *testing.T) { + jwt, err := New( + WithExpiry(time.Second), + WithNonce(), + ) + if err != nil { + t.Error(err) + } + token := NewToken(claims.Copy(), nil) + err = jwt.Sign(token) + if err != nil { + t.Error(err) + } + nt, err := DecodeToken(token.String()) + if err != nil { + t.Error(err) + } + err = jwt.Verify(nt) + if err != nil { + t.Error(err) + } + nonce := nt.Claims.GetString(NonceClaim) + if nonce == "" { + t.Error(errors.New("nonce is empty")) + } +} + +func TestNoBlacklist(t *testing.T) { + jwt, err := New(WithExpiry(time.Second)) + if err != nil { + t.Error(err) + } + token := NewToken(claims.Copy(), nil) + err = jwt.Sign(token) + if err != nil { + t.Error(err) + } + _, err = DecodeToken("") + if err == nil { + t.Error(errors.New("token is empty")) + } + nt, err := DecodeToken(token.String()) + if err != nil { + t.Error(err) + } + err = jwt.Verify(nt) + if err != nil { + t.Error(err) + } + time.Sleep(time.Second * 2) + err = jwt.Verify(nt) + if err == nil { + t.Error(errors.New("token should be expired")) + } + + err = jwt.Invalidate(nt) + if err == nil { + t.Error(errors.New("blacklisting should be disabled")) + } +} + +func TestEmptySecretReader(t *testing.T) { + _, err := New( + WithExpiry(time.Second), + WithSecret(new(bytes.Buffer)), + ) + if err == nil { + t.Error(errors.New("error should be secret reader error")) + } +} + +func TestInvalidSecretReader(t *testing.T) { + _, err := New( + WithExpiry(time.Second), + WithSecret(bytes.NewBufferString("123")), + ) + if err != ErrInvalidKeySize { + t.Error(errors.New("error should be invalid key size")) + } +} + +func BenchmarkDecodeToken(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), nil) + jwt.Sign(t) + token := t.String() + b.ReportAllocs() + for i := 0; i < b.N; i++ { + DecodeToken(token) + } +} + +func BenchmarkNew(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + New() + } +} + +func BenchmarkNewWithBlacklist(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + New(WithBlacklist(NewMemBlacklist())) + } +} + +func BenchmarkSignHS256(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), nil) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkSignHS384(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), NewHash(HS384Name)) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkSignHS512(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), NewHash(HS512Name)) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Sign(t) + } +} + +func BenchmarkVerifyHS256(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), nil) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } +} + +func BenchmarkVerifyHS384(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), NewHash(HS384Name)) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } +} + +func BenchmarkVerifyHS512(b *testing.B) { + jwt, _ := New() + t := NewToken(claims.Copy(), NewHash(HS512Name)) + jwt.Sign(t) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + jwt.Verify(t) + } } diff --git a/rand.go b/rand.go new file mode 100644 index 0000000..bd1b191 --- /dev/null +++ b/rand.go @@ -0,0 +1,11 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import "crypto/rand" + +func key16() []byte { + b := make([]byte, 16) + rand.Read(b) + return b +} diff --git a/time.go b/time.go new file mode 100644 index 0000000..77e151f --- /dev/null +++ b/time.go @@ -0,0 +1,24 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import "time" + +// Now returns the current time in UTC Unix format +func Now() int64 { + return NewNbf(time.Now()) +} + +// NewNbf returns a new 'not before' date +func NewNbf(t time.Time) int64 { + return t.UTC().Unix() +} + +// NewExp returns a new expiration date +func NewExp(d time.Duration) int64 { + return newExp(time.Now(), d) +} + +func newExp(t time.Time, d time.Duration) int64 { + return t.UTC().Add(d).Unix() +} diff --git a/token.go b/token.go new file mode 100644 index 0000000..772415c --- /dev/null +++ b/token.go @@ -0,0 +1,98 @@ +// Copyright (C) 2025 Marius Schellenberger + +package jwt + +import ( + "encoding/json" + "strings" +) + +// Token represents a JWT token +type Token struct { + raw string + data string + signature string + rawSignature []byte + Header Claims + Claims Claims +} + +// 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 +} + +// Data returns the first two URLEncoded token fields +func (t *Token) Data() string { + return t.data +} + +// NewToken returns a new *Token using the provided hash algorithm and claims +// If claims is nil, an empty map is used +// If hash is nil, HS256 is used +func NewToken(claims Claims, hash Hash) *Token { + if claims == nil { + claims = make(Claims) + } + if hash == nil { + hash = NewHash(HS256Name) + } + return &Token{ + Header: Claims{ + AlgClaim: hash.Alg(), + TypClaim: Typ, + }, + Claims: claims, + } +} + +// 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, TokenSeparator) + 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 := enc.DecodeString(parts[0]) + if err != nil { + return + } + err = json.Unmarshal(header, &t.Header) + if err != nil { + return + } + claims, err := enc.DecodeString(parts[1]) + if err != nil { + return + } + err = json.Unmarshal(claims, &t.Claims) + if err != nil { + return + } + t.rawSignature, err = enc.DecodeString(parts[2]) + t.signature = parts[2] + t.data = parts[0] + TokenSeparator + parts[1] + t.raw = token + return +}