From 5a6c4277ad313f39560e012e2bce9f6ba7fad1a6 Mon Sep 17 00:00:00 2001 From: binwiederhier Date: Sat, 1 Aug 2026 16:10:31 +0200 Subject: [PATCH] WIP: Clsuter support (cross node delivery, leader election) --- cluster/cluster.go | 139 ++++++++ cluster/cluster_load_test.go | 133 +++++++ cluster/cluster_mesh.go | 450 +++++++++++++++++++++++ cluster/cluster_mesh_test.go | 572 ++++++++++++++++++++++++++++++ cluster/cluster_nop.go | 24 ++ cluster/cluster_test.go | 83 +++++ cluster/registry/registry.go | 149 ++++++++ cluster/registry/registry_test.go | 222 ++++++++++++ cluster/types.go | 47 +++ cluster/util.go | 69 ++++ cmd/serve.go | 40 +++ cmd/serve_test.go | 35 ++ db/pg/leader.go | 83 +++++ db/pg/leader_test.go | 84 +++++ db/pg/pg.go | 1 + db/schema/schema_test.go | 4 +- db/test/test.go | 4 +- message/cache.go | 36 ++ message/cache_postgres.go | 8 + message/cache_test.go | 38 ++ metrics/metrics.go | 36 ++ metrics/metrics_test.go | 9 + server/config.go | 16 +- server/config_test.go | 1 + server/server.go | 232 +++++++++--- server/server.yml | 28 ++ server/server_account.go | 3 +- server/server_cluster_test.go | 318 +++++++++++++++++ server/server_manager.go | 14 +- server/topic.go | 15 +- util/bloom.go | 103 ++++++ util/bloom_test.go | 61 ++++ util/linger_queue.go | 144 ++++++++ util/linger_queue_test.go | 90 +++++ webpush/store_test.go | 4 +- 35 files changed, 3232 insertions(+), 63 deletions(-) create mode 100644 cluster/cluster.go create mode 100644 cluster/cluster_load_test.go create mode 100644 cluster/cluster_mesh.go create mode 100644 cluster/cluster_mesh_test.go create mode 100644 cluster/cluster_nop.go create mode 100644 cluster/cluster_test.go create mode 100644 cluster/registry/registry.go create mode 100644 cluster/registry/registry_test.go create mode 100644 cluster/types.go create mode 100644 cluster/util.go create mode 100644 db/pg/leader.go create mode 100644 db/pg/leader_test.go create mode 100644 server/server_cluster_test.go create mode 100644 util/bloom.go create mode 100644 util/bloom_test.go create mode 100644 util/linger_queue.go create mode 100644 util/linger_queue_test.go diff --git a/cluster/cluster.go b/cluster/cluster.go new file mode 100644 index 00000000..8c6dd27d --- /dev/null +++ b/cluster/cluster.go @@ -0,0 +1,139 @@ +// Package cluster implements cross-node message delivery for a multi-node ntfy cluster. Nodes +// register themselves in a PostgreSQL node registry (control plane) and fan published messages +// out to each other directly over HTTP (data plane); PostgreSQL is never on the message path. +// The single-node default is the nop cluster, which does nothing. +package cluster + +import ( + "errors" + "net/http" + "time" + + "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/model" +) + +// The internal peer API: every kind of node-to-node communication is a path under +// /v1/internal/, served only on the dedicated cluster listener. Future concerns (rate limit +// counters, stats) become new paths or new sections of the state envelope. +const ( + // MessagePath receives batches of published messages (NDJSON, one apiMessage per line). + MessagePath = "/v1/internal/message" + // StatePath receives peer state (JSON apiState): full subscription snapshots and + // incremental updates. + StatePath = "/v1/internal/state" +) + +// NodeID identifies a cluster node; it keys the registry, the per-peer queues, and the peer +// state table. +// +// Naming convention: a "node" is any cluster member in the absolute sense (identity, registry, +// config); a "peer" is another node as seen from this one (Peers, peerQueue, peerState). A peer +// IS a node, which is why peer values carry a NodeID. +type NodeID string + +const ( + // secretHeader carries the shared secret authenticating node-to-node fan-out requests. + secretHeader = "X-Cluster-Secret" + + // originHeader carries the sending node's ID on fan-out requests, so a node can skip + // requests that carry its own broadcasts (loop prevention). + originHeader = "X-Cluster-Origin" +) + +// Content types of the peer API: message bodies are NDJSON (one JSON message per line, matching +// the framing of ntfy's own /topic/json subscribe stream), state bodies are plain JSON. Future +// node-to-node request types get their own paths on the cluster listener; an old node answering +// 404 on an unknown path keeps mixed-version clusters working during rolling deploys. +const ( + contentTypeNDJSON = "application/x-ndjson" + contentTypeJSON = "application/json" +) + +const ( + defaultHeartbeatInterval = 3 * time.Second // How often a node refreshes its registry heartbeat + defaultNodeTTL = 30 * time.Second // A node counts as live if its heartbeat is newer than this; generous to avoid false-dead flapping (see plans) + defaultStateInterval = 15 * time.Second // How often the full subscription state is pushed to peers + + // DefaultBatchLinger is how long a fan-out message may wait in a peer's queue for more + // messages to arrive, so they are delivered as one batch. It trades up to this much + // cross-node latency for a bounded request rate per peer. + DefaultBatchLinger = 500 * time.Millisecond +) + +// Config configures the cluster. It is assembled by the server from its own config, which keeps +// this package free of server types. +type Config struct { + Enabled bool // Master switch; when false, New returns the nop cluster + NodeID NodeID // Stable per-node identifier; required + AdvertiseURL string // Base URL peers use to reach this node's fan-out endpoint + Secret string // Shared secret authenticating node-to-node fan-out requests + HeartbeatInterval time.Duration // How often the node registry heartbeat is refreshed + NodeTTL time.Duration // Registry rows older than this do not count as live peers + BatchLinger time.Duration // How long messages wait in a peer queue to form a batch; 0 = send immediately + StateInterval time.Duration // How often the full subscription state is pushed to peers + MaxMessageBytes int64 // Upper bound for a single message on the wire (batch limits derive from this) +} + +// DeliverFunc hands a message received from a peer node to this node's local subscribers. The +// server supplies it, which inverts the dependency: this package never imports the server. +type DeliverFunc func(m *model.Message) + +// TopicsFunc returns the topics that currently have at least one live subscriber, computed +// fresh on every call: membership is never tracked as a list, so topics "leave" simply by not +// appearing in the next snapshot. The server supplies it (same inversion as DeliverFunc). +type TopicsFunc func() []string + +// Cluster fans published messages out to peer cluster nodes and receives their fan-out requests. +// Local delivery to a node's own subscribers still happens inline in the server; the cluster +// only covers the cross-node hop. +type Cluster interface { + // The internal peer API (/v1/internal/*), served on the dedicated cluster listener. The + // nop cluster answers 404. + http.Handler + // Relay sends a locally published message on to the peer nodes that may have subscribers + // for its topic (all of them, when subscription knowledge is missing or stale). It is + // fire-and-forget and must not block the caller's request path. + Relay(m *model.Message) error + // AnnounceTopics tells peers that these topics just gained their first local subscriber, + // closing the routing-knowledge window to ~one round trip. Nop in single-node mode. + AnnounceTopics(topics []string) + // IsLeader reports whether this node holds the cluster leader lock. Singleton background + // jobs (e.g. the Firebase keepaliver) are gated on the leader. + IsLeader() bool + // TODO(T8): add Healthy() bool reflecting registration health (unhealthy when the last + // successful Register is older than NodeTTL, i.e. when peers stop relaying to this node), + // so /v1/health can pull a DB-partitioned node out of DNS rotation. Needs a fail-open + // answer at the health checker first: during a full DB outage ALL nodes fail to register + // while the mesh keeps delivering on stale peer caches, and naively pulling every node + // would turn a control-plane blip into a total outage. + // Close stops the cluster and releases its resources. + Close() error +} + +// New creates the cluster for the given config: the nop cluster when clustering is disabled (the +// single-node default), or the peer-mesh cluster otherwise. +func New(conf *Config, pool *db.DB, deliver DeliverFunc, topics TopicsFunc) (Cluster, error) { + if !conf.Enabled { + return &nopCluster{}, nil + } + if pool == nil { + return nil, errors.New("cluster mode requires a PostgreSQL database (set database-url)") + } + if conf.AdvertiseURL == "" { + return nil, errors.New("cluster mode requires an advertise URL (set cluster-advertise-url)") + } + if conf.NodeID == "" { + return nil, errors.New("cluster mode requires a stable node ID (set cluster-node-id)") + } + if conf.HeartbeatInterval == 0 { + conf.HeartbeatInterval = defaultHeartbeatInterval + } + if conf.NodeTTL == 0 { + conf.NodeTTL = defaultNodeTTL + } + if conf.StateInterval == 0 { + conf.StateInterval = defaultStateInterval + } + return newMeshCluster(conf, pool, deliver, topics) +} diff --git a/cluster/cluster_load_test.go b/cluster/cluster_load_test.go new file mode 100644 index 00000000..a2689603 --- /dev/null +++ b/cluster/cluster_load_test.go @@ -0,0 +1,133 @@ +package cluster + +import ( + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + dbtest "heckel.io/ntfy/v2/db/test" + "heckel.io/ntfy/v2/model" +) + +// TestMesh_Soak floods the mesh with concurrent publishers and asserts exact delivery: every +// message reaches the peer exactly once, nothing is dropped, and batching keeps the request +// count far below the message count. Skipped unless NTFY_TEST_SOAK is set (it takes a few +// seconds and is meant for pre-deploy verification, not the regular suite). +func TestMesh_Soak(t *testing.T) { + if os.Getenv("NTFY_TEST_SOAK") == "" { + t.Skip("NTFY_TEST_SOAK not set") + } + // ~1000 msg/s aggregate (10x the ntfy.sh peak of ~88 msg/s): each publisher paces itself to + // 100 msg/s. Unthrottled publishing intentionally overruns the bounded per-peer queue (load + // shedding by design), so a zero-drop assertion only holds below the drain ceiling. + const ( + publishers = 10 + messagesPerPublisher = 300 + publishInterval = 10 * time.Millisecond + total = publishers * messagesPerPublisher + ) + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + received := make(map[string]int, total) // message body -> count, to catch duplicates + requests := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.Nil(t, err) + messages, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + mu.Lock() + requests++ + for _, m := range messages { + received[m.Message]++ + } + mu.Unlock() + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + conf := newTestMeshConfig("node-a", "http://127.0.0.1:1") + conf.BatchLinger = 50 * time.Millisecond + conf.NodeTTL = time.Minute // The fake peer never heartbeats; liveness is not under test here + registerFakePeer(t, pool, "node-peer", srv.URL) + mesh, err := newMeshCluster(conf, pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + start := time.Now() + var wg sync.WaitGroup + for p := 0; p < publishers; p++ { + wg.Add(1) + go func(p int) { + defer wg.Done() + ticker := time.NewTicker(publishInterval) + defer ticker.Stop() + for i := 0; i < messagesPerPublisher; i++ { + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", fmt.Sprintf("p%d-m%d", p, i)))) + <-ticker.C + } + }(p) + } + wg.Wait() + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(received) == total + }) + elapsed := time.Since(start) + mu.Lock() + defer mu.Unlock() + for body, count := range received { + require.Equalf(t, 1, count, "message %s delivered %d times", body, count) + } + require.Less(t, requests, total/10, "expected strong batching under load") + t.Logf("soak: %d messages, %d requests (%.1f msgs/request), %.0f msgs/s", + total, requests, float64(total)/float64(requests), float64(total)/elapsed.Seconds()) +} + +// BenchmarkRelay measures the publish-path cost of Relay: marshal + peer lookup (cached) +// + enqueue. The peer never drains, so enqueued fragments are dropped once the queue fills; +// the benchmark measures the hot path, not HTTP delivery. +func BenchmarkRelay(b *testing.B) { + if os.Getenv("NTFY_TEST_DATABASE_URL") == "" { + b.Skip("NTFY_TEST_DATABASE_URL not set") + } + schemaDSN := dbtest.CreateTestPostgresSchema(b) + pool := openTestPool(b, schemaDSN) + conf := newTestMeshConfig("node-a", "http://127.0.0.1:1") + conf.BatchLinger = time.Minute // Never flush; we measure enqueue only + mesh, err := newMeshCluster(conf, pool, nil, nil) + require.Nil(b, err) + defer mesh.Close() + registerFakePeer(b, pool, "node-peer", "http://127.0.0.1:1") + m := model.NewDefaultMessage("mytopic", "benchmark message body of typical size for a push") + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := mesh.Relay(m); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkDecodeFanout measures the receive-path cost of decoding a 100-message NDJSON body. +func BenchmarkDecodeFanout(b *testing.B) { + frags := make([][]byte, 100) + for i := range frags { + frag, err := marshalMessage(model.NewDefaultMessage("mytopic", fmt.Sprintf("benchmark message %d", i))) + require.Nil(b, err) + frags[i] = frag + } + body := assembleMessageBody(frags) + b.SetBytes(int64(len(body))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + messages, err := unmarshalMessageBody(body, 1<<20) + if err != nil || len(messages) != 100 { + b.Fatal("decode failed") + } + } +} diff --git a/cluster/cluster_mesh.go b/cluster/cluster_mesh.go new file mode 100644 index 00000000..0404b3c0 --- /dev/null +++ b/cluster/cluster_mesh.go @@ -0,0 +1,450 @@ +package cluster + +import ( + "bytes" + "context" + "crypto/subtle" + "encoding/json" + "io" + "net/http" + "sync" + "time" + + "heckel.io/ntfy/v2/cluster/registry" + "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/pg" + "heckel.io/ntfy/v2/log" + "heckel.io/ntfy/v2/metrics" + "heckel.io/ntfy/v2/model" + "heckel.io/ntfy/v2/util" +) + +const ( + meshHTTPTimeout = 5 * time.Second + peerQueueSize = 1024 // Bounded per-peer fan-out queue (drop on overflow) + batchMaxMessages = 100 // Flush a batch early when it reaches this many messages + batchMaxBytes = 256 * 1024 // Flush a batch early when it reaches this size + stateMaxBytes = 4 * 1024 * 1024 // Upper bound for inbound state bodies (filter over ~1M topics) + stateFilterFPRate = 0.01 // Bloom false-positive rate; a false positive is one wasted send + tag = "cluster" +) + +// meshCluster fans messages out directly to peer nodes over HTTP (the data plane), using +// PostgreSQL only as a control plane: the node_registry table for membership/discovery, and a +// Postgres advisory lock for singleton-job leader election. Fan-out never touches the database on +// the message path (only the cached peer list does). See plans/260715-scale-out-mesh.md. +// +// Each peer has its own bounded send queue and delivery worker, so a slow or wedged peer only +// backs up (and eventually drops) its own queue and never delays delivery to healthy peers. +type meshCluster struct { + conf *Config + deliver DeliverFunc + topics TopicsFunc + registry *registry.Registry + leader *pg.Leader + httpClient *http.Client + mux *http.ServeMux // The internal peer API; Cluster is an http.Handler + queues map[NodeID]*peerQueue // per-peer send queues; reconciled against the registry + closed bool // Guards against Relay spawning new workers after Close + states map[NodeID]*peerState // what each peer last told us (subscription knowledge) + lastStatePush time.Time // Only touched by the heartbeat goroutine + ctx context.Context + cancel context.CancelFunc + wg sync.WaitGroup + mu sync.Mutex // Protects queues and closed + statesMu sync.Mutex // Protects states +} + +// newMeshCluster creates the mesh cluster: it sets up the registry schema, registers this node +// (synchronously, so it is discoverable before New returns), and starts the heartbeat loop. +// Peer delivery workers are started lazily as peers appear in the registry. +func newMeshCluster(conf *Config, pool *db.DB, deliver DeliverFunc, topics TopicsFunc) (*meshCluster, error) { + if topics == nil { + topics = func() []string { return nil } // No known topics; peers will broadcast to us + } + reg, err := registry.New(pool, string(conf.NodeID), conf.AdvertiseURL, conf.NodeTTL) + if err != nil { + return nil, err + } + // Register synchronously so the node is discoverable before the constructor returns; the + // heartbeat loop refreshes the registration from here on + if err := reg.Register(); err != nil { + return nil, err + } + ctx, cancel := context.WithCancel(context.Background()) + c := &meshCluster{ + conf: conf, + deliver: deliver, + topics: topics, + registry: reg, + leader: pg.NewLeader(pool.Primary(), pg.LeaderLockKey), + httpClient: &http.Client{Timeout: meshHTTPTimeout}, + queues: make(map[NodeID]*peerQueue), + states: make(map[NodeID]*peerState), + ctx: ctx, + cancel: cancel, + } + c.mux = http.NewServeMux() + c.mux.HandleFunc("POST "+MessagePath, c.authenticated(c.handleMessage)) + c.mux.HandleFunc("POST "+StatePath, c.authenticated(c.handleState)) + c.wg.Add(1) + go c.heartbeatLoop() + return c, nil +} + +// ServeHTTP serves the internal peer API. Auth lives in the authenticated middleware, so every +// endpoint gets the same shared-secret and origin handling. +func (c *meshCluster) ServeHTTP(w http.ResponseWriter, r *http.Request) { + c.mux.ServeHTTP(w, r) +} + +// authenticated wraps a peer API handler with the checks every endpoint needs: the shared +// secret (constant-time compare, rejected before any body is read), a present origin, and the +// origin self-skip (a request carrying this node's own traffic is acknowledged but ignored). +func (c *meshCluster) authenticated(h func(origin NodeID, w http.ResponseWriter, r *http.Request)) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + if c.conf.Secret == "" || subtle.ConstantTimeCompare([]byte(r.Header.Get(secretHeader)), []byte(c.conf.Secret)) != 1 { + w.WriteHeader(http.StatusUnauthorized) + return + } + origin := NodeID(r.Header.Get(originHeader)) + if origin == "" { + w.WriteHeader(http.StatusBadRequest) + return + } + if origin == c.conf.NodeID { + w.WriteHeader(http.StatusOK) // Our own traffic; nothing to do + return + } + h(origin, w, r) + } +} + +// heartbeatLoop runs one heartbeat immediately (the ticker first fires a full interval after +// startup, and a fresh node should be leader-capable and state-visible right away), then one per +// interval until shutdown. +func (c *meshCluster) heartbeatLoop() { + defer c.wg.Done() + ticker := time.NewTicker(c.conf.HeartbeatInterval) + defer ticker.Stop() + if err := c.heartbeat(); err != nil { + log.Tag(tag).Err(err).Warn("Cluster heartbeat failed") + } + for { + select { + case <-c.ctx.Done(): + return + case <-ticker.C: + if err := c.heartbeat(); err != nil { + log.Tag(tag).Err(err).Warn("Cluster heartbeat failed") + } + } + } +} + +// heartbeat is one control-plane tick: refresh this node's registry row, retry/confirm the +// leader lock, prune long-dead registry rows (as leader), reconcile the per-peer queues, and +// periodically push our subscription state to peers. +// +// A node that cannot even register itself aborts the tick: the remaining database work would +// fail against the same database, and everything downstream degrades safely without it -- Relay +// serves the stale peer cache on its own, and peers fall back to broadcasting to us once our +// last pushed state expires. +func (c *meshCluster) heartbeat() error { + // TODO(T8): record the last successful Register here to back the future Healthy() method + // (see the Cluster interface) + if err := c.registry.Register(); err != nil { + return err + } + if c.leader.TryAcquire(c.ctx) { + metrics.ClusterLeader.Set(1) + if err := c.registry.Prune(); err != nil { + log.Tag(tag).Err(err).Warn("Failed to prune stale nodes") // Housekeeping only; not fatal for the tick + } + } else { + metrics.ClusterLeader.Set(0) + } + peers, err := c.registry.Peers() + if err != nil { + return err + } + c.reconcilePeers(peers) + if time.Since(c.lastStatePush) >= c.conf.StateInterval { + c.pushState(peers) + c.lastStatePush = time.Now() + } + return nil +} + +// reconcilePeers aligns this node's per-peer attachments with the live peer set: it retires the +// queues (and workers) of peers that have left the registry or re-registered under a new +// advertise URL (the retired queue's remainder was headed for a dead address anyway), and prunes +// the stale state of departed peers. New and replacement queues are created lazily by Relay, not +// here, so a freshly joined peer is reachable immediately. +func (c *meshCluster) reconcilePeers(peers []*registry.Peer) { + metrics.ClusterPeers.Set(float64(len(peers))) + alive := make(map[NodeID]string, len(peers)) // node ID -> advertise URL + for _, p := range peers { + alive[NodeID(p.NodeID)] = p.AdvertiseURL + } + c.mu.Lock() + for nodeID, q := range c.queues { + if url, ok := alive[nodeID]; !ok || q.advertiseURL != url { + q.queue.Close() // Flushes the remainder; the worker exits when the queue is drained + delete(c.queues, nodeID) + } + } + c.mu.Unlock() + // Prune the state of departed peers, but only once stale: state is push-driven and can + // arrive before a new peer is visible in the (up to NodeTTL stale) registry view, so fresh + // state must survive even when its peer is not in the live set. Without this, the states of + // long-gone nodes would accumulate forever. + c.statesMu.Lock() + for nodeID, state := range c.states { + if _, ok := alive[nodeID]; !ok && time.Since(state.updatedAt) > 3*c.conf.StateInterval { + delete(c.states, nodeID) + } + } + c.statesMu.Unlock() +} + +// queueFor returns the send queue for the given peer, creating it (and its delivery worker) if it +// does not exist yet. The caller must hold c.mu. +func (c *meshCluster) queueFor(p *registry.Peer) *peerQueue { + nodeID := NodeID(p.NodeID) + q, ok := c.queues[nodeID] + if ok { + return q + } + q = &peerQueue{ + advertiseURL: p.AdvertiseURL, + queue: util.NewLingerQueue(peerQueueSize, batchMaxMessages, batchMaxBytes, + func(frag []byte) int { return len(frag) }, c.conf.BatchLinger), + } + c.queues[nodeID] = q + c.wg.Add(1) + go c.peerWorker(nodeID, q) + return q +} + +// Relay enqueues the message for delivery to every live peer node that may have subscribers for +// its topic (all of them, absent fresh knowledge). Delivery is fire-and-forget via each peer's +// bounded batching queue; if a peer's queue is full the message is dropped for that peer +// (subscribers reconnect and re-poll history from the database). +func (c *meshCluster) Relay(msg *model.Message) error { + peers, err := c.registry.Peers() + if err != nil { + return err + } + if len(peers) == 0 { + return nil // Cluster of one; skip the marshal + } + frag, err := marshalMessage(msg) + if err != nil { + return err + } + metrics.ClusterMessagesRelayed.Inc() + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return nil // Shutting down; the message is dropped like any other in-flight fan-out + } + for _, p := range peers { + // Route around peers whose fresh state provably excludes this topic; anything less + // certain (no state, stale state) falls back to broadcasting + if !c.mayNeed(NodeID(p.NodeID), msg.Topic) { + metrics.ClusterRouteSkipped.Inc() + continue + } + if !c.queueFor(p).queue.TryEnqueue(frag) { + metrics.ClusterQueueDropped.Inc() + log.Tag(tag).Warn("Fan-out queue for peer %s full, dropping message %s", p.NodeID, msg.ID) + } + } + return nil +} + +// mayNeed reports whether the peer may have a subscriber for the topic. Conservative by +// construction: it returns false only when a fresh state snapshot provably excludes the topic. +// A false positive costs one wasted send; a false negative would lose a message and cannot +// happen for topics a peer has reported (Bloom filters have no false negatives). +func (c *meshCluster) mayNeed(peer NodeID, topic string) bool { + c.statesMu.Lock() + defer c.statesMu.Unlock() + state, ok := c.states[peer] + if !ok || time.Since(state.updatedAt) > 3*c.conf.StateInterval { + return true // No knowledge, or too old to trust for skipping + } + return state.topics.Contains(topic) +} + +// peerWorker delivers batches of queued fan-out messages to a single peer. Batches form in the +// peer's LingerQueue (up to BatchLinger delay, flushed early on size/count caps); the worker +// exits when the queue is closed (peer left the registry, or mesh shutdown) and drained. +func (c *meshCluster) peerWorker(nodeID NodeID, q *peerQueue) { + defer c.wg.Done() + for frags := range q.queue.Dequeue() { + c.postToPeer(nodeID, messageURL(q.advertiseURL), contentTypeNDJSON, assembleMessageBody(frags)) + metrics.ClusterBatchesSent.Inc() + } +} + +// postToPeer POSTs a peer API payload, authenticated with the shared cluster secret. Failures +// are logged and counted, never retried: peer traffic is best-effort by design (messages are +// recovered via since= replay, state via the next periodic push). +func (c *meshCluster) postToPeer(nodeID NodeID, url, contentType string, payload []byte) { + req, err := http.NewRequestWithContext(c.ctx, http.MethodPost, url, bytes.NewReader(payload)) + if err != nil { + metrics.ClusterSendErrors.Inc() + log.Tag(tag).Err(err).Warn("Failed to build request for peer %s", nodeID) + return + } + req.Header.Set("Content-Type", contentType) + req.Header.Set(secretHeader, c.conf.Secret) + req.Header.Set(originHeader, string(c.conf.NodeID)) + resp, err := c.httpClient.Do(req) + if err != nil { + if c.ctx.Err() == nil { + metrics.ClusterSendErrors.Inc() + log.Tag(tag).Err(err).Warn("Failed to send to peer %s (%s)", nodeID, url) + } + return + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + metrics.ClusterSendErrors.Inc() + log.Tag(tag).Warn("Peer %s (%s) rejected request with HTTP %d", nodeID, url, resp.StatusCode) + } +} + +// handleMessage receives a batch of peer messages (NDJSON) and streams them to local +// subscribers line by line, delivering each message as it is decoded. +func (c *meshCluster) handleMessage(_ NodeID, w http.ResponseWriter, r *http.Request) { + // A batch can exceed its byte cap by one message, plus framing overhead + maxBodyBytes := int64(batchMaxBytes) + c.conf.MaxMessageBytes + 1024 + if err := decodeMessageBody(io.LimitReader(r.Body, maxBodyBytes), int(c.conf.MaxMessageBytes), c.deliver); err != nil { + w.WriteHeader(http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusOK) +} + +// handleState receives a peer's state envelope and applies each section it carries. +func (c *meshCluster) handleState(origin NodeID, w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(io.LimitReader(r.Body, stateMaxBytes)) + if err != nil { + w.WriteHeader(http.StatusBadRequest) + return + } + var state apiState + if err := json.Unmarshal(body, &state); err != nil { + w.WriteHeader(http.StatusBadRequest) + return + } + if state.Topics != nil { + if err := c.applyTopicState(origin, state.Topics); err != nil { + w.WriteHeader(http.StatusBadRequest) + return + } + } + w.WriteHeader(http.StatusOK) +} + +// applyTopicState updates what we know about a peer's subscriptions: a full snapshot replaces +// all prior knowledge, an incremental add merges into it. Increments without a baseline are +// ignored on purpose -- without a snapshot the peer is broadcast to anyway. +func (c *meshCluster) applyTopicState(origin NodeID, topics *apiStateTopics) error { + c.statesMu.Lock() + defer c.statesMu.Unlock() + if len(topics.Filter) > 0 { + filter, err := util.UnmarshalBloomFilter(topics.Filter) + if err != nil { + return err + } + c.states[origin] = &peerState{topics: filter, updatedAt: time.Now()} + return nil + } + if state, ok := c.states[origin]; ok { + for _, topic := range topics.Added { + state.topics.Add(topic) + } + state.updatedAt = time.Now() + } + return nil +} + +// pushState sends a full state snapshot to every live peer: a Bloom filter over the topics that +// currently have local subscribers. Sent directly (not via the linger queues -- state must not +// wait behind message batches); a lost push self-heals at the next interval. Topics without +// subscribers disappear simply by not being in the next snapshot. +func (c *meshCluster) pushState(peers []*registry.Peer) { + if len(peers) == 0 { + return + } + topics := c.topics() + filter := util.NewBloomFilter(len(topics), stateFilterFPRate) + for _, topic := range topics { + filter.Add(topic) + } + data, err := filter.MarshalBinary() + if err != nil { + return + } + body, err := json.Marshal(&apiState{Topics: &apiStateTopics{Filter: data}}) + if err != nil { + return + } + for _, p := range peers { + go c.postToPeer(NodeID(p.NodeID), stateURL(p.AdvertiseURL), contentTypeJSON, body) + } + metrics.ClusterStatePushes.Inc() +} + +// AnnounceTopics immediately tells all live peers that these topics gained their first local +// subscriber, shrinking the window in which a publisher could wrongly skip this node from a +// full state interval down to about one round trip. +func (c *meshCluster) AnnounceTopics(topics []string) { + if len(topics) == 0 { + return + } + peers, err := c.registry.Peers() + if err != nil || len(peers) == 0 { + return + } + body, err := json.Marshal(&apiState{Topics: &apiStateTopics{Added: topics}}) + if err != nil { + return + } + for _, p := range peers { + go c.postToPeer(NodeID(p.NodeID), stateURL(p.AdvertiseURL), contentTypeJSON, body) + } +} + +// IsLeader reports whether this node currently holds singleton-job leadership. +func (c *meshCluster) IsLeader() bool { + return c.leader.IsLeader() +} + +// Close stops the mesh: it deregisters this node, releases leadership, stops all peer workers, +// and waits for them to exit. +func (c *meshCluster) Close() error { + c.cancel() // Stops the heartbeat loop and aborts in-flight peer deliveries + // Close the peer queues so their workers flush and exit; final sends are best-effort since + // the context is already canceled (parity with fire-and-forget delivery) + c.mu.Lock() + c.closed = true + for nodeID, q := range c.queues { + q.queue.Close() + delete(c.queues, nodeID) + } + c.mu.Unlock() + // Wait for the loops BEFORE deregistering: an in-flight heartbeat's Register would otherwise + // re-insert our row right after Deregister deleted it + c.wg.Wait() + if err := c.registry.Deregister(); err != nil { + log.Tag(tag).Err(err).Warn("Failed to deregister node") + } + c.leader.Release() + metrics.ClusterLeader.Set(0) + return nil +} diff --git a/cluster/cluster_mesh_test.go b/cluster/cluster_mesh_test.go new file mode 100644 index 00000000..49f19774 --- /dev/null +++ b/cluster/cluster_mesh_test.go @@ -0,0 +1,572 @@ +package cluster + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/cluster/registry" + "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/pg" + dbtest "heckel.io/ntfy/v2/db/test" + "heckel.io/ntfy/v2/model" + "heckel.io/ntfy/v2/util" +) + +const ( + testSecret = "s3cret" +) + +// openTestPool opens a dedicated connection pool to the given test schema, so that each simulated +// node has its own pool like real nodes would. +func openTestPool(t testing.TB, dsn string) *db.DB { + host, err := pg.Open(dsn) + require.Nil(t, err) + d := db.New(host, nil) + t.Cleanup(func() { d.Close() }) + return d +} + +func newTestMeshConfig(nodeID, advertiseURL string) *Config { + return &Config{ + Enabled: true, + NodeID: NodeID(nodeID), + AdvertiseURL: advertiseURL, + Secret: testSecret, + HeartbeatInterval: 100 * time.Millisecond, + NodeTTL: time.Second, // Also the peer cache bound; short so fake peers registered mid-test are seen quickly + MaxMessageBytes: 1 << 20, + StateInterval: time.Minute, // Individual tests lower this to exercise state pushes + } +} + +// registerFakePeer registers a fake peer via the registry (creating the table if the mesh has +// not been constructed yet): tests register fakes before the mesh boots, since its first +// heartbeat caches the peer list. The fake never refreshes its heartbeat. +func registerFakePeer(t testing.TB, pool *db.DB, nodeID NodeID, url string) { + t.Helper() + reg, err := registry.New(pool, string(nodeID), url, time.Minute) + require.Nil(t, err) + require.Nil(t, reg.Register()) +} + +func waitFor(t *testing.T, f func() bool) { + t.Helper() + for i := 0; i < 100; i++ { + if f() { + return + } + time.Sleep(50 * time.Millisecond) + } + t.Fatal("timed out waiting for condition") +} + +func TestMesh_CrossNodeDelivery(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + poolA, poolB := openTestPool(t, schemaDSN), openTestPool(t, schemaDSN) + var mu sync.Mutex + var received []*model.Message + var meshB *meshCluster + srvB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + meshB.ServeHTTP(w, r) + })) + defer srvB.Close() + meshB, err := newMeshCluster(newTestMeshConfig("node-b", srvB.URL), poolB, func(m *model.Message) { + mu.Lock() + defer mu.Unlock() + received = append(received, m) + }, nil) + require.Nil(t, err) + defer meshB.Close() + meshA, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), poolA, func(m *model.Message) { + t.Error("node A must not receive its own relayed message") + }, nil) + require.Nil(t, err) + defer meshA.Close() + msg := model.NewDefaultMessage("mytopic", "hello cross-node") + require.Nil(t, meshA.Relay(msg)) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(received) == 1 + }) + mu.Lock() + defer mu.Unlock() + require.Equal(t, "mytopic", received[0].Topic) + require.Equal(t, "hello cross-node", received[0].Message) +} + +func TestMesh_PeerAPI_Auth(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var delivered int + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, func(m *model.Message) { + delivered++ + }, nil) + require.Nil(t, err) + defer mesh.Close() + frag, err := marshalMessage(model.NewDefaultMessage("mytopic", "hi")) + require.Nil(t, err) + payload := assembleMessageBody([][]byte{frag}) + + // Wrong secret -> 401, not delivered + rr := httptest.NewRecorder() + req := httptest.NewRequest("POST", MessagePath, strings.NewReader(string(payload))) + req.Header.Set(secretHeader, "wrong") + req.Header.Set(originHeader, "node-b") + mesh.ServeHTTP(rr, req) + require.Equal(t, 401, rr.Code) + + // Missing secret -> 401, not delivered + rr = httptest.NewRecorder() + mesh.ServeHTTP(rr, httptest.NewRequest("POST", MessagePath, strings.NewReader(string(payload)))) + require.Equal(t, 401, rr.Code) + require.Equal(t, 0, delivered) + + // Missing origin -> 400, not delivered + rr = httptest.NewRecorder() + req = httptest.NewRequest("POST", MessagePath, strings.NewReader(string(payload))) + req.Header.Set(secretHeader, testSecret) + mesh.ServeHTTP(rr, req) + require.Equal(t, 400, rr.Code) + require.Equal(t, 0, delivered) + + // Correct secret and origin -> 200, delivered + rr = httptest.NewRecorder() + req = httptest.NewRequest("POST", MessagePath, strings.NewReader(string(payload))) + req.Header.Set(secretHeader, testSecret) + req.Header.Set(originHeader, "node-b") + mesh.ServeHTTP(rr, req) + require.Equal(t, 200, rr.Code) + require.Equal(t, 1, delivered) +} + +func TestMesh_PeerAPI_SelfOrigin(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var delivered int + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, func(m *model.Message) { + delivered++ + }, nil) + require.Nil(t, err) + defer mesh.Close() + + // A request that carries this node's own broadcasts must not be re-delivered (loop prevention) + frag, err := marshalMessage(model.NewDefaultMessage("mytopic", "loop")) + require.Nil(t, err) + payload := assembleMessageBody([][]byte{frag}) + rr := httptest.NewRecorder() + req := httptest.NewRequest("POST", MessagePath, strings.NewReader(string(payload))) + req.Header.Set(secretHeader, testSecret) + req.Header.Set(originHeader, "node-a") // Same as the receiving node's ID + mesh.ServeHTTP(rr, req) + require.Equal(t, 200, rr.Code) + require.Equal(t, 0, delivered) +} + +func TestMesh_SlowPeerIsolation(t *testing.T) { + // A wedged peer must not delay delivery to healthy peers: each peer has its own queue and + // delivery worker. With a shared send queue (the design this replaces), the slow peer's + // requests would occupy all delivery workers and starve the fast peer. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + fastReceived := 0 // Messages, not requests: with batching, one request can carry many + srvFast := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.Nil(t, err) + messages, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + mu.Lock() + fastReceived += len(messages) + mu.Unlock() + w.WriteHeader(http.StatusOK) + })) + defer srvFast.Close() + release := make(chan struct{}) + srvSlow := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + <-release // Wedged until the end of the test + w.WriteHeader(http.StatusOK) + })) + defer srvSlow.Close() + defer close(release) + // Register the fake peers before the mesh boots; its first heartbeat caches the peer list + for i, url := range []string{srvFast.URL, srvSlow.URL} { + registerFakePeer(t, pool, NodeID(fmt.Sprintf("node-fake-%d", i)), url) + } + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + const n = 20 + for i := 0; i < n; i++ { + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", fmt.Sprintf("message %d", i)))) + } + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return fastReceived == n + }) +} + +func TestMesh_BatchCoalescing(t *testing.T) { + // Messages published within the linger window arrive as batches: fewer HTTP requests than + // messages, with nothing lost. Fails against a one-request-per-message sender. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + requests, messages := 0, 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.Nil(t, err) + decoded, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + mu.Lock() + requests++ + messages += len(decoded) + mu.Unlock() + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + registerFakePeer(t, pool, "node-fake", srv.URL) + conf := newTestMeshConfig("node-a", "http://127.0.0.1:1") + conf.BatchLinger = 150 * time.Millisecond + mesh, err := newMeshCluster(conf, pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + const n = 20 + for i := 0; i < n; i++ { + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", fmt.Sprintf("message %d", i)))) + } + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return messages == n + }) + mu.Lock() + defer mu.Unlock() + require.Less(t, requests, 5, "expected %d messages coalesced into few requests, got %d", n, requests) +} + +func TestMesh_DeadPeerRemovedAndRejoin(t *testing.T) { + // A peer that dies ungracefully (no Deregister) stops refreshing its heartbeat: after the + // TTL it no longer counts as live (no more sends), its queue/worker are reconciled away, the + // leader prunes its registry row, and a re-registered peer starts receiving again. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + received := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.Nil(t, err) + messages, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + mu.Lock() + received += len(messages) + mu.Unlock() + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + conf := newTestMeshConfig("node-a", "http://127.0.0.1:1") + conf.NodeTTL = 300 * time.Millisecond // Fast expiry so the test observes TTL-based removal + mesh, err := newMeshCluster(conf, pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + // The fake peer registers once and then "dies": its heartbeat is never refreshed + registerFakePeer(t, pool, "node-dead", srv.URL) + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", "while alive"))) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return received == 1 + }) + // After the TTL, the peer is no longer live: its queue is reconciled away and its registry + // row is pruned by the leader (this mesh is the only real node, so it holds the lock) + waitFor(t, func() bool { + mesh.mu.Lock() + defer mesh.mu.Unlock() + return len(mesh.queues) == 0 + }) + waitFor(t, func() bool { + var count int + require.Nil(t, pool.QueryRow(`SELECT COUNT(*) FROM node_registry WHERE node_id = 'node-dead'`).Scan(&count)) + return count == 0 + }) + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", "while dead"))) + time.Sleep(250 * time.Millisecond) // Give a wrong implementation time to deliver anyway + mu.Lock() + require.Equal(t, 1, received) // Only the first message arrived + mu.Unlock() + // The peer comes back (same node ID, fresh heartbeat) and receives messages again; the + // relay retries because the peer list is cached for up to the node TTL + registerFakePeer(t, pool, "node-dead", srv.URL) + waitFor(t, func() bool { + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", "after rejoin"))) + mu.Lock() + defer mu.Unlock() + return received > 1 + }) +} + +func TestMesh_RelayAfterClose(t *testing.T) { + // A Relay racing shutdown (e.g. an in-flight publish during server Stop) must not spawn + // a new peer queue and worker after Close: the worker would never exit (its queue is never + // closed) and nothing waits for it. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + registerFakePeer(t, pool, "node-peer", "http://127.0.0.1:1") + require.Nil(t, mesh.Close()) + require.Nil(t, mesh.Relay(model.NewDefaultMessage("mytopic", "too late"))) // Dropped silently + mesh.mu.Lock() + defer mesh.mu.Unlock() + require.Empty(t, mesh.queues) +} + +func TestMesh_LeaderFailover(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + poolA, poolB := openTestPool(t, schemaDSN), openTestPool(t, schemaDSN) + meshA, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), poolA, nil, nil) + require.Nil(t, err) + defer meshA.Close() + meshB, err := newMeshCluster(newTestMeshConfig("node-b", "http://127.0.0.1:1"), poolB, nil, nil) + require.Nil(t, err) + defer meshB.Close() + // Exactly one node becomes leader + waitFor(t, func() bool { + return meshA.IsLeader() != meshB.IsLeader() // Exactly one + }) + // The leader steps down; the follower takes over + leader, follower := meshA, meshB + if meshB.IsLeader() { + leader, follower = meshB, meshA + } + require.Nil(t, leader.Close()) + waitFor(t, follower.IsLeader) +} + +func TestMesh_CloseDeregisters(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + var count int + require.Nil(t, pool.QueryRow(`SELECT COUNT(*) FROM node_registry WHERE node_id = 'node-a'`).Scan(&count)) + require.Equal(t, 1, count) + require.Nil(t, mesh.Close()) + require.Nil(t, pool.QueryRow(`SELECT COUNT(*) FROM node_registry WHERE node_id = 'node-a'`).Scan(&count)) + require.Equal(t, 0, count) +} + +// postState delivers a state envelope to a mesh's peer API, as a peer would. +func postState(c *meshCluster, origin NodeID, state *apiState) *httptest.ResponseRecorder { + body, err := json.Marshal(state) + if err != nil { + panic(err) + } + rr := httptest.NewRecorder() + req := httptest.NewRequest("POST", StatePath, bytes.NewReader(body)) + req.Header.Set(secretHeader, testSecret) + req.Header.Set(originHeader, string(origin)) + c.ServeHTTP(rr, req) + return rr +} + +// topicFilter builds a marshaled Bloom filter over the given topics. +func topicFilter(t *testing.T, topics ...string) []byte { + t.Helper() + filter := util.NewBloomFilter(len(topics), 0.01) + for _, topic := range topics { + filter.Add(topic) + } + data, err := filter.MarshalBinary() + require.Nil(t, err) + return data +} + +func TestMesh_RouteSkipsUnsubscribedPeer(t *testing.T) { + // A peer whose fresh state provably excludes a topic is not contacted for it; a topic in its + // state is delivered as usual. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + received := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.Nil(t, err) + messages, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + mu.Lock() + received += len(messages) + mu.Unlock() + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + registerFakePeer(t, pool, "node-b", srv.URL) + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + // node-b reports subscribers only for "subscribed-topic" + rr := postState(mesh, "node-b", &apiState{Topics: &apiStateTopics{Filter: topicFilter(t, "subscribed-topic")}}) + require.Equal(t, 200, rr.Code) + // A topic outside the peer's state is skipped + require.Nil(t, mesh.Relay(model.NewDefaultMessage("other-topic", "skipped"))) + time.Sleep(300 * time.Millisecond) // Give a wrong implementation time to deliver anyway + mu.Lock() + require.Equal(t, 0, received) + mu.Unlock() + // A topic inside the peer's state is delivered + require.Nil(t, mesh.Relay(model.NewDefaultMessage("subscribed-topic", "delivered"))) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return received == 1 + }) +} + +func TestMesh_RouteBroadcastsOnStaleState(t *testing.T) { + // State too old to trust cannot justify skipping: the peer is broadcast to as if unknown. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + var mu sync.Mutex + received := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == MessagePath { // The mesh also pushes state here; count only messages + mu.Lock() + received++ + mu.Unlock() + } + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + registerFakePeer(t, pool, "node-b", srv.URL) + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + rr := postState(mesh, "node-b", &apiState{Topics: &apiStateTopics{Filter: topicFilter(t, "subscribed-topic")}}) + require.Equal(t, 200, rr.Code) + // Age the state beyond the trust window + mesh.statesMu.Lock() + mesh.states["node-b"].updatedAt = time.Now().Add(-time.Hour) + mesh.statesMu.Unlock() + require.Nil(t, mesh.Relay(model.NewDefaultMessage("other-topic", "broadcast anyway"))) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return received == 1 + }) +} + +func TestMesh_StatePushReplacesAndRemoves(t *testing.T) { + // Node A periodically pushes a full snapshot of its live topics to node B; each snapshot + // REPLACES B's knowledge, so topics that lost their subscribers disappear without any + // explicit removal protocol. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + poolA, poolB := openTestPool(t, schemaDSN), openTestPool(t, schemaDSN) + var topicsMu sync.Mutex + topicsA := []string{"topic-1"} + source := func() []string { + topicsMu.Lock() + defer topicsMu.Unlock() + return append([]string{}, topicsA...) + } + var meshB *meshCluster + srvB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + meshB.ServeHTTP(w, r) + })) + defer srvB.Close() + meshB, err := newMeshCluster(newTestMeshConfig("node-b", srvB.URL), poolB, nil, nil) + require.Nil(t, err) + defer meshB.Close() + confA := newTestMeshConfig("node-a", "http://127.0.0.1:1") + confA.StateInterval = 200 * time.Millisecond + meshA, err := newMeshCluster(confA, poolA, nil, source) + require.Nil(t, err) + defer meshA.Close() + // B learns A's topics via the periodic push + knows := func(topic string) func() bool { + return func() bool { + meshB.statesMu.Lock() + defer meshB.statesMu.Unlock() + state, ok := meshB.states["node-a"] + return ok && state.topics.Contains(topic) + } + } + waitFor(t, knows("topic-1")) + // A's subscribers change; the next snapshot replaces the old knowledge entirely + topicsMu.Lock() + topicsA = []string{"topic-2"} + topicsMu.Unlock() + waitFor(t, knows("topic-2")) + waitFor(t, func() bool { return !knows("topic-1")() }) +} + +func TestMesh_AnnounceClosesWindow(t *testing.T) { + // A topic gaining its first subscriber is announced immediately, so peers learn about it + // without waiting for the next full state push. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + poolA, poolB := openTestPool(t, schemaDSN), openTestPool(t, schemaDSN) + var meshB *meshCluster + srvB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + meshB.ServeHTTP(w, r) + })) + defer srvB.Close() + meshB, err := newMeshCluster(newTestMeshConfig("node-b", srvB.URL), poolB, nil, nil) + require.Nil(t, err) + defer meshB.Close() + confA := newTestMeshConfig("node-a", "http://127.0.0.1:1") + confA.StateInterval = 200 * time.Millisecond // One full push establishes the baseline + meshA, err := newMeshCluster(confA, poolA, nil, func() []string { return []string{"existing"} }) + require.Nil(t, err) + defer meshA.Close() + waitFor(t, func() bool { + meshB.statesMu.Lock() + defer meshB.statesMu.Unlock() + _, ok := meshB.states["node-a"] + return ok + }) + // Announcements merge into the baseline right away + meshA.AnnounceTopics([]string{"fresh-topic"}) + waitFor(t, func() bool { + meshB.statesMu.Lock() + defer meshB.statesMu.Unlock() + state, ok := meshB.states["node-a"] + return ok && state.topics.Contains("fresh-topic") + }) +} + +func TestMesh_StateOfDepartedPeerPruned(t *testing.T) { + // peerState is push-driven and can arrive before the peer is visible in the registry, so it + // must survive reconcile while fresh -- but a departed peer's state must not leak forever: + // once it is both absent from the registry and stale past the trust window, it is pruned. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + mesh, err := newMeshCluster(newTestMeshConfig("node-a", "http://127.0.0.1:1"), pool, nil, nil) + require.Nil(t, err) + defer mesh.Close() + rr := postState(mesh, "node-gone", &apiState{Topics: &apiStateTopics{Filter: topicFilter(t, "some-topic")}}) + require.Equal(t, 200, rr.Code) + // Fresh state of an unknown peer survives reconcile (the new-node visibility window) + mesh.reconcilePeers(nil) + mesh.statesMu.Lock() + _, ok := mesh.states["node-gone"] + mesh.statesMu.Unlock() + require.True(t, ok) + // Stale state of an absent peer is pruned + mesh.statesMu.Lock() + mesh.states["node-gone"].updatedAt = time.Now().Add(-time.Hour) + mesh.statesMu.Unlock() + mesh.reconcilePeers(nil) + mesh.statesMu.Lock() + _, ok = mesh.states["node-gone"] + mesh.statesMu.Unlock() + require.False(t, ok) +} diff --git a/cluster/cluster_nop.go b/cluster/cluster_nop.go new file mode 100644 index 00000000..06231f52 --- /dev/null +++ b/cluster/cluster_nop.go @@ -0,0 +1,24 @@ +package cluster + +import ( + "net/http" + + "heckel.io/ntfy/v2/model" +) + +// nopCluster is the single-node default: it drops all relayed messages, rejects peer API requests, and +// reports this node as leader (a single node is trivially the leader, so leader-gated jobs need +// no special-casing in single-node mode). +type nopCluster struct{} + +func (c *nopCluster) Relay(_ *model.Message) error { return nil } + +func (c *nopCluster) ServeHTTP(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotFound) +} + +func (c *nopCluster) AnnounceTopics(_ []string) {} + +func (c *nopCluster) IsLeader() bool { return true } + +func (c *nopCluster) Close() error { return nil } diff --git a/cluster/cluster_test.go b/cluster/cluster_test.go new file mode 100644 index 00000000..b5ae2290 --- /dev/null +++ b/cluster/cluster_test.go @@ -0,0 +1,83 @@ +package cluster + +import ( + "bytes" + "net/http/httptest" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/model" +) + +func TestDeliver_RoundTrip(t *testing.T) { + // The fan-out body is NDJSON: one apiDeliverMessage per line, joined from pre-marshaled + // fragments; the origin travels in a header, not the body + m1 := model.NewDefaultMessage("mytopic", "my message") + m1.Sender = netip.MustParseAddr("1.2.3.4") + m1.User = "u_abc" + m2 := model.NewDefaultMessage("othertopic", "other message") + frag1, err := marshalMessage(m1) + require.Nil(t, err) + frag2, err := marshalMessage(m2) + require.Nil(t, err) + messages, err := unmarshalMessageBody(assembleMessageBody([][]byte{frag1, frag2}), 1<<20) + require.Nil(t, err) + require.Len(t, messages, 2) + require.Equal(t, "mytopic", messages[0].Topic) + require.Equal(t, "my message", messages[0].Message) + // Sender and User are json:"-" on model.Message; the lines must carry and reattach them + require.Equal(t, netip.MustParseAddr("1.2.3.4"), messages[0].Sender) + require.Equal(t, "u_abc", messages[0].User) + require.Equal(t, "othertopic", messages[1].Topic) + require.False(t, messages[1].Sender.IsValid()) +} + +func TestDeliver_SingleMessage(t *testing.T) { + // A single message is just a one-line body; there is no separate single-message format + frag, err := marshalMessage(model.NewDefaultMessage("mytopic", "hi")) + require.Nil(t, err) + messages, err := unmarshalMessageBody(assembleMessageBody([][]byte{frag}), 1<<20) + require.Nil(t, err) + require.Len(t, messages, 1) +} + +func TestDeliver_MalformedLinesSkipped(t *testing.T) { + // Fan-out is fire-and-forget: a malformed or message-less line is skipped (and logged), the + // remaining lines are still delivered + frag, err := marshalMessage(model.NewDefaultMessage("mytopic", "good")) + require.Nil(t, err) + body := []byte("this is not json\n{\"sender\":\"1.2.3.4\"}\n" + string(frag) + "\n\n") + messages, err := unmarshalMessageBody(body, 1<<20) + require.Nil(t, err) + require.Len(t, messages, 1) + require.Equal(t, "good", messages[0].Message) +} + +// unmarshalMessageBody is a test helper collecting the messages of an NDJSON message body. +func unmarshalMessageBody(body []byte, maxLineBytes int) ([]*model.Message, error) { + var messages []*model.Message + err := decodeMessageBody(bytes.NewReader(body), maxLineBytes, func(m *model.Message) { + messages = append(messages, m) + }) + return messages, err +} + +func TestNop(t *testing.T) { + b, err := New(&Config{}, nil, nil, nil) // not enabled -> nop cluster, no database required + require.Nil(t, err) + require.IsType(t, &nopCluster{}, b) + require.Nil(t, b.Relay(model.NewDefaultMessage("mytopic", "hi"))) + // A single node is trivially the leader, so leader-gated jobs run without special-casing + require.True(t, b.IsLeader()) + rr := httptest.NewRecorder() + b.ServeHTTP(rr, httptest.NewRequest("POST", MessagePath, nil)) + require.Equal(t, 404, rr.Code) + require.Nil(t, b.Close()) +} + +func TestNew_EnabledRequiresDatabase(t *testing.T) { + _, err := New(&Config{Enabled: true, Secret: "secret"}, nil, nil, nil) + require.Error(t, err) + require.Contains(t, err.Error(), "database") +} diff --git a/cluster/registry/registry.go b/cluster/registry/registry.go new file mode 100644 index 00000000..52daf5c2 --- /dev/null +++ b/cluster/registry/registry.go @@ -0,0 +1,149 @@ +// Package registry implements cluster membership: each node upserts its own row into the +// node_registry table with a fresh heartbeat, and discovers its peers by reading the other +// fresh rows. Node IDs are plain strings here; the cluster package layers its NodeID type on +// top. +package registry + +import ( + "sync" + "time" + + "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/schema" +) + +// Registry queries +const ( + upsertNodeQuery = ` + INSERT INTO node_registry (node_id, advertise_url, last_heartbeat) + VALUES ($1, $2, $3) + ON CONFLICT (node_id) DO UPDATE SET advertise_url = EXCLUDED.advertise_url, last_heartbeat = EXCLUDED.last_heartbeat + ` + selectPeersQuery = `SELECT node_id, advertise_url FROM node_registry WHERE last_heartbeat >= $1 AND node_id != $2` + pruneStaleNodesQuery = `DELETE FROM node_registry WHERE last_heartbeat < $1` + deleteNodeQuery = `DELETE FROM node_registry WHERE node_id = $1` +) + +// Schema version and queries + +const ( + schemaVersion = 1 + schemaStoreKey = "node_registry" +) + +var ( + createTable = schema.AsMigrateFunc(` + CREATE TABLE IF NOT EXISTS node_registry ( + node_id TEXT PRIMARY KEY, + advertise_url TEXT NOT NULL, + last_heartbeat BIGINT NOT NULL + ) + `) +) + +// Peer is a live remote node as read from the registry. +type Peer struct { + NodeID string + AdvertiseURL string +} + +// Registry is the node membership table (control plane): each node upserts its own row with a +// fresh heartbeat every few seconds, and peers are the other rows with a heartbeat newer than +// the TTL. Stale rows are pruned by the leader. The TTL bounds membership staleness in BOTH +// directions: how long a silent node still counts as live, and how long the cached peer list is +// served before a re-read -- so a new node may take up to a TTL to become visible. +type Registry struct { + pool *db.DB + nodeID string + advertiseURL string + ttl time.Duration + peers []*Peer // cached peer list + peersFetched time.Time + mu sync.Mutex // Protects peers and peersFetched +} + +// New creates or migrates the registry schema and returns this node's membership handle. It +// does NOT register the node: joining the cluster is an explicit Register call, owned by the +// caller, so read-only uses of the registry stay side-effect free. +func New(pool *db.DB, nodeID, advertiseURL string, ttl time.Duration) (*Registry, error) { + if err := schema.Migrate(pool.Primary(), schema.Postgres, schemaStoreKey, schemaVersion, createTable, nil); err != nil { + return nil, err + } + return &Registry{ + pool: pool, + nodeID: nodeID, + advertiseURL: advertiseURL, + ttl: ttl, + }, nil +} + +// Register upserts this node into the registry with a fresh heartbeat. It is a pure write: it +// does not touch the peer cache, because our own row is excluded from Peers() anyway. +func (r *Registry) Register() error { + _, err := r.pool.Exec(upsertNodeQuery, r.nodeID, r.advertiseURL, time.Now().Unix()) + return err +} + +// Peers returns the current set of live peer nodes (all registry rows with a fresh heartbeat, +// excluding this node), cached for the TTL. +func (r *Registry) Peers() ([]*Peer, error) { + r.mu.Lock() + if r.peers != nil && time.Since(r.peersFetched) < r.ttl { + peers := r.peers + r.mu.Unlock() + return peers, nil + } + r.mu.Unlock() + peers, err := r.queryPeers() + if err != nil { + // Serve the last-known peer list during database hiccups: fan-out keeps flowing to + // known peers instead of erroring (and logging) once per published message for the + // duration of the outage. Dead peers in the stale list only cost failed sends. + r.mu.Lock() + defer r.mu.Unlock() + if r.peers != nil { + return r.peers, nil + } + return nil, err + } + r.mu.Lock() + r.peers = peers + r.peersFetched = time.Now() + r.mu.Unlock() + return peers, nil +} + +// Prune deletes registry rows whose heartbeat is long expired. Only the leader calls this; the +// grace period of 3x the TTL avoids deleting rows of nodes that are merely slow to heartbeat. +func (r *Registry) Prune() error { + _, err := r.pool.Exec(pruneStaleNodesQuery, time.Now().Add(-3*r.ttl).Unix()) + return err +} + +// Deregister deletes this node's registry row; called on shutdown. +func (r *Registry) Deregister() error { + _, err := r.pool.Exec(deleteNodeQuery, r.nodeID) + return err +} + +// queryPeers reads the current live peer set from the registry table. +func (r *Registry) queryPeers() ([]*Peer, error) { + cutoff := time.Now().Add(-r.ttl).Unix() + rows, err := r.pool.Query(selectPeersQuery, cutoff, r.nodeID) + if err != nil { + return nil, err + } + defer rows.Close() + peers := make([]*Peer, 0) + for rows.Next() { + p := &Peer{} + if err := rows.Scan(&p.NodeID, &p.AdvertiseURL); err != nil { + return nil, err + } + peers = append(peers, p) + } + if err := rows.Err(); err != nil { + return nil, err + } + return peers, nil +} diff --git a/cluster/registry/registry_test.go b/cluster/registry/registry_test.go new file mode 100644 index 00000000..c4fb6610 --- /dev/null +++ b/cluster/registry/registry_test.go @@ -0,0 +1,222 @@ +package registry + +import ( + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/pg" + dbtest "heckel.io/ntfy/v2/db/test" +) + +func openTestPool(t *testing.T, dsn string) *db.DB { + t.Helper() + host, err := pg.Open(dsn) + require.Nil(t, err) + d := db.New(host, nil) + t.Cleanup(func() { d.Close() }) + return d +} + +func TestRegistry_NewDoesNotRegister(t *testing.T) { + // New only sets up the schema and the identity handle; joining the cluster is an explicit + // Register call, owned by the caller (the mesh registers synchronously at construction). + // This keeps read-only uses (ops tooling, future admin endpoints) side-effect free. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + require.Equal(t, 0, countRows(t, pool, "node-1")) + require.Nil(t, r1.Register()) + require.Equal(t, 1, countRows(t, pool, "node-1")) +} + +func TestRegistry_RegisterAndPeers(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, r1.Register()) + r2, err := New(pool, "node-2", "http://10.0.0.2:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, r2.Register()) + // Each node sees the other, never itself + peers, err := r1.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) + require.Equal(t, "node-2", peers[0].NodeID) + require.Equal(t, "http://10.0.0.2:2587", peers[0].AdvertiseURL) + peers, err = r2.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) + require.Equal(t, "node-1", peers[0].NodeID) +} + +func TestRegistry_ReRegisterUpdatesAdvertiseURL(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + // The same node comes back under a new address; the upsert replaces the row + old, err := New(pool, "node-2", "http://old:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, old.Register()) + renewed, err := New(pool, "node-2", "http://new:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, renewed.Register()) + expireCache(r1) + peers, err := r1.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) + require.Equal(t, "http://new:2587", peers[0].AdvertiseURL) +} + +func TestRegistry_PeersCachedForTTL(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + peers, err := r1.Peers() + require.Nil(t, err) + require.Empty(t, peers) + // A node joining after the cache was populated is invisible until the cache expires + r2, err := New(pool, "node-2", "http://10.0.0.2:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, r2.Register()) + peers, err = r1.Peers() + require.Nil(t, err) + require.Empty(t, peers) + expireCache(r1) + peers, err = r1.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) +} + +func TestRegistry_TTLExcludesSilentNodes(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + // A node whose heartbeat is older than the TTL does not count as live + _, err = pool.Exec(upsertNodeQuery, "node-silent", "http://10.0.0.9:2587", time.Now().Add(-2*time.Minute).Unix()) + require.Nil(t, err) + peers, err := r1.Peers() + require.Nil(t, err) + require.Empty(t, peers) +} + +func TestRegistry_PruneDeletesLongDeadOnly(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + // One node beyond the 3x TTL grace period, one merely stale + _, err = pool.Exec(upsertNodeQuery, "node-long-dead", "http://10.0.0.8:2587", time.Now().Add(-4*time.Minute).Unix()) + require.Nil(t, err) + _, err = pool.Exec(upsertNodeQuery, "node-slow", "http://10.0.0.9:2587", time.Now().Add(-2*time.Minute).Unix()) + require.Nil(t, err) + require.Nil(t, r1.Prune()) + require.Equal(t, 0, countRows(t, pool, "node-long-dead")) + require.Equal(t, 1, countRows(t, pool, "node-slow")) // Slow, not dead: kept +} + +func TestRegistry_Deregister(t *testing.T) { + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, r1.Register()) + require.Equal(t, 1, countRows(t, pool, "node-1")) + require.Nil(t, r1.Deregister()) + require.Equal(t, 0, countRows(t, pool, "node-1")) +} + +func TestRegistry_PeersStaleCacheOnError(t *testing.T) { + // During a database hiccup, Peers serves the last-known peer list instead of erroring: + // fan-out keeps flowing to known peers, and the publish path does not log a warning per + // message for the duration of the outage. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + r1, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + r2, err := New(pool, "node-2", "http://10.0.0.2:2587", time.Minute) + require.Nil(t, err) + require.Nil(t, r2.Register()) + peers, err := r1.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) + // Expire the cache and break the database; the stale list must still be served + expireCache(r1) + require.Nil(t, pool.Close()) + peers, err = r1.Peers() + require.Nil(t, err) + require.Len(t, peers, 1) + require.Equal(t, "node-2", peers[0].NodeID) +} + +func TestRegistry_ConcurrentCreate(t *testing.T) { + // Multiple nodes cold-booting on a fresh database must not race on table creation: CREATE + // TABLE IF NOT EXISTS is not atomic in PostgreSQL, so creation is serialized via an advisory + // lock. Without it, this test fails sporadically with a duplicate-key error on pg_class. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + const n = 8 + errs := make(chan error, n) + for i := 0; i < n; i++ { + go func(i int) { + pool, err := pg.Open(schemaDSN) + if err != nil { + errs <- err + return + } + defer pool.DB.Close() + _, err = New(db.New(pool, nil), fmt.Sprintf("node-%d", i), "http://127.0.0.1:1", time.Second) + errs <- err + }(i) + } + for i := 0; i < n; i++ { + require.Nil(t, <-errs) + } +} + +func TestRegistry_SchemaVersionWritten(t *testing.T) { + // The registry participates in the shared schema_version framework like every other store, + // so future table changes can be applied as migrations. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + _, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + var version int + require.Nil(t, pool.QueryRow(`SELECT version FROM schema_version WHERE store = $1`, schemaStoreKey).Scan(&version)) + require.Equal(t, schemaVersion, version) + // Setup is idempotent: a second node boots against the migrated schema + _, err = New(pool, "node-2", "http://10.0.0.2:2587", time.Minute) + require.Nil(t, err) +} + +func TestRegistry_SchemaVersionFromTheFuture(t *testing.T) { + // A node running older code must refuse to touch a schema migrated by newer code + schemaDSN := dbtest.CreateTestPostgresSchema(t) + pool := openTestPool(t, schemaDSN) + _, err := New(pool, "node-1", "http://10.0.0.1:2587", time.Minute) + require.Nil(t, err) + _, err = pool.Exec(`UPDATE schema_version SET version = 99 WHERE store = $1`, schemaStoreKey) + require.Nil(t, err) + _, err = New(pool, "node-2", "http://10.0.0.2:2587", time.Minute) + require.Error(t, err) +} + +// expireCache forces the next Peers() call to re-read the registry table. +func expireCache(r *Registry) { + r.mu.Lock() + r.peersFetched = time.Time{} + r.mu.Unlock() +} + +func countRows(t *testing.T, pool *db.DB, nodeID string) int { + t.Helper() + var count int + require.Nil(t, pool.QueryRow(`SELECT COUNT(*) FROM node_registry WHERE node_id = $1`, nodeID).Scan(&count)) + return count +} diff --git a/cluster/types.go b/cluster/types.go new file mode 100644 index 00000000..b85247da --- /dev/null +++ b/cluster/types.go @@ -0,0 +1,47 @@ +package cluster + +import ( + "time" + + "heckel.io/ntfy/v2/model" + "heckel.io/ntfy/v2/util" +) + +// apiMessage is one line of a message request body (NDJSON: one message per line; a single +// message is just a one-line body). It carries the two fields that model.Message does not +// serialize to JSON (Sender and User), which are needed to reconstruct the visitor on the +// receiving node. The origin node travels in a request header, not in the body. +type apiMessage struct { + Sender string `json:"sender,omitempty"` + User string `json:"user,omitempty"` + Message *model.Message `json:"message"` +} + +// apiState is the peer state-exchange envelope. Each concern is an optional section; future +// concerns (rate limit counters, stats) become siblings of Topics. +type apiState struct { + Topics *apiStateTopics `json:"topics,omitempty"` +} + +// apiStateTopics carries a peer's subscription knowledge: either a full snapshot (Filter, a +// marshaled Bloom filter over the topics with live subscribers) replacing all prior knowledge, +// or an incremental update (Added) merged into it. +type apiStateTopics struct { + Filter []byte `json:"filter,omitempty"` + Added []string `json:"added,omitempty"` +} + +// peerState is what a peer last told us about itself; Relay routes around peers whose +// fresh state provably excludes a topic. +type peerState struct { + topics *util.BloomFilter + updatedAt time.Time +} + +// peerQueue is the bounded, batching send queue for a single peer, pinned to the advertise URL +// the peer was created with: a peer re-registering under a different advertise URL is treated +// as a replacement (reconcile retires the old queue; Relay creates a fresh one on demand). +type peerQueue struct { + advertiseURL string + queue *util.LingerQueue[[]byte] // pre-marshaled apiMessage fragments +} diff --git a/cluster/util.go b/cluster/util.go new file mode 100644 index 00000000..8897ed5a --- /dev/null +++ b/cluster/util.go @@ -0,0 +1,69 @@ +package cluster + +import ( + "bufio" + "bytes" + "encoding/json" + "io" + "net/netip" + "strings" + + "heckel.io/ntfy/v2/log" + "heckel.io/ntfy/v2/model" +) + +// messageURL derives the peer's message endpoint URL from its advertise URL. +func messageURL(advertiseURL string) string { + return strings.TrimRight(advertiseURL, "/") + MessagePath +} + +// stateURL derives the peer's state endpoint URL from its advertise URL. +func stateURL(advertiseURL string) string { + return strings.TrimRight(advertiseURL, "/") + StatePath +} + +// marshalMessage serializes one message and its non-JSON fields (Sender, User) as an +// apiMessage line. Lines are marshaled once per publish and shared across all per-peer +// queues; assembleMessageBody joins them without re-marshaling. +func marshalMessage(m *model.Message) ([]byte, error) { + apiMsg := &apiMessage{User: m.User, Message: m} + if m.Sender.IsValid() { + apiMsg.Sender = m.Sender.String() + } + return json.Marshal(apiMsg) +} + +// assembleMessageBody builds an NDJSON fan-out request body from pre-marshaled apiMessage +// lines, avoiding a second JSON marshal of the messages. +func assembleMessageBody(frags [][]byte) []byte { + return append(bytes.Join(frags, []byte("\n")), '\n') +} + +// decodeMessageBody reads NDJSON apiMessage lines from r, reattaches the non-JSON fields +// (Sender, User) onto each message, and hands them to deliver. Malformed or message-less lines +// are skipped and logged, not fatal: fan-out is fire-and-forget, so the valid remainder of a +// request is still delivered. It returns an error only for stream-level failures (e.g. a line +// exceeding maxLineBytes). +func decodeMessageBody(r io.Reader, maxLineBytes int, deliver DeliverFunc) error { + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 64*1024), maxLineBytes) + for scanner.Scan() { + line := bytes.TrimSpace(scanner.Bytes()) + if len(line) == 0 { + continue + } + var apiMsg apiMessage + if err := json.Unmarshal(line, &apiMsg); err != nil || apiMsg.Message == nil { + log.Tag(tag).Warn("Skipping malformed fan-out line") + continue + } + apiMsg.Message.User = apiMsg.User + if apiMsg.Sender != "" { + if addr, err := netip.ParseAddr(apiMsg.Sender); err == nil { + apiMsg.Message.Sender = addr + } + } + deliver(apiMsg.Message) + } + return scanner.Err() +} diff --git a/cmd/serve.go b/cmd/serve.go index f0198d5f..a8d748f7 100644 --- a/cmd/serve.go +++ b/cmd/serve.go @@ -19,6 +19,7 @@ import ( "github.com/urfave/cli/v2" "github.com/urfave/cli/v2/altsrc" "heckel.io/ntfy/v2/ban" + "heckel.io/ntfy/v2/cluster" "heckel.io/ntfy/v2/log" "heckel.io/ntfy/v2/payments" "heckel.io/ntfy/v2/server" @@ -43,6 +44,11 @@ var flagsServe = append( altsrc.NewStringFlag(&cli.StringFlag{Name: "firebase-key-file", Aliases: []string{"firebase_key_file", "F"}, EnvVars: []string{"NTFY_FIREBASE_KEY_FILE"}, Usage: "Firebase credentials file; if set additionally publish to FCM topic"}), altsrc.NewStringFlag(&cli.StringFlag{Name: "database-url", Aliases: []string{"database_url"}, EnvVars: []string{"NTFY_DATABASE_URL"}, Usage: "PostgreSQL connection string for database-backed stores (e.g. postgres://user:pass@host:5432/ntfy)"}), altsrc.NewStringSliceFlag(&cli.StringSliceFlag{Name: "database-replica-urls", Aliases: []string{"database_replica_urls"}, EnvVars: []string{"NTFY_DATABASE_REPLICA_URLS"}, Usage: "PostgreSQL read replica connection strings for offloading read queries"}), + altsrc.NewStringFlag(&cli.StringFlag{Name: "cluster-node-id", Aliases: []string{"cluster_node_id"}, EnvVars: []string{"NTFY_CLUSTER_NODE_ID"}, Usage: "stable per-node identifier for the cluster node registry (required in cluster mode)"}), + altsrc.NewStringFlag(&cli.StringFlag{Name: "cluster-listen", Aliases: []string{"cluster_listen"}, EnvVars: []string{"NTFY_CLUSTER_LISTEN"}, Usage: "ip:port for the dedicated cluster fan-out listener; bind it to the private network (e.g. 10.0.0.5:2587)"}), + altsrc.NewStringFlag(&cli.StringFlag{Name: "cluster-advertise-url", Aliases: []string{"cluster_advertise_url"}, EnvVars: []string{"NTFY_CLUSTER_ADVERTISE_URL"}, Usage: "base URL peer nodes use to reach this node's fan-out listener (defaults to http://)"}), + altsrc.NewStringFlag(&cli.StringFlag{Name: "cluster-secret", Aliases: []string{"cluster_secret"}, EnvVars: []string{"NTFY_CLUSTER_SECRET"}, Usage: "shared secret authenticating node-to-node fan-out requests"}), + altsrc.NewStringFlag(&cli.StringFlag{Name: "cluster-batch-linger", Aliases: []string{"cluster_batch_linger"}, EnvVars: []string{"NTFY_CLUSTER_BATCH_LINGER"}, Value: util.FormatDuration(cluster.DefaultBatchLinger), Usage: "how long fan-out messages wait to form a batch per peer node (0 = send immediately)"}), altsrc.NewStringFlag(&cli.StringFlag{Name: "cache-file", Aliases: []string{"cache_file", "C"}, EnvVars: []string{"NTFY_CACHE_FILE"}, Usage: "cache file used for message caching"}), altsrc.NewStringFlag(&cli.StringFlag{Name: "cache-duration", Aliases: []string{"cache_duration", "b"}, EnvVars: []string{"NTFY_CACHE_DURATION"}, Value: util.FormatDuration(server.DefaultCacheDuration), Usage: "buffer messages for this time to allow `since` requests"}), altsrc.NewIntFlag(&cli.IntFlag{Name: "cache-batch-size", Aliases: []string{"cache_batch_size"}, EnvVars: []string{"NTFY_BATCH_SIZE"}, Usage: "max size of messages to batch together when writing to message cache (if zero, writes are synchronous)"}), @@ -157,6 +163,11 @@ func execServe(c *cli.Context) error { firebaseKeyFile := c.String("firebase-key-file") databaseURL := c.String("database-url") databaseReplicaURLs := c.StringSlice("database-replica-urls") + clusterNodeID := c.String("cluster-node-id") + clusterListen := c.String("cluster-listen") + clusterAdvertiseURL := c.String("cluster-advertise-url") + clusterSecret := c.String("cluster-secret") + clusterBatchLingerStr := c.String("cluster-batch-linger") webPushPrivateKey := c.String("web-push-private-key") webPushPublicKey := c.String("web-push-public-key") webPushFile := c.String("web-push-file") @@ -252,6 +263,10 @@ func execServe(c *cli.Context) error { if err != nil { return fmt.Errorf("invalid keepalive interval: %s", keepaliveIntervalStr) } + clusterBatchLinger, err := util.ParseDuration(clusterBatchLingerStr) + if err != nil || clusterBatchLinger < 0 { + return fmt.Errorf("invalid cluster batch linger: %s", clusterBatchLingerStr) + } managerInterval, err := util.ParseDuration(managerIntervalStr) if err != nil { return fmt.Errorf("invalid manager interval: %s", managerIntervalStr) @@ -322,6 +337,16 @@ func execServe(c *cli.Context) error { return errors.New("if database-url is set, auth-file, cache-file, and web-push-file must not be set") } else if len(databaseReplicaURLs) > 0 && databaseURL == "" { return errors.New("database-replica-urls can only be used if database-url is also set") + } else if clusterListen != "" && databaseURL == "" { + return errors.New("cluster-listen requires database-url to be set") + } else if clusterListen != "" && clusterSecret == "" { + return errors.New("cluster-listen requires cluster-secret to be set") + } else if clusterListen != "" && clusterNodeID == "" { + return errors.New("cluster-listen requires cluster-node-id to be set") + } else if clusterListen == "" && clusterSecret != "" { + return errors.New("cluster-secret can only be used if cluster-listen is set") + } else if clusterListen != "" && clusterAdvertiseURL == "" && wildcardAddr(clusterListen) { + return errors.New("cluster-advertise-url must be set if cluster-listen binds a wildcard address") } else if firebaseKeyFile != "" && !util.FileExists(firebaseKeyFile) { return errors.New("if set, FCM key file must exist") } else if firebaseKeyFile != "" && !server.FirebaseAvailable { @@ -559,6 +584,11 @@ func execServe(c *cli.Context) error { conf.ProfileListenHTTP = profileListenHTTP conf.DatabaseURL = databaseURL conf.DatabaseReplicaURLs = databaseReplicaURLs + conf.ClusterNodeID = clusterNodeID + conf.ClusterListen = clusterListen + conf.ClusterAdvertiseURL = clusterAdvertiseURL + conf.ClusterSecret = clusterSecret + conf.ClusterBatchLinger = clusterBatchLinger conf.WebPushPrivateKey = webPushPrivateKey conf.WebPushPublicKey = webPushPublicKey conf.WebPushFile = webPushFile @@ -737,3 +767,13 @@ func maybeFromMetadata(m map[string]any, key string) string { } return s } + +// wildcardAddr reports whether the given listen address binds all interfaces (e.g. ":2587", +// "0.0.0.0:2587", "[::]:2587"), in which case peers cannot derive a reachable URL from it. +func wildcardAddr(addr string) bool { + host, _, err := net.SplitHostPort(addr) + if err != nil { + return true // Unparseable -> cannot derive a URL either + } + return host == "" || host == "0.0.0.0" || host == "::" +} diff --git a/cmd/serve_test.go b/cmd/serve_test.go index b89efa8a..50cff723 100644 --- a/cmd/serve_test.go +++ b/cmd/serve_test.go @@ -536,6 +536,41 @@ func TestIP_Host_Parsing(t *testing.T) { } } +func TestCLI_Serve_ClusterValidation(t *testing.T) { + configFile := newEmptyFile(t) // Avoid issues with existing server.yml file on system + // Setting cluster-listen implicitly enables clustering, which requires database-url; all + // validation must fail before any database connection is attempted + app, _, _, _ := newTestApp() + err := app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-listen=127.0.0.1:2587"}) + require.Error(t, err) + require.Contains(t, err.Error(), "database-url") + // cluster-listen requires cluster-secret + app, _, _, _ = newTestApp() + err = app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-listen=127.0.0.1:2587", "--database-url=postgres://user:pass@localhost:1/na"}) + require.Error(t, err) + require.Contains(t, err.Error(), "cluster-secret") + // cluster-listen requires an explicit stable node ID + app, _, _, _ = newTestApp() + err = app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-listen=127.0.0.1:2587", "--database-url=postgres://user:pass@localhost:1/na", "--cluster-secret=s3cret"}) + require.Error(t, err) + require.Contains(t, err.Error(), "cluster-node-id") + // cluster-secret without cluster-listen is a config error (clustering would silently be off) + app, _, _, _ = newTestApp() + err = app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-secret=s3cret"}) + require.Error(t, err) + require.Contains(t, err.Error(), "cluster-listen") + // A wildcard cluster-listen bind cannot derive an advertise URL + app, _, _, _ = newTestApp() + err = app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-listen=:2587", "--database-url=postgres://user:pass@localhost:1/na", "--cluster-secret=s3cret", "--cluster-node-id=node-a"}) + require.Error(t, err) + require.Contains(t, err.Error(), "cluster-advertise-url") + // cluster-batch-linger must not be negative + app, _, _, _ = newTestApp() + err = app.Run([]string{"ntfy", "serve", "--config=" + configFile, "--cluster-batch-linger=-1s"}) + require.Error(t, err) + require.Contains(t, err.Error(), "cluster batch linger") +} + func newEmptyFile(t *testing.T) string { filename := filepath.Join(t.TempDir(), "empty") require.Nil(t, os.WriteFile(filename, []byte{}, 0600)) diff --git a/db/pg/leader.go b/db/pg/leader.go new file mode 100644 index 00000000..8bf88369 --- /dev/null +++ b/db/pg/leader.go @@ -0,0 +1,83 @@ +package pg + +import ( + "context" + "database/sql" + "sync" +) + +const ( + tryAdvisoryLockQuery = `SELECT pg_try_advisory_lock($1)` + advisoryUnlockQuery = `SELECT pg_advisory_unlock($1)` +) + +// Leader implements singleton-job leader election via a Postgres advisory lock held on a pinned +// connection. The lock auto-releases if the holding connection dies, so a crashed leader is +// replaced without manual fencing or lease bookkeeping. Multiple independent leaderships can +// coexist by using distinct lock keys. +type Leader struct { + db *sql.DB + key int64 + conn *sql.Conn // holds the advisory lock while this process is leader + held bool + mu sync.Mutex // Protects conn and held +} + +// NewLeader creates a Leader competing for the advisory lock identified by key. It does not +// attempt to acquire the lock; call TryAcquire periodically. +func NewLeader(db *sql.DB, key int64) *Leader { + return &Leader{db: db, key: key} +} + +// TryAcquire attempts to grab (or confirm) the advisory lock on a pinned connection, without +// blocking. It returns whether this process is the leader after the attempt, i.e. it acquired +// the lock now or still holds it. It is meant to be called periodically; on a healthy leader it +// is a cheap ping, on a follower it retries the lock. +func (l *Leader) TryAcquire(ctx context.Context) bool { + l.mu.Lock() + conn, held := l.conn, l.held + l.mu.Unlock() + if held { + if conn != nil && conn.PingContext(ctx) == nil { + return true // Still leader, connection healthy + } + l.Release() // Connection died; the lock is already gone, re-acquire below + } + newConn, err := l.db.Conn(ctx) + if err != nil { + return false + } + var acquired bool + if err := newConn.QueryRowContext(ctx, tryAdvisoryLockQuery, l.key).Scan(&acquired); err != nil || !acquired { + newConn.Close() + return false + } + l.mu.Lock() + l.conn = newConn + l.held = true + l.mu.Unlock() + return true +} + +// Release unlocks the advisory lock and returns the pinned connection to the pool. +func (l *Leader) Release() { + l.mu.Lock() + conn := l.conn + l.conn = nil + l.held = false + l.mu.Unlock() + if conn != nil { + // Unlock explicitly before returning the connection to the pool: sql.Conn.Close() returns + // the physical connection to the pool rather than closing it, so the session-scoped lock + // would otherwise stay held. + conn.ExecContext(context.Background(), advisoryUnlockQuery, l.key) + conn.Close() + } +} + +// IsLeader reports whether this process currently holds the advisory lock. +func (l *Leader) IsLeader() bool { + l.mu.Lock() + defer l.mu.Unlock() + return l.held +} diff --git a/db/pg/leader_test.go b/db/pg/leader_test.go new file mode 100644 index 00000000..9f5ecbe2 --- /dev/null +++ b/db/pg/leader_test.go @@ -0,0 +1,84 @@ +package pg_test + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/db/pg" + dbtest "heckel.io/ntfy/v2/db/test" +) + +func TestLeader_AcquireAndFailover(t *testing.T) { + testDB := dbtest.CreateTestPostgres(t) // skips if NTFY_TEST_DATABASE_URL is unset + const key = int64(42) + ctx := context.Background() + l1 := pg.NewLeader(testDB.Primary(), key) + l2 := pg.NewLeader(testDB.Primary(), key) + defer l1.Release() + defer l2.Release() + // First to try wins; the second stays follower + require.True(t, l1.TryAcquire(ctx)) + require.False(t, l2.TryAcquire(ctx)) + require.True(t, l1.IsLeader()) + require.False(t, l2.IsLeader()) + // Repeated TryAcquire on the leader is a no-op ping and reports leadership + require.True(t, l1.TryAcquire(ctx)) + // Release -> the follower can take over + l1.Release() + require.False(t, l1.IsLeader()) + require.True(t, l2.TryAcquire(ctx)) + require.True(t, l2.IsLeader()) +} + +func TestLeader_ConnectionLossFailover(t *testing.T) { + // A leader that dies without calling Release (crash, network loss) must not wedge the + // cluster: the advisory lock is session-scoped, so Postgres releases it when the pinned + // connection dies, and a follower can take over. Simulated by terminating the lock-holding + // backend server-side. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + hostA, err := pg.Open(schemaDSN) + require.Nil(t, err) + defer hostA.DB.Close() + hostB, err := pg.Open(schemaDSN) + require.Nil(t, err) + defer hostB.DB.Close() + const key = int64(43) + ctx := context.Background() + l1 := pg.NewLeader(hostA.DB, key) + l2 := pg.NewLeader(hostB.DB, key) + defer l1.Release() + defer l2.Release() + require.True(t, l1.TryAcquire(ctx)) + require.False(t, l2.TryAcquire(ctx)) + // Kill the backend holding the lock (advisory lock keys map to classid/objid) + _, err = hostB.DB.Exec(`SELECT pg_terminate_backend(pid) FROM pg_locks WHERE locktype = 'advisory' AND objid = $1 AND granted`, key) + require.Nil(t, err) + // The lock is auto-released; the follower takes over + waitForCond(t, func() bool { return l2.TryAcquire(ctx) }) + // The old leader discovers its dead connection on the next attempt and stays follower + require.False(t, l1.TryAcquire(ctx)) +} + +func waitForCond(t *testing.T, f func() bool) { + t.Helper() + for i := 0; i < 100; i++ { + if f() { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatal("timed out waiting for condition") +} + +func TestLeader_DistinctKeysAreIndependent(t *testing.T) { + testDB := dbtest.CreateTestPostgres(t) + ctx := context.Background() + l1 := pg.NewLeader(testDB.Primary(), 1) + l2 := pg.NewLeader(testDB.Primary(), 2) + defer l1.Release() + defer l2.Release() + require.True(t, l1.TryAcquire(ctx)) + require.True(t, l2.TryAcquire(ctx)) // Different keys do not compete +} diff --git a/db/pg/pg.go b/db/pg/pg.go index 015910d6..52a96f32 100644 --- a/db/pg/pg.go +++ b/db/pg/pg.go @@ -17,6 +17,7 @@ import ( // ntfy key is defined here, following the "ntfy"+2586+letter scheme const ( SchemaLockKey = int64(0x6e7466792586a) // Schema setup serialization (transaction-scoped, see db/schema) + LeaderLockKey = int64(0x6e7466792586b) // Cluster singleton-job leader (session-scoped, held for process lifetime) ) // Open opens a PostgreSQL connection pool for a primary database. It pings the database diff --git a/db/schema/schema_test.go b/db/schema/schema_test.go index 28e6c376..cd5c38a0 100644 --- a/db/schema/schema_test.go +++ b/db/schema/schema_test.go @@ -14,7 +14,9 @@ import ( _ "github.com/mattn/go-sqlite3" ) -const testCreateQuery = `CREATE TABLE IF NOT EXISTS things (id TEXT PRIMARY KEY, name TEXT NOT NULL)` +const ( + testCreateQuery = `CREATE TABLE IF NOT EXISTS things (id TEXT PRIMARY KEY, name TEXT NOT NULL)` +) func testCreate(tx *sql.Tx) error { _, err := tx.Exec(testCreateQuery) diff --git a/db/test/test.go b/db/test/test.go index 8d3f329b..d9db8f61 100644 --- a/db/test/test.go +++ b/db/test/test.go @@ -17,7 +17,7 @@ const testPoolMaxConns = "2" // CreateTestPostgresSchema creates a temporary PostgreSQL schema and returns the DSN pointing to it. // It registers a cleanup function to drop the schema when the test finishes. // If NTFY_TEST_DATABASE_URL is not set, the test is skipped. -func CreateTestPostgresSchema(t *testing.T) string { +func CreateTestPostgresSchema(t testing.TB) string { t.Helper() dsn := os.Getenv("NTFY_TEST_DATABASE_URL") if dsn == "" { @@ -51,7 +51,7 @@ func CreateTestPostgresSchema(t *testing.T) string { // CreateTestPostgres creates a temporary PostgreSQL schema and returns an open *db.DB connection to it. // It registers cleanup functions to close the DB and drop the schema when the test finishes. // If NTFY_TEST_DATABASE_URL is not set, the test is skipped. -func CreateTestPostgres(t *testing.T) *db.DB { +func CreateTestPostgres(t testing.TB) *db.DB { t.Helper() schemaDSN := CreateTestPostgresSchema(t) testHost, err := pg.Open(schemaDSN) diff --git a/message/cache.go b/message/cache.go index 73aaa076..58aacabc 100644 --- a/message/cache.go +++ b/message/cache.go @@ -35,6 +35,7 @@ type queries struct { selectMessagesSinceIDScheduled string selectMessagesLatest string selectMessagesDue string + selectMessagesDueForUpdate string // Postgres-only: claims due rows via FOR UPDATE SKIP LOCKED; empty for SQLite/mem deleteExpiredMessages string updateMessagePublished string selectMessagesCount string @@ -238,6 +239,15 @@ func (c *Cache) messagesLatest(topic string) ([]*model.Message, error) { // MessagesDue returns all messages that are due for publishing func (c *Cache) MessagesDue() ([]*model.Message, error) { + // On Postgres (cluster mode), claim due rows atomically so that concurrent delayed senders + // on other nodes cannot pick up the same message. We SELECT ... FOR UPDATE SKIP LOCKED and + // mark the claimed rows published in the same transaction; each row is thus handed to exactly + // one node. We deliberately mark published at claim time (not after delivery) to keep the row + // lock short: holding a transaction open across Firebase/WebPush/email delivery would be far + // worse than the small at-most-once window if a node crashes between claim and delivery. + if c.queries.selectMessagesDueForUpdate != "" { + return c.claimMessagesDue() + } rows, err := c.db.Query(c.queries.selectMessagesDue, time.Now().Unix()) if err != nil { return nil, err @@ -245,6 +255,32 @@ func (c *Cache) MessagesDue() ([]*model.Message, error) { return readMessages(rows) } +// claimMessagesDue is the Postgres claiming path for MessagesDue (see its comment). +func (c *Cache) claimMessagesDue() ([]*model.Message, error) { + tx, err := c.db.Begin() + if err != nil { + return nil, err + } + defer tx.Rollback() + rows, err := tx.Query(c.queries.selectMessagesDueForUpdate, time.Now().Unix()) + if err != nil { + return nil, err + } + messages, err := readMessages(rows) // reads all rows and closes them + if err != nil { + return nil, err + } + for _, m := range messages { + if _, err := tx.Exec(c.queries.updateMessagePublished, m.ID); err != nil { + return nil, err + } + } + if err := tx.Commit(); err != nil { + return nil, err + } + return messages, nil +} + // DeleteExpiredMessages deletes up to `limit` expired messages in a single query // and returns the number of deleted rows. func (c *Cache) DeleteExpiredMessages(limit int) (int64, error) { diff --git a/message/cache_postgres.go b/message/cache_postgres.go index e588f8c4..19b67173 100644 --- a/message/cache_postgres.go +++ b/message/cache_postgres.go @@ -61,6 +61,13 @@ const ( WHERE time <= $1 AND published = FALSE ORDER BY time, id ` + postgresSelectMessagesDueForUpdateQuery = ` + SELECT mid, sequence_id, time, event, expires, topic, message, title, priority, tags, click, icon, actions, attachment_name, attachment_type, attachment_size, attachment_expires, attachment_url, sender, user_id, content_type, encoding + FROM message + WHERE time <= $1 AND published = FALSE + ORDER BY time, id + FOR UPDATE SKIP LOCKED + ` postgresUpdateMessagePublishedQuery = `UPDATE message SET published = TRUE WHERE mid = $1` postgresSelectMessagesCountQuery = `SELECT COUNT(*) FROM message` postgresSelectTopicsQuery = `SELECT topic FROM message GROUP BY topic` @@ -88,6 +95,7 @@ var postgresQueries = queries{ selectMessagesSinceIDScheduled: postgresSelectMessagesSinceIDIncludeScheduledQuery, selectMessagesLatest: postgresSelectMessagesLatestQuery, selectMessagesDue: postgresSelectMessagesDueQuery, + selectMessagesDueForUpdate: postgresSelectMessagesDueForUpdateQuery, deleteExpiredMessages: postgresDeleteExpiredMessagesQuery, updateMessagePublished: postgresUpdateMessagePublishedQuery, selectMessagesCount: postgresSelectMessagesCountQuery, diff --git a/message/cache_test.go b/message/cache_test.go index 059a1f62..c169a0bd 100644 --- a/message/cache_test.go +++ b/message/cache_test.go @@ -1,6 +1,7 @@ package message_test import ( + "fmt" "net/netip" "path/filepath" "sync" @@ -556,6 +557,43 @@ func TestStore_MarkPublished(t *testing.T) { }) } +func TestStore_MessagesDue_ClaimExactlyOnce(t *testing.T) { + // Postgres-only: exercises "FOR UPDATE SKIP LOCKED" claiming so that concurrent delayed + // senders running on different cluster nodes never pick up (and deliver) the same due + // message twice. With the non-claiming implementation, every concurrent caller sees every + // due row, so this test fails; with claiming, each row is returned to exactly one caller. + s := newTestPostgresStore(t) // skips if NTFY_TEST_DATABASE_URL is unset + const n = 40 + for i := 0; i < n; i++ { + m := model.NewDefaultMessage("mytopic", fmt.Sprintf("scheduled %d", i)) + m.Time = time.Now().Add(time.Hour).Unix() // future -> stored as published=FALSE + require.Nil(t, s.AddMessage(m)) + // Move the time into the past so the message is due now (but still unpublished) + require.Nil(t, s.UpdateMessageTime(m.ID, time.Now().Add(-time.Minute).Unix())) + } + var mu sync.Mutex + seen := make(map[string]int) + var wg sync.WaitGroup + for c := 0; c < 6; c++ { + wg.Add(1) + go func() { + defer wg.Done() + due, err := s.MessagesDue() + require.Nil(t, err) + mu.Lock() + defer mu.Unlock() + for _, m := range due { + seen[m.ID]++ + } + }() + } + wg.Wait() + require.Len(t, seen, n) // every due message was claimed + for id, count := range seen { + require.Equalf(t, 1, count, "message %s was claimed %d times, want exactly 1", id, count) + } +} + func TestStore_ExpireMessages(t *testing.T) { forEachBackend(t, func(t *testing.T, s *message.Cache) { // Add messages to two topics diff --git a/metrics/metrics.go b/metrics/metrics.go index 1de5903c..00b9741c 100644 --- a/metrics/metrics.go +++ b/metrics/metrics.go @@ -75,6 +75,33 @@ var ( HTTPRequests = prometheus.NewCounterVec(prometheus.CounterOpts{ Name: "ntfy_http_requests_total", }, []string{"http_code", "ntfy_code", "http_method"}) + ClusterPeers = prometheus.NewGauge(prometheus.GaugeOpts{ + Name: "ntfy_cluster_peers", + }) + ClusterMessagesRelayed = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_messages_relayed_total", + }) + ClusterSendErrors = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_send_errors_total", + }) + ClusterQueueDropped = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_queue_dropped_total", + }) + ClusterBatchesSent = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_batches_sent_total", + }) + ClusterMessagesWasted = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_messages_wasted_total", + }) + ClusterRouteSkipped = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_route_skipped_total", + }) + ClusterStatePushes = prometheus.NewCounter(prometheus.CounterOpts{ + Name: "ntfy_cluster_state_pushes_total", + }) + ClusterLeader = prometheus.NewGauge(prometheus.GaugeOpts{ + Name: "ntfy_cluster_leader", + }) ) // init registers all collectors with the default Prometheus registry. Registration is @@ -103,5 +130,14 @@ func init() { Subscribers, Topics, HTTPRequests, + ClusterPeers, + ClusterMessagesRelayed, + ClusterSendErrors, + ClusterQueueDropped, + ClusterBatchesSent, + ClusterMessagesWasted, + ClusterRouteSkipped, + ClusterStatePushes, + ClusterLeader, ) } diff --git a/metrics/metrics_test.go b/metrics/metrics_test.go index 70c322c5..6bc9e223 100644 --- a/metrics/metrics_test.go +++ b/metrics/metrics_test.go @@ -15,6 +15,15 @@ var expectedMetricNames = []string{ "ntfy_attachments_total_size", "ntfy_calls_made_failure", "ntfy_calls_made_success", + "ntfy_cluster_batches_sent_total", + "ntfy_cluster_leader", + "ntfy_cluster_messages_relayed_total", + "ntfy_cluster_messages_wasted_total", + "ntfy_cluster_peers", + "ntfy_cluster_queue_dropped_total", + "ntfy_cluster_route_skipped_total", + "ntfy_cluster_send_errors_total", + "ntfy_cluster_state_pushes_total", "ntfy_emails_received_failure", "ntfy_emails_received_success", "ntfy_emails_sent_failure", diff --git a/server/config.go b/server/config.go index b2c169a4..46cae5fb 100644 --- a/server/config.go +++ b/server/config.go @@ -10,6 +10,8 @@ import ( "text/template" "time" + "heckel.io/ntfy/v2/cluster" + "heckel.io/ntfy/v2/ban" "heckel.io/ntfy/v2/user" ) @@ -119,8 +121,13 @@ type Config struct { ListenUnixMode fs.FileMode KeyFile string CertFile string - DatabaseURL string // PostgreSQL connection string (e.g. "postgres://user:pass@host:5432/ntfy") - DatabaseReplicaURLs []string // PostgreSQL read replica connection strings + DatabaseURL string // PostgreSQL connection string (e.g. "postgres://user:pass@host:5432/ntfy") + DatabaseReplicaURLs []string // PostgreSQL read replica connection strings + ClusterNodeID string // Stable per-node identifier used to skip a node's own fan-out; required in cluster mode + ClusterListen string // ip:port the dedicated cluster fan-out listener binds to (private network, e.g. "10.0.0.5:2587") + ClusterAdvertiseURL string // Base URL peers use to reach this node's fan-out listener (defaults to "http://") + ClusterSecret string `hash:"-"` // Shared secret authenticating node-to-node fan-out requests + ClusterBatchLinger time.Duration // How long fan-out messages wait to form a batch per peer; 0 sends immediately FirebaseKeyFile string CacheFile string CacheDuration time.Duration @@ -236,6 +243,11 @@ func NewConfig() *Config { KeyFile: "", CertFile: "", DatabaseURL: "", + ClusterNodeID: "", + ClusterListen: "", + ClusterAdvertiseURL: "", + ClusterSecret: "", + ClusterBatchLinger: cluster.DefaultBatchLinger, FirebaseKeyFile: "", CacheFile: "", CacheDuration: DefaultCacheDuration, diff --git a/server/config_test.go b/server/config_test.go index e789d55d..c4d36172 100644 --- a/server/config_test.go +++ b/server/config_test.go @@ -25,6 +25,7 @@ func TestConfig_HashExcludesSecrets(t *testing.T) { conf2.UpstreamAccessToken = "tk_upstream" conf2.WebPushPrivateKey = "web-push-private-key" conf2.SMTPSenderPass = "hunter2" + conf2.ClusterSecret = "cluster-secret" conf2.AuthUsers = []*user.User{{Name: "phil", Hash: "$2a$10$somebcrypthash"}} conf2.AuthTokens = map[string][]*user.Token{"phil": {{Value: "tk_secrettoken"}}} assert.Equal(t, conf1.Hash(), conf2.Hash()) diff --git a/server/server.go b/server/server.go index dca96747..27b8b1b1 100644 --- a/server/server.go +++ b/server/server.go @@ -32,6 +32,7 @@ import ( "heckel.io/ntfy/v2/action" "heckel.io/ntfy/v2/attachment" "heckel.io/ntfy/v2/ban" + "heckel.io/ntfy/v2/cluster" "heckel.io/ntfy/v2/db" "heckel.io/ntfy/v2/db/pg" "heckel.io/ntfy/v2/log" @@ -72,6 +73,8 @@ type Server struct { stripe stripeAPI // Stripe API, can be replaced with a mock priceCache *util.LookupCache[map[string]int64] // Stripe price ID -> price as cents (USD implied!) metricsHandler http.Handler // Handles /metrics if enable-metrics set, and listen-metrics-http not set + cluster cluster.Cluster // Fans messages out to peer cluster nodes (nop when not clustered) + httpClusterServer *http.Server // Dedicated private listener for node-to-node fan-out (cluster-listen) closeChan chan bool mu sync.RWMutex } @@ -334,9 +337,90 @@ func New(conf *Config) (*Server, error) { stripe: stripe, } s.priceCache = util.NewLookupCache(s.fetchStripePrices, conf.StripePriceCacheDuration) + // Cross-node cluster; delivery of peer messages to local subscribers is injected + // as a callback, so the cluster package never depends on the server. Peers talk to each + // other only via the dedicated cluster listener, never via the public listeners. + advertiseURL := conf.ClusterAdvertiseURL + if advertiseURL == "" && conf.ClusterListen != "" { + advertiseURL = "http://" + conf.ClusterListen + } + s.cluster, err = cluster.New(&cluster.Config{ + Enabled: conf.ClusterListen != "", // Setting cluster-listen implicitly enables clustering + NodeID: cluster.NodeID(conf.ClusterNodeID), + AdvertiseURL: advertiseURL, + Secret: conf.ClusterSecret, + BatchLinger: conf.ClusterBatchLinger, + MaxMessageBytes: int64(conf.MessageSizeLimit)*4 + 1024, // Envelope overhead over the raw message + }, pool, s.deliverFromBus, s.liveTopics) + if err != nil { + return nil, err + } + // The cluster routes messages by subscription knowledge: peers learn this node's live topics + // via periodic state pushes (liveTopics) and immediate announcements on a topic's first + // subscriber (the hook below; also set in topicsFromIDs for topics created later) + for _, t := range s.topics { + t.onFirstSubscriber = s.topicAnnouncer(t.ID) + } return s, nil } +// clusterHandler returns the handler served on the dedicated cluster listener (cluster-listen). +// It serves the internal peer API (owned and routed by the cluster itself, including auth) plus +// a health endpoint; the public listeners never expose these paths, so internal traffic cannot +// be reached from the outside even before any firewalling. +func (s *Server) clusterHandler() http.Handler { + mux := http.NewServeMux() + // TODO(T8): reflect s.cluster.Healthy() here (and in handleHealth) instead of a static + // response, so health checks can pull a node that lost its registry heartbeat + mux.HandleFunc(apiHealthPath, func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + io.WriteString(w, `{"healthy":true}`+"\n") + }) + mux.Handle("/", s.cluster) + return mux +} + +// liveTopics returns the topics that currently have at least one subscriber, computed fresh on +// every call. Deliberately NOT the whole topics map: it also holds subscriber-less topics +// rebuilt from the message cache, which would gut the routing filter's selectivity. +func (s *Server) liveTopics() []string { + s.mu.RLock() + defer s.mu.RUnlock() + topics := make([]string, 0, len(s.topics)) + for _, t := range s.topics { + if subscribers, _ := t.Stats(); subscribers > 0 { + topics = append(topics, t.ID) + } + } + return topics +} + +// topicAnnouncer returns the first-subscriber hook for a topic: it tells peer nodes right away +// that this node now wants messages for it (see Cluster.AnnounceTopics). +func (s *Server) topicAnnouncer(id string) func() { + return func() { + s.cluster.AnnounceTopics([]string{id}) + } +} + +// deliverFromBus delivers a message received from a peer node (via the cluster) to this +// node's local subscribers. It is the receive-side counterpart to Cluster.Relay: local +// delivery and all global side effects (Firebase, email, web push, upstream) already ran on the +// origin node, so this only publishes to the local topic, and never re-relays. +func (s *Server) deliverFromBus(m *model.Message) { + s.mu.RLock() + t, ok := s.topics[m.Topic] + s.mu.RUnlock() + if !ok { + metrics.ClusterMessagesWasted.Inc() // Relayed here needlessly: this node had no use for the message + return + } + v := s.visitor(m.Sender, nil) + if err := t.Publish(v, m); err != nil { + logvm(v, m).Err(err).Warn("Cluster: unable to deliver fan-out message to local subscribers") + } +} + func createMessageCache(conf *Config, pool *db.DB) (*message.Cache, error) { if conf.CacheDuration == 0 { return message.NewNopStore() @@ -379,6 +463,9 @@ func (s *Server) Run() error { if s.config.ProfileListenHTTP != "" { listenStr += fmt.Sprintf(" %s[http/profile]", s.config.ProfileListenHTTP) } + if s.config.ClusterListen != "" { + listenStr += fmt.Sprintf(" %s[http/cluster]", s.config.ClusterListen) + } log.Tag(tagStartup).Info("Listening on%s, ntfy %s, log level is %s", listenStr, s.config.BuildVersion, log.CurrentLevel().String()) if log.IsFile() { fmt.Fprintf(os.Stderr, "Listening on%s, ntfy %s\n", listenStr, s.config.BuildVersion) @@ -433,6 +520,12 @@ func (s *Server) Run() error { } else if s.config.EnableMetrics { s.metricsHandler = promhttp.Handler() } + if s.config.ClusterListen != "" { + s.httpClusterServer = &http.Server{Addr: s.config.ClusterListen, Handler: s.clusterHandler()} + go func() { + errChan <- s.httpClusterServer.ListenAndServe() + }() + } if s.config.ProfileListenHTTP != "" { profileMux := http.NewServeMux() profileMux.HandleFunc("/debug/pprof/", pprof.Index) @@ -478,6 +571,12 @@ func (s *Server) Stop() { if s.attachment != nil { s.attachment.Close() } + if s.httpClusterServer != nil { + s.httpClusterServer.Close() + } + if s.cluster != nil { + s.cluster.Close() + } s.closeDatabases() if s.ban != nil { s.ban.Close() @@ -730,6 +829,8 @@ func (s *Server) handleTopicAuth(w http.ResponseWriter, _ *http.Request, _ *visi } func (s *Server) handleHealth(w http.ResponseWriter, _ *http.Request, _ *visitor) error { + // TODO(T8): in cluster mode, reflect s.cluster.Healthy() so DNS/LB health checks stop + // routing NEW clients to a node whose peers no longer relay to it (see Cluster interface) response := &apiHealthResponse{ Healthy: true, } @@ -833,6 +934,60 @@ func (s *Server) handleMatrixDiscovery(w http.ResponseWriter) error { return writeMatrixDiscoveryResponse(w) } +// dispatchOpts selects which delivery targets fire for a published message, beyond delivery to +// local subscribers and the cross-node broadcast (which always happen). +type dispatchOpts struct { + firebase bool // Send to Firebase (if configured) + email string // Send an email to this address (if a mailer is configured) + call string // Call this phone number (if Twilio is configured) + upstream bool // Forward a poll request to the upstream server (if configured) + webPush bool // Publish to web push endpoints (if configured) + async bool // Deliver to local subscribers in a goroutine, logging errors instead of returning them +} + +// dispatch delivers m to local subscribers, relays it to peer cluster nodes, and fires the +// requested side-effect targets. It is the single choke point through which every published +// message must pass; t may be nil when the topic has no local subscribers (delayed sender). +// +// The side-effect targets (Firebase, email, calls, upstream, web push) are global, not per-node: +// they fire only here on the origin node, and never again when a peer node receives the message +// via the cluster (see deliverFromBus). +func (s *Server) dispatch(v *visitor, t *topic, m *model.Message, opts dispatchOpts) error { + // Deliver to local subscribers + if t != nil { + if opts.async { + go func() { + if err := t.Publish(v, m); err != nil { + logvm(v, m).Err(err).Warn("Unable to publish message") + } + }() + } else if err := t.Publish(v, m); err != nil { + return err + } + } + // Relay to peer cluster nodes, whose subscribers do not show up in this node's topics map + if err := s.cluster.Relay(m); err != nil { + logvm(v, m).Err(err).Warn("Cluster: unable to relay message to peer nodes") + } + // Fire the requested side-effect targets + if s.firebaseClient != nil && opts.firebase { + go s.sendToFirebase(v, m) + } + if s.mailer != nil && opts.email != "" { + go s.sendEmail(v, m, opts.email) + } + if s.config.TwilioAccount != "" && opts.call != "" { + go s.callPhone(v, m, opts.call) + } + if s.config.UpstreamBaseURL != "" && opts.upstream { + go s.forwardPollRequest(v, m) + } + if s.config.WebPushPublicKey != "" && opts.webPush { + go s.publishToWebPushEndpoints(v, m) + } + return nil +} + func (s *Server) handlePublishInternal(r *http.Request, v *visitor) (*model.Message, error) { start := time.Now() t, err := fromContext[*topic](r, contextTopic) @@ -911,24 +1066,16 @@ func (s *Server) handlePublishInternal(r *http.Request, v *visitor) (*model.Mess ev.Debug("Received message") } if !delayed { - if err := t.Publish(v, m); err != nil { + err := s.dispatch(v, t, m, dispatchOpts{ + firebase: firebase, + email: email, + call: call, + upstream: !unifiedpush, // UP messages are not sent to upstream + webPush: true, + }) + if err != nil { return nil, err } - if s.firebaseClient != nil && firebase { - go s.sendToFirebase(v, m) - } - if s.mailer != nil && email != "" { - go s.sendEmail(v, m, email) - } - if s.config.TwilioAccount != "" && call != "" { - go s.callPhone(v, m, call) - } - if s.config.UpstreamBaseURL != "" && !unifiedpush { // UP messages are not sent to upstream - go s.forwardPollRequest(v, m) - } - if s.config.WebPushPublicKey != "" { - go s.publishToWebPushEndpoints(v, m) - } } else { logvrm(v, r, m).Tag(tagPublish).Debug("Message delayed, will process later") } @@ -1027,18 +1174,10 @@ func (s *Server) handleActionMessage(w http.ResponseWriter, r *http.Request, v * m.Sender = v.IP() m.User = v.MaybeUserID() m.Expires = time.Unix(m.Time, 0).Add(v.Limits().MessageExpiryDuration).Unix() - // Publish to subscribers - if err := t.Publish(v, m); err != nil { + // Publish to subscribers, peer nodes, Firebase (for Android clients), and web push endpoints + if err := s.dispatch(v, t, m, dispatchOpts{firebase: true, webPush: true}); err != nil { return err } - // Send to Firebase for Android clients - if s.firebaseClient != nil { - go s.sendToFirebase(v, m) - } - // Send to web push endpoints - if s.config.WebPushPublicKey != "" { - go s.publishToWebPushEndpoints(v, m) - } if event == model.MessageDeleteEvent { // Delete any existing scheduled message with the same sequence ID deletedIDs, err := s.messageCache.DeleteScheduledBySequenceID(t.ID, sequenceID) @@ -1832,7 +1971,9 @@ func (s *Server) topicsFromIDs(v *visitor, ids ...string) ([]*topic, error) { if v != nil && !v.TopicCreationAllowed() { return nil, errHTTPTooManyRequestsLimitTopicCreation } - s.topics[id] = newTopic(id) + t := newTopic(id) + t.onFirstSubscriber = s.topicAnnouncer(id) + s.topics[id] = t } topics = append(topics, s.topics[id]) } @@ -1919,7 +2060,8 @@ func (s *Server) resetStats() { for _, v := range s.visitors { v.ResetStats() } - if s.userManager != nil { + // The user database is shared; only the cluster leader resets it (always true single-node) + if s.userManager != nil && s.cluster.IsLeader() { if err := s.userManager.ResetStats(); err != nil { log.Tag(tagResetter).Warn("Failed to write to database: %s", err.Error()) } @@ -1934,7 +2076,11 @@ func (s *Server) runFirebaseKeepaliver() { for { select { case <-time.After(s.config.FirebaseKeepaliveInterval): - s.sendToFirebase(v, model.NewKeepaliveMessage(firebaseControlTopic)) + // Leader only: every FCM keepalive wakes all subscribed phones, so a cluster must + // send it exactly once, not once per node (checked per tick to survive failover) + if s.cluster.IsLeader() { + s.sendToFirebase(v, model.NewKeepaliveMessage(firebaseControlTopic)) + } /* FIXME: Disable iOS polling entirely for now due to thundering herd problem (see #677) To solve this, we'd have to shard the iOS poll topics to spread out the polling evenly. @@ -1987,24 +2133,18 @@ func (s *Server) sendDelayedMessages() error { func (s *Server) sendDelayedMessage(v *visitor, m *model.Message) error { logvm(v, m).Debug("Sending delayed message") s.mu.RLock() - t, ok := s.topics[m.Topic] // If no subscribers, just mark message as published + t := s.topics[m.Topic] // May be nil if there are no local subscribers; dispatch handles that s.mu.RUnlock() - if ok { - go func() { - // We do not rate-limit messages here, since we've rate limited them in the PUT/POST handler - if err := t.Publish(v, m); err != nil { - logvm(v, m).Err(err).Warn("Unable to publish message") - } - }() - } - if s.firebaseClient != nil { // Firebase subscribers may not show up in topics map - go s.sendToFirebase(v, m) - } - if s.config.UpstreamBaseURL != "" { - go s.forwardPollRequest(v, m) - } - if s.config.WebPushPublicKey != "" { - go s.publishToWebPushEndpoints(v, m) + // We do not rate-limit messages here, since we've rate limited them in the PUT/POST handler. + // Firebase subscribers may not show up in the topics map, so side effects fire regardless. + err := s.dispatch(v, t, m, dispatchOpts{ + firebase: true, + upstream: true, + webPush: true, + async: true, + }) + if err != nil { + return err } if err := s.messageCache.MarkPublished(m); err != nil { return err diff --git a/server/server.yml b/server/server.yml index 4b3acc46..45d0160a 100644 --- a/server/server.yml +++ b/server/server.yml @@ -61,6 +61,34 @@ # # database-url: +# If "cluster-listen" is set, clustering is implicitly enabled: this node registers itself in +# the PostgreSQL node registry and fans published messages out to the other cluster nodes over +# HTTP, so subscribers connected to any node receive messages published to any other node. +# Requires "database-url" and "cluster-secret". +# +# - cluster-listen is the ip:port of the dedicated fan-out listener that peer nodes talk to. +# Bind it to a private network interface (e.g. "10.0.0.5:2587"); the public listeners never +# serve the fan-out endpoint. +# - cluster-node-id is a stable per-node identifier (e.g. the hostname); required. +# - cluster-advertise-url is the base URL peer nodes use to reach this node's fan-out listener; +# it defaults to "http://". It must be set explicitly if cluster-listen binds +# a wildcard address (e.g. ":2587"). +# - cluster-secret authenticates node-to-node fan-out requests; it must be identical on all +# nodes. +# - cluster-batch-linger is how long fan-out messages may wait to form a batch per peer node, +# trading up to that much cross-node delivery latency for a bounded request rate between +# nodes. Set to 0 to send each message immediately. +# +# SECURITY: The fan-out endpoint injects messages into arbitrary topics. The shared secret +# protects it, and it is only served on the dedicated cluster listener -- but you should still +# make sure that listener is reachable only from the private network (firewall/VPC rules). +# +# cluster-listen: +# cluster-node-id: +# cluster-advertise-url: "http://" +# cluster-secret: +# cluster-batch-linger: 500ms + # If "cache-file" is set, messages are cached in a local SQLite database instead of only in-memory. # This allows for service restarts without losing messages in support of the since= parameter. # Not required if "database-url" is set (messages are stored in PostgreSQL instead). diff --git a/server/server_account.go b/server/server_account.go index d53313a2..cfd57af1 100644 --- a/server/server_account.go +++ b/server/server_account.go @@ -985,7 +985,8 @@ func (s *Server) publishSyncEventForUser(v *visitor, u *user.User) error { return err } m := model.NewDefaultMessage(syncTopic.ID, string(messageBytes)) - if err := syncTopic.Publish(v, m); err != nil { + // Dispatch so the sync event also reaches the user's devices connected to peer cluster nodes + if err := s.dispatch(v, syncTopic, m, dispatchOpts{}); err != nil { return err } return nil diff --git a/server/server_cluster_test.go b/server/server_cluster_test.go new file mode 100644 index 00000000..0daf20cf --- /dev/null +++ b/server/server_cluster_test.go @@ -0,0 +1,318 @@ +package server + +import ( + "database/sql" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + dbtest "heckel.io/ntfy/v2/db/test" + "heckel.io/ntfy/v2/model" + "heckel.io/ntfy/v2/user" +) + +// fakeCluster records relayed messages and topic announcements so tests can assert that every +// publish path passes through the cluster exactly once, and that subscription hooks fire. +type fakeCluster struct { + mu sync.Mutex + messages []*model.Message + announced []string + notLeader bool +} + +func (b *fakeCluster) Relay(m *model.Message) error { + b.mu.Lock() + defer b.mu.Unlock() + b.messages = append(b.messages, m) + return nil +} + +func (b *fakeCluster) ServeHTTP(_ http.ResponseWriter, _ *http.Request) {} + +func (b *fakeCluster) AnnounceTopics(topics []string) { + b.mu.Lock() + defer b.mu.Unlock() + b.announced = append(b.announced, topics...) +} + +func (b *fakeCluster) IsLeader() bool { + b.mu.Lock() + defer b.mu.Unlock() + return !b.notLeader +} + +func (b *fakeCluster) setLeader(leader bool) { + b.mu.Lock() + defer b.mu.Unlock() + b.notLeader = !leader +} + +func (b *fakeCluster) Close() error { return nil } + +func (b *fakeCluster) Messages() []*model.Message { + b.mu.Lock() + defer b.mu.Unlock() + return append([]*model.Message{}, b.messages...) +} + +func (b *fakeCluster) Announced() []string { + b.mu.Lock() + defer b.mu.Unlock() + return append([]string{}, b.announced...) +} + +func TestServer_Cluster_PublishRelaysOnce(t *testing.T) { + s := newTestServer(t, newTestConfig(t, "")) + b := &fakeCluster{} + s.cluster = b + response := request(t, s, "PUT", "/mytopic", "hi there", nil) + require.Equal(t, 200, response.Code) + messages := b.Messages() + require.Len(t, messages, 1) + require.Equal(t, "mytopic", messages[0].Topic) + require.Equal(t, "hi there", messages[0].Message) +} + +func TestServer_Cluster_SyncEventRelays(t *testing.T) { + // Account sync events are delivered via the user's st_... sync topic; without relaying + // them, cross-device account sync silently breaks when a user's devices land on different + // cluster nodes. + s := newTestServer(t, newTestConfig(t, "")) + b := &fakeCluster{} + s.cluster = b + u := &user.User{ID: "u_abc", Name: "phil", SyncTopic: "st_1234"} + v := s.visitor(netip.MustParseAddr("1.2.3.4"), nil) + require.Nil(t, s.publishSyncEventForUser(v, u)) + messages := b.Messages() + require.Len(t, messages, 1) + require.Equal(t, "st_1234", messages[0].Topic) +} + +func TestServer_Cluster_DeliverNotOnPublicHandler(t *testing.T) { + // The fan-out endpoint lives only on the dedicated cluster listener; the public handler must + // not serve it, even with cluster mode on and a valid secret. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + conf := newTestConfig(t, schemaDSN) + conf.ClusterNodeID = "node-a" + conf.ClusterListen = "127.0.0.1:1" // Enables clustering; not bound since Run() is not called + conf.ClusterSecret = "s3cret" + conf.ClusterAdvertiseURL = "http://127.0.0.1:1" + s := newTestServer(t, conf) + topics, err := s.topicsFromIDs(nil, "mytopic") + require.Nil(t, err) + var mu sync.Mutex + var received []*model.Message + topics[0].Subscribe(func(_ *visitor, m *model.Message) error { + mu.Lock() + defer mu.Unlock() + received = append(received, m) + return nil + }, "", func() {}) + // A valid fan-out request against the PUBLIC handler must not deliver + response := request(t, s, "POST", "/v1/internal/message", + `{"message":{"id":"x1","time":1,"event":"message","topic":"mytopic","message":"sneaky"}}`, + map[string]string{"X-Cluster-Secret": "s3cret", "X-Cluster-Origin": "node-b"}) + require.Equal(t, 404, response.Code) + time.Sleep(250 * time.Millisecond) // Delivery is async; give a wrong implementation time to fail + mu.Lock() + require.Empty(t, received) + mu.Unlock() + // The same request against the cluster listener handler DOES deliver + rr := httptest.NewRecorder() + req, err := http.NewRequest("POST", "/v1/internal/message", + strings.NewReader(`{"message":{"id":"x2","time":1,"event":"message","topic":"mytopic","message":"legit"}}`)) + require.Nil(t, err) + req.Header.Set("X-Cluster-Secret", "s3cret") + req.Header.Set("X-Cluster-Origin", "node-b") + s.clusterHandler().ServeHTTP(rr, req) + require.Equal(t, 200, rr.Code) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(received) == 1 + }) +} + +func TestServer_Cluster_EndToEnd(t *testing.T) { + // Two full servers sharing one Postgres schema: a message published to node A over HTTP must + // reach a subscriber connected to node B, via the node registry and the fan-out endpoint. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + // Node B: create the listener first so its advertise URL is known before the server exists + listenerB, err := net.Listen("tcp", "127.0.0.1:0") + require.Nil(t, err) + confB := newTestConfig(t, schemaDSN) + confB.ClusterNodeID = "node-b" + confB.ClusterListen = listenerB.Addr().String() // Enables clustering; the test serves it below + confB.ClusterSecret = "s3cret" + confB.ClusterAdvertiseURL = "http://" + listenerB.Addr().String() + sB := newTestServer(t, confB) + srvB := &http.Server{Handler: sB.clusterHandler()} + go srvB.Serve(listenerB) + defer srvB.Close() + // Node A: publish-only in this test, so its advertise URL is never called + confA := newTestConfig(t, schemaDSN) + confA.ClusterNodeID = "node-a" + confA.ClusterListen = "127.0.0.1:1" // Enables clustering; not bound since Run() is not called + confA.ClusterSecret = "s3cret" + confA.ClusterAdvertiseURL = "http://127.0.0.1:1" + sA := newTestServer(t, confA) + // Subscribe on node B + topics, err := sB.topicsFromIDs(nil, "mytopic") + require.Nil(t, err) + var mu sync.Mutex + var received []*model.Message + topics[0].Subscribe(func(_ *visitor, m *model.Message) error { + mu.Lock() + defer mu.Unlock() + received = append(received, m) + return nil + }, "", func() {}) + // Publish on node A + response := request(t, sA, "PUT", "/mytopic", "hello cluster", nil) + require.Equal(t, 200, response.Code) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(received) == 1 + }) + mu.Lock() + defer mu.Unlock() + require.Equal(t, "hello cluster", received[0].Message) +} + +func TestServer_Cluster_DeliverFromBus(t *testing.T) { + // deliverFromBus is the receive side of the broadcaster: a message that originated on a peer + // node must reach this node's local subscribers, but must NOT be re-broadcast (loop) nor + // re-trigger origin-only side effects. + s := newTestServer(t, newTestConfig(t, "")) + b := &fakeCluster{} + s.cluster = b + topics, err := s.topicsFromIDs(nil, "mytopic") + require.Nil(t, err) + var mu sync.Mutex + var received []*model.Message + topics[0].Subscribe(func(_ *visitor, m *model.Message) error { + mu.Lock() + defer mu.Unlock() + received = append(received, m) + return nil + }, "", func() {}) + m := model.NewDefaultMessage("mytopic", "from peer") + m.Sender = netip.MustParseAddr("5.6.7.8") + s.deliverFromBus(m) + waitFor(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(received) == 1 + }) + require.Empty(t, b.Messages()) // Peer messages are never re-relayed +} + +func TestServer_Cluster_FirstSubscriberAnnounces(t *testing.T) { + // A topic gaining its FIRST subscriber is announced to peers exactly once, so publishers on + // other nodes stop skipping this node for it without waiting for the next state push. + s := newTestServer(t, newTestConfig(t, "")) + b := &fakeCluster{} + s.cluster = b + topics, err := s.topicsFromIDs(nil, "mytopic") + require.Nil(t, err) + subscriber := func(_ *visitor, _ *model.Message) error { return nil } + topics[0].Subscribe(subscriber, "", func() {}) + waitFor(t, func() bool { + return len(b.Announced()) == 1 && b.Announced()[0] == "mytopic" + }) + // A second subscriber does not re-announce + topics[0].Subscribe(subscriber, "", func() {}) + time.Sleep(250 * time.Millisecond) + require.Len(t, b.Announced(), 1) +} + +func TestServer_Cluster_ManagerPrunesOnlyOnLeader(t *testing.T) { + c := newTestConfig(t, "") + s := newTestServer(t, c) + cl := &fakeCluster{notLeader: true} + s.cluster = cl + + // Publish and expire a message + rr := request(t, s, "POST", "/mytopic", "hi", nil) + require.Equal(t, 200, rr.Code) + m := toMessage(t, rr.Body.String()) + require.Nil(t, s.messageCache.ExpireMessages("mytopic")) + + // A non-leader node leaves shared-database pruning to the leader + s.execManager() + _, err := s.messageCache.Message(m.ID) + require.Nil(t, err) + + // Once this node is the leader, the same run prunes + cl.setLeader(true) + s.execManager() + _, err = s.messageCache.Message(m.ID) + require.Equal(t, model.ErrMessageNotFound, err) +} + +func TestServer_Cluster_StatsResetOnlyOnLeader(t *testing.T) { + c := newTestConfigWithAuthFile(t, "") + s := newTestServer(t, c) + cl := &fakeCluster{notLeader: true} + s.cluster = cl + + // An anonymous visitor with an in-memory message count + v := newVisitor(c, s.messageCache, s.userManager, netip.MustParseAddr("1.2.3.4"), nil) + require.True(t, v.MessageAllowed()) + s.mu.Lock() + s.visitors["ip:1.2.3.4"] = v + s.mu.Unlock() + require.Equal(t, int64(1), v.Stats().Messages) + + // A user with persisted stats in the (shared) user database + require.Nil(t, s.userManager.AddUser("phil", "phil1234", user.RoleUser, false)) + authDB, err := sql.Open("sqlite3", c.AuthFile) + require.Nil(t, err) + defer authDB.Close() + _, err = authDB.Exec(`UPDATE user SET stats_messages = 5 WHERE user = 'phil'`) + require.Nil(t, err) + + // A non-leader node resets its own in-memory visitor stats, but leaves the user database + // to the leader + s.resetStats() + require.Equal(t, int64(0), v.Stats().Messages) + u, err := s.userManager.User("phil") + require.Nil(t, err) + require.Equal(t, int64(5), u.Stats.Messages) + + // The leader resets the user database too + cl.setLeader(true) + s.resetStats() + u, err = s.userManager.User("phil") + require.Nil(t, err) + require.Equal(t, int64(0), u.Stats.Messages) +} + +func TestServer_Cluster_FirebaseKeepaliverOnlyOnLeader(t *testing.T) { + // Every FCM keepalive wakes all subscribed phones, so only the leader may send them; + // N nodes sending N keepalives would multiply the battery cost for every user + c := newTestConfig(t, "") + c.FirebaseKeepaliveInterval = 20 * time.Millisecond + s := newTestServer(t, c) + sender := newTestFirebaseSender(100) + s.firebaseClient = newFirebaseClient(sender, &testAuther{Allow: true}) + cl := &fakeCluster{notLeader: true} + s.cluster = cl + s.closeChan = make(chan bool) // Closed by Stop() in the test cleanup + go s.runFirebaseKeepaliver() + + // A non-leader node stays silent + time.Sleep(150 * time.Millisecond) + require.Empty(t, sender.Messages()) + + // The leader sends keepalives + cl.setLeader(true) + waitFor(t, func() bool { return len(sender.Messages()) > 0 }) +} diff --git a/server/server_manager.go b/server/server_manager.go index 204f5fbe..e15883af 100644 --- a/server/server_manager.go +++ b/server/server_manager.go @@ -10,12 +10,16 @@ func (s *Server) execManager() { // WARNING: Make sure to only selectively lock with the mutex, and be aware that this // there is no mutex for the entire function. - // Prune all the things + // Prune all the things. In-memory state is pruned on every node; jobs touching shared + // databases (and the web push job, which also sends expiry-warning notifications) run on + // the cluster leader only. In a single-node setup, IsLeader is always true. s.pruneVisitors() - s.pruneTokens() - s.pruneAttachments() - s.pruneMessages() - s.pruneAndNotifyWebPushSubscriptions() + if s.cluster.IsLeader() { + s.pruneTokens() + s.pruneAttachments() + s.pruneMessages() + s.pruneAndNotifyWebPushSubscriptions() + } // Message count messagesCached, err := s.messageCache.MessagesCount() diff --git a/server/topic.go b/server/topic.go index f373a5e6..54a9d24f 100644 --- a/server/topic.go +++ b/server/topic.go @@ -20,11 +20,12 @@ const ( // topic represents a channel to which subscribers can subscribe, and publishers // can publish a message type topic struct { - ID string - subscribers map[int]*topicSubscriber - rateVisitor *visitor - lastAccess time.Time - mu sync.RWMutex + ID string + subscribers map[int]*topicSubscriber + rateVisitor *visitor + lastAccess time.Time + onFirstSubscriber func() // Fired (async) when the subscriber count goes 0 -> 1; may be nil + mu sync.RWMutex } type topicSubscriber struct { @@ -56,6 +57,10 @@ func (t *topic) Subscribe(s subscriber, userID string, cancel func()) (subscribe break } } + if len(t.subscribers) == 0 && t.onFirstSubscriber != nil { + // Fired async so cluster announcements never run under the topic lock + go t.onFirstSubscriber() + } t.subscribers[subscriberID] = &topicSubscriber{ userID: userID, // May be empty subscriber: s, diff --git a/util/bloom.go b/util/bloom.go new file mode 100644 index 00000000..b48a3a90 --- /dev/null +++ b/util/bloom.go @@ -0,0 +1,103 @@ +package util + +import ( + "errors" + "hash/fnv" + "math" +) + +// BloomFilter is a fixed-size probabilistic set: Contains may return false positives (bounded by +// the target rate the filter was sized for) but never false negatives. Elements cannot be +// removed; rebuild the filter from scratch instead. +type BloomFilter struct { + bits []uint64 + k int // number of hash probes per element, derived via double hashing +} + +// NewBloomFilter creates a filter sized for n elements at the given false-positive rate. +func NewBloomFilter(n int, fpRate float64) *BloomFilter { + if n < 1 { + n = 1 + } + if fpRate <= 0 || fpRate >= 1 { + fpRate = 0.01 + } + m := int(math.Ceil(-float64(n) * math.Log(fpRate) / (math.Ln2 * math.Ln2))) // bits + k := int(math.Round(float64(m) / float64(n) * math.Ln2)) // probes + if k < 1 { + k = 1 + } + return &BloomFilter{ + bits: make([]uint64, (m+63)/64), + k: k, + } +} + +// Add inserts an element into the filter. +func (b *BloomFilter) Add(s string) { + h1, h2 := hashPair(s) + m := uint64(len(b.bits)) * 64 + for i := 0; i < b.k; i++ { + bit := (h1 + uint64(i)*h2) % m + b.bits[bit/64] |= 1 << (bit % 64) + } +} + +// Contains reports whether the element may be in the set. A false result is definitive: the +// element was never added. +func (b *BloomFilter) Contains(s string) bool { + h1, h2 := hashPair(s) + m := uint64(len(b.bits)) * 64 + for i := 0; i < b.k; i++ { + bit := (h1 + uint64(i)*h2) % m + if b.bits[bit/64]&(1<<(bit%64)) == 0 { + return false + } + } + return true +} + +// MarshalBinary serializes the filter as [k][bits...], with 64-bit little-endian words. +func (b *BloomFilter) MarshalBinary() ([]byte, error) { + data := make([]byte, 1+len(b.bits)*8) + data[0] = byte(b.k) + for i, word := range b.bits { + for j := 0; j < 8; j++ { + data[1+i*8+j] = byte(word >> (8 * j)) + } + } + return data, nil +} + +// UnmarshalBloomFilter deserializes a filter produced by MarshalBinary. +func UnmarshalBloomFilter(data []byte) (*BloomFilter, error) { + if len(data) < 9 || (len(data)-1)%8 != 0 { + return nil, errors.New("invalid bloom filter data") + } + b := &BloomFilter{ + bits: make([]uint64, (len(data)-1)/8), + k: int(data[0]), + } + if b.k < 1 { + return nil, errors.New("invalid bloom filter hash count") + } + for i := range b.bits { + var word uint64 + for j := 0; j < 8; j++ { + word |= uint64(data[1+i*8+j]) << (8 * j) + } + b.bits[i] = word + } + return b, nil +} + +// hashPair derives the two independent hash values used for double hashing (probe i uses +// h1 + i*h2), from FNV-1a over the element and a domain-separated variant of it. +func hashPair(s string) (uint64, uint64) { + f := fnv.New64a() + f.Write([]byte(s)) + h1 := f.Sum64() + f.Write([]byte{0xff}) // Domain separation for the second hash + h2 := f.Sum64() | 1 // Odd, so probes cycle through all bit positions + return h1, h2 +} diff --git a/util/bloom_test.go b/util/bloom_test.go new file mode 100644 index 00000000..92512f79 --- /dev/null +++ b/util/bloom_test.go @@ -0,0 +1,61 @@ +package util_test + +import ( + "fmt" + "testing" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/util" +) + +func TestBloomFilter_NoFalseNegatives(t *testing.T) { + // The property routing correctness relies on: an added element is ALWAYS reported present + b := util.NewBloomFilter(1000, 0.01) + for i := 0; i < 1000; i++ { + b.Add(fmt.Sprintf("topic-%d", i)) + } + for i := 0; i < 1000; i++ { + require.True(t, b.Contains(fmt.Sprintf("topic-%d", i))) + } +} + +func TestBloomFilter_FalsePositiveRate(t *testing.T) { + b := util.NewBloomFilter(10000, 0.01) + for i := 0; i < 10000; i++ { + b.Add(fmt.Sprintf("added-%d", i)) + } + falsePositives := 0 + for i := 0; i < 10000; i++ { + if b.Contains(fmt.Sprintf("absent-%d", i)) { + falsePositives++ + } + } + require.Less(t, falsePositives, 300, "expected ~1%% false positives, got %d/10000", falsePositives) +} + +func TestBloomFilter_EmptyContainsNothing(t *testing.T) { + b := util.NewBloomFilter(100, 0.01) + require.False(t, b.Contains("anything")) +} + +func TestBloomFilter_MarshalRoundTrip(t *testing.T) { + b := util.NewBloomFilter(500, 0.01) + for i := 0; i < 500; i++ { + b.Add(fmt.Sprintf("topic-%d", i)) + } + data, err := b.MarshalBinary() + require.Nil(t, err) + b2, err := util.UnmarshalBloomFilter(data) + require.Nil(t, err) + for i := 0; i < 500; i++ { + require.True(t, b2.Contains(fmt.Sprintf("topic-%d", i))) + } + require.False(t, b2.Contains("never-added-topic")) +} + +func TestBloomFilter_UnmarshalGarbage(t *testing.T) { + _, err := util.UnmarshalBloomFilter([]byte{}) + require.Error(t, err) + _, err = util.UnmarshalBloomFilter([]byte{1, 2, 3}) + require.Error(t, err) +} diff --git a/util/linger_queue.go b/util/linger_queue.go new file mode 100644 index 00000000..efa620dd --- /dev/null +++ b/util/linger_queue.go @@ -0,0 +1,144 @@ +package util + +import ( + "sync" + "time" +) + +// LingerQueue is a bounded, non-blocking batching queue: enqueued elements are emitted as +// batches once a linger window expires, a batch count cap is reached, or a batch size cap is +// reached, whichever comes first. Unlike BatchingQueue, producers never block: TryEnqueue drops +// (returns false) when the queue is full, and Close flushes the remainder and terminates the +// consumer channel, so per-entity queues can be created and destroyed dynamically. +// +// Example: +// +// q := NewLingerQueue[int](64, 10, 0, nil, 500*time.Millisecond) +// go func() { +// for batch := range q.Dequeue() { +// send(batch) +// } +// }() +// q.TryEnqueue(1) +// q.TryEnqueue(2) // emitted together as [1, 2] after <= 500ms +type LingerQueue[T any] struct { + in chan T + out chan []T + max int // max elements per batch + maxSize int // max cumulative size per batch; 0 = no size cap + size func(T) int // element size function; nil = no size cap + linger time.Duration // max time the first element of a batch waits; 0 = emit immediately + closed bool + mu sync.Mutex // Protects closed, and guards TryEnqueue's send against Close's close(in) +} + +// NewLingerQueue creates a LingerQueue holding up to capacity queued elements, emitting batches +// of up to max elements or maxSize cumulative size (as measured by size; pass 0/nil for no size +// cap) after at most linger. +func NewLingerQueue[T any](capacity, max, maxSize int, size func(T) int, linger time.Duration) *LingerQueue[T] { + q := &LingerQueue[T]{ + in: make(chan T, capacity), + out: make(chan []T), + max: max, + maxSize: maxSize, + size: size, + linger: linger, + } + go q.run() + return q +} + +// TryEnqueue enqueues an element without blocking. It returns false if the queue is full or +// closed; the caller decides how to account for the drop. +func (q *LingerQueue[T]) TryEnqueue(t T) bool { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return false + } + select { + case q.in <- t: + return true + default: + return false + } +} + +// Dequeue returns the channel emitting batches. It is closed after Close, once the remaining +// elements have been flushed. +func (q *LingerQueue[T]) Dequeue() <-chan []T { + return q.out +} + +// Close stops the queue: remaining elements are flushed as final batches, then the Dequeue +// channel is closed. TryEnqueue returns false after Close. Close is idempotent. +func (q *LingerQueue[T]) Close() { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return + } + q.closed = true + close(q.in) +} + +// run is the batching loop: it blocks for the first element of a batch, then collects more until +// the linger expires or a cap is hit, and emits the batch. It exits once the queue is closed and +// drained. Note that receiving from the closed in channel still yields the buffered remainder +// before reporting closed, which is what flushes on Close. +func (q *LingerQueue[T]) run() { + defer close(q.out) + for { + first, ok := <-q.in + if !ok { + return + } + batch := []T{first} + bytes := q.sizeOf(first) + var timeout <-chan time.Time + if q.linger > 0 { + timeout = time.After(q.linger) + } + closed := false + collect: + for len(batch) < q.max && (q.maxSize <= 0 || bytes < q.maxSize) { + if timeout == nil { + // Zero linger: greedily drain what is immediately available, never wait + select { + case t, ok := <-q.in: + if !ok { + closed = true + break collect + } + batch = append(batch, t) + bytes += q.sizeOf(t) + default: + break collect + } + } else { + select { + case t, ok := <-q.in: + if !ok { + closed = true + break collect + } + batch = append(batch, t) + bytes += q.sizeOf(t) + case <-timeout: + break collect + } + } + } + q.out <- batch + if closed { + return + } + } +} + +func (q *LingerQueue[T]) sizeOf(t T) int { + if q.size == nil { + return 0 + } + return q.size(t) +} diff --git a/util/linger_queue_test.go b/util/linger_queue_test.go new file mode 100644 index 00000000..43b02bde --- /dev/null +++ b/util/linger_queue_test.go @@ -0,0 +1,90 @@ +package util_test + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/util" +) + +func TestLingerQueue_LingerFlush(t *testing.T) { + // Items enqueued within the linger window are emitted as a single batch when it expires + q := util.NewLingerQueue[int](16, 100, 0, nil, 50*time.Millisecond) + defer q.Close() + require.True(t, q.TryEnqueue(1)) + require.True(t, q.TryEnqueue(2)) + require.True(t, q.TryEnqueue(3)) + start := time.Now() + batch := <-q.Dequeue() + require.Equal(t, []int{1, 2, 3}, batch) + require.GreaterOrEqual(t, time.Since(start), 30*time.Millisecond) // Waited out the linger +} + +func TestLingerQueue_MaxBatchFlush(t *testing.T) { + // Hitting the count cap flushes early, before the linger expires + q := util.NewLingerQueue[int](16, 5, 0, nil, time.Minute) + defer q.Close() + for i := 0; i < 12; i++ { + require.True(t, q.TryEnqueue(i)) + } + require.Len(t, <-q.Dequeue(), 5) + require.Len(t, <-q.Dequeue(), 5) + q.Close() // Flushes the remainder + require.Len(t, <-q.Dequeue(), 2) +} + +func TestLingerQueue_SizeCapFlush(t *testing.T) { + // Hitting the byte cap flushes early, before count cap or linger + q := util.NewLingerQueue(16, 100, 10, func(s string) int { return len(s) }, time.Minute) + defer q.Close() + require.True(t, q.TryEnqueue("aaaa")) + require.True(t, q.TryEnqueue("bbbb")) + require.True(t, q.TryEnqueue("cccc")) // 12 bytes >= 10 -> flush + batch := <-q.Dequeue() + require.Equal(t, []string{"aaaa", "bbbb", "cccc"}, batch) +} + +func TestLingerQueue_TryEnqueueFull(t *testing.T) { + // A full queue drops (returns false) instead of blocking the producer + q := util.NewLingerQueue[int](1, 1, 0, nil, 0) + defer q.Close() + require.True(t, q.TryEnqueue(1)) // Taken by the batcher, blocks emitting (no consumer) + waitForCond(t, func() bool { return q.TryEnqueue(2) }) // Fills the buffer once slot frees + require.False(t, q.TryEnqueue(3)) // Buffer full, batcher blocked -> drop +} + +func TestLingerQueue_CloseFlushesAndCloses(t *testing.T) { + q := util.NewLingerQueue[int](16, 100, 0, nil, time.Minute) + require.True(t, q.TryEnqueue(1)) + require.True(t, q.TryEnqueue(2)) + q.Close() + require.Equal(t, []int{1, 2}, <-q.Dequeue()) // Remainder flushed without waiting out the linger + _, ok := <-q.Dequeue() + require.False(t, ok) // Channel closed + require.False(t, q.TryEnqueue(3)) + q.Close() // Idempotent +} + +func TestLingerQueue_LingerZeroImmediate(t *testing.T) { + q := util.NewLingerQueue[int](16, 100, 0, nil, 0) + defer q.Close() + require.True(t, q.TryEnqueue(1)) + select { + case batch := <-q.Dequeue(): + require.Contains(t, batch, 1) + case <-time.After(time.Second): + t.Fatal("expected immediate flush with zero linger") + } +} + +func waitForCond(t *testing.T, f func() bool) { + t.Helper() + for i := 0; i < 100; i++ { + if f() { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatal("timed out waiting for condition") +} diff --git a/webpush/store_test.go b/webpush/store_test.go index 65dd817f..124f15d8 100644 --- a/webpush/store_test.go +++ b/webpush/store_test.go @@ -13,7 +13,9 @@ import ( "heckel.io/ntfy/v2/webpush" ) -const testWebPushEndpoint = "https://updates.push.services.mozilla.com/wpush/v1/AAABBCCCDDEEEFFF" +const ( + testWebPushEndpoint = "https://updates.push.services.mozilla.com/wpush/v1/AAABBCCCDDEEEFFF" +) // Schema layout as written by ntfy releases before the db/schema framework; used to verify // that existing databases open cleanly without an adoption step