// 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() }