feat: add backend jwt middleware
This commit is contained in:
@@ -0,0 +1,350 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/hmac"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidToken = errors.New("invalid token")
|
||||
ErrExpiredToken = errors.New("expired token")
|
||||
)
|
||||
|
||||
type Validator struct {
|
||||
issuer string
|
||||
jwksURL string
|
||||
hmacKey []byte
|
||||
client *http.Client
|
||||
now func() time.Time
|
||||
|
||||
mu sync.Mutex
|
||||
cachedAt time.Time
|
||||
cacheTTL time.Duration
|
||||
cachedSet keySet
|
||||
}
|
||||
|
||||
type Option func(*Validator)
|
||||
|
||||
func WithHTTPClient(client *http.Client) Option {
|
||||
return func(v *Validator) {
|
||||
v.client = client
|
||||
}
|
||||
}
|
||||
|
||||
func WithNow(now func() time.Time) Option {
|
||||
return func(v *Validator) {
|
||||
v.now = now
|
||||
}
|
||||
}
|
||||
|
||||
func WithCacheTTL(ttl time.Duration) Option {
|
||||
return func(v *Validator) {
|
||||
v.cacheTTL = ttl
|
||||
}
|
||||
}
|
||||
|
||||
func WithHMACSecret(secret string) Option {
|
||||
return func(v *Validator) {
|
||||
secret = strings.TrimSpace(secret)
|
||||
if secret != "" {
|
||||
v.hmacKey = []byte(secret)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func NewValidator(issuer string, jwksURL string, opts ...Option) (*Validator, error) {
|
||||
issuer = strings.TrimSpace(issuer)
|
||||
jwksURL = strings.TrimSpace(jwksURL)
|
||||
if issuer == "" {
|
||||
return nil, fmt.Errorf("issuer is required")
|
||||
}
|
||||
|
||||
validator := &Validator{
|
||||
issuer: issuer,
|
||||
jwksURL: jwksURL,
|
||||
client: http.DefaultClient,
|
||||
now: time.Now,
|
||||
cacheTTL: 5 * time.Minute,
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(validator)
|
||||
}
|
||||
if validator.jwksURL == "" && len(validator.hmacKey) == 0 {
|
||||
return nil, fmt.Errorf("jwks url or hmac secret is required")
|
||||
}
|
||||
|
||||
return validator, nil
|
||||
}
|
||||
|
||||
func (v *Validator) Validate(ctx context.Context, token string) (Principal, error) {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != 3 {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
var header jwtHeader
|
||||
if err := decodeJSONPart(parts[0], &header); err != nil {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
if header.Alg == "" {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
var claims jwtClaims
|
||||
if err := decodeJSONPart(parts[1], &claims); err != nil {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(claims.Issuer), []byte(v.issuer)) != 1 {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
if claims.Subject == "" {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
now := v.now().Unix()
|
||||
if claims.ExpiresAt <= now {
|
||||
return Principal{}, ErrExpiredToken
|
||||
}
|
||||
if claims.NotBefore != 0 && claims.NotBefore > now {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
signingInput := parts[0] + "." + parts[1]
|
||||
signature, err := base64.RawURLEncoding.DecodeString(parts[2])
|
||||
if err != nil {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
if header.Alg == "HS256" {
|
||||
if err := v.verifyHMAC([]byte(signingInput), signature); err != nil {
|
||||
return Principal{}, err
|
||||
}
|
||||
return Principal{
|
||||
Subject: claims.Subject,
|
||||
Email: claims.Email,
|
||||
Role: claims.Role,
|
||||
}, nil
|
||||
}
|
||||
|
||||
if header.KID == "" {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
keys, err := v.keys(ctx)
|
||||
if err != nil {
|
||||
return Principal{}, err
|
||||
}
|
||||
|
||||
key, ok := keys[header.KID]
|
||||
if !ok {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
if key.Alg != "" && key.Alg != header.Alg {
|
||||
return Principal{}, ErrInvalidToken
|
||||
}
|
||||
if err := verifySignature(header.Alg, key, []byte(signingInput), signature); err != nil {
|
||||
return Principal{}, err
|
||||
}
|
||||
|
||||
return Principal{
|
||||
Subject: claims.Subject,
|
||||
Email: claims.Email,
|
||||
Role: claims.Role,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (v *Validator) verifyHMAC(signingInput []byte, signature []byte) error {
|
||||
if len(v.hmacKey) == 0 {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
|
||||
mac := hmac.New(sha256.New, v.hmacKey)
|
||||
_, _ = mac.Write(signingInput)
|
||||
if !hmac.Equal(signature, mac.Sum(nil)) {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *Validator) keys(ctx context.Context) (keySet, error) {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.cachedSet != nil && v.now().Sub(v.cachedAt) < v.cacheTTL {
|
||||
return v.cachedSet, nil
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, v.jwksURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp, err := v.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("fetch jwks: status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var raw jwksResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&raw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
keys := make(keySet, len(raw.Keys))
|
||||
for _, key := range raw.Keys {
|
||||
if key.KID != "" {
|
||||
keys[key.KID] = key
|
||||
}
|
||||
}
|
||||
if len(keys) == 0 {
|
||||
return nil, fmt.Errorf("jwks contains no usable keys")
|
||||
}
|
||||
|
||||
v.cachedSet = keys
|
||||
v.cachedAt = v.now()
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func verifySignature(alg string, key jwk, signingInput []byte, signature []byte) error {
|
||||
digest := sha256.Sum256(signingInput)
|
||||
|
||||
switch alg {
|
||||
case "RS256":
|
||||
publicKey, err := rsaPublicKey(key)
|
||||
if err != nil {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
if err := rsa.VerifyPKCS1v15(publicKey, crypto.SHA256, digest[:], signature); err != nil {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
return nil
|
||||
case "ES256":
|
||||
publicKey, err := ecdsaPublicKey(key)
|
||||
if err != nil {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
if len(signature) != 64 {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
r := new(big.Int).SetBytes(signature[:32])
|
||||
s := new(big.Int).SetBytes(signature[32:])
|
||||
if !ecdsa.Verify(publicKey, digest[:], r, s) {
|
||||
return ErrInvalidToken
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return ErrInvalidToken
|
||||
}
|
||||
}
|
||||
|
||||
func decodeJSONPart(part string, out any) error {
|
||||
data, err := base64.RawURLEncoding.DecodeString(part)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(data, out)
|
||||
}
|
||||
|
||||
func rsaPublicKey(key jwk) (*rsa.PublicKey, error) {
|
||||
if key.Kty != "RSA" {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
nBytes, err := base64.RawURLEncoding.DecodeString(key.N)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
eBytes, err := base64.RawURLEncoding.DecodeString(key.E)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
e := 0
|
||||
for _, b := range eBytes {
|
||||
e = e<<8 + int(b)
|
||||
}
|
||||
if e == 0 {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
|
||||
return &rsa.PublicKey{
|
||||
N: new(big.Int).SetBytes(nBytes),
|
||||
E: e,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ecdsaPublicKey(key jwk) (*ecdsa.PublicKey, error) {
|
||||
if key.Kty != "EC" || key.Crv != "P-256" {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
|
||||
xBytes, err := base64.RawURLEncoding.DecodeString(key.X)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
yBytes, err := base64.RawURLEncoding.DecodeString(key.Y)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
publicKey := &ecdsa.PublicKey{
|
||||
Curve: elliptic.P256(),
|
||||
X: new(big.Int).SetBytes(xBytes),
|
||||
Y: new(big.Int).SetBytes(yBytes),
|
||||
}
|
||||
if !publicKey.Curve.IsOnCurve(publicKey.X, publicKey.Y) {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
return publicKey, nil
|
||||
}
|
||||
|
||||
type keySet map[string]jwk
|
||||
|
||||
type jwksResponse struct {
|
||||
Keys []jwk `json:"keys"`
|
||||
}
|
||||
|
||||
type jwk struct {
|
||||
KID string `json:"kid"`
|
||||
Kty string `json:"kty"`
|
||||
Alg string `json:"alg"`
|
||||
N string `json:"n"`
|
||||
E string `json:"e"`
|
||||
Crv string `json:"crv"`
|
||||
X string `json:"x"`
|
||||
Y string `json:"y"`
|
||||
}
|
||||
|
||||
type jwtHeader struct {
|
||||
Alg string `json:"alg"`
|
||||
KID string `json:"kid"`
|
||||
Typ string `json:"typ"`
|
||||
}
|
||||
|
||||
type jwtClaims struct {
|
||||
Issuer string `json:"iss"`
|
||||
Subject string `json:"sub"`
|
||||
Email string `json:"email"`
|
||||
Role string `json:"role"`
|
||||
ExpiresAt int64 `json:"exp"`
|
||||
NotBefore int64 `json:"nbf"`
|
||||
IssuedAt int64 `json:"iat"`
|
||||
}
|
||||
Reference in New Issue
Block a user