351 lines
7.1 KiB
Go
351 lines
7.1 KiB
Go
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"`
|
|
}
|