mirror of
https://github.com/multipleof4/ntfy.git
synced 2026-10-08 21:05:21 +00:00
WIP: Clsuter support (cross node delivery, leader election)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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 == "::"
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 })
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user