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"` }