state: refresh policy nodes inside the peer map build

One peer build per tag/user/IP/route write; callers detect policy moves
via NodesGeneration. Per-node caches only store results for the node
pm holds, so a mapper reading mid-build cannot pin a stale filter.
This commit is contained in:
Kristoffer Dalby
2026-09-25 17:36:26 +00:00
parent 311d9323e0
commit 21f6e46fb8
17 changed files with 1085 additions and 148 deletions
+420 -18
View File
@@ -15,6 +15,7 @@ import (
"github.com/juanfont/headscale/hscontrol/db"
policyv2 "github.com/juanfont/headscale/hscontrol/policy/v2"
"github.com/juanfont/headscale/hscontrol/types"
"github.com/juanfont/headscale/hscontrol/types/change"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"pgregory.net/rapid"
@@ -1694,9 +1695,7 @@ func nodeStoreWithPolicy(t fatalfer, pol string, users []types.User, nodes types
t.Fatalf("policy: %v", err)
}
store := NewNodeStore(nodes, func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
return pm.BuildPeerMap(views.SliceOf(ns))
}, TestBatchSize, TestBatchTimeout)
store := NewNodeStore(nodes, policyPeersFunc(pm), TestBatchSize, TestBatchTimeout)
store.Start()
return store, pm
@@ -1718,12 +1717,21 @@ func syncPolicy(t fatalfer, store *NodeStore, pm *policyv2.PolicyManager) {
// checkAdjacencyMatchesFullBuild compares the NodeStore's current
// adjacency — however it got there, including the reused-from-previous-
// snapshot path taken for payload-only writes — against a from-scratch
// [policyv2.PolicyManager.BuildPeerMap] over the same nodes. Divergence
// means the reuse path served stale adjacency.
func checkAdjacencyMatchesFullBuild(t fatalfer, store *NodeStore, pm *policyv2.PolicyManager) {
// snapshot path taken for payload-only writes — against
// [policyv2.PolicyManager.BuildPeerMap] from a fresh policy manager over
// the same nodes. The fresh manager keeps the oracle independent of the
// one the NodeStore writer updates. Divergence means the store served
// stale adjacency.
func checkAdjacencyMatchesFullBuild(t fatalfer, store *NodeStore, pol string, users []types.User) {
snap := store.data.Load()
want := pm.BuildPeerMap(views.SliceOf(snap.allNodes))
nodes := views.SliceOf(snap.allNodes)
fresh, err := policyv2.NewPolicyManager([]byte(pol), users, nodes)
if err != nil {
t.Fatalf("fresh policy: %v", err)
}
want := fresh.BuildPeerMap(nodes)
for id := range snap.nodesByID {
got := slices.Sorted(slices.Values(snap.peersByNode[id]))
@@ -1822,17 +1830,16 @@ func TestNodeStoreAdjacencyMatchesFullBuild(t *testing.T) {
})
}
// Check before syncPolicy: at this point pm still has the
// pre-write nodes, so BuildPeerMap(new nodes) against pm's
// old matchers is exactly what the write's own reuse-vs-
// recompute decision (updateChanges) should have produced.
// Checking only after syncPolicy would let a wrong
// updateChanges classification hide behind the
// RebuildPeerMaps that SetNodes triggers on its own.
checkAdjacencyMatchesFullBuild(rt, store, pm)
// Check before syncPolicy: the write's own snapshot must
// already be right, both on the reuse path (updateChanges
// said payload-only) and on the recompute path (the
// peersFunc refreshed pm before building). Checking only
// after syncPolicy would let either mistake hide behind a
// RebuildPeerMaps.
checkAdjacencyMatchesFullBuild(rt, store, tc.pol, users)
syncPolicy(rt, store, pm)
checkAdjacencyMatchesFullBuild(rt, store, pm)
checkAdjacencyMatchesFullBuild(rt, store, tc.pol, users)
}
})
})
@@ -2007,10 +2014,11 @@ func BenchmarkNodeStoreWrite(b *testing.B) {
pm, err := policyv2.NewPolicyManager([]byte(policyGlobal), users, nodes.ViewSlice())
require.NoError(b, err)
inner := policyPeersFunc(pm)
peersFunc := func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
calls.Add(1)
return pm.BuildPeerMap(views.SliceOf(ns))
return inner(ns)
}
store := NewNodeStore(nodes, peersFunc, TestBatchSize, TestBatchTimeout)
@@ -2037,3 +2045,397 @@ func BenchmarkNodeStoreWrite(b *testing.B) {
})
}
}
// countStatePeerBuilds swaps s's NodeStore for one wired the same way
// but counting peersFunc runs, so a test can see how many O(n^2) peer
// builds a State write costs. Call before anything else uses s.
func countStatePeerBuilds(t *testing.T, s *State) *atomic.Int64 {
t.Helper()
nodes := make(types.Nodes, 0, s.nodeStore.ListNodes().Len())
for _, nv := range s.nodeStore.ListNodes().All() {
nodes = append(nodes, nv.AsStruct())
}
s.nodeStore.Stop()
var calls atomic.Int64
inner := policyPeersFunc(s.polMan)
s.nodeStore = NewNodeStore(nodes, func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
calls.Add(1)
return inner(ns)
}, TestBatchSize, TestBatchTimeout)
s.nodeStore.Start()
calls.Store(0)
return &calls
}
// peerBuildTestPolicy is the policy newPeerBuildTestState installs.
const peerBuildTestPolicy = `{
"tagOwners": {"tag:a": ["pb-user@"], "tag:b": ["pb-user@"]},
"acls": [
{"action": "accept", "src": ["tag:a"], "dst": ["tag:b:*"]},
{"action": "accept", "src": ["pb-user@"], "dst": ["10.55.0.0/24:*"]}
]}`
// checkStateAdjacencyMatchesFullBuild is checkAdjacencyMatchesFullBuild
// for a State from newPeerBuildTestState.
func checkStateAdjacencyMatchesFullBuild(t *testing.T, s *State) {
t.Helper()
users, err := s.ListAllUsers()
require.NoError(t, err)
checkAdjacencyMatchesFullBuild(t, s.nodeStore, peerBuildTestPolicy, users)
}
// newPeerBuildTestState returns a State over three user-owned nodes, the
// first announcing a subnet, under a policy where both a tag and that
// subnet decide who sees whom.
func newPeerBuildTestState(t *testing.T) (*State, []types.NodeID, *atomic.Int64) {
t.Helper()
dbPath := t.TempDir() + "/headscale.db"
cfg := persistTestConfig(dbPath)
database, err := db.NewHeadscaleDatabase(cfg)
require.NoError(t, err)
user := database.CreateUserForTest("pb-user")
nodes := database.CreateRegisteredNodesForTest(user, 3, "pb-node")
require.NoError(t, database.Close())
s, err := NewState(cfg)
require.NoError(t, err)
t.Cleanup(func() { _ = s.Close() })
_, err = s.SetPolicy([]byte(peerBuildTestPolicy))
require.NoError(t, err)
ids := make([]types.NodeID, 0, len(nodes))
for _, n := range nodes {
ids = append(ids, n.ID)
}
_, ok := s.nodeStore.UpdateNode(ids[0], func(n *types.Node) {
n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: []netip.Prefix{netip.MustParsePrefix("10.55.0.0/24")}}
})
require.True(t, ok)
return s, ids, countStatePeerBuilds(t, s)
}
// TestStatePolicyWriteBuildsPeersOnce pins that a policy-relevant State
// write costs one peer build: the NodeStore writer's own build must
// already use the matchers the written node implies, not the old ones
// followed by a second rebuild once the policy manager catches up.
func TestStatePolicyWriteBuildsPeersOnce(t *testing.T) {
tests := []struct {
name string
write func(t *testing.T, s *State, id types.NodeID) change.Change
}{
{name: "tag", write: func(t *testing.T, s *State, id types.NodeID) change.Change {
t.Helper()
_, c, err := s.SetNodeTags(id, []string{"tag:a"})
require.NoError(t, err)
return c
}},
{name: "route", write: func(t *testing.T, s *State, id types.NodeID) change.Change {
t.Helper()
_, c, err := s.SetApprovedRoutes(id, []netip.Prefix{netip.MustParsePrefix("10.55.0.0/24")})
require.NoError(t, err)
return c
}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s, ids, builds := newPeerBuildTestState(t)
c := tt.write(t, s, ids[0])
assert.Equal(t, "policy", c.Type(), "a policy-relevant write must still report a policy change")
assert.Equal(t, int64(1), builds.Load(), "peer builds for one policy-relevant write")
checkStateAdjacencyMatchesFullBuild(t, s)
})
}
}
// TestStateConcurrentTagWritesEachReportPolicyChange runs two SetNodeTags
// on different nodes at once. The NodeStore may apply both in one batch,
// so the policy manager sees both tags in a single SetNodes; each caller
// must still report a policy change for its own write, and adjacency
// must end up matching a full build. A report from only one of them could
// be sent before the other's snapshot is published.
func TestStateConcurrentTagWritesEachReportPolicyChange(t *testing.T) {
s, ids, _ := newPeerBuildTestState(t)
tags := [2]string{"tag:a", "tag:b"}
for round := range 20 {
var (
wg sync.WaitGroup
changes [2]change.Change
errs [2]error
)
for i := range 2 {
wg.Go(func() {
tag := tags[(i+round)%2]
_, changes[i], errs[i] = s.SetNodeTags(ids[1+i], []string{tag})
})
}
wg.Wait()
for i := range 2 {
require.NoError(t, errs[i])
require.True(t, changes[i].RequiresRuntimePeerComputation,
"round %d writer %d: %s must report a policy change", round, i, changes[i].Type())
require.Equal(t, ids[1+i], changes[i].OriginNode, "round %d writer %d", round, i)
}
checkStateAdjacencyMatchesFullBuild(t, s)
}
}
// TestPolicyWriteReportsAfterPublish holds the NodeStore writer between the
// policy manager's SetNodes and the snapshot swap, and lets another caller
// report in that window. The writer's own caller must still report a policy
// change: one reported in the window is computed against the old snapshot,
// so peers would keep the adjacency the write replaced.
func TestPolicyWriteReportsAfterPublish(t *testing.T) {
s, ids, _ := newPeerBuildTestState(t)
nodes := make(types.Nodes, 0, s.nodeStore.ListNodes().Len())
for _, nv := range s.nodeStore.ListNodes().All() {
nodes = append(nodes, nv.AsStruct())
}
s.nodeStore.Stop()
var (
armed atomic.Bool
reached = make(chan struct{})
release = make(chan struct{})
)
inner := policyPeersFunc(s.polMan)
s.nodeStore = NewNodeStore(nodes, func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
if armed.CompareAndSwap(true, false) {
_, err := s.polMan.SetNodes(views.SliceOf(ns))
assert.NoError(t, err)
close(reached)
<-release
}
return inner(ns)
}, TestBatchSize, TestBatchTimeout)
s.nodeStore.Start()
other := s.polMan.NodesGeneration()
armed.Store(true)
var (
wg sync.WaitGroup
tagC change.Change
err error
)
wg.Go(func() {
_, tagC, err = s.SetNodeTags(ids[1], []string{"tag:a"})
})
<-reached
published, ok := s.GetNodeByID(ids[1])
require.True(t, ok)
require.False(t, published.IsTagged(), "the tag must not be published yet")
otherC := s.policyChangeSince(other)
close(release)
wg.Wait()
require.NoError(t, err)
assert.True(t, otherC.IncludePolicy, "a caller whose window saw the move may report it early")
assert.True(t, tagC.IncludePolicy,
"the writer must report the policy change once its snapshot is published")
assert.Equal(t, ids[1], tagC.OriginNode)
}
// TestBackfillNodeIPsReportsPolicyChange pins that assigning a missing
// address reports a policy change: the address is a policy input, and
// without the change clients only learned it from whichever unrelated
// write next refreshed the policy.
func TestBackfillNodeIPsReportsPolicyChange(t *testing.T) {
dbPath := t.TempDir() + "/headscale.db"
cfg := persistTestConfig(dbPath)
database, err := db.NewHeadscaleDatabase(cfg)
require.NoError(t, err)
user := database.CreateUserForTest("bf-user")
nodes := database.CreateRegisteredNodesForTest(user, 2, "bf-node")
require.NoError(t, database.DB.Model(&types.Node{}).Where("id = ?", nodes[0].ID).Update("ipv4", nil).Error)
// Backfill copies the stored Hostinfo, which a registered client always has.
require.NoError(t, database.DB.Model(&types.Node{}).Where("1 = 1").Update("host_info", "{}").Error)
require.NoError(t, database.Close())
s, err := NewState(cfg)
require.NoError(t, err)
t.Cleanup(func() { _ = s.Close() })
backfilled, cs, err := s.BackfillNodeIPs()
require.NoError(t, err)
require.NotEmpty(t, backfilled)
assert.True(t, slices.ContainsFunc(cs, change.Change.IsBroadcastPolicyChange),
"backfill must report a policy change: %v", cs)
}
// TestPolicyCachesSurviveOldViewDuringBuild covers the window between the
// writer's SetNodes and the snapshot swap: a mapper still holding the
// written node's old view can ask for its filter or SSH policy then. The
// answer for that old view must not be cached under the node's ID, or
// the node keeps it after the swap, since nothing invalidates it again.
func TestPolicyCachesSurviveOldViewDuringBuild(t *testing.T) {
users := []types.User{{ID: 1, Name: "u1"}, {ID: 2, Name: "u2"}}
subnet := netip.MustParsePrefix("10.33.0.0/24")
pol := `{
"tagOwners": {"tag:srv": ["u1@"]},
"acls": [{"action": "accept", "src": ["u2@"], "dst": ["10.33.0.0/24:*"]}],
"ssh": [{"action": "accept", "src": ["autogroup:member"], "dst": ["autogroup:self"], "users": ["root"]}]
}`
tests := []struct {
name string
mutate func(n *types.Node)
// probe reports a property of node 1's cached artefact that the
// write flips from !want to want. It takes no *testing.T because it
// also runs on the NodeStore writer, where FailNow would hang.
probe func(pm *policyv2.PolicyManager, view types.NodeView) (bool, error)
want bool
}{
{
name: "filter after route approval",
mutate: func(n *types.Node) {
n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: []netip.Prefix{subnet}}
n.ApprovedRoutes = []netip.Prefix{subnet}
},
probe: func(pm *policyv2.PolicyManager, view types.NodeView) (bool, error) {
rules, err := pm.FilterForNode(view)
if err != nil {
return false, err
}
for _, r := range rules {
for _, d := range r.DstPorts {
if d.IP == subnet.String() {
return true, nil
}
}
}
return false, nil
},
want: true,
},
{
// A tagged node is outside autogroup:self, so tagging it must
// drop its SSH rules.
name: "ssh after tagging",
mutate: func(n *types.Node) {
n.Tags = []string{"tag:srv"}
n.UserID, n.User = nil, nil
},
probe: func(pm *policyv2.PolicyManager, view types.NodeView) (bool, error) {
sshPol, err := pm.SSHPolicy("", view)
if err != nil {
return false, err
}
return sshPol != nil && len(sshPol.Rules) > 0, nil
},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Node 3 shares node 1's user so autogroup:self has a source
// for node 1 once node 1 itself is tagged away.
owners := []int{0, 1, 0}
nodes := make(types.Nodes, 0, len(owners))
for i, o := range owners {
id := i + 1
n := createTestNode(types.NodeID(id), users[o].ID, users[o].Name, fmt.Sprintf("n%d", id)) //nolint:gosec
ip4 := netip.AddrFrom4([4]byte{100, 64, 0, byte(id)}) //nolint:gosec
n.IPv4, n.IPv6 = &ip4, nil
n.User = &users[o]
nodes = append(nodes, &n)
}
pm, err := policyv2.NewPolicyManager([]byte(pol), users, nodes.ViewSlice())
require.NoError(t, err)
var (
oldView atomic.Pointer[types.NodeView]
buildErr atomic.Pointer[error]
)
inner := policyPeersFunc(pm)
store := NewNodeStore(nodes, func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
// Stand in for a mapper that read the snapshot just before
// this write and asks between SetNodes and the swap.
if v := oldView.Load(); v != nil {
_, err := pm.SetNodes(views.SliceOf(ns))
if err == nil {
_, err = tt.probe(pm, *v)
}
if err != nil {
buildErr.CompareAndSwap(nil, &err)
}
}
return inner(ns)
}, TestBatchSize, TestBatchTimeout)
store.Start()
defer store.Stop()
before, ok := store.GetNode(1)
require.True(t, ok)
got, err := tt.probe(pm, before)
require.NoError(t, err)
require.NotEqual(t, tt.want, got, "precondition: the write must flip the probed artefact")
oldView.Store(&before)
after, ok := store.UpdateNode(1, tt.mutate)
require.True(t, ok)
oldView.Store(nil)
if e := buildErr.Load(); e != nil {
require.NoError(t, *e, "probe during peer map build")
}
got, err = tt.probe(pm, after)
require.NoError(t, err)
assert.Equal(t, tt.want, got, "cached artefact must reflect the written node")
})
}
}