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
This commit is contained in:
Kristoffer Dalby
2026-10-02 09:58:42 +00:00
committed by Kristoffer Dalby
parent 41c137bb3b
commit 785ba22c65
7 changed files with 250 additions and 2 deletions
+3
View File
@@ -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.
+101
View File
@@ -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)
}
}
+83
View File
@@ -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")
}
+2
View File
@@ -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)
}
+2 -1
View File
@@ -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]
//
+23 -1
View File
@@ -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)
}
}
}
+36
View File
@@ -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()