diff --git a/CHANGELOG.md b/CHANGELOG.md index 1f1cca3a5..f94853e6e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -131,6 +131,7 @@ clients, and how to run the same setup without Nix. - Fix SSH check accepting a repeated follow-up for an already-decided session, even after a rejection [#3526](https://github.com/juanfont/headscale/pull/3526) - A node re-registering with a spent, expired or revoked pre-auth key is now rejected if it expired or changed node key while the re-registration was in flight [#3525](https://github.com/juanfont/headscale/pull/3525) +- Fix a node ping being lost when a full map update is queued at the same time ## 0.29.5 (202x-xx-xx) diff --git a/hscontrol/mapper/batcher.go b/hscontrol/mapper/batcher.go index 71f7eac2b..321814212 100644 --- a/hscontrol/mapper/batcher.go +++ b/hscontrol/mapper/batcher.go @@ -687,17 +687,15 @@ func (b *Batcher) addToBatch(changes ...change.Change) { } } - // Short circuit if any of the changes is a full update, which - // means we can skip sending individual changes. + // A full update supersedes every state change pending or in this call, + // but not the pings addressed to a node; those follow the full. if change.HasFull(changes) { b.nodes.Range(func(_ types.NodeID, nc *multiChannelNodeConn) bool { if nc == nil { return true } - nc.pendingMu.Lock() - nc.pending = []change.Change{change.FullUpdate()} - nc.pendingMu.Unlock() + nc.collapsePendingToFull(changes) return true }) diff --git a/hscontrol/mapper/batcher_concurrency_test.go b/hscontrol/mapper/batcher_concurrency_test.go index 3c21f122a..dad31e321 100644 --- a/hscontrol/mapper/batcher_concurrency_test.go +++ b/hscontrol/mapper/batcher_concurrency_test.go @@ -275,6 +275,96 @@ func TestAddToBatch_FullUpdateOverrides(t *testing.T) { }) } +// TestAddToBatch_FullUpdateKeepsPingRequests verifies that a full update +// keeps the pings it cannot carry: one pending before the full, and one in +// the same call. A ping lost here times out its caller. +func TestAddToBatch_FullUpdateKeepsPingRequests(t *testing.T) { + lb := setupLightweightBatcher(t, 3, 10) + defer lb.cleanup() + + prPending := &tailcfg.PingRequest{URL: "https://example.com/ping/pending"} + prSameCall := &tailcfg.PingRequest{URL: "https://example.com/ping/same-call"} + + lb.b.addToBatch(change.PingNode(1, prPending)) + lb.b.addToBatch(change.NodeOnline(3)) + lb.b.addToBatch(change.PingNode(2, prSameCall), change.UserRemoved()) + + want := map[types.NodeID][]change.Change{ + 1: {change.FullUpdate(), change.PingNode(1, prPending)}, + 2: {change.FullUpdate(), change.PingNode(2, prSameCall)}, + 3: {change.FullUpdate()}, + } + + for id, w := range want { + assert.Equal(t, w, getPendingForNode(lb.b, id), "node %d", id) + } + + // A second full neither stacks nor drops the rescued ping. + lb.b.addToBatch(change.FullUpdate()) + + for id, w := range want { + assert.Equal(t, w, getPendingForNode(lb.b, id), "node %d after second full", id) + } +} + +// TestAddToBatch_ConcurrentPingAndFullUpdate_NoPingLoss races pings against +// fulls: every ping must end up pending exactly once, on its own node, behind +// a single full. +func TestAddToBatch_ConcurrentPingAndFullUpdate_NoPingLoss(t *testing.T) { + const ( + nodes = 4 + pings = 64 + fulls = 16 + ) + + lb := setupLightweightBatcher(t, nodes, 10) + defer lb.cleanup() + + target := func(i int) types.NodeID { + return types.NodeID(i%nodes + 1) //nolint:gosec // test with small values + } + + panics := runConcurrentlyWithTimeout(t, pings+fulls, 10*time.Second, func(i int) { + if i < pings { + lb.b.addToBatch(change.PingNode(target(i), &tailcfg.PingRequest{ + URL: fmt.Sprintf("ping-%d", i), + })) + + return + } + + lb.b.addToBatch(change.FullUpdate()) + }) + require.Zero(t, panics) + + seen := make(map[string]types.NodeID, pings) + + for id := range lb.channels { + pending := getPendingForNode(lb.b, id) + require.NotEmpty(t, pending, "node %d", id) + assert.True(t, pending[0].IsFull(), "node %d: full must lead", id) + + for _, c := range pending[1:] { + require.NotNil(t, c.PingRequest, "node %d: only pings follow the full", id) + assert.Equal(t, change.PingNode(id, c.PingRequest), c, "node %d: ping-only", id) + + prev, dup := seen[c.PingRequest.URL] + assert.False(t, dup, "%s pending on node %d and %d", c.PingRequest.URL, prev, id) + + seen[c.PingRequest.URL] = id + } + } + + for i := range pings { + url := fmt.Sprintf("ping-%d", i) + got, ok := seen[url] + + if assert.True(t, ok, "%s lost", url) { + assert.Equal(t, target(i), got, "%s on wrong node", url) + } + } +} + // TestAddToBatch_NodeRemovalCleanup verifies that a permanent node deletion // cleans up the node from the batcher's internal state. func TestAddToBatch_NodeRemovalCleanup(t *testing.T) { diff --git a/hscontrol/mapper/batcher_test.go b/hscontrol/mapper/batcher_test.go index 2559d1199..0ff1a5de6 100644 --- a/hscontrol/mapper/batcher_test.go +++ b/hscontrol/mapper/batcher_test.go @@ -1171,6 +1171,76 @@ func TestBatcherCoalescesPolicyRecomputesPerTick(t *testing.T) { } } +// TestBatcherPingSurvivesFullUpdate queues a ping and a full in one AddWork +// call against real connections: the target renders the full and then a +// ping-only frame, and no other node sees the ping. +func TestBatcherPingSurvivesFullUpdate(t *testing.T) { + for _, bf := range allBatcherFunctions { + t.Run(bf.name, func(t *testing.T) { + testData, cleanup := setupBatcherWithTestData(t, bf.fn, 1, 2, normalBufferSize) + defer cleanup() + + batcher := testData.Batcher + target, other := &testData.Nodes[0], &testData.Nodes[1] + + for _, n := range []*node{target, other} { + require.NoError(t, batcher.AddNode(n.n.ID, n.ch, tailcfg.CapabilityVersion(100), nil)) + } + + // Settle initial maps and online patches so only this call's + // frames remain. + drainChannelTimeout(target.ch, 300*time.Millisecond) + drainChannelTimeout(other.ch, 300*time.Millisecond) + + pr := &tailcfg.PingRequest{URL: "https://example.com/ping", Log: true} + batcher.AddWork(change.PingNode(target.n.ID, pr), change.UserRemoved()) + + collect := func(ch <-chan *tailcfg.MapResponse) []*tailcfg.MapResponse { + var frames []*tailcfg.MapResponse + + deadline := time.After(updateTimeout) + quiet := time.NewTimer(time.Hour) + + defer quiet.Stop() + + for { + select { + case resp := <-ch: + frames = append(frames, resp) + + quiet.Reset(300 * time.Millisecond) + case <-quiet.C: + return frames + case <-deadline: + return frames + } + } + } + + isFull := func(r *tailcfg.MapResponse) bool { + return r.Node != nil && r.DERPMap != nil && len(r.Peers) > 0 + } + + frames := collect(target.ch) + require.Len(t, frames, 2, "target: full then ping") + assert.True(t, isFull(frames[0]), "target: first frame is the full") + assert.Nil(t, frames[0].PingRequest, "target: the full carries no ping") + + ping := frames[1] + assert.Equal(t, pr, ping.PingRequest) + assert.Nil(t, ping.Node, "ping frame is ping-only") + assert.Nil(t, ping.DERPMap, "ping frame is ping-only") + assert.Empty(t, ping.Peers, "ping frame is ping-only") + assert.Empty(t, ping.PeersChangedPatch, "ping frame is ping-only") + + frames = collect(other.ch) + require.Len(t, frames, 1, "other: only the full") + assert.True(t, isFull(frames[0]), "other: frame is the full") + assert.Nil(t, frames[0].PingRequest, "other: no ping") + }) + } +} + // TestBatcherWorkerChannelSafety tests that worker goroutines handle closed // channels safely without panicking when processing work items. // diff --git a/hscontrol/mapper/node_conn.go b/hscontrol/mapper/node_conn.go index acbab5361..3ef68b67b 100644 --- a/hscontrol/mapper/node_conn.go +++ b/hscontrol/mapper/node_conn.go @@ -262,6 +262,14 @@ func (mc *multiChannelNodeConn) prependPending(changes ...change.Change) { mc.pendingMu.Unlock() } +// collapsePendingToFull replaces pending, together with incoming, by a +// single full update and the pings it cannot carry ([change.CollapseToFull]). +func (mc *multiChannelNodeConn) collapsePendingToFull(incoming []change.Change) { + mc.pendingMu.Lock() + mc.pending = change.CollapseToFull(mc.id, slices.Concat(mc.pending, incoming)) + mc.pendingMu.Unlock() +} + // drainPending atomically removes and returns all pending changes. // Returns nil if there are no pending changes. func (mc *multiChannelNodeConn) drainPending() []change.Change { diff --git a/hscontrol/servertest/ping_test.go b/hscontrol/servertest/ping_test.go index eba4add67..245ccd75e 100644 --- a/hscontrol/servertest/ping_test.go +++ b/hscontrol/servertest/ping_test.go @@ -136,6 +136,61 @@ func TestPingTwoSameNode(t *testing.T) { } } +// TestPingSurvivesFullUpdate verifies that a full update queued alongside a +// ping, in the same call or right after it, does not swallow the ping. +func TestPingSurvivesFullUpdate(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + send func(h *servertest.TestHarness, ping change.Change) + }{ + { + name: "same call", + send: func(h *servertest.TestHarness, ping change.Change) { + h.Server.App.Change(ping, change.UserRemoved()) + }, + }, + { + name: "full after pending ping", + send: func(h *servertest.TestHarness, ping change.Change) { + h.Server.App.Change(ping) + h.Server.App.Change(change.UserRemoved()) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + h := servertest.NewHarness(t, 1) + + nm := h.Client(0).Netmap() + require.NotNil(t, nm) + + nodeID := types.NodeID(nm.SelfNode.ID()) //nolint:gosec + + st := h.Server.State() + pingID, responseCh := st.RegisterPing(nodeID) + + defer st.CancelPing(pingID) + + tt.send(h, change.PingNode(nodeID, &tailcfg.PingRequest{ + URL: h.Server.URL + "/machine/ping-response?id=" + pingID, + Log: true, + })) + + select { + case latency := <-responseCh: + assert.GreaterOrEqual(t, latency, time.Duration(0)) + case <-time.After(15 * time.Second): + t.Fatal("ping lost to the full update") + } + }) + } +} + // TestPingResolveByHostname verifies that [state.State.ResolveNode] can find a node // by hostname and that the resolved node can be pinged. func TestPingResolveByHostname(t *testing.T) { diff --git a/hscontrol/types/change/change.go b/hscontrol/types/change/change.go index d3cc8001b..3c5dfe5bc 100644 --- a/hscontrol/types/change/change.go +++ b/hscontrol/types/change/change.go @@ -230,6 +230,23 @@ func HasFull(rs []Change) bool { return slices.ContainsFunc(rs, Change.IsFull) } +// CollapseToFull returns what nodeID receives when a full update supersedes +// changes: one [FullUpdate], then a ping-only [PingNode] for every +// [Change.PingRequest] addressed to nodeID, in order. A full renders state at +// drain time, so it covers every state change; a ping is a one-shot command +// the full cannot carry. +func CollapseToFull(nodeID types.NodeID, changes []Change) []Change { + out := []Change{FullUpdate()} + + for _, c := range changes { + if c.PingRequest != nil && c.ShouldSendToNode(nodeID) { + out = append(out, PingNode(nodeID, c.PingRequest)) + } + } + + return out +} + // SplitTargetedAndBroadcast separates responses into targeted (to specific node) and broadcast. func SplitTargetedAndBroadcast(rs []Change) ([]Change, []Change) { var broadcast, targeted []Change diff --git a/hscontrol/types/change/change_test.go b/hscontrol/types/change/change_test.go index abae55a7d..aa2dfc115 100644 --- a/hscontrol/types/change/change_test.go +++ b/hscontrol/types/change/change_test.go @@ -32,6 +32,82 @@ func TestChange_FieldSync(t *testing.T) { } } +// TestChange_FullUpdateSubsumesAllButPing classifies every [Change] field by +// whether a full update covers it. A full is rendered from state at drain +// time, so it covers anything describing state; a one-shot command is not +// state and must survive the collapse. A new field fails here until it is +// classified. +func TestChange_FullUpdateSubsumesAllButPing(t *testing.T) { + const self = types.NodeID(1) + + // true: the field carries a command the full cannot re-render. + carried := map[string]bool{ + "Reason": false, // logging only + "TargetNode": false, // routing; the full is queued per node + "OriginNode": false, // self detection; the full includes self + "IncludeSelf": false, + "IncludeDERPMap": false, + "IncludeDNS": false, + "IncludeDomain": false, + "IncludePolicy": false, + "PeersChanged": false, // SendAllPeers re-lists every peer + "PeersRemoved": false, // diffed from the full's peer list + "PeerPatches": false, // peers rendered from current state + "SendAllPeers": false, + "DeletedNodes": false, // torn down in addToBatch before the collapse + "RequiresRuntimePeerComputation": false, // SendAllPeers recomputes visibility + "PingRequest": true, + } + + typ := reflect.TypeFor[Change]() + require.Len(t, carried, typ.NumField(), "classify every Change field") + + for field := range typ.Fields() { + t.Run(field.Name, func(t *testing.T) { + isCarried, ok := carried[field.Name] + require.True(t, ok, "field %s is not classified", field.Name) + + var c Change + + v := reflect.ValueOf(&c).Elem().FieldByIndex(field.Index) + setNonZero(t, v, self) + require.False(t, v.IsZero()) + + got := CollapseToFull(self, []Change{c, FullUpdate()}) + + want := []Change{FullUpdate()} + if isCarried { + want = append(want, PingNode(self, c.PingRequest)) + } + + assert.Equal(t, want, got) + }) + } +} + +// setNonZero sets v to a non-zero value; node IDs get id so the change is +// addressed to the node under test. +func setNonZero(t *testing.T, v reflect.Value, id types.NodeID) { + t.Helper() + + switch v.Kind() { //nolint:exhaustive // only the kinds Change uses; default fails the test + case reflect.Bool: + v.SetBool(true) + case reflect.String: + v.SetString("set") + case reflect.Uint64: + v.SetUint(id.Uint64()) + case reflect.Pointer: + v.Set(reflect.New(v.Type().Elem())) + case reflect.Slice: + s := reflect.MakeSlice(v.Type(), 1, 1) + setNonZero(t, s.Index(0), id) + v.Set(s) + default: + t.Fatalf("setNonZero: unhandled kind %s", v.Kind()) + } +} + func TestChange_IsEmpty(t *testing.T) { tests := []struct { name string @@ -587,6 +663,71 @@ func TestPingNode(t *testing.T) { assert.Equal(t, "ping", r.Type()) } +func TestCollapseToFullKeepsPings(t *testing.T) { + const self = types.NodeID(1) + + prA := &tailcfg.PingRequest{URL: "https://example.com/ping/a"} + prB := &tailcfg.PingRequest{URL: "https://example.com/ping/b"} + prOther := &tailcfg.PingRequest{URL: "https://example.com/ping/other"} + + tests := []struct { + name string + changes []Change + want []Change + }{ + { + name: "nil yields a lone full", + changes: nil, + want: []Change{FullUpdate()}, + }, + { + name: "state changes are subsumed", + changes: []Change{NodeOnline(2), PolicyChange(), SelfUpdate(self), UserRemoved()}, + want: []Change{FullUpdate()}, + }, + { + name: "pings follow the full in order", + changes: []Change{PingNode(self, prA), NodeOnline(2), PingNode(self, prB)}, + want: []Change{FullUpdate(), PingNode(self, prA), PingNode(self, prB)}, + }, + { + name: "ping for another node is not taken", + changes: []Change{PingNode(2, prOther), PingNode(self, prA)}, + want: []Change{FullUpdate(), PingNode(self, prA)}, + }, + { + name: "merged ping is rescued ping-only", + changes: []Change{PingNode(self, prA).Merge(NodeOnline(3)).Merge(SelfUpdate(self))}, + want: []Change{FullUpdate(), PingNode(self, prA)}, + }, + { + name: "untargeted ping is addressed to every node", + changes: []Change{{Reason: "broadcast ping", PingRequest: prA}}, + want: []Change{FullUpdate(), PingNode(self, prA)}, + }, + { + name: "fulls never stack", + changes: []Change{FullUpdate(), PingNode(self, prA), UserAdded(), FullUpdate()}, + want: []Change{FullUpdate(), PingNode(self, prA)}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := CollapseToFull(self, tt.changes) + assert.Equal(t, tt.want, got) + + require.NotEmpty(t, got) + assert.True(t, got[0].IsFull(), "first entry must be the full") + + for _, c := range got[1:] { + assert.False(t, c.IsFull(), "only one full per collapse") + assert.Equal(t, PingNode(self, c.PingRequest), c, "rescued entries are ping-only") + } + }) + } +} + func TestUniqueNodeIDs(t *testing.T) { tests := []struct { name string