mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-06 14:50:07 +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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -47,7 +47,7 @@ func TestPersistNodeDoesNotClobberConcurrentAdminWrite(t *testing.T) {
|
||||
"precondition: admin SetNodeTags must have written the tag to the DB")
|
||||
|
||||
// (3) Map-request persists its stale snapshot.
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(staleView)
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(staleView, s.polMan.NodesGeneration())
|
||||
require.NoError(t, err)
|
||||
|
||||
// The admin write must survive.
|
||||
|
||||
@@ -3,6 +3,7 @@ package state
|
||||
import (
|
||||
"errors"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -141,7 +142,7 @@ func TestPersistEmptyTags(t *testing.T) {
|
||||
seeded, ok := s.nodeStore.GetNode(nodeID)
|
||||
require.True(t, ok)
|
||||
|
||||
_, _, err := s.persistNodeAndRefreshPolicy(seeded)
|
||||
_, _, err := s.persistNodeAndRefreshPolicy(seeded, s.polMan.NodesGeneration())
|
||||
require.NoError(t, err)
|
||||
|
||||
gotAfterSeed, err := s.DB().GetNodeByID(nodeID)
|
||||
@@ -154,7 +155,7 @@ func TestPersistEmptyTags(t *testing.T) {
|
||||
})
|
||||
require.True(t, ok)
|
||||
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(cleared)
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(cleared, s.polMan.NodesGeneration())
|
||||
require.NoError(t, err)
|
||||
|
||||
gotAfterClear, err := s.DB().GetNodeByID(nodeID)
|
||||
@@ -189,7 +190,7 @@ func TestPersistEmptyEndpoints(t *testing.T) {
|
||||
seeded, ok := s.nodeStore.GetNode(nodeID)
|
||||
require.True(t, ok)
|
||||
|
||||
_, _, err := s.persistNodeAndRefreshPolicy(seeded)
|
||||
_, _, err := s.persistNodeAndRefreshPolicy(seeded, s.polMan.NodesGeneration())
|
||||
require.NoError(t, err)
|
||||
|
||||
gotAfterSeed, err := s.DB().GetNodeByID(nodeID)
|
||||
@@ -202,7 +203,7 @@ func TestPersistEmptyEndpoints(t *testing.T) {
|
||||
})
|
||||
require.True(t, ok)
|
||||
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(cleared)
|
||||
_, _, err = s.persistNodeAndRefreshPolicy(cleared, s.polMan.NodesGeneration())
|
||||
require.NoError(t, err)
|
||||
|
||||
gotAfterClear, err := s.DB().GetNodeByID(nodeID)
|
||||
@@ -724,12 +725,14 @@ func TestPersistNodeAndRefreshPolicyEmptyForPayloadOnlyChange(t *testing.T) {
|
||||
_, s, nodeID := persistTestSetup(t)
|
||||
t.Cleanup(func() { _ = s.Close() })
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
view, ok := s.nodeStore.UpdateNode(nodeID, func(n *types.Node) {
|
||||
n.Hostinfo = &tailcfg.Hostinfo{Hostname: "payload-only"}
|
||||
})
|
||||
require.True(t, ok)
|
||||
|
||||
_, c, err := s.persistNodeAndRefreshPolicy(view)
|
||||
_, c, err := s.persistNodeAndRefreshPolicy(view, genBefore)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, c.IsEmpty(), "a payload-only write must not fabricate a change")
|
||||
}
|
||||
@@ -871,3 +874,182 @@ func TestPersistCallerChangeDecisions(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestRetryAfterFailedPersistReportsPolicyChange proves a policy move made by
|
||||
// a write whose database persist failed is still reported. The NodeStore write
|
||||
// already fed the policy manager, so an identical retry sees nothing new; if
|
||||
// neither call reported the move, clients would keep the filter and SSH
|
||||
// policy the write revoked.
|
||||
func TestRetryAfterFailedPersistReportsPolicyChange(t *testing.T) {
|
||||
_, s, nodeID := persistTestSetup(t)
|
||||
t.Cleanup(func() { _ = s.Close() })
|
||||
|
||||
_, err := s.SetPolicy([]byte(`{
|
||||
"tagOwners": {"tag:ci": ["persist-user@"]},
|
||||
"acls": [{"action": "accept", "src": ["autogroup:member"], "dst": ["autogroup:self:*"]}]
|
||||
}`))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, s.db.DB.Callback().Update().Before("gorm:update").
|
||||
Register("fail_node_update", func(tx *gorm.DB) {
|
||||
if tx.Statement.Table == "nodes" {
|
||||
_ = tx.AddError(errInjectedNodeUpdate)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(func() { _ = s.db.DB.Callback().Update().Remove("fail_node_update") })
|
||||
|
||||
_, failedC, err := s.SetNodeTags(nodeID, []string{"tag:ci"})
|
||||
require.ErrorIs(t, err, errInjectedNodeUpdate)
|
||||
require.NoError(t, s.db.DB.Callback().Update().Remove("fail_node_update"))
|
||||
|
||||
_, c, err := s.SetNodeTags(nodeID, []string{"tag:ci"})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, failedC.IncludePolicy || c.IncludePolicy,
|
||||
"the failed call or its retry must report the policy move: failed %s, retry %s",
|
||||
failedC.Type(), c.Type())
|
||||
assert.Equal(t, nodeID, c.OriginNode)
|
||||
|
||||
if !failedC.IsEmpty() {
|
||||
assert.Equal(t, nodeID, failedC.OriginNode,
|
||||
"a failed call's change must still refresh the tagged node's self view")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFailedPersistPolicyChangeSurvivesDroppedChange covers a policy move
|
||||
// whose write failed to persist, followed by an unrelated map request whose
|
||||
// change never reaches the batcher (its initial map failed) and then a
|
||||
// successful retry. The changes that do get published must still carry the
|
||||
// policy, or clients keep the filter and SSH policy the write revoked.
|
||||
func TestFailedPersistPolicyChangeSurvivesDroppedChange(t *testing.T) {
|
||||
_, s, nodeID := persistTestSetup(t)
|
||||
t.Cleanup(func() { _ = s.Close() })
|
||||
|
||||
_, err := s.SetPolicy([]byte(`{
|
||||
"tagOwners": {"tag:ci": ["persist-user@"]},
|
||||
"acls": [{"action": "accept", "src": ["autogroup:member"], "dst": ["autogroup:self:*"]}]
|
||||
}`))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, s.db.DB.Callback().Update().Before("gorm:update").
|
||||
Register("fail_node_update", func(tx *gorm.DB) {
|
||||
if tx.Statement.Table == "nodes" {
|
||||
_ = tx.AddError(errInjectedNodeUpdate)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(func() { _ = s.db.DB.Callback().Update().Remove("fail_node_update") })
|
||||
|
||||
_, failedC, err := s.SetNodeTags(nodeID, []string{"tag:ci"})
|
||||
require.ErrorIs(t, err, errInjectedNodeUpdate)
|
||||
require.NoError(t, s.db.DB.Callback().Update().Remove("fail_node_update"))
|
||||
|
||||
nv, ok := s.GetNodeByID(nodeID)
|
||||
require.True(t, ok)
|
||||
|
||||
// The map request's change is dropped, as when its initial map fails.
|
||||
_, err = s.UpdateNodeFromMapRequest(nodeID, tailcfg.MapRequest{
|
||||
NodeKey: nv.NodeKey(),
|
||||
DiscoKey: nv.DiscoKey(),
|
||||
Hostinfo: &tailcfg.Hostinfo{Hostname: nv.Hostname()},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, retryC, err := s.SetNodeTags(nodeID, []string{"tag:ci"})
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, failedC.IncludePolicy || retryC.IncludePolicy,
|
||||
"published changes must carry the policy move: failed %s, retry %s",
|
||||
failedC.Type(), retryC.Type())
|
||||
}
|
||||
|
||||
// TestReloadPolicyReturnsChangesOnAutoApproveFailure covers a policy reload
|
||||
// whose route auto-approval fails to persist. The NodeStore write already
|
||||
// fed the approved routes to the policy manager, so no later write sees the
|
||||
// policy move again; the reload must return its changes with the error for
|
||||
// the caller to publish.
|
||||
func TestReloadPolicyReturnsChangesOnAutoApproveFailure(t *testing.T) {
|
||||
_, s, nodeID := persistTestSetup(t)
|
||||
t.Cleanup(func() { _ = s.Close() })
|
||||
|
||||
route := netip.MustParsePrefix("10.9.0.0/24")
|
||||
_, ok := s.nodeStore.UpdateNode(nodeID, func(n *types.Node) {
|
||||
n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: []netip.Prefix{route}}
|
||||
})
|
||||
require.True(t, ok)
|
||||
|
||||
_, err := s.db.SetPolicy(`{
|
||||
"acls": [{"action": "accept", "src": ["*"], "dst": ["*:*"]}],
|
||||
"autoApprovers": {"routes": {"10.9.0.0/24": ["persist-user@"]}}
|
||||
}`)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, s.db.DB.Callback().Update().Before("gorm:update").
|
||||
Register("fail_node_update", func(tx *gorm.DB) {
|
||||
if tx.Statement.Table == "nodes" {
|
||||
_ = tx.AddError(errInjectedNodeUpdate)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(func() { _ = s.db.DB.Callback().Update().Remove("fail_node_update") })
|
||||
|
||||
cs, err := s.ReloadPolicy()
|
||||
require.ErrorIs(t, err, errInjectedNodeUpdate)
|
||||
|
||||
approved, ok := s.GetNodeByID(nodeID)
|
||||
require.True(t, ok)
|
||||
require.Contains(t, approved.ApprovedRoutes().AsSlice(), route,
|
||||
"the NodeStore holds the approval the database write lost")
|
||||
|
||||
assert.True(t, slices.ContainsFunc(cs, func(c change.Change) bool { return c.IncludePolicy }),
|
||||
"the reload must return its policy change with the error: %v", cs)
|
||||
}
|
||||
|
||||
// TestNodeWriteChangeWhenPolicyRefreshFails fails both the NodeStore
|
||||
// writer's SetNodes and the caller's. The write still reached the
|
||||
// NodeStore, so the change returned with the error must resend the node to
|
||||
// itself and its peers rather than be empty.
|
||||
func TestNodeWriteChangeWhenPolicyRefreshFails(t *testing.T) {
|
||||
_, s, nodeID := persistTestSetup(t)
|
||||
t.Cleanup(func() { _ = s.Close() })
|
||||
|
||||
_, err := s.SetPolicy([]byte(`{
|
||||
"tagOwners": {"tag:ci": ["persist-user@"]},
|
||||
"acls": [{"action": "accept", "src": ["*"], "dst": ["*:*"]}]
|
||||
}`))
|
||||
require.NoError(t, err)
|
||||
|
||||
nodes := make(types.Nodes, 0, s.nodeStore.ListNodes().Len())
|
||||
for _, nv := range s.nodeStore.ListNodes().All() {
|
||||
nodes = append(nodes, nv.AsStruct())
|
||||
}
|
||||
|
||||
s.nodeStore.Stop()
|
||||
s.polMan = failingSetNodesPolicyManager{PolicyManager: s.polMan}
|
||||
s.nodeStore = NewNodeStore(nodes, policyPeersFunc(s.polMan), TestBatchSize, TestBatchTimeout)
|
||||
s.nodeStore.Start()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
write func() (change.Change, error)
|
||||
}{
|
||||
{name: "SetNodeTags", write: func() (change.Change, error) {
|
||||
_, c, err := s.SetNodeTags(nodeID, []string{"tag:ci"})
|
||||
return c, err
|
||||
}},
|
||||
{name: "RenameNode", write: func() (change.Change, error) {
|
||||
_, c, err := s.RenameNode(nodeID, "renamed")
|
||||
return c, err
|
||||
}},
|
||||
{name: "SetNodeExpiry", write: func() (change.Change, error) {
|
||||
_, c, err := s.SetNodeExpiry(nodeID, nil)
|
||||
return c, err
|
||||
}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
c, err := tt.write()
|
||||
require.ErrorIs(t, err, errInjectedPolicyNodeUpdate)
|
||||
assert.Equal(t, nodeID, c.OriginNode, "change: %s", c.Type())
|
||||
assert.Contains(t, c.PeersChanged, nodeID, "change: %s", c.Type())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+176
-74
@@ -268,16 +268,7 @@ func NewState(cfg *types.Config) (*State, error) {
|
||||
|
||||
batchTimeout := cmp.Or(cfg.Tuning.NodeStoreBatchTimeout, defaultNodeStoreBatchTimeout)
|
||||
|
||||
// [policy.PolicyManager.BuildPeerMap] handles both global and per-node filter complexity.
|
||||
// This moves the complex peer relationship logic into the policy package where it belongs.
|
||||
nodeStore := NewNodeStore(
|
||||
nodes,
|
||||
func(nodes []types.NodeView) map[types.NodeID][]types.NodeID {
|
||||
return polMan.BuildPeerMap(views.SliceOf(nodes))
|
||||
},
|
||||
batchSize,
|
||||
batchTimeout,
|
||||
)
|
||||
nodeStore := NewNodeStore(nodes, policyPeersFunc(polMan), batchSize, batchTimeout)
|
||||
nodeStore.Start()
|
||||
|
||||
s := &State{
|
||||
@@ -302,6 +293,31 @@ func NewState(cfg *types.Config) (*State, error) {
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// policyPeersFunc is the [PeersFunc] [State] runs its [NodeStore] with.
|
||||
// It feeds the nodes being built to the policy manager first, so the
|
||||
// build already uses the matchers they imply; building with the old
|
||||
// matchers would serve stale adjacency until a second full rebuild.
|
||||
// [policy.PolicyManager.NodesGeneration] tells the writing caller the
|
||||
// SetNodes happened here.
|
||||
//
|
||||
// It runs on the NodeStore writer goroutine and takes the policy
|
||||
// manager's lock, which is safe only while the policy manager never
|
||||
// waits on a NodeStore write while holding it.
|
||||
func policyPeersFunc(pm policy.PolicyManager) PeersFunc {
|
||||
return func(nodes []types.NodeView) map[types.NodeID][]types.NodeID {
|
||||
slice := views.SliceOf(nodes)
|
||||
|
||||
// A failed recompile keeps the old nodes, so a caller that
|
||||
// refreshes the policy retries it and returns the error.
|
||||
_, err := pm.SetNodes(slice)
|
||||
if err != nil {
|
||||
log.Error().Err(err).Msg("refreshing policy nodes before peer map build")
|
||||
}
|
||||
|
||||
return pm.BuildPeerMap(slice)
|
||||
}
|
||||
}
|
||||
|
||||
// Close gracefully shuts down the [State] instance and releases all resources.
|
||||
func (s *State) Close() error {
|
||||
s.pings.drain()
|
||||
@@ -326,7 +342,8 @@ func (s *State) DERPMap() tailcfg.DERPMapView {
|
||||
}
|
||||
|
||||
// ReloadPolicy reloads the access control policy and triggers auto-approval if changed.
|
||||
// Returns the resulting [change.Change] slice when the policy or routes changed.
|
||||
// Returns the resulting [change.Change] slice when the policy or routes changed,
|
||||
// also alongside an error once the policy is swapped.
|
||||
func (s *State) ReloadPolicy() ([]change.Change, error) {
|
||||
pol, err := hsdb.PolicyBytes(s.db.DB, s.cfg)
|
||||
if err != nil {
|
||||
@@ -369,7 +386,9 @@ func (s *State) ReloadPolicy() ([]change.Change, error) {
|
||||
// with the current policy.
|
||||
rcs, err := s.autoApproveNodes()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("auto approving nodes: %w", err)
|
||||
// The policy is already swapped and the approvals already in the
|
||||
// NodeStore; callers publish these before handling the error.
|
||||
return append(cs, rcs...), fmt.Errorf("auto approving nodes: %w", err)
|
||||
}
|
||||
|
||||
// TODO(kradalby): These changes can probably be safely ignored.
|
||||
@@ -565,16 +584,17 @@ func (s *State) persistNode(node types.NodeView) (types.NodeView, error) {
|
||||
// persistNodeAndRefreshPolicy saves the given node state to the database and refreshes the
|
||||
// policy manager. The exact row written comes from [NodeStore]; see
|
||||
// [State.persistNode].
|
||||
func (s *State) persistNodeAndRefreshPolicy(node types.NodeView) (types.NodeView, change.Change, error) {
|
||||
// genBefore is as for [State.updatePolicyManagerNodes].
|
||||
func (s *State) persistNodeAndRefreshPolicy(node types.NodeView, genBefore uint64) (types.NodeView, change.Change, error) {
|
||||
fresh, err := s.persistNode(node)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, nodeWriteFailed(node.ID(), s.policyChangeSince(genBefore)), err
|
||||
}
|
||||
|
||||
// Check if policy manager needs updating
|
||||
c, err := s.updatePolicyManagerNodes()
|
||||
c, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return fresh, change.Change{}, fmt.Errorf("updating policy manager after node save: %w", err)
|
||||
return fresh, nodeWriteFailed(node.ID(), c), fmt.Errorf("updating policy manager after node save: %w", err)
|
||||
}
|
||||
|
||||
return fresh, c, nil
|
||||
@@ -584,10 +604,11 @@ func (s *State) SaveNode(node types.NodeView) (types.NodeView, change.Change, er
|
||||
// Update [NodeStore] first
|
||||
nodePtr := node.AsStruct()
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
resultNode := s.nodeStore.PutNode(*nodePtr)
|
||||
|
||||
// Then save to database using the result from [NodeStore.PutNode]
|
||||
return s.persistNodeAndRefreshPolicy(resultNode)
|
||||
return s.persistNodeAndRefreshPolicy(resultNode, genBefore)
|
||||
}
|
||||
|
||||
// DeleteNode permanently removes a node and cleans up associated resources.
|
||||
@@ -596,6 +617,8 @@ func (s *State) SaveNode(node types.NodeView) (types.NodeView, change.Change, er
|
||||
// publish a non-empty change before handling the error so live sessions are
|
||||
// still torn down after a committed deletion.
|
||||
func (s *State) DeleteNode(node types.NodeView) (change.Change, error) {
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
s.persistMu.Lock()
|
||||
|
||||
err := s.db.DeleteNode(node.AsStruct())
|
||||
@@ -616,9 +639,9 @@ func (s *State) DeleteNode(node types.NodeView) (change.Change, error) {
|
||||
c := change.NodeRemoved(node.ID())
|
||||
|
||||
// Check if policy manager needs updating after node deletion
|
||||
policyChange, err := s.updatePolicyManagerNodes()
|
||||
policyChange, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return c, fmt.Errorf("updating policy manager after node deletion: %w", err)
|
||||
return c.Merge(policyChange), fmt.Errorf("updating policy manager after node deletion: %w", err)
|
||||
}
|
||||
|
||||
if !policyChange.IsEmpty() {
|
||||
@@ -908,11 +931,12 @@ func (s *State) ListEphemeralNodes() views.Slice[types.NodeView] {
|
||||
func (s *State) SetNodeExpiry(nodeID types.NodeID, expiry *time.Time) (types.NodeView, change.Change, error) {
|
||||
var onlineChanged bool
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
// Update [NodeStore] before database to ensure consistency. The [NodeStore] update
|
||||
// is blocking and will be the source of truth for the batcher. The database update
|
||||
// must make the exact same change. If the database update fails, the [NodeStore]
|
||||
// change will remain, but since we return an error, no change notification will be
|
||||
// sent to the batcher, preventing inconsistent state propagation.
|
||||
// change will remain, and the change describing it is returned with the error.
|
||||
n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) {
|
||||
wasOnline := node.Online()
|
||||
node.Expiry = expiry
|
||||
@@ -926,24 +950,30 @@ func (s *State) SetNodeExpiry(nodeID types.NodeID, expiry *time.Time) (types.Nod
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("%w: %d", ErrNodeNotInNodeStore, nodeID)
|
||||
}
|
||||
|
||||
// The online flip already re-elected primaries in the NodeStore, so
|
||||
// peers need it even when the database write below fails.
|
||||
var recompute change.Change
|
||||
if onlineChanged && s.polMan.NodeNeedsPeerRecompute(n) {
|
||||
recompute = change.PolicyChange()
|
||||
}
|
||||
|
||||
// Persist expiry change to database directly since persistNodeAndRefreshPolicy omits expiry.
|
||||
err := s.db.NodeSetExpiry(nodeID, expiry)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("setting node expiry in database: %w", err)
|
||||
c := nodeWriteFailed(nodeID, s.policyChangeSince(genBefore).Merge(recompute))
|
||||
|
||||
return types.NodeView{}, c, fmt.Errorf("setting node expiry in database: %w", err)
|
||||
}
|
||||
|
||||
// Update policy manager and generate change notification.
|
||||
c, err := s.updatePolicyManagerNodes()
|
||||
c, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return n, change.Change{}, fmt.Errorf("updating policy manager after setting expiry: %w", err)
|
||||
return n, nodeWriteFailed(nodeID, c.Merge(recompute)), fmt.Errorf("updating policy manager after setting expiry: %w", err)
|
||||
}
|
||||
|
||||
// Resolve expiry and online status together from the current snapshot
|
||||
// when the mapper sends the change, including after a rapid restoration.
|
||||
c = c.Merge(change.NodeAdded(n.ID()))
|
||||
if onlineChanged && s.polMan.NodeNeedsPeerRecompute(n) {
|
||||
c = c.Merge(change.PolicyChange())
|
||||
}
|
||||
c = c.Merge(change.NodeAdded(n.ID())).Merge(recompute)
|
||||
|
||||
return n, c, nil
|
||||
}
|
||||
@@ -986,6 +1016,8 @@ func (s *State) SetNodeTags(nodeID types.NodeID, tags []string) (types.NodeView,
|
||||
// Log the operation
|
||||
logTagOperation(existingNode, validatedTags)
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
// Update [NodeStore] before database to ensure consistency. The [NodeStore] update
|
||||
// is blocking and will be the source of truth for the batcher. The database update
|
||||
// must make the exact same change.
|
||||
@@ -1000,7 +1032,7 @@ func (s *State) SetNodeTags(nodeID types.NodeID, tags []string) (types.NodeView,
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("%w: %d", ErrNodeNotInNodeStore, nodeID)
|
||||
}
|
||||
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(n)
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(n, genBefore)
|
||||
if err != nil {
|
||||
return nodeView, c, err
|
||||
}
|
||||
@@ -1026,6 +1058,7 @@ func (s *State) SetApprovedRoutes(nodeID types.NodeID, routes []netip.Prefix) (t
|
||||
// because even if the CLI removes an auto-approved route, it will be added
|
||||
// back automatically.
|
||||
prevRoutes := s.nodeStore.PrimaryRoutes()
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) {
|
||||
node.ApprovedRoutes = routes
|
||||
@@ -1042,9 +1075,9 @@ func (s *State) SetApprovedRoutes(nodeID types.NodeID, routes []netip.Prefix) (t
|
||||
}
|
||||
|
||||
// Persist the node changes to the database
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(n)
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(n, genBefore)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return nodeView, c, err
|
||||
}
|
||||
|
||||
// PolicyChange fans out a fresh netmap whenever the new approved
|
||||
@@ -1070,6 +1103,8 @@ func (s *State) RenameNode(nodeID types.NodeID, newName string) (types.NodeView,
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("%w: %w", ErrGivenNameInvalid, err)
|
||||
}
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
view, err := s.nodeStore.SetGivenName(nodeID, newName)
|
||||
if err != nil {
|
||||
switch {
|
||||
@@ -1082,7 +1117,7 @@ func (s *State) RenameNode(nodeID types.NodeID, newName string) (types.NodeView,
|
||||
}
|
||||
}
|
||||
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(view)
|
||||
nodeView, c, err := s.persistNodeAndRefreshPolicy(view, genBefore)
|
||||
if err != nil {
|
||||
return nodeView, c, err
|
||||
}
|
||||
@@ -1095,23 +1130,34 @@ func (s *State) RenameNode(nodeID types.NodeID, newName string) (types.NodeView,
|
||||
return nodeView, c, nil
|
||||
}
|
||||
|
||||
// BackfillNodeIPs assigns IP addresses to nodes that don't have them.
|
||||
func (s *State) BackfillNodeIPs() ([]string, error) {
|
||||
// BackfillNodeIPs assigns IP addresses to nodes that don't have them. The
|
||||
// returned changes tell clients about the new addresses.
|
||||
// Like the other writes, it returns the changes alongside an error once the
|
||||
// NodeStore holds new addresses; callers publish them before handling it.
|
||||
func (s *State) BackfillNodeIPs() ([]string, []change.Change, error) {
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
changes, err := s.db.BackfillNodeIPs(s.ipAlloc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
var readdressed []types.NodeID
|
||||
|
||||
// Refresh [NodeStore] after IP changes to ensure consistency
|
||||
if len(changes) > 0 {
|
||||
nodes, err := s.db.ListNodes()
|
||||
if err != nil {
|
||||
return changes, fmt.Errorf("refreshing NodeStore after IP backfill: %w", err)
|
||||
return changes, nil, fmt.Errorf("refreshing NodeStore after IP backfill: %w", err)
|
||||
}
|
||||
|
||||
for _, node := range nodes {
|
||||
// Preserve online status and NetInfo when refreshing from database
|
||||
existingNode, exists := s.nodeStore.GetNode(node.ID)
|
||||
if !exists || !slices.Equal(existingNode.IPs(), node.IPs()) {
|
||||
readdressed = append(readdressed, node.ID)
|
||||
}
|
||||
|
||||
if exists && existingNode.Valid() {
|
||||
node.IsOnline = new(existingNode.IsOnline().Get())
|
||||
|
||||
@@ -1129,7 +1175,22 @@ func (s *State) BackfillNodeIPs() ([]string, error) {
|
||||
}
|
||||
}
|
||||
|
||||
return changes, nil
|
||||
// IPs are policy inputs: without this, clients only learned the new
|
||||
// addresses from whichever unrelated write next refreshed the policy.
|
||||
c, err := s.updatePolicyManagerNodes(genBefore)
|
||||
|
||||
// A policy change carries no self node, so a readdressed node would
|
||||
// not learn its own new addresses from it.
|
||||
cs := make([]change.Change, 0, len(readdressed)+1)
|
||||
if !c.IsEmpty() {
|
||||
cs = append(cs, c)
|
||||
}
|
||||
|
||||
for _, id := range readdressed {
|
||||
cs = append(cs, change.NodeAdded(id))
|
||||
}
|
||||
|
||||
return changes, cs, err
|
||||
}
|
||||
|
||||
// ExpireExpiredNodes finds and processes expired nodes since the last check.
|
||||
@@ -1273,7 +1334,7 @@ func (s *State) AutoApproveRoutes(nv types.NodeView) (change.Change, error) {
|
||||
Err(err).
|
||||
Msg("Failed to persist auto-approved routes")
|
||||
|
||||
return change.Change{}, err
|
||||
return c, err
|
||||
}
|
||||
|
||||
log.Info().EmbedObject(nv).Strs(zf.RoutesApproved, util.PrefixesToString(approved)).Msg("routes approved")
|
||||
@@ -2295,16 +2356,19 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
expiry *time.Time,
|
||||
registrationMethod string,
|
||||
) (types.NodeView, change.Change, error) {
|
||||
// Read before any NodeStore write below; see updatePolicyManagerNodes.
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
// Get the registration entry from cache
|
||||
regEntry, ok := s.GetAuthCacheEntry(authID)
|
||||
if !ok {
|
||||
return types.NodeView{}, change.Change{}, hsdb.ErrNodeNotFoundRegistrationCache
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), hsdb.ErrNodeNotFoundRegistrationCache
|
||||
}
|
||||
|
||||
// Get the user
|
||||
user, err := s.db.GetUserByID(userID)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("finding user: %w", err)
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("finding user: %w", err)
|
||||
}
|
||||
|
||||
regData := regEntry.RegistrationData()
|
||||
@@ -2352,7 +2416,7 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
// present the machine key is in a corrupt/ambiguous state; reject rather
|
||||
// than converting an arbitrary node and orphaning the other.
|
||||
if existingNodeIsTagged && (nodeExistsForSameUser || existingNodeOwnedByOtherUser) {
|
||||
return types.NodeView{}, change.Change{}, ErrAmbiguousNodeOwnership
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), ErrAmbiguousNodeOwnership
|
||||
}
|
||||
|
||||
// Create logger with common fields for all auth operations
|
||||
@@ -2380,7 +2444,7 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
|
||||
finalNode, err = s.applyAuthNodeUpdate(updateParams)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
} else if existingNodeIsTagged {
|
||||
updateParams.ExistingNode = taggedNode
|
||||
@@ -2388,7 +2452,7 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
|
||||
finalNode, err = s.applyAuthNodeUpdate(updateParams)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
} else if existingNodeOwnedByOtherUser {
|
||||
oldUser := existingNodeOtherUser.User()
|
||||
@@ -2409,7 +2473,7 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
expiry, registrationMethod, existingNodeOtherUser,
|
||||
)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
} else {
|
||||
finalNode, err = s.createNewNodeFromAuth(
|
||||
@@ -2417,7 +2481,7 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
expiry, registrationMethod, types.NodeView{},
|
||||
)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2430,12 +2494,12 @@ func (s *State) HandleNodeFromAuthPath(
|
||||
// Update policy managers
|
||||
usersChange, err := s.updatePolicyManagerUsers()
|
||||
if err != nil {
|
||||
return finalNode, change.NodeAdded(finalNode.ID()), fmt.Errorf("updating policy manager users: %w", err)
|
||||
return finalNode, change.NodeAdded(finalNode.ID()).Merge(s.policyChangeSince(genBefore)), fmt.Errorf("updating policy manager users: %w", err)
|
||||
}
|
||||
|
||||
nodesChange, err := s.updatePolicyManagerNodes()
|
||||
nodesChange, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return finalNode, change.NodeAdded(finalNode.ID()), fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
return finalNode, change.NodeAdded(finalNode.ID()).Merge(nodesChange), fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
}
|
||||
|
||||
policyChanged := !usersChange.IsEmpty() || !nodesChange.IsEmpty()
|
||||
@@ -2550,9 +2614,12 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
// to a single node rather than racing the find-then-create section.
|
||||
defer s.lockRegistration(machineKey)()
|
||||
|
||||
// Read before any NodeStore write below; see updatePolicyManagerNodes.
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
pak, err := s.GetPreAuthKey(regReq.Auth.AuthKey)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
|
||||
// A pre-auth key node's tags come from the key, never from RequestTags.
|
||||
@@ -2570,7 +2637,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
}
|
||||
|
||||
if len(extraTags) > 0 {
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("%w %v are invalid or not permitted", ErrRequestedTagsInvalidOrNotPermitted, extraTags)
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("%w %v are invalid or not permitted", ErrRequestedTagsInvalidOrNotPermitted, extraTags)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2585,7 +2652,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
|
||||
existingNodeSameUser, existsSameUser, err := s.findExistingNodeForPAK(machineKey, pak)
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
|
||||
// For existing nodes, skip validation if:
|
||||
@@ -2657,7 +2724,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
// New node or NodeKey rotation: require valid auth key.
|
||||
err = pak.Validate()
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, err
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2705,7 +2772,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
// NodeStore NodeKey index, denying the victim service.
|
||||
if existing, ok := s.nodeStore.GetNodeByNodeKey(regReq.NodeKey); ok &&
|
||||
existing.MachineKey() != machineKey {
|
||||
return types.NodeView{}, change.Change{}, ErrNodeKeyInUse
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), ErrNodeKeyInUse
|
||||
}
|
||||
|
||||
// Snapshot the pre-update node so the NodeStore can be rolled back if
|
||||
@@ -2799,7 +2866,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
})
|
||||
|
||||
if !ok {
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("%w: %d", ErrNodeNotInNodeStore, existingNodeSameUser.ID())
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("%w: %d", ErrNodeNotInNodeStore, existingNodeSameUser.ID())
|
||||
}
|
||||
|
||||
_, err = hsdb.Write(s.db.DB, func(tx *gorm.DB) (*types.Node, error) {
|
||||
@@ -2839,7 +2906,7 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
s.nodeStore.PutNode(*priorNode)
|
||||
}
|
||||
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("writing node to database: %w", err)
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("writing node to database: %w", err)
|
||||
}
|
||||
|
||||
log.Trace().
|
||||
@@ -2928,19 +2995,19 @@ func (s *State) HandleNodeFromPreAuthKey(
|
||||
ExistingNodeForNetinfo: differentUserNode,
|
||||
})
|
||||
if err != nil {
|
||||
return types.NodeView{}, change.Change{}, fmt.Errorf("creating new node: %w", err)
|
||||
return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("creating new node: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Update policy managers
|
||||
usersChange, err := s.updatePolicyManagerUsers()
|
||||
if err != nil {
|
||||
return finalNode, change.NodeAdded(finalNode.ID()), fmt.Errorf("updating policy manager users: %w", err)
|
||||
return finalNode, change.NodeAdded(finalNode.ID()).Merge(s.policyChangeSince(genBefore)), fmt.Errorf("updating policy manager users: %w", err)
|
||||
}
|
||||
|
||||
nodesChange, err := s.updatePolicyManagerNodes()
|
||||
nodesChange, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return finalNode, change.NodeAdded(finalNode.ID()), fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
return finalNode, change.NodeAdded(finalNode.ID()).Merge(nodesChange), fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
}
|
||||
|
||||
policyChanged := !usersChange.IsEmpty() || !nodesChange.IsEmpty()
|
||||
@@ -3006,28 +3073,60 @@ func (s *State) UpdatePolicyManagerUsersForTest() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// updatePolicyManagerNodes updates the policy manager with current nodes.
|
||||
// Returns true if the policy changed and notifications should be sent.
|
||||
// updatePolicyManagerNodes refreshes the policy manager with current node
|
||||
// data and returns a PolicyChange when a node write since genBefore moved
|
||||
// the policy. genBefore is [policy.PolicyManager.NodesGeneration] read
|
||||
// before the caller's NodeStore write: the writer's [policyPeersFunc]
|
||||
// usually applies the change, so the SetNodes here alone would miss it.
|
||||
// On error the change is still returned; see [State.policyChangeSince].
|
||||
// TODO(kradalby): This is a temporary stepping stone, ultimately we should
|
||||
// have the list already available so it could go much quicker. Alternatively
|
||||
// the policy manager could have a remove or add list for nodes.
|
||||
// updatePolicyManagerNodes refreshes the policy manager with current node data.
|
||||
func (s *State) updatePolicyManagerNodes() (change.Change, error) {
|
||||
func (s *State) updatePolicyManagerNodes(genBefore uint64) (change.Change, error) {
|
||||
nodes := s.ListNodes()
|
||||
|
||||
changed, err := s.polMan.SetNodes(nodes)
|
||||
if err != nil {
|
||||
return change.Change{}, fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
return s.policyChangeSince(genBefore), fmt.Errorf("updating policy manager nodes: %w", err)
|
||||
}
|
||||
|
||||
if changed {
|
||||
// Rebuild peer maps because policy-affecting node changes (tags, user, IPs)
|
||||
// affect ACL visibility. Without this, cached peer relationships use stale data.
|
||||
// The writer refreshes the policy before every relation build, so
|
||||
// a change here means this snapshot raced another writer and moved
|
||||
// the policy manager away from what adjacency was built with.
|
||||
s.nodeStore.RebuildPeerMaps()
|
||||
return change.PolicyChange(), nil
|
||||
}
|
||||
|
||||
return change.Change{}, nil
|
||||
return s.policyChangeSince(genBefore), nil
|
||||
}
|
||||
|
||||
// nodeWriteFailed is the change a write returns with its error once its
|
||||
// NodeStore write happened: the node's new state is live, so it goes out
|
||||
// to the node itself and its peers even when the policy did not move.
|
||||
func nodeWriteFailed(id types.NodeID, c change.Change) change.Change {
|
||||
if c.IsEmpty() {
|
||||
return change.NodeAdded(id)
|
||||
}
|
||||
|
||||
c.OriginNode = id
|
||||
|
||||
return c
|
||||
}
|
||||
|
||||
// policyChangeSince returns a PolicyChange when a SetNodes since genBefore
|
||||
// moved the policy. A caller whose NodeStore write already happened returns
|
||||
// it even alongside an error, and callers publish it before handling the
|
||||
// error: the writer applied the move to the policy manager, and no later
|
||||
// caller will see it move again, so dropping it would leave clients on the
|
||||
// filter and SSH policy the write replaced. Each caller reports its own
|
||||
// window and never consumes another's, so a concurrent or failing caller
|
||||
// can only add a report, not take one away.
|
||||
func (s *State) policyChangeSince(genBefore uint64) change.Change {
|
||||
if s.polMan.NodesGeneration() != genBefore {
|
||||
return change.PolicyChange()
|
||||
}
|
||||
|
||||
return change.Change{}
|
||||
}
|
||||
|
||||
// PingDB checks if the database connection is healthy.
|
||||
@@ -3071,6 +3170,8 @@ func (s *State) autoApproveNodes() ([]change.Change, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
updates := make(map[types.NodeID]UpdateNodeFunc, len(approvedByID))
|
||||
for id, approved := range approvedByID {
|
||||
updates[id] = func(n *types.Node) {
|
||||
@@ -3094,13 +3195,13 @@ func (s *State) autoApproveNodes() ([]change.Change, error) {
|
||||
|
||||
_, err := s.persistNode(fresh)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return []change.Change{s.policyChangeSince(genBefore)}, err
|
||||
}
|
||||
}
|
||||
|
||||
c, err := s.updatePolicyManagerNodes()
|
||||
c, err := s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return []change.Change{c}, err
|
||||
}
|
||||
|
||||
if c.IsEmpty() {
|
||||
@@ -3159,6 +3260,7 @@ func (s *State) UpdateNodeFromMapRequest(id types.NodeID, req tailcfg.MapRequest
|
||||
// Snapshot the primary assignment so we can tell whether the
|
||||
// Hostinfo + auto-approval that follows shifted any prefix.
|
||||
prevRoutes := s.nodeStore.PrimaryRoutes()
|
||||
genBefore := s.polMan.NodesGeneration()
|
||||
|
||||
// We need to ensure we update the node as it is in the [NodeStore] at
|
||||
// the time of the request.
|
||||
@@ -3374,16 +3476,16 @@ func (s *State) UpdateNodeFromMapRequest(id types.NodeID, req tailcfg.MapRequest
|
||||
|
||||
updatedNode, err = s.persistNode(updatedNode)
|
||||
if err != nil {
|
||||
return change.Change{}, fmt.Errorf("saving to database: %w", err)
|
||||
return nodeWriteFailed(id, s.policyChangeSince(genBefore).Merge(nodeRouteChange)), fmt.Errorf("saving to database: %w", err)
|
||||
}
|
||||
|
||||
// Only refresh the policy manager when something it depends on
|
||||
// might have moved. Endpoint/key/DERP/LastSeen-only updates do not
|
||||
// affect policy evaluation and are deliberately skipped here.
|
||||
if delta.peerHostinfoChanged || delta.routesChanged {
|
||||
policyChange, err = s.updatePolicyManagerNodes()
|
||||
policyChange, err = s.updatePolicyManagerNodes(genBefore)
|
||||
if err != nil {
|
||||
return change.Change{}, fmt.Errorf("updating policy manager after node save: %w", err)
|
||||
return nodeWriteFailed(id, policyChange.Merge(nodeRouteChange)), fmt.Errorf("updating policy manager after node save: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user