feat: add backend jwt middleware

This commit is contained in:
Philipp
2026-06-10 16:23:55 +02:00
parent d87d8c6be7
commit 9947286b1e
12 changed files with 795 additions and 20 deletions
+350
View File
@@ -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"`
}