feat(11-04): add flare Web Push with a stdlib VAPID driver
- 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
This commit is contained in:
423
modules/flare/flare.go
Normal file
423
modules/flare/flare.go
Normal file
@@ -0,0 +1,423 @@
|
||||
// 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()
|
||||
}
|
||||
Reference in New Issue
Block a user