mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-06 23:00:09 +09:00
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:
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user