From 785ba22c65e1f4ca948901ab04d556305a0387c2 Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Fri, 2 Oct 2026 09:58:42 +0000 Subject: [PATCH] mapper: send SSHPolicy only when it changes for the connection Any non-nil SSHPolicy forces a full client netmap rebuild; the empty policy now sent to every node made each policy change one. Updates #3508 --- hscontrol/mapper/batcher.go | 3 + hscontrol/mapper/batcher_test.go | 101 ++++++++++++++++++++++++++ hscontrol/mapper/batcher_unit_test.go | 83 +++++++++++++++++++++ hscontrol/mapper/builder.go | 2 + hscontrol/mapper/mapper.go | 3 +- hscontrol/mapper/node_conn.go | 24 +++++- hscontrol/servertest/policy_test.go | 36 +++++++++ 7 files changed, 250 insertions(+), 2 deletions(-) diff --git a/hscontrol/mapper/batcher.go b/hscontrol/mapper/batcher.go index c70a4197f..985bf1f40 100644 --- a/hscontrol/mapper/batcher.go +++ b/hscontrol/mapper/batcher.go @@ -342,6 +342,9 @@ func (b *Batcher) AddNode( // and we want to avoid the race condition where the receiver isn't ready yet select { case c <- initialMap: + // Still pendingInitial, so no broadcast can race this. + newEntry.lastSSHPolicy.Store(initialMap.SSHPolicy) + // Record sent peers only after confirmed delivery, mirroring the async // path, and under workMu so a concurrent async bundle for this node // cannot interleave its own lastSentPeers update. diff --git a/hscontrol/mapper/batcher_test.go b/hscontrol/mapper/batcher_test.go index 55b4a555a..798b1e34a 100644 --- a/hscontrol/mapper/batcher_test.go +++ b/hscontrol/mapper/batcher_test.go @@ -2826,3 +2826,104 @@ func TestSelfSentOnlyWhenChanged(t *testing.T) { assert.Contains(t, frames[0].Node.Name, "renamed-again") assert.Nil(t, frames[1].Node, "second connection already holds the new self") } + +// TestSSHPolicySentOnlyWhenChanged pins that policy frames drop an SSHPolicy +// the client already holds, since any non-nil SSHPolicy forces a full client +// netmap rebuild, while initial maps and real changes still carry it. +// https://github.com/juanfont/headscale/issues/3508 +func TestSSHPolicySentOnlyWhenChanged(t *testing.T) { + testData, cleanup := setupBatcherWithTestData(t, NewBatcherAndMapper, 1, 2, normalBufferSize) + defer cleanup() + + b := testData.Batcher.Batcher + self := &testData.Nodes[0] + + policy := func(sshUser, group string) []byte { + ssh := "" + if sshUser != "" { + ssh = fmt.Sprintf(`, "ssh": [{ + "action": "accept", + "src": ["autogroup:member"], + "dst": ["autogroup:self"], + "users": [%q] + }]`, sshUser) + } + + return fmt.Appendf(nil, `{ + "groups": {%q: []}, + "acls": [{"action": "accept", "src": ["*"], "dst": ["*:*"]}]%s + }`, "group:"+group, ssh) + } + + _, err := testData.State.SetPolicy(policy("root", "a")) + require.NoError(t, err) + + require.NoError(t, b.AddNode(self.n.ID, self.ch, tailcfg.CapabilityVersion(100), nil)) + + initial := expectReceive(t, self.ch, "initial map") + require.NotNil(t, initial.SSHPolicy) + require.NotEmpty(t, initial.SSHPolicy.Rules) + + nc, ok := b.nodes.Load(self.n.ID) + require.True(t, ok) + + policyFrame := func(chs ...chan *tailcfg.MapResponse) []*tailcfg.MapResponse { + nc.workMu.Lock() + defer nc.workMu.Unlock() + + require.NoError(t, handleNodeChange(nc, b.mapper, change.PolicyChange())) + + frames := make([]*tailcfg.MapResponse, 0, len(chs)) + for _, ch := range chs { + frames = append(frames, expectReceive(t, ch, "policy frame")) + } + + return frames + } + + steps := []struct { + name string + policy []byte // nil: no policy change + wantSent bool + wantSSH bool + }{ + {"nothing changed", nil, false, false}, + {"acl-only change", policy("root", "b"), false, false}, + {"ssh users changed", policy("alice", "b"), true, true}, + {"ssh removed", policy("", "b"), true, false}, + {"acl-only change without ssh", policy("", "c"), false, false}, + } + + for _, step := range steps { + if step.policy != nil { + _, err := testData.State.SetPolicy(step.policy) + require.NoError(t, err) + } + + frame := policyFrame(self.ch)[0] + + if !step.wantSent { + assert.Nil(t, frame.SSHPolicy, "%s: unchanged SSHPolicy must be dropped", step.name) + + _, delta := netmap.MutationsFromMapResponse(frame, time.Now()) + assert.True(t, delta, "%s: frame must apply as a delta", step.name) + + continue + } + + require.NotNil(t, frame.SSHPolicy, "%s: changed SSHPolicy must be sent", step.name) + assert.Equal(t, step.wantSSH, len(frame.SSHPolicy.Rules) > 0, step.name) + } + + // A new connection gets the policy in its initial map; afterwards both + // connections hold it and drop it from the next policy frame. + ch2 := make(chan *tailcfg.MapResponse, normalBufferSize) + require.NoError(t, b.AddNode(self.n.ID, ch2, tailcfg.CapabilityVersion(100), nil)) + + initial2 := expectReceive(t, ch2, "second connection's initial map") + require.NotNil(t, initial2.SSHPolicy, "every initial map must carry SSHPolicy") + + for i, frame := range policyFrame(self.ch, ch2) { + assert.Nil(t, frame.SSHPolicy, "connection %d already holds the policy", i) + } +} diff --git a/hscontrol/mapper/batcher_unit_test.go b/hscontrol/mapper/batcher_unit_test.go index 5337ab179..7995e346b 100644 --- a/hscontrol/mapper/batcher_unit_test.go +++ b/hscontrol/mapper/batcher_unit_test.go @@ -1194,3 +1194,86 @@ func TestRemoveConnectionAtIndex_NilsTrailingSlot(t *testing.T) { mc.mutex.Unlock() } + +// ============================================================================ +// SSHPolicy delta Tests +// ============================================================================ + +func sshPolicyForUser(user string) *tailcfg.SSHPolicy { + return &tailcfg.SSHPolicy{Rules: []*tailcfg.SSHRule{{ + Principals: []*tailcfg.SSHPrincipal{{NodeIP: "100.64.0.1"}}, + SSHUsers: map[string]string{user: user}, + Action: &tailcfg.SSHAction{Accept: true}, + }}} +} + +func emptySSHPolicy() *tailcfg.SSHPolicy { + return &tailcfg.SSHPolicy{Rules: []*tailcfg.SSHRule{}} +} + +func TestMultiChannelSend_SSHPolicyOnlyWhenChanged(t *testing.T) { + tests := []struct { + name string + last, next *tailcfg.SSHPolicy + wantSent bool + }{ + {"unchanged empty", emptySSHPolicy(), emptySSHPolicy(), false}, + {"unchanged rules", sshPolicyForUser("root"), sshPolicyForUser("root"), false}, + {"rules removed", sshPolicyForUser("root"), emptySSHPolicy(), true}, + {"rules changed", sshPolicyForUser("root"), sshPolicyForUser("alice"), true}, + {"none delivered yet", nil, emptySSHPolicy(), true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + mc := newMultiChannelNodeConn(1, nil) + ch := make(chan *tailcfg.MapResponse, 1) + entry := makeConnectionEntry("conn", ch) + entry.lastSSHPolicy.Store(tt.last) + mc.addConnection(entry) + + data := testMapResponse() + data.PacketFilters = map[string][]tailcfg.FilterRule{"base": nil} + data.SSHPolicy = tt.next + + require.NoError(t, mc.send(data)) + + got := expectReceive(t, ch, "connection should receive the response") + assert.Equal(t, data.PacketFilters, got.PacketFilters, "other fields must be kept") + + if tt.wantSent { + assert.Same(t, tt.next, got.SSHPolicy) + } else { + assert.Nil(t, got.SSHPolicy) + } + + assert.Same(t, tt.next, entry.lastSSHPolicy.Load(), "latest copy must be stored") + assert.Same(t, tt.next, data.SSHPolicy, "shared response must not be modified") + }) + } +} + +// TestMultiChannelSend_SSHPolicyPerConnection pins per-connection tracking: +// connections of one node can hold different policies, e.g. one that missed +// deltas while its initial map was in flight. +func TestMultiChannelSend_SSHPolicyPerConnection(t *testing.T) { + mc := newMultiChannelNodeConn(1, nil) + + chA := make(chan *tailcfg.MapResponse, 1) + chB := make(chan *tailcfg.MapResponse, 1) + a := makeConnectionEntry("a", chA) + b := makeConnectionEntry("b", chB) + + a.lastSSHPolicy.Store(sshPolicyForUser("alice")) + b.lastSSHPolicy.Store(sshPolicyForUser("root")) + mc.addConnection(a) + mc.addConnection(b) + + data := testMapResponse() + data.SSHPolicy = sshPolicyForUser("alice") + + require.NoError(t, mc.send(data)) + + assert.Nil(t, expectReceive(t, chA, "a").SSHPolicy, "a already holds this policy") + assert.Same(t, data.SSHPolicy, expectReceive(t, chB, "b").SSHPolicy, "b holds an older policy") +} diff --git a/hscontrol/mapper/builder.go b/hscontrol/mapper/builder.go index 68f03161e..e3f794d58 100644 --- a/hscontrol/mapper/builder.go +++ b/hscontrol/mapper/builder.go @@ -347,6 +347,8 @@ func (b *MapResponseBuilder) Build() (*tailcfg.MapResponse, error) { return nil, multierr.New(b.errs...) } + // Dumps show the generated response; an unchanged SSHPolicy is + // dropped later, per connection, before it reaches the wire. if debugDumpMapResponsePath != "" { writeDebugMapResponse(b.resp, b.debugType, b.nodeID) } diff --git a/hscontrol/mapper/mapper.go b/hscontrol/mapper/mapper.go index eb96090fd..361178033 100644 --- a/hscontrol/mapper/mapper.go +++ b/hscontrol/mapper/mapper.go @@ -308,7 +308,8 @@ func (m *mapper) selfMapResponse( // [handleNodeChange]. It sends: // - PeersChanged for remaining peers (their AllowedIPs may have changed due to policy) // - Updated PacketFilters -// - Updated SSHPolicy (SSH rules may reference users/groups that changed) +// - Updated SSHPolicy (SSH rules may reference users/groups that changed); +// dropped per connection when unchanged, see [connectionEntry.withSSHPolicyDelta] // - The node's own self info, which renders from the same state as peers; // dropped per connection when unchanged, see [connectionEntry.withSelfDelta] // diff --git a/hscontrol/mapper/node_conn.go b/hscontrol/mapper/node_conn.go index 386a6862c..a43138a9a 100644 --- a/hscontrol/mapper/node_conn.go +++ b/hscontrol/mapper/node_conn.go @@ -3,6 +3,7 @@ package mapper import ( "errors" "fmt" + "reflect" "slices" "strconv" "sync" @@ -56,6 +57,9 @@ type connectionEntry struct { // established connection must go through [multiChannelNodeConn.send] // to keep it current. lastSelf atomic.Pointer[tailcfg.Node] + + // lastSSHPolicy is the last non-nil policy delivered to this connection. + lastSSHPolicy atomic.Pointer[tailcfg.SSHPolicy] } // withSelfDelta returns data without its Node when this client already @@ -71,6 +75,21 @@ func (entry *connectionEntry) withSelfDelta(data *tailcfg.MapResponse) *tailcfg. return &stripped } +// withSSHPolicyDelta returns data without its SSHPolicy when this client +// already holds an equal one: any non-nil SSHPolicy forces a full client +// netmap rebuild. Equal content arrives in fresh pointers after every +// policy reload, so compare deeply. +func (entry *connectionEntry) withSSHPolicyDelta(data *tailcfg.MapResponse) *tailcfg.MapResponse { + if data.SSHPolicy == nil || !reflect.DeepEqual(entry.lastSSHPolicy.Load(), data.SSHPolicy) { + return data + } + + stripped := *data + stripped.SSHPolicy = nil + + return &stripped +} + // multiChannelNodeConn manages multiple concurrent connections for a single node. type multiChannelNodeConn struct { id types.NodeID @@ -360,7 +379,7 @@ func (mc *multiChannelNodeConn) send(data *tailcfg.MapResponse) error { ) for _, conn := range snapshot { - err := conn.send(conn.withSelfDelta(data)) + err := conn.send(conn.withSSHPolicyDelta(conn.withSelfDelta(data))) if err != nil { lastErr = err @@ -375,6 +394,9 @@ func (mc *multiChannelNodeConn) send(data *tailcfg.MapResponse) error { if data.Node != nil { conn.lastSelf.Store(data.Node) } + if data.SSHPolicy != nil { + conn.lastSSHPolicy.Store(data.SSHPolicy) + } } } diff --git a/hscontrol/servertest/policy_test.go b/hscontrol/servertest/policy_test.go index 569328cc5..45b07d45f 100644 --- a/hscontrol/servertest/policy_test.go +++ b/hscontrol/servertest/policy_test.go @@ -1,6 +1,7 @@ package servertest_test import ( + "fmt" "net/netip" "testing" "time" @@ -157,6 +158,41 @@ func TestPolicyChanges(t *testing.T) { }) }) + // An unchanged SSHPolicy is dropped from policy frames; the client + // must keep the rules it holds. + t.Run("unrelated_policy_change_keeps_client_ssh_rules", func(t *testing.T) { + t.Parallel() + h := servertest.NewHarness(t, 2) + + withSSH := func(group string) []byte { + return fmt.Appendf(nil, `{ + "groups": {%q: []}, + "acls": [{"action": "accept", "src": ["*"], "dst": ["*:*"]}], + "ssh": [{ + "action": "accept", + "src": ["autogroup:member"], + "dst": ["autogroup:self"], + "users": ["root"] + }] + }`, "group:"+group) + } + + h.ChangePolicy(t, withSSH("a")) + h.Client(0).WaitForCondition(t, "SSH rules present", 10*time.Second, + func(nm *netmap.NetworkMap) bool { + return nm.SSHPolicy != nil && len(nm.SSHPolicy.Rules) > 0 + }) + + countBefore := h.Client(0).UpdateCount() + + h.ChangePolicy(t, withSSH("b")) + h.Client(0).WaitForCondition(t, "policy update with SSH rules kept", 10*time.Second, + func(nm *netmap.NetworkMap) bool { + return h.Client(0).UpdateCount() > countBefore && + nm.SSHPolicy != nil && len(nm.SSHPolicy.Rules) > 0 + }) + }) + t.Run("policy_with_multiple_users", func(t *testing.T) { t.Parallel()