package flare import ( "crypto/aes" "crypto/cipher" "crypto/ecdh" "crypto/hkdf" "crypto/rand" "crypto/sha256" "encoding/base64" "encoding/binary" "errors" "fmt" "strings" ) // ContentEncoding is the only content coding a push message may use // (RFC 8291 Section 4). const ContentEncoding = "aes128gcm" // MaxPayloadSize is the largest plaintext Encrypt accepts: a push service // need not accept more than 4096 octets of body, and the aes128gcm header // (86 octets), the padding delimiter (1) and the GCM tag (16) take the rest. const MaxPayloadSize = 3993 // recordSize is the "rs" field of the aes128gcm header. A push message is a // single record, so any value above the record length works; 4096 is the // value RFC 8291 uses. const recordSize = 4096 const ( saltSize = 16 authSize = 16 publicKeySize = 65 headerSize = saltSize + 4 + 1 + publicKeySize ) // Key derivation labels from RFC 8291 Section 3.4 (the RFC 8188 labels end // with a zero octet). const ( webPushInfo = "WebPush: info\x00" cekInfo = "Content-Encoding: aes128gcm\x00" nonceInfo = "Content-Encoding: nonce\x00" ) // ErrPayloadTooLarge is returned by Encrypt for a payload above // MaxPayloadSize. var ErrPayloadTooLarge = errors.New("flare: payload exceeds 3993 bytes") // Encrypt encrypts payload for the user agent that owns sub, following // RFC 8291: an ephemeral P-256 key agreement with sub.P256dh, mixed with the // sub.Auth secret, then one aes128gcm record (RFC 8188). The result is the // request body: salt || record size || key id length || the ephemeral // public key || ciphertext. func Encrypt(payload []byte, sub Subscription) ([]byte, error) { asKey, err := ecdh.P256().GenerateKey(rand.Reader) if err != nil { return nil, fmt.Errorf("flare: generate ephemeral key: %w", err) } salt := make([]byte, saltSize) if _, err := rand.Read(salt); err != nil { return nil, fmt.Errorf("flare: generate salt: %w", err) } return encrypt(payload, sub, asKey, salt) } // encrypt is Encrypt with the application server key and the salt given, // so the RFC 8291 Appendix A vector can fix both. func encrypt(payload []byte, sub Subscription, asKey *ecdh.PrivateKey, salt []byte) ([]byte, error) { if len(payload) > MaxPayloadSize { return nil, ErrPayloadTooLarge } if len(salt) != saltSize { return nil, fmt.Errorf("flare: salt must be %d bytes", saltSize) } uaRaw, err := decodeBase64URL(sub.P256dh) if err != nil { return nil, fmt.Errorf("flare: subscription p256dh is not base64url") } uaPublic, err := ecdh.P256().NewPublicKey(uaRaw) if err != nil { return nil, fmt.Errorf("flare: subscription p256dh is not an uncompressed P-256 point") } authSecret, err := decodeBase64URL(sub.Auth) if err != nil { return nil, fmt.Errorf("flare: subscription auth is not base64url") } if len(authSecret) != authSize { return nil, fmt.Errorf("flare: subscription auth must be %d bytes, got %d", authSize, len(authSecret)) } ecdhSecret, err := asKey.ECDH(uaPublic) if err != nil { return nil, fmt.Errorf("flare: key agreement: %w", err) } asPublic := asKey.PublicKey().Bytes() cek, nonce, err := deriveKeys(ecdhSecret, authSecret, salt, uaPublic.Bytes(), asPublic) if err != nil { return nil, err } block, err := aes.NewCipher(cek) if err != nil { return nil, fmt.Errorf("flare: %w", err) } gcm, err := cipher.NewGCM(block) if err != nil { return nil, fmt.Errorf("flare: %w", err) } // A single, final record: the plaintext followed by the 0x02 padding // delimiter and no padding. plain := make([]byte, 0, len(payload)+1) plain = append(plain, payload...) plain = append(plain, 0x02) out := make([]byte, headerSize, headerSize+len(plain)+gcm.Overhead()) copy(out, salt) binary.BigEndian.PutUint32(out[saltSize:], recordSize) out[saltSize+4] = byte(len(asPublic)) copy(out[saltSize+5:], asPublic) return gcm.Seal(out, nonce, plain, nil), nil } // deriveKeys runs the RFC 8291 Section 3.4 key schedule and returns the // content encryption key and the nonce. func deriveKeys(ecdhSecret, authSecret, salt, uaPublic, asPublic []byte) (cek, nonce []byte, err error) { prkKey, err := hkdf.Extract(sha256.New, ecdhSecret, authSecret) if err != nil { return nil, nil, fmt.Errorf("flare: %w", err) } keyInfo := make([]byte, 0, len(webPushInfo)+len(uaPublic)+len(asPublic)) keyInfo = append(keyInfo, webPushInfo...) keyInfo = append(keyInfo, uaPublic...) keyInfo = append(keyInfo, asPublic...) ikm, err := hkdf.Expand(sha256.New, prkKey, string(keyInfo), 32) if err != nil { return nil, nil, fmt.Errorf("flare: %w", err) } prk, err := hkdf.Extract(sha256.New, ikm, salt) if err != nil { return nil, nil, fmt.Errorf("flare: %w", err) } cek, err = hkdf.Expand(sha256.New, prk, cekInfo, 16) if err != nil { return nil, nil, fmt.Errorf("flare: %w", err) } nonce, err = hkdf.Expand(sha256.New, prk, nonceInfo, 12) if err != nil { return nil, nil, fmt.Errorf("flare: %w", err) } return cek, nonce, nil } // decodeBase64URL decodes base64url with or without padding. Browsers hand // out unpadded keys; stored copies sometimes carry padding or use the // standard alphabet, which is accepted too. func decodeBase64URL(s string) ([]byte, error) { s = strings.TrimRight(strings.TrimSpace(s), "=") if b, err := base64.RawURLEncoding.DecodeString(s); err == nil { return b, nil } return base64.RawStdEncoding.DecodeString(s) }