diff --git a/cmd/beacon/main.go b/cmd/beacon/main.go index 2f17043..8e322e6 100644 --- a/cmd/beacon/main.go +++ b/cmd/beacon/main.go @@ -12,6 +12,7 @@ import ( "os" "os/signal" "strconv" + "sync" "syscall" "time" @@ -284,8 +285,9 @@ func main() { broker2.SetCacheInvalidators(cr.InvalidateNode, cr.InvalidateObserver) } - go broker1.Start(ctx) - go broker2.Start(ctx) + var ingestWorkers sync.WaitGroup + ingestWorkers.Go(func() { broker1.Start(ctx) }) + ingestWorkers.Go(func() { broker2.Start(ctx) }) tasks := []background.Task{ background.ViewRefreshTask(store, resolved.ViewRefreshInterval), @@ -343,6 +345,7 @@ func main() { if err := srv.Shutdown(shutdownCtx); err != nil { slog.Error("server shutdown error", "component", "startup", "error", err) } + ingestWorkers.Wait() coalescer.Flush(shutdownCtx) } diff --git a/internal/ingest/capability.go b/internal/ingest/capability.go index 240c381..22f7ce5 100644 --- a/internal/ingest/capability.go +++ b/internal/ingest/capability.go @@ -5,28 +5,47 @@ package ingest import ( "context" + "sync" "github.com/google/uuid" ) -// runCapabilityDetection checks hash sizes and flips firmware capability flags. -// Called only when the observation INSERT succeeded (no dedup conflict). -// -// Rules (from design doc): -// - hash_size == 1: do nothing (proves nothing about firmware) -// - duplicate hash prefixes within the path: skip entirely -// - non-trace + hash_size 2 or 3 → supports_multibyte_paths = TRUE -// - trace (0x09) + hash_size 2 or 4 → supports_multibyte_traces = TRUE +const capabilityCacheLimit = 4096 + +type capabilityCache struct { + mu sync.Mutex + nodes map[uuid.UUID]uint8 +} + func (w *Worker) runCapabilityDetection(ctx context.Context, payloadType uint8, hashSize uint8, resolvedNodeIDs []uuid.UUID) { - if hashSize < 2 { + var flag uint8 + switch { + case payloadType != 0x09 && (hashSize == 2 || hashSize == 3): + flag = 1 + case payloadType == 0x09 && (hashSize == 2 || hashSize == 4): + flag = 2 + default: return } for _, nodeID := range resolvedNodeIDs { - switch { - case payloadType != 0x09 && (hashSize == 2 || hashSize == 3): - _ = w.db.SetNodeCapability(ctx, nodeID, true, false) - case payloadType == 0x09 && (hashSize == 2 || hashSize == 4): - _ = w.db.SetNodeCapability(ctx, nodeID, false, true) + w.capabilities.mu.Lock() + recorded := w.capabilities.nodes[nodeID]&flag != 0 + w.capabilities.mu.Unlock() + if recorded { + continue + } + if err := w.db.SetNodeCapability(ctx, nodeID, flag == 1, flag == 2); err != nil { + continue + } + w.capabilities.mu.Lock() + if w.capabilities.nodes == nil { + w.capabilities.nodes = make(map[uuid.UUID]uint8) + } + // Eviction only costs another write; capabilities never downgrade. + if len(w.capabilities.nodes) >= capabilityCacheLimit { + clear(w.capabilities.nodes) } + w.capabilities.nodes[nodeID] |= flag + w.capabilities.mu.Unlock() } } diff --git a/internal/ingest/capability_test.go b/internal/ingest/capability_test.go new file mode 100644 index 0000000..2ebfb42 --- /dev/null +++ b/internal/ingest/capability_test.go @@ -0,0 +1,72 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ingest + +import ( + "context" + "errors" + "testing" + + "github.com/google/uuid" +) + +func TestCapabilityAlreadyRecorded(t *testing.T) { + w, db := newTestWorker() + id := uuid.New() + for range 3 { + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{id, id}) + } + if got := len(db.setCapabilityCalls); got != 1 { + t.Fatalf("recorded the same capability %d times, want 1", got) + } + w.runCapabilityDetection(context.Background(), 9, 2, []uuid.UUID{id}) + w.runCapabilityDetection(context.Background(), 9, 2, []uuid.UUID{id}) + if got := len(db.setCapabilityCalls); got != 2 || !db.setCapabilityCalls[1].traces { + t.Fatalf("trace capability was not recorded separately: %+v", db.setCapabilityCalls) + } +} + +type failingCapabilityDB struct { + *stubDB + fail bool +} + +func (d *failingCapabilityDB) SetNodeCapability(ctx context.Context, id uuid.UUID, paths, traces bool) error { + _ = d.stubDB.SetNodeCapability(ctx, id, paths, traces) + if d.fail { + return errors.New("write failed") + } + return nil +} + +func TestCapabilityRetriesFailedWrite(t *testing.T) { + w, base := newTestWorker() + db := &failingCapabilityDB{stubDB: base, fail: true} + w.db = db + id := uuid.New() + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{id}) + db.fail = false + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{id}) + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{id}) + if got := len(db.setCapabilityCalls); got != 2 { + t.Fatalf("got %d writes, want the failed write and one successful retry", got) + } +} + +func TestCapabilityCacheEviction(t *testing.T) { + w, db := newTestWorker() + first := uuid.New() + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{first}) + for range capabilityCacheLimit { + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{uuid.New()}) + } + if len(w.capabilities.nodes) > capabilityCacheLimit { + t.Fatal("cache grew beyond limit") + } + before := len(db.setCapabilityCalls) + w.runCapabilityDetection(context.Background(), 5, 2, []uuid.UUID{first}) + if len(db.setCapabilityCalls) != before+1 { + t.Fatal("evicted capability was not recorded again") + } +} diff --git a/internal/ingest/endpoint_matching_test.go b/internal/ingest/endpoint_matching_test.go index 1b60062..52d1219 100644 --- a/internal/ingest/endpoint_matching_test.go +++ b/internal/ingest/endpoint_matching_test.go @@ -35,6 +35,7 @@ func (s *endpointRoutingDB) ResolvePathHashes(_ context.Context, iata string, ha func TestHandlePacketSeparatesEndpointAndRelayMatching(t *testing.T) { for _, kind := range []uint8{meshcore.PayloadTypeReq, meshcore.PayloadTypeResponse, meshcore.PayloadTypeTxtMsg, meshcore.PayloadTypePath} { w, base := newTestWorker() + base.observationInserted = true db := &endpointRoutingDB{stubDB: base} w.db = db packet := &meshcore.Packet{Header: meshcore.MakeHeader(meshcore.RouteTypeFlood, kind, 0), diff --git a/internal/ingest/endpoint_work_test.go b/internal/ingest/endpoint_work_test.go new file mode 100644 index 0000000..5030291 --- /dev/null +++ b/internal/ingest/endpoint_work_test.go @@ -0,0 +1,54 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ingest + +import ( + "context" + "testing" + + "github.com/MeshCore-Beacon/beacon-server/internal/api" + "github.com/meshcore-go/meshcore-go" +) + +type endpointLookupDB struct { + *stubDB + calls [][][]byte +} + +func (d *endpointLookupDB) ResolveEndpointHashes(_ context.Context, _ string, hashes [][]byte) (map[string][]api.ResolvedPathEntry, error) { + d.calls = append(d.calls, hashes) + return nil, nil +} +func endpointTestPacket() *meshcore.Packet { + return &meshcore.Packet{Header: meshcore.MakeHeader(meshcore.RouteTypeFlood, meshcore.PayloadTypeTxtMsg, 0), Payload: append([]byte{0xbb, 0xaa, 0, 0}, make([]byte, 16)...)} +} +func TestEndpointLookupsOnlyForLiveHearings(t *testing.T) { + r := newRepeatHarness(t, true) + d := &endpointLookupDB{stubDB: r.db.stubDB} + r.w.db = d + packet := endpointTestPacket() + hear := func(inserted bool, path byte) { + d.observationInserted = inserted + packet.Path = []byte{path} + packet.PathLength = 1 + r.w.handlePacket(r.ctx, "YOW", "0102", packetEnvelope(t, packet)) + } + hear(true, 0x11) + first := len(d.calls) + if first == 0 { + t.Fatal("first hearing missing endpoint lookup") + } + hear(false, 0x11) + if len(d.calls) != first { + t.Fatal("suppressed copy resolved endpoints") + } + hear(false, 0x22) + if len(d.calls) != 2*first { + t.Fatal("new path missing endpoint lookup") + } + hear(false, 0x22) + if len(d.calls) != 2*first { + t.Fatal("suppressed repeat resolved endpoints") + } +} diff --git a/internal/ingest/ingest.go b/internal/ingest/ingest.go index f19e403..998c205 100644 --- a/internal/ingest/ingest.go +++ b/internal/ingest/ingest.go @@ -227,6 +227,7 @@ type Worker struct { keys ChannelKeyStore scopes ScopeStore client mqtt.Client + capabilities capabilityCache onNodeUpsert func(ctx context.Context, nodeID uuid.UUID) onObserverUpsert func(ctx context.Context, observerID uuid.UUID) } @@ -243,6 +244,28 @@ func New(cfg Config, db DB, h *hub.Hub, keys ChannelKeyStore, scopes ScopeStore) // // Intended usage: go worker.Start(ctx) func (w *Worker) Start(ctx context.Context) { + if w.cfg.URL == "" { + return + } + queue := newMessageQueue(8, 2048, 32<<20, func(ctx context.Context, msg mqtt.Message) { + w.handleMessageContext(ctx, msg) + msg.Ack() + }) + reportDrops := func() { + if n := queue.takeDropped(); n > 0 { + w.log.Error("ingest queue full, messages dropped", "count", n) + } + } + defer func() { + drainCtx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + defer cancel() + if err := queue.close(drainCtx); err != nil { + w.log.Warn("ingest shutdown timed out", "error", err) + } + reportDrops() + w.log.Info("stopped") + }() + // Isolate workers across deployments; Paho reuses this ID on reconnect. // Keep it alphanumeric and within MQTT 3.1's 23-character client ID limit. opts := mqtt.NewClientOptions(). @@ -251,6 +274,7 @@ func (w *Worker) Start(ctx context.Context) { SetUsername(w.cfg.Username). SetPassword(w.cfg.Password). SetAutoReconnect(true). + SetAutoAckDisabled(true). SetMaxReconnectInterval(30 * time.Second). SetKeepAlive(30 * time.Second). SetPingTimeout(10 * time.Second). @@ -260,21 +284,33 @@ func (w *Worker) Start(ctx context.Context) { SetConnectRetryInterval(5 * time.Second). SetOnConnectHandler(func(c mqtt.Client) { w.log.Info("connected") - w.subscribe(c) + w.subscribe(c, queue) }). SetConnectionLostHandler(func(_ mqtt.Client, err error) { w.log.Warn("connection lost, will reconnect", "error", err) }) w.client = mqtt.NewClient(opts) - if tok := w.client.Connect(); tok.Wait() && tok.Error() != nil { - w.log.Error("initial connect failed", "error", tok.Error()) - // paho will retry; we fall through and wait for ctx + + token := w.client.Connect() + ticker := time.NewTicker(5 * time.Second) + defer ticker.Stop() + connected := token.Done() + for { + select { + case <-connected: + if err := token.Error(); err != nil { + w.log.Error("initial connect failed", "error", err) + } + connected = nil + case <-ticker.C: + reportDrops() + case <-ctx.Done(): + w.client.Disconnect(500) + return + } } - <-ctx.Done() - w.client.Disconnect(500) - w.log.Info("stopped") } func (w *Worker) BrokerName() string { @@ -294,12 +330,15 @@ func (w *Worker) SetCacheInvalidators(onNode, onObserver func(ctx context.Contex } // subscribe registers the wildcard topic handler after (re)connect. -func (w *Worker) subscribe(client mqtt.Client) { +func (w *Worker) subscribe(client mqtt.Client, queue *messageQueue) { // meshcore/{IATA}/{pubkey}/packets // meshcore/{IATA}/{pubkey}/status // We do NOT subscribe to /internal (Role 2 access). tok := client.Subscribe("meshcore/#", 1, func(_ mqtt.Client, msg mqtt.Message) { - w.handleMessage(msg) + if !queue.enqueue(msg) { + // Counted drops must release the broker's inflight slot. + msg.Ack() + } }) if tok.Wait() && tok.Error() != nil { w.log.Error("subscribe error", "error", tok.Error()) @@ -323,8 +362,12 @@ func isValidIATA(s string) bool { // handleMessage dispatches incoming MQTT messages by subtopic. // Each message is processed with a 30s timeout to prevent slow DB calls -// from blocking the MQTT receive goroutine indefinitely. +// from holding a processing worker indefinitely. func (w *Worker) handleMessage(msg mqtt.Message) { + w.handleMessageContext(context.Background(), msg) +} + +func (w *Worker) handleMessageContext(parent context.Context, msg mqtt.Message) { // Topic shape: meshcore/{IATA}/{pubkey}/{subtopic} parts := strings.SplitN(msg.Topic(), "/", 4) if len(parts) != 4 || parts[0] != "meshcore" { @@ -349,7 +392,7 @@ func (w *Worker) handleMessage(msg mqtt.Message) { } } - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + ctx, cancel := context.WithTimeout(parent, 30*time.Second) defer cancel() switch subtopic { diff --git a/internal/ingest/packet.go b/internal/ingest/packet.go index ad1bfcb..92f3da7 100644 --- a/internal/ingest/packet.go +++ b/internal/ingest/packet.go @@ -766,29 +766,6 @@ func (w *Worker) handlePacket(ctx context.Context, iata, pubkeyHex string, raw [ if err != nil { w.log.Error(fmt.Sprintf("db: get observer radio failed for %s", pubkeyHex), "error", err) } - // Endpoints for the live event only; stored rows resolve them at read time. - var resolvedSource, resolvedDestination *api.ResolvedHop - if packet.PayloadType() == meshcore.PayloadTypeAdvert && originPubkey != nil { - // Exact match: ADVERT carries the sender's real identity pubkey, not a - // short ambiguous hash prefix like the other resolvable payload types. - if nodeID, err := w.db.GetNodeByPubkey(ctx, originPubkey); err == nil { - if nodes, err := w.db.GetNodesByIDs(ctx, []uuid.UUID{nodeID}); err == nil { - hop := api.ResolveExactNode(nodes[nodeID]) - resolvedSource = &hop - } - } - } else if len(sourceHashByte) == 1 { - if r, err := w.db.ResolveEndpointHashes(ctx, iata, [][]byte{sourceHashByte}); err == nil { - hop := api.BuildResolvedPath([][]byte{sourceHashByte}, r)[0] - resolvedSource = &hop - } - } - if len(destHashByte) == 1 { - if r, err := w.db.ResolveEndpointHashes(ctx, iata, [][]byte{destHashByte}); err == nil { - hop := api.BuildResolvedPath([][]byte{destHashByte}, r)[0] - resolvedDestination = &hop - } - } // Airtime is costed from the frame as received; zero radio columns mean the // observer never reported its settings, so there is nothing to cost. var airtimeMs *float32 @@ -885,6 +862,29 @@ func (w *Worker) handlePacket(ctx context.Context, iata, pubkeyHex string, raw [ // A duplicate is streamed, never stored, to includeRepeats clients when its path is new. repeat := !inserted && w.hub.RepeatsWanted() && w.hub.MarkSent(packetHash[:], id[:], packet.Path) if inserted || repeat { + // Endpoints for the live event only; stored rows resolve them at read time. + var resolvedSource, resolvedDestination *api.ResolvedHop + if packet.PayloadType() == meshcore.PayloadTypeAdvert && originPubkey != nil { + // Exact match: ADVERT carries the sender's real identity pubkey, not a + // short ambiguous hash prefix like the other resolvable payload types. + if nodeID, err := w.db.GetNodeByPubkey(ctx, originPubkey); err == nil { + if nodes, err := w.db.GetNodesByIDs(ctx, []uuid.UUID{nodeID}); err == nil { + hop := api.ResolveExactNode(nodes[nodeID]) + resolvedSource = &hop + } + } + } else if len(sourceHashByte) == 1 { + if r, err := w.db.ResolveEndpointHashes(ctx, iata, [][]byte{sourceHashByte}); err == nil { + hop := api.BuildResolvedPath([][]byte{sourceHashByte}, r)[0] + resolvedSource = &hop + } + } + if len(destHashByte) == 1 { + if r, err := w.db.ResolveEndpointHashes(ctx, iata, [][]byte{destHashByte}); err == nil { + hop := api.BuildResolvedPath([][]byte{destHashByte}, r)[0] + resolvedDestination = &hop + } + } if inserted { w.handlePayloadTypeSideEffects(ctx, packet, iata, packetHash[:], radio, scopeID, matchedScope, pubkeyBytes, float32(parseNumber(envelope.SNR))) if w.hub.RepeatsWanted() { diff --git a/internal/ingest/queue.go b/internal/ingest/queue.go new file mode 100644 index 0000000..ff7e6d5 --- /dev/null +++ b/internal/ingest/queue.go @@ -0,0 +1,108 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ingest + +import ( + "context" + "strings" + "sync" + + mqtt "github.com/eclipse/paho.mqtt.golang" +) + +type messageQueue struct { + mu sync.Mutex + lanes []chan mqtt.Message + bytes, maxBytes int + dropped uint64 + closed bool + done chan struct{} + cancel context.CancelFunc +} + +func newMessageQueue(workers, capacity, maxBytes int, handle func(context.Context, mqtt.Message)) *messageQueue { + ctx, cancel := context.WithCancel(context.Background()) + q := &messageQueue{lanes: make([]chan mqtt.Message, workers), maxBytes: maxBytes, done: make(chan struct{}), cancel: cancel} + var wg sync.WaitGroup + for i := range q.lanes { + lane := make(chan mqtt.Message, capacity) + q.lanes[i] = lane + wg.Go(func() { + for m := range lane { + if ctx.Err() != nil { + return + } + handle(ctx, m) + q.mu.Lock() + q.bytes -= len(m.Topic()) + len(m.Payload()) + q.mu.Unlock() + } + }) + } + go func() { wg.Wait(); close(q.done) }() + return q +} + +func (q *messageQueue) enqueue(m mqtt.Message) bool { + // Keep status and packets from each observer in arrival order. + parts := strings.SplitN(m.Topic(), "/", 4) + key := m.Topic() + if len(parts) == 4 { + key = parts[2] + } + hash := uint32(2166136261) + for i := range len(key) { + c := key[i] + if c >= 'A' && c <= 'Z' { + c += 'a' - 'A' + } + hash = (hash ^ uint32(c)) * 16777619 + } + size := len(m.Topic()) + len(m.Payload()) + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return false + } + if size > q.maxBytes-q.bytes { + q.dropped++ + return false + } + select { + case q.lanes[int(hash%uint32(len(q.lanes)))] <- m: + q.bytes += size + return true + default: + q.dropped++ + return false + } +} + +func (q *messageQueue) takeDropped() uint64 { + q.mu.Lock() + defer q.mu.Unlock() + n := q.dropped + q.dropped = 0 + return n +} + +func (q *messageQueue) close(ctx context.Context) error { + q.mu.Lock() + if !q.closed { + q.closed = true + for _, lane := range q.lanes { + close(lane) + } + } + q.mu.Unlock() + select { + case <-q.done: + q.cancel() + return nil + case <-ctx.Done(): + q.cancel() + <-q.done + return ctx.Err() + } +} diff --git a/internal/ingest/queue_test.go b/internal/ingest/queue_test.go new file mode 100644 index 0000000..130a17a --- /dev/null +++ b/internal/ingest/queue_test.go @@ -0,0 +1,194 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ingest + +import ( + "context" + "fmt" + "reflect" + "sync" + "testing" + "time" + + mqtt "github.com/eclipse/paho.mqtt.golang" +) + +type queuedTestMessage struct { + malformedTopicMessage + payload []byte +} + +func (m queuedTestMessage) Payload() []byte { return m.payload } +func queueMessage(observer, kind string, n int) mqtt.Message { + return queuedTestMessage{malformedTopicMessage{"meshcore/SEA/" + observer + "/" + kind}, []byte(fmt.Sprint(n))} +} + +func TestQueueKeepsObserverOrderWhileDatabaseIsBusy(t *testing.T) { + entered, release := make(chan struct{}), make(chan struct{}) + var mu sync.Mutex + var got []string + q := newMessageQueue(2, 8, 4096, func(ctx context.Context, m mqtt.Message) { + if string(m.Payload()) == "0" { + close(entered) + select { + case <-release: + case <-ctx.Done(): + return + } + } + mu.Lock() + got = append(got, string(m.Payload())) + mu.Unlock() + }) + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _ = q.close(ctx) + }) + if !q.enqueue(queueMessage("aa", "status", 0)) { + t.Fatal("first message rejected") + } + <-entered + for n := 1; n <= 3; n++ { + if !q.enqueue(queueMessage("AA", "packets", n)) { + t.Fatal("callback blocked or message rejected") + } + } + close(release) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := q.close(ctx); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(got, []string{"0", "1", "2", "3"}) { + t.Fatalf("observer order: %v", got) + } + if q.enqueue(queueMessage("aa", "packets", 4)) { + t.Fatal("accepted after shutdown") + } +} + +func TestQueueRunsOtherObserversAndReportsOverflow(t *testing.T) { + entered, other := make(chan struct{}), make(chan struct{}) + q := newMessageQueue(2, 1, 4096, func(ctx context.Context, m mqtt.Message) { + if m.Topic() == "meshcore/SEA/a/packets" { + close(entered) + <-ctx.Done() + } else { + close(other) + } + }) + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _ = q.close(ctx) + }) + q.enqueue(queueMessage("a", "packets", 0)) + <-entered + if !q.enqueue(queueMessage("b", "packets", 0)) { + t.Fatal("other observer rejected") + } + select { + case <-other: + case <-time.After(time.Second): + t.Fatal("other observer blocked") + } + if !q.enqueue(queueMessage("a", "packets", 1)) { + t.Fatal("buffered message rejected") + } + if q.enqueue(queueMessage("a", "packets", 2)) { + t.Fatal("accepted beyond queue capacity") + } + if q.takeDropped() != 1 { + t.Fatal("overflow was not counted") + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if q.close(ctx) == nil { + t.Fatal("expected shutdown deadline") + } +} + +func TestQueueBoundsPayloadMemory(t *testing.T) { + q := newMessageQueue(1, 8, 64, func(context.Context, mqtt.Message) {}) + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _ = q.close(ctx) + }) + m := queuedTestMessage{malformedTopicMessage{"meshcore/SEA/a/packets"}, make([]byte, 65)} + if q.enqueue(m) { + t.Fatal("accepted oversized payload") + } + if q.takeDropped() != 1 { + t.Fatal("memory overflow was not counted") + } +} + +type subscribeCapture struct { + mqtt.Client + callback mqtt.MessageHandler +} + +func (c *subscribeCapture) Subscribe(_ string, _ byte, callback mqtt.MessageHandler) mqtt.Token { + c.callback = callback + return completedSubscribe{} +} + +type completedSubscribe struct{ mqtt.Token } + +func (completedSubscribe) Wait() bool { return true } +func (completedSubscribe) Error() error { return nil } + +type ackMessage struct { + mqtt.Message + acked bool +} + +func (m *ackMessage) Ack() { m.acked = true } + +func TestSubscribeReleasesRejectedMessage(t *testing.T) { + w, _ := newTestWorker() + q := newMessageQueue(1, 1, 1, func(context.Context, mqtt.Message) { t.Error("rejected message processed") }) + defer q.close(context.Background()) + client := &subscribeCapture{} + w.subscribe(client, q) + m := &ackMessage{Message: queueMessage("aa", "packets", 0)} + client.callback(client, m) + if !m.acked { + t.Fatal("rejected message still occupies the broker inflight window") + } + if q.takeDropped() != 1 { + t.Fatal("rejected message was not counted") + } +} + +func TestQueueShutdownWaitsForCanceledHandler(t *testing.T) { + entered, canceled, release := make(chan struct{}), make(chan struct{}), make(chan struct{}) + q := newMessageQueue(1, 1, 4096, func(ctx context.Context, _ mqtt.Message) { + close(entered) + <-ctx.Done() + close(canceled) + <-release + }) + q.enqueue(queueMessage("aa", "packets", 0)) + <-entered + ctx, cancel := context.WithCancel(context.Background()) + cancel() + stopped := make(chan error, 1) + go func() { stopped <- q.close(ctx) }() + <-canceled + select { + case <-stopped: + close(release) + t.Fatal("shutdown returned while handler was still running") + case <-time.After(20 * time.Millisecond): + } + close(release) + select { + case <-stopped: + case <-time.After(time.Second): + t.Fatal("shutdown did not finish") + } +}