- RFC 8291 aes128gcm encryption from crypto/ecdh, crypto/hkdf and AES-GCM, matching the RFC 8291 Appendix A vector byte for byte - RFC 8292 vapid t=<ES256 JWT>, k=<key> header (aud origin, exp +12h, sub) - Pusher, Subscription, SendOptions, SubscriptionSource, Service and From reading push.* (enabled, keys, subject, ttl, allowed_hosts) - sends only to https endpoints on push.allowed_hosts, checked before dialing, and never follows redirects; 404/410 map to ErrSubscriptionGone - module README and root modules row
424 lines
12 KiB
Go
424 lines
12 KiB
Go
// Package flare delivers Web Push notifications: a small Pusher interface
|
|
// and a VAPID driver that encrypts every payload with aes128gcm (RFC 8291)
|
|
// and signs every request with an ES256 VAPID token (RFC 8292). Push is a
|
|
// separate channel from realtime; the application owns the subscription
|
|
// store and hands subscriptions to flare through a SubscriptionSource.
|
|
package flare
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"git.golem15.com/golem15/summercms/modules/backpack"
|
|
"git.golem15.com/golem15/summercms/modules/compass"
|
|
)
|
|
|
|
// DefaultTTL is how long a push service keeps an undelivered message when
|
|
// push.ttl is not set: four weeks.
|
|
const DefaultTTL = 2419200 * time.Second
|
|
|
|
// DefaultTimeout bounds every request to a push service.
|
|
const DefaultTimeout = 10 * time.Second
|
|
|
|
// Errors returned by Pusher.Send and SubscriptionSource.
|
|
var (
|
|
// ErrPushDisabled is returned while push.enabled is false.
|
|
ErrPushDisabled = errors.New("flare: push is disabled (push.enabled is false)")
|
|
// ErrEndpointNotAllowed is returned, before any connection is made,
|
|
// for an endpoint that is not https or whose host is not in
|
|
// push.allowed_hosts.
|
|
ErrEndpointNotAllowed = errors.New("flare: push endpoint not allowed")
|
|
// ErrSubscriptionGone is returned when the push service answers 404 or
|
|
// 410: the subscription expired and should be deleted.
|
|
ErrSubscriptionGone = errors.New("flare: push subscription is gone")
|
|
// ErrUserNotFound is returned by a SubscriptionSource for an unknown
|
|
// user.
|
|
ErrUserNotFound = errors.New("flare: user not found")
|
|
)
|
|
|
|
// StatusError is a push service answer other than 2xx, 404 or 410. It
|
|
// carries the status code, never the response body.
|
|
type StatusError struct {
|
|
Code int
|
|
}
|
|
|
|
func (e *StatusError) Error() string {
|
|
return "flare: push service answered HTTP " + strconv.Itoa(e.Code)
|
|
}
|
|
|
|
// Subscription is a browser push subscription as PushSubscription.toJSON
|
|
// returns it: the endpoint URL and the user agent's P-256 public key and
|
|
// 16-byte authentication secret, both base64url.
|
|
type Subscription struct {
|
|
Endpoint string `json:"endpoint"`
|
|
P256dh string `json:"p256dh"`
|
|
Auth string `json:"auth"`
|
|
}
|
|
|
|
// SendOptions tune one push. A zero TTL uses push.ttl; Urgency ("very-low",
|
|
// "low", "normal" or "high") and Topic are sent only when set.
|
|
type SendOptions struct {
|
|
TTL time.Duration
|
|
Urgency string
|
|
Topic string
|
|
}
|
|
|
|
// Pusher sends one encrypted push message to one subscription.
|
|
type Pusher interface {
|
|
Send(ctx context.Context, sub Subscription, payload []byte, opts SendOptions) error
|
|
}
|
|
|
|
// SubscriptionInfo is a stored subscription with the details the
|
|
// websockets:test-push command reports.
|
|
type SubscriptionInfo struct {
|
|
Subscription
|
|
ID uint
|
|
UserAgent string
|
|
SubscribedAt *time.Time
|
|
LastUsedAt *time.Time
|
|
}
|
|
|
|
// SubscriptionSource is implemented by the application that stores push
|
|
// subscriptions, and published on the app with backpack.App.Publish as a
|
|
// flare.SubscriptionSource. Subscriptions returns ErrUserNotFound for an
|
|
// unknown user and an empty list for a user without subscriptions.
|
|
type SubscriptionSource interface {
|
|
Subscriptions(ctx context.Context, userID uint) ([]SubscriptionInfo, error)
|
|
}
|
|
|
|
// Config is the push.* configuration. String and GoString never print
|
|
// PrivateKey.
|
|
type Config struct {
|
|
// Enabled is push.enabled; nothing is sent while it is false.
|
|
Enabled bool
|
|
// PublicKey and PrivateKey are the VAPID key pair, base64url.
|
|
PublicKey string
|
|
PrivateKey string
|
|
// Subject is the VAPID sub claim, a mailto: or https: URI.
|
|
Subject string
|
|
// TTL is the default TTL header.
|
|
TTL time.Duration
|
|
// AllowedHosts are the push service hosts an endpoint may point at.
|
|
// "*.example.com" matches any subdomain of example.com.
|
|
AllowedHosts []string
|
|
}
|
|
|
|
// String prints the configuration without the private key.
|
|
func (c Config) String() string {
|
|
priv := "unset"
|
|
if c.PrivateKey != "" {
|
|
priv = "[redacted]"
|
|
}
|
|
return fmt.Sprintf("flare.Config{Enabled: %t, PublicKey: %s, PrivateKey: %s, Subject: %s, TTL: %s, AllowedHosts: %v}",
|
|
c.Enabled, c.PublicKey, priv, c.Subject, c.TTL, c.AllowedHosts)
|
|
}
|
|
|
|
// GoString is String.
|
|
func (c Config) GoString() string { return c.String() }
|
|
|
|
// LogValue keeps the private key out of structured logs.
|
|
func (c Config) LogValue() slog.Value { return slog.StringValue(c.String()) }
|
|
|
|
// Keys returns the configured VAPID key pair.
|
|
func (c Config) Keys() VAPIDKeys {
|
|
return VAPIDKeys{PublicKey: c.PublicKey, PrivateKey: c.PrivateKey}
|
|
}
|
|
|
|
// DefaultAllowedHosts are the push services endpoints may point at when
|
|
// push.allowed_hosts is not set: Firebase Cloud Messaging (Chrome, Edge on
|
|
// Android), Mozilla autopush (Firefox), Apple and Windows push.
|
|
func DefaultAllowedHosts() []string {
|
|
return []string{
|
|
"fcm.googleapis.com",
|
|
"updates.push.services.mozilla.com",
|
|
"*.push.apple.com",
|
|
"*.notify.windows.com",
|
|
}
|
|
}
|
|
|
|
// LoadConfig reads push.* from c, filling the defaults. push.ttl is an
|
|
// integer number of seconds or a duration string; push.allowed_hosts is a
|
|
// list or a comma-separated string.
|
|
func LoadConfig(c *compass.Config) Config {
|
|
cfg := Config{TTL: DefaultTTL, AllowedHosts: DefaultAllowedHosts()}
|
|
if c == nil {
|
|
return cfg
|
|
}
|
|
cfg.Enabled = c.Bool("push.enabled")
|
|
cfg.PublicKey = strings.TrimSpace(c.String("push.public_key"))
|
|
cfg.PrivateKey = strings.TrimSpace(c.String("push.private_key"))
|
|
cfg.Subject = strings.TrimSpace(c.String("push.subject"))
|
|
if d := durationSetting(c, "push.ttl"); d > 0 {
|
|
cfg.TTL = d
|
|
}
|
|
if v, ok := c.Lookup("push.allowed_hosts"); ok {
|
|
if hosts := stringList(v); len(hosts) > 0 {
|
|
cfg.AllowedHosts = hosts
|
|
}
|
|
}
|
|
return cfg
|
|
}
|
|
|
|
func durationSetting(c *compass.Config, key string) time.Duration {
|
|
raw := strings.TrimSpace(c.String(key))
|
|
if raw == "" {
|
|
return 0
|
|
}
|
|
if n, err := strconv.ParseInt(raw, 10, 64); err == nil {
|
|
return time.Duration(n) * time.Second
|
|
}
|
|
if d, err := time.ParseDuration(raw); err == nil {
|
|
return d
|
|
}
|
|
return 0
|
|
}
|
|
|
|
func stringList(v any) []string {
|
|
var parts []string
|
|
switch t := v.(type) {
|
|
case string:
|
|
parts = strings.Split(t, ",")
|
|
case []string:
|
|
parts = t
|
|
case []any:
|
|
for _, e := range t {
|
|
if s, ok := e.(string); ok {
|
|
parts = append(parts, s)
|
|
}
|
|
}
|
|
}
|
|
out := make([]string, 0, len(parts))
|
|
for _, p := range parts {
|
|
if p = strings.ToLower(strings.TrimSpace(p)); p != "" {
|
|
out = append(out, p)
|
|
}
|
|
}
|
|
return out
|
|
}
|
|
|
|
// Service is the app-scoped push service. Get it with From.
|
|
type Service struct {
|
|
cfg Config
|
|
log *slog.Logger
|
|
|
|
mu sync.RWMutex
|
|
pusher *VAPIDPusher
|
|
}
|
|
|
|
// From returns the app's Service, building it from push.* and publishing it
|
|
// on first use.
|
|
func From(app *backpack.App) (*Service, error) {
|
|
if app == nil {
|
|
return nil, fmt.Errorf("flare: app is nil")
|
|
}
|
|
if s, ok := app.Lookup[*Service](); ok && s != nil {
|
|
return s, nil
|
|
}
|
|
cfg := LoadConfig(app.Config)
|
|
svc := &Service{cfg: cfg, log: loggerFromApp(app), pusher: NewVAPIDPusher(cfg, nil)}
|
|
if err := app.Publish(svc); err != nil {
|
|
if existing, ok := app.Lookup[*Service](); ok && existing != nil {
|
|
return existing, nil
|
|
}
|
|
return nil, fmt.Errorf("flare: %w", err)
|
|
}
|
|
return svc, nil
|
|
}
|
|
|
|
// Config returns the push configuration the service was built with.
|
|
func (s *Service) Config() Config {
|
|
if s == nil {
|
|
return LoadConfig(nil)
|
|
}
|
|
return s.cfg
|
|
}
|
|
|
|
// Enabled reports push.enabled.
|
|
func (s *Service) Enabled() bool { return s != nil && s.cfg.Enabled }
|
|
|
|
// Pusher returns the VAPID driver.
|
|
func (s *Service) Pusher() Pusher {
|
|
if s == nil {
|
|
return NewVAPIDPusher(LoadConfig(nil), nil)
|
|
}
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
return s.pusher
|
|
}
|
|
|
|
// SetHTTPClient replaces the HTTP client of the VAPID driver, for example
|
|
// with an httptest TLS client in tests. Redirects stay refused.
|
|
func (s *Service) SetHTTPClient(hc *http.Client) {
|
|
if s == nil {
|
|
return
|
|
}
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
s.pusher = NewVAPIDPusher(s.cfg, hc)
|
|
}
|
|
|
|
// Logger returns the app logger.
|
|
func (s *Service) Logger() *slog.Logger {
|
|
if s == nil || s.log == nil {
|
|
return slog.Default()
|
|
}
|
|
return s.log
|
|
}
|
|
|
|
// VAPIDPusher is the Pusher that talks to push services directly with
|
|
// VAPID authentication. It is safe for concurrent use.
|
|
type VAPIDPusher struct {
|
|
cfg Config
|
|
hc *http.Client
|
|
now func() time.Time
|
|
}
|
|
|
|
var _ Pusher = (*VAPIDPusher)(nil)
|
|
|
|
// NewVAPIDPusher returns a driver for cfg. hc may be nil for a client with
|
|
// a DefaultTimeout timeout; a given client is copied. Either way the driver
|
|
// never follows redirects, so a push service cannot send it to a host
|
|
// outside push.allowed_hosts.
|
|
func NewVAPIDPusher(cfg Config, hc *http.Client) *VAPIDPusher {
|
|
var client http.Client
|
|
if hc != nil {
|
|
client = *hc
|
|
}
|
|
if client.Timeout == 0 {
|
|
client.Timeout = DefaultTimeout
|
|
}
|
|
client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }
|
|
return &VAPIDPusher{cfg: cfg, hc: &client, now: time.Now}
|
|
}
|
|
|
|
// Send encrypts payload for sub and POSTs it to sub.Endpoint. It returns
|
|
// ErrPushDisabled while push is disabled and ErrEndpointNotAllowed, before
|
|
// dialing, for an endpoint outside the allowlist. A 2xx answer is success,
|
|
// 404 and 410 are ErrSubscriptionGone, any other status a *StatusError.
|
|
func (p *VAPIDPusher) Send(ctx context.Context, sub Subscription, payload []byte, opts SendOptions) error {
|
|
if p == nil || !p.cfg.Enabled {
|
|
return ErrPushDisabled
|
|
}
|
|
if err := p.checkEndpoint(sub.Endpoint); err != nil {
|
|
return err
|
|
}
|
|
body, err := Encrypt(payload, sub)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
auth, err := VAPIDHeader(sub.Endpoint, p.cfg.Subject, p.cfg.Keys(), p.now())
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if ctx == nil {
|
|
ctx = context.Background()
|
|
}
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, sub.Endpoint, bytes.NewReader(body))
|
|
if err != nil {
|
|
return fmt.Errorf("flare: build request: %w", err)
|
|
}
|
|
ttl := opts.TTL
|
|
if ttl <= 0 {
|
|
ttl = p.cfg.TTL
|
|
}
|
|
req.Header.Set("TTL", strconv.FormatInt(int64(ttl/time.Second), 10))
|
|
req.Header.Set("Content-Encoding", ContentEncoding)
|
|
req.Header.Set("Content-Type", "application/octet-stream")
|
|
req.Header.Set("Authorization", auth)
|
|
if opts.Urgency != "" {
|
|
req.Header.Set("Urgency", opts.Urgency)
|
|
}
|
|
if opts.Topic != "" {
|
|
req.Header.Set("Topic", opts.Topic)
|
|
}
|
|
resp, err := p.hc.Do(req)
|
|
if err != nil {
|
|
// *url.Error repeats the full endpoint; keep only the cause.
|
|
var uerr *url.Error
|
|
if errors.As(err, &uerr) && uerr.Err != nil {
|
|
err = uerr.Err
|
|
}
|
|
return fmt.Errorf("flare: push to %s: %w", hostOf(sub.Endpoint), err)
|
|
}
|
|
defer resp.Body.Close()
|
|
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64<<10))
|
|
switch {
|
|
case resp.StatusCode >= 200 && resp.StatusCode <= 299:
|
|
return nil
|
|
case resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusGone:
|
|
return ErrSubscriptionGone
|
|
default:
|
|
return &StatusError{Code: resp.StatusCode}
|
|
}
|
|
}
|
|
|
|
// checkEndpoint allows only https URLs, without user info, whose host is in
|
|
// push.allowed_hosts. The error names the host, never the full endpoint,
|
|
// whose path is a capability.
|
|
func (p *VAPIDPusher) checkEndpoint(endpoint string) error {
|
|
u, err := url.Parse(endpoint)
|
|
if err != nil || u.Host == "" {
|
|
return fmt.Errorf("%w: not an absolute URL", ErrEndpointNotAllowed)
|
|
}
|
|
if !strings.EqualFold(u.Scheme, "https") {
|
|
return fmt.Errorf("%w: %s is not https", ErrEndpointNotAllowed, strings.ToLower(u.Scheme))
|
|
}
|
|
if u.User != nil {
|
|
return fmt.Errorf("%w: user info in URL", ErrEndpointNotAllowed)
|
|
}
|
|
host := strings.ToLower(u.Hostname())
|
|
if !HostAllowed(host, p.cfg.AllowedHosts) {
|
|
return fmt.Errorf("%w: host %s is not in push.allowed_hosts", ErrEndpointNotAllowed, host)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// HostAllowed reports whether host matches an entry of allowed. An entry
|
|
// "*.example.com" matches any subdomain of example.com but not example.com
|
|
// itself; any other entry must equal host. Matching is case-insensitive.
|
|
func HostAllowed(host string, allowed []string) bool {
|
|
host = strings.TrimSuffix(strings.ToLower(host), ".")
|
|
if host == "" {
|
|
return false
|
|
}
|
|
for _, entry := range allowed {
|
|
entry = strings.ToLower(strings.TrimSpace(entry))
|
|
if suffix, ok := strings.CutPrefix(entry, "*."); ok {
|
|
if suffix != "" && strings.HasSuffix(host, "."+suffix) {
|
|
return true
|
|
}
|
|
continue
|
|
}
|
|
if entry != "" && host == entry {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func hostOf(endpoint string) string {
|
|
if u, err := url.Parse(endpoint); err == nil {
|
|
return u.Host
|
|
}
|
|
return "endpoint"
|
|
}
|
|
|
|
func loggerFromApp(app *backpack.App) *slog.Logger {
|
|
if app != nil {
|
|
if log, ok := app.Lookup[*slog.Logger](); ok && log != nil {
|
|
return log
|
|
}
|
|
}
|
|
return slog.Default()
|
|
}
|