mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-10 08:40:07 +09:00
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:
committed by
Kristoffer Dalby
parent
41c137bb3b
commit
785ba22c65
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
//
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user