WIP: Clsuter support (cross node delivery, leader election)

This commit is contained in:
binwiederhier
2026-08-01 16:10:31 +02:00
parent dc11655153
commit 5a6c4277ad
35 changed files with 3232 additions and 63 deletions
+139
View File
@@ -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)
}
+133
View File
@@ -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")
}
}
}
+450
View File
@@ -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
}
+572
View File
@@ -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)
}
+24
View File
@@ -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 }
+83
View File
@@ -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")
}
+149
View File
@@ -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
}
+222
View File
@@ -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
}
+47
View File
@@ -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
}
+69
View File
@@ -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()
}
+40
View File
@@ -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://<cluster-listen>)"}),
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 == "::"
}
+35
View File
@@ -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))
+83
View File
@@ -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
}
+84
View File
@@ -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
}
+1
View File
@@ -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
+3 -1
View File
@@ -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)
+2 -2
View File
@@ -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)
+36
View File
@@ -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) {
+8
View File
@@ -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,
+38
View File
@@ -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
+36
View File
@@ -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,
)
}
+9
View File
@@ -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",
+14 -2
View File
@@ -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://<cluster-listen>")
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,
+1
View File
@@ -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())
+186 -46
View File
@@ -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
+28
View File
@@ -61,6 +61,34 @@
#
# database-url: <connection-string>
# 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://<cluster-listen>". 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: <ip:port>
# cluster-node-id: <hostname>
# cluster-advertise-url: "http://<cluster-listen>"
# cluster-secret: <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).
+2 -1
View File
@@ -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
+318
View File
@@ -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 })
}
+9 -5
View File
@@ -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()
+10 -5
View File
@@ -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,
+103
View File
@@ -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
}
+61
View File
@@ -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)
}
+144
View File
@@ -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)
}
+90
View File
@@ -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")
}
+3 -1
View File
@@ -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