package flare import ( "crypto/ecdh" "crypto/ecdsa" "crypto/elliptic" "crypto/rand" "crypto/subtle" "encoding/base64" "errors" "fmt" "net/url" "strings" "time" "github.com/golang-jwt/jwt/v5" ) // VAPIDTokenLifetime is how far ahead of the request the exp claim of a // VAPID token lies. RFC 8292 allows at most 24 hours. const VAPIDTokenLifetime = 12 * time.Hour // Lengths of unpadded base64url VAPID keys: a 65-byte uncompressed P-256 // point and a 32-byte scalar. const ( PublicKeyLength = 87 PrivateKeyLength = 43 ) // ErrInvalidVAPIDKeys is returned when the configured key pair cannot be // parsed or does not match. It never carries key material. var ErrInvalidVAPIDKeys = errors.New("flare: invalid VAPID key pair") // ErrInvalidSubject is returned when the VAPID subject is not a mailto: or // https: URI. var ErrInvalidSubject = errors.New("flare: VAPID subject must start with mailto: or https:") // VAPIDKeys is an application server key pair (RFC 8292) as unpadded // base64url: PublicKey is the 65-byte uncompressed P-256 point, PrivateKey // the 32-byte private scalar. String and GoString never print PrivateKey. type VAPIDKeys struct { PublicKey string PrivateKey string } // String prints the public key and hides the private key. func (k VAPIDKeys) String() string { return "flare.VAPIDKeys{PublicKey: " + k.PublicKey + ", PrivateKey: [redacted]}" } // GoString is String, so %#v does not print the private key either. func (k VAPIDKeys) GoString() string { return k.String() } // GenerateVAPIDKeys returns a fresh P-256 key pair. func GenerateVAPIDKeys() (VAPIDKeys, error) { priv, err := ecdh.P256().GenerateKey(rand.Reader) if err != nil { return VAPIDKeys{}, fmt.Errorf("flare: generate VAPID keys: %w", err) } return VAPIDKeys{ PublicKey: base64.RawURLEncoding.EncodeToString(priv.PublicKey().Bytes()), PrivateKey: base64.RawURLEncoding.EncodeToString(priv.Bytes()), }, nil } // ParseVAPIDKeys decodes a key pair given as base64url, with or without // padding, and checks that the public key belongs to the private key. func ParseVAPIDKeys(public, private string) (*ecdh.PrivateKey, error) { privRaw, err := decodeBase64URL(private) if err != nil || len(privRaw) != 32 { return nil, fmt.Errorf("%w: the private key must be 32 bytes of base64url", ErrInvalidVAPIDKeys) } priv, err := ecdh.P256().NewPrivateKey(privRaw) if err != nil { return nil, fmt.Errorf("%w: the private key is not a P-256 scalar", ErrInvalidVAPIDKeys) } pubRaw, err := decodeBase64URL(public) if err != nil || len(pubRaw) != publicKeySize { return nil, fmt.Errorf("%w: the public key must be 65 bytes of base64url", ErrInvalidVAPIDKeys) } if subtle.ConstantTimeCompare(pubRaw, priv.PublicKey().Bytes()) != 1 { return nil, fmt.Errorf("%w: the public key does not match the private key", ErrInvalidVAPIDKeys) } return priv, nil } // VAPIDHeader returns the RFC 8292 Authorization header value for a push to // endpoint: "vapid t=, k=". The ES256 token carries aud // (the endpoint's origin), exp (now + VAPIDTokenLifetime) and sub. func VAPIDHeader(endpoint, subject string, keys VAPIDKeys, now time.Time) (string, error) { if !validSubject(subject) { return "", ErrInvalidSubject } aud, err := origin(endpoint) if err != nil { return "", err } priv, err := ParseVAPIDKeys(keys.PublicKey, keys.PrivateKey) if err != nil { return "", err } signer, err := ecdsa.ParseRawPrivateKey(elliptic.P256(), priv.Bytes()) if err != nil { return "", ErrInvalidVAPIDKeys } token := jwt.NewWithClaims(jwt.SigningMethodES256, jwt.MapClaims{ "aud": aud, "exp": now.Add(VAPIDTokenLifetime).Unix(), "sub": subject, }) signed, err := token.SignedString(signer) if err != nil { return "", fmt.Errorf("flare: sign VAPID token: %w", err) } k := base64.RawURLEncoding.EncodeToString(priv.PublicKey().Bytes()) return "vapid t=" + signed + ", k=" + k, nil } func validSubject(subject string) bool { return strings.HasPrefix(subject, "mailto:") || strings.HasPrefix(subject, "https:") } // origin is the RFC 6454 serialization of the endpoint's origin: scheme and // host, with the port only when it is not the scheme's default. func origin(endpoint string) (string, error) { u, err := url.Parse(endpoint) if err != nil || u.Scheme == "" || u.Host == "" { return "", fmt.Errorf("%w: not an absolute URL", ErrEndpointNotAllowed) } scheme := strings.ToLower(u.Scheme) host := strings.ToLower(u.Hostname()) if strings.Contains(host, ":") { host = "[" + host + "]" } port := u.Port() if port != "" && !(scheme == "https" && port == "443") && !(scheme == "http" && port == "80") { host += ":" + port } return scheme + "://" + host, nil }