diff --git a/hscontrol/policy/pm.go b/hscontrol/policy/pm.go index 04269a672..78bf62438 100644 --- a/hscontrol/policy/pm.go +++ b/hscontrol/policy/pm.go @@ -30,6 +30,9 @@ type PolicyManager interface { // adjacency may have changed. Both are false when the list is unchanged. SetUsers(users []types.User) (policyChanged, peerMapChanged bool, err error) SetNodes(nodes views.Slice[types.NodeView]) (bool, error) + // NodesGeneration counts SetNodes calls that reported a change; see + // [policyv2.PolicyManager.NodesGeneration]. + NodesGeneration() uint64 // NodeCanHaveTag reports whether the given node can have the given tag. NodeCanHaveTag(node types.NodeView, tag string) bool diff --git a/hscontrol/policy/v2/policy.go b/hscontrol/policy/v2/policy.go index 4939c27fb..8ca9819fe 100644 --- a/hscontrol/policy/v2/policy.go +++ b/hscontrol/policy/v2/policy.go @@ -10,6 +10,7 @@ import ( "slices" "strings" "sync" + "sync/atomic" "time" "github.com/juanfont/headscale/hscontrol/policy/matcher" @@ -88,6 +89,11 @@ type PolicyManager struct { nodeAttrsMap map[types.NodeID]tailcfg.NodeCapMap nodeAttrsHashes map[types.NodeID]deephash.Sum nodeAttrsChanged []types.NodeID + + // nodesGen counts SetNodes calls that reported a change, so a caller + // can tell the policy moved even when another goroutine (the + // NodeStore writer) applied the SetNodes. + nodesGen atomic.Uint64 } // filterAndPolicy combines the compiled filter rules with policy content for hashing. @@ -917,12 +923,25 @@ func (pm *PolicyManager) SetNodes(nodes views.Slice[types.NodeView]) (bool, erro } // Always return true when nodes changed, even if filter hash didn't change // (can happen with autogroup:self or when nodes are added but don't affect rules) + pm.nodesGen.Add(1) + return true, nil } return false, nil } +// NodesGeneration returns how many SetNodes calls have reported a change. +// A value past the last one a caller acted on means some SetNodes since +// then, possibly run on another goroutine, moved the policy. +func (pm *PolicyManager) NodesGeneration() uint64 { + if pm == nil { + return 0 + } + + return pm.nodesGen.Load() +} + // nodeIDViewMap indexes a slice of node views by node ID. On duplicate IDs the // last view wins, matching the open-coded loops it replaces. func nodeIDViewMap(s views.Slice[types.NodeView]) map[types.NodeID]types.NodeView { diff --git a/hscontrol/policy/v2/policy_test.go b/hscontrol/policy/v2/policy_test.go index 77c49a2b0..87b9ab3df 100644 --- a/hscontrol/policy/v2/policy_test.go +++ b/hscontrol/policy/v2/policy_test.go @@ -2869,3 +2869,33 @@ func BenchmarkSetNodes(b *testing.B) { }) } } + +// TestNodesGenerationCountsChangingSetNodes pins that NodesGeneration +// moves exactly when SetNodes reports a change, so a caller that did not +// run the SetNodes itself can still tell its write moved the policy. +func TestNodesGenerationCountsChangingSetNodes(t *testing.T) { + users := types.Users{{ID: 1, Name: "user1"}} + + nodes := types.Nodes{node("n1", "100.64.0.1", "fd7a:115c:a1e0::1", users[0])} + nodes[0].ID = 1 + + pm, err := NewPolicyManager([]byte(`{ + "acls": [{"action": "accept", "src": ["user1@"], "dst": ["user1@:*"]}] + }`), users, nodes.ViewSlice()) + require.NoError(t, err) + + gen := pm.NodesGeneration() + + changed, err := pm.SetNodes(nodes.ViewSlice()) + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, gen, pm.NodesGeneration(), "unchanged nodes must not advance the generation") + + added := node("n2", "100.64.0.2", "fd7a:115c:a1e0::2", users[0]) + added.ID = 2 + + changed, err = pm.SetNodes(append(nodes, added).ViewSlice()) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, gen+1, pm.NodesGeneration(), "a changing SetNodes must advance the generation once") +}