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
+4 -2
View File
@@ -91,16 +91,18 @@ func registerAuth(api huma.API, b Backend) {
util.RegisterMethodCLI, util.RegisterMethodCLI,
) )
if err != nil { if err != nil {
b.Change(nodeChange)
return nil, mapError("registering node", err) return nil, mapError("registering node", err)
} }
routeChange, err := b.State.AutoApproveRoutes(node) routeChange, err := b.State.AutoApproveRoutes(node)
b.Change(nodeChange, routeChange)
if err != nil { if err != nil {
return nil, huma.Error500InternalServerError("auto approving routes", err) return nil, huma.Error500InternalServerError("auto approving routes", err)
} }
b.Change(nodeChange, routeChange)
out := &authRegisterOutput{} out := &authRegisterOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
+19 -14
View File
@@ -309,12 +309,12 @@ func registerNodeWriteOps(api huma.API, b Backend) {
switch { switch {
case disableExpiry: case disableExpiry:
node, nodeChange, expErr := b.State.SetNodeExpiry(nodeID, nil) node, nodeChange, expErr := b.State.SetNodeExpiry(nodeID, nil)
b.Change(nodeChange)
if expErr != nil { if expErr != nil {
return nil, mapError("expiring node", expErr) return nil, mapError("expiring node", expErr)
} }
b.Change(nodeChange)
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
@@ -324,12 +324,12 @@ func registerNodeWriteOps(api huma.API, b Backend) {
} }
node, nodeChange, err := b.State.SetNodeExpiry(nodeID, &expiry) node, nodeChange, err := b.State.SetNodeExpiry(nodeID, &expiry)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("expiring node", err) return nil, mapError("expiring node", err)
} }
b.Change(nodeChange)
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
@@ -350,12 +350,12 @@ func registerNodeWriteOps(api huma.API, b Backend) {
} }
node, nodeChange, err := b.State.RenameNode(nodeID, in.NewName) node, nodeChange, err := b.State.RenameNode(nodeID, in.NewName)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("renaming node", err) return nil, mapError("renaming node", err)
} }
b.Change(nodeChange)
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
@@ -396,12 +396,12 @@ func registerNodeWriteOps(api huma.API, b Backend) {
} }
node, nodeChange, err := b.State.SetNodeTags(nodeID, in.Body.Tags) node, nodeChange, err := b.State.SetNodeTags(nodeID, in.Body.Tags)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, huma.Error400BadRequest("setting tags", err) return nil, huma.Error400BadRequest("setting tags", err)
} }
b.Change(nodeChange)
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
@@ -444,12 +444,12 @@ func registerNodeAdminOps(api huma.API, b Backend) {
newApproved = slices.Compact(newApproved) newApproved = slices.Compact(newApproved)
node, nodeChange, err := b.State.SetApprovedRoutes(nodeID, newApproved) node, nodeChange, err := b.State.SetApprovedRoutes(nodeID, newApproved)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("setting approved routes", err) return nil, mapError("setting approved routes", err)
} }
b.Change(nodeChange)
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
// SubnetRoutes here excludes exit routes, unlike the list handler. // SubnetRoutes here excludes exit routes, unlike the list handler.
@@ -485,17 +485,20 @@ func registerNodeAdminOps(api huma.API, b Backend) {
util.RegisterMethodCLI, util.RegisterMethodCLI,
) )
if err != nil { if err != nil {
b.Change(nodeChange)
return nil, mapError("registering node", err) return nil, mapError("registering node", err)
} }
routeChange, err := b.State.AutoApproveRoutes(node) routeChange, err := b.State.AutoApproveRoutes(node)
if err != nil {
return nil, huma.Error500InternalServerError("auto approving routes", err)
}
// Empty changes are ignored by the change sink. // Empty changes are ignored by the change sink.
b.Change(nodeChange, routeChange) b.Change(nodeChange, routeChange)
if err != nil {
return nil, huma.Error500InternalServerError("auto approving routes", err)
}
out := &nodeOutput{} out := &nodeOutput{}
out.Body.Node = nodeFromView(node) out.Body.Node = nodeFromView(node)
@@ -514,7 +517,9 @@ func registerNodeAdminOps(api huma.API, b Backend) {
return nil, huma.Error400BadRequest("backfilling node IPs", errBackfillNotConfirmed) return nil, huma.Error400BadRequest("backfilling node IPs", errBackfillNotConfirmed)
} }
changes, err := b.State.BackfillNodeIPs() changes, cs, err := b.State.BackfillNodeIPs()
b.Change(cs...)
if err != nil { if err != nil {
return nil, huma.Error500InternalServerError("backfilling node IPs", err) return nil, huma.Error500InternalServerError("backfilling node IPs", err)
} }
+2 -4
View File
@@ -143,14 +143,12 @@ func registerPolicy(api huma.API, b Backend) {
// Reload even when content is unchanged: routes manually disabled before // Reload even when content is unchanged: routes manually disabled before
// may now qualify for auto-approval, so they must be re-evaluated. // may now qualify for auto-approval, so they must be re-evaluated.
cs, err := b.State.ReloadPolicy() cs, err := b.State.ReloadPolicy()
b.Change(cs...)
if err != nil { if err != nil {
return nil, huma.Error500InternalServerError("reloading policy", err) return nil, huma.Error500InternalServerError("reloading policy", err)
} }
if len(cs) > 0 {
b.Change(cs...)
}
out := &setPolicyOutput{} out := &setPolicyOutput{}
out.Body.Policy = updated.Data out.Body.Policy = updated.Data
out.Body.UpdatedAt = updated.UpdatedAt out.Body.UpdatedAt = updated.UpdatedAt
+2 -4
View File
@@ -134,14 +134,12 @@ func registerACL(api huma.API, b Backend) {
} }
cs, err := b.State.ReloadPolicy() cs, err := b.State.ReloadPolicy()
b.Change(cs...)
if err != nil { if err != nil {
return nil, huma.Error500InternalServerError("reloading policy", err) return nil, huma.Error500InternalServerError("reloading policy", err)
} }
if len(cs) > 0 {
b.Change(cs...)
}
return streamPolicy([]byte(updated.Data), aclContentType(in.Accept)), nil return streamPolicy([]byte(updated.Data), aclContentType(in.Accept)), nil
}) })
} }
+8 -8
View File
@@ -232,12 +232,12 @@ func registerDevices(api huma.API, b Backend) {
} }
_, nodeChange, err := b.State.RenameNode(node.ID(), in.Body.Name) _, nodeChange, err := b.State.RenameNode(node.ID(), in.Body.Name)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("renaming device", err) return nil, mapError("renaming device", err)
} }
b.Change(nodeChange)
return &emptyOutput{}, nil return &emptyOutput{}, nil
}) })
@@ -278,12 +278,12 @@ func registerDevices(api huma.API, b Backend) {
} }
_, nodeChange, err := b.State.SetNodeTags(node.ID(), in.Body.Tags) _, nodeChange, err := b.State.SetNodeTags(node.ID(), in.Body.Tags)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("setting device tags", err) return nil, mapError("setting device tags", err)
} }
b.Change(nodeChange)
return &emptyOutput{}, nil return &emptyOutput{}, nil
}) })
@@ -311,12 +311,12 @@ func registerDevices(api huma.API, b Backend) {
} }
_, nodeChange, err := b.State.SetNodeExpiry(node.ID(), nil) _, nodeChange, err := b.State.SetNodeExpiry(node.ID(), nil)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("setting device key expiry", err) return nil, mapError("setting device key expiry", err)
} }
b.Change(nodeChange)
return &emptyOutput{}, nil return &emptyOutput{}, nil
}) })
@@ -341,12 +341,12 @@ func registerDevices(api huma.API, b Backend) {
} }
updated, nodeChange, err := b.State.SetApprovedRoutes(node.ID(), approved) updated, nodeChange, err := b.State.SetApprovedRoutes(node.ID(), approved)
b.Change(nodeChange)
if err != nil { if err != nil {
return nil, mapError("setting device routes", err) return nil, mapError("setting device routes", err)
} }
b.Change(nodeChange)
return &deviceRoutesOutput{Body: routesFromView(updated)}, nil return &deviceRoutesOutput{Body: routesFromView(updated)}, nil
}) })
+2 -2
View File
@@ -862,13 +862,13 @@ func (h *Headscale) Serve() error {
} }
changes, err := h.state.ReloadPolicy() changes, err := h.state.ReloadPolicy()
h.Change(changes...)
if err != nil { if err != nil {
log.Error().Err(err).Msgf("reloading policy") log.Error().Err(err).Msgf("reloading policy")
continue continue
} }
h.Change(changes...)
default: default:
info := func(msg string) { log.Info().Msg(msg) } info := func(msg string) { log.Info().Msg(msg) }
+8 -5
View File
@@ -263,12 +263,12 @@ func (h *Headscale) handleLogout(
} }
updatedNode, c, err := h.state.SetNodeExpiry(node.ID(), &expiry) updatedNode, c, err := h.state.SetNodeExpiry(node.ID(), &expiry)
h.Change(c)
if err != nil { if err != nil {
return nil, fmt.Errorf("setting node expiry: %w", err) return nil, fmt.Errorf("setting node expiry: %w", err)
} }
h.Change(c)
return nodeToRegisterResponse(updatedNode), nil return nodeToRegisterResponse(updatedNode), nil
} }
@@ -420,6 +420,8 @@ func (h *Headscale) handleRegisterWithAuthKey(
machineKey, machineKey,
) )
if err != nil { if err != nil {
h.Change(changed)
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, NewHTTPError(http.StatusUnauthorized, "invalid pre auth key", nil) return nil, NewHTTPError(http.StatusUnauthorized, "invalid pre auth key", nil)
} }
@@ -451,13 +453,14 @@ func (h *Headscale) handleRegisterWithAuthKey(
// TODO(kradalby): This needs to be ran as part of the batcher maybe? // TODO(kradalby): This needs to be ran as part of the batcher maybe?
// now since we dont update the node/pol here anymore // now since we dont update the node/pol here anymore
routesChange, err := h.state.AutoApproveRoutes(node) routesChange, err := h.state.AutoApproveRoutes(node)
if err != nil {
return nil, fmt.Errorf("auto approving routes: %w", err)
}
// Send both changes. Empty changes are ignored by Change(). // Send both changes. Empty changes are ignored by Change().
h.Change(changed, routesChange) h.Change(changed, routesChange)
if err != nil {
return nil, fmt.Errorf("auto approving routes: %w", err)
}
resp := &tailcfg.RegisterResponse{ resp := &tailcfg.RegisterResponse{
MachineAuthorized: true, MachineAuthorized: true,
NodeKeyExpired: node.IsExpired(), NodeKeyExpired: node.IsExpired(),
+168
View File
@@ -1,6 +1,7 @@
package mapper package mapper
import ( import (
"errors"
"fmt" "fmt"
"net/netip" "net/netip"
"strings" "strings"
@@ -15,6 +16,7 @@ import (
"github.com/juanfont/headscale/hscontrol/types/change" "github.com/juanfont/headscale/hscontrol/types/change"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"gorm.io/gorm"
"tailscale.com/tailcfg" "tailscale.com/tailcfg"
"tailscale.com/types/dnstype" "tailscale.com/types/dnstype"
) )
@@ -811,3 +813,169 @@ func TestNoSelfAsPeerDuringRealNodeChurn(t *testing.T) {
}) })
} }
} }
// TestBackfillNodeIPsReachesBackfilledNode proves the node that receives a
// backfilled address learns it too. Peers pick it up from the policy
// refresh, but that refresh carries no self node, so the node itself kept
// serving its old addresses until an unrelated self update.
func TestBackfillNodeIPsReachesBackfilledNode(t *testing.T) {
tmp := t.TempDir()
p4 := netip.MustParsePrefix("100.64.0.0/10")
p6 := netip.MustParsePrefix("fd7a:115c:a1e0::/48")
cfg := &types.Config{
Database: types.DatabaseConfig{
Type: types.DatabaseSqlite,
Sqlite: types.SqliteConfig{Path: tmp + "/h.db"},
},
PrefixV4: &p4,
PrefixV6: &p6,
IPAllocation: types.IPAllocationStrategySequential,
BaseDomain: "headscale.test",
Policy: types.PolicyConfig{Mode: types.PolicyModeDB},
DERP: types.DERPConfig{
DERPMap: &tailcfg.DERPMap{
Regions: map[tailcfg.DERPRegionID]*tailcfg.DERPRegion{999: {RegionID: 999}},
},
},
Tuning: types.Tuning{
NodeStoreBatchSize: state.TestBatchSize,
NodeStoreBatchTimeout: state.TestBatchTimeout,
},
}
database, err := db.NewHeadscaleDatabase(cfg)
require.NoError(t, err)
user := database.CreateUserForTest("u1")
nodes := database.CreateRegisteredNodesForTest(user, 2, "bf")
target := nodes[0].ID
require.NoError(t, database.DB.Model(&types.Node{}).Where("id = ?", target).Update("ipv6", 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 := state.NewState(cfg)
require.NoError(t, err)
t.Cleanup(func() { _ = s.Close() })
backfilled, cs, err := s.BackfillNodeIPs()
require.NoError(t, err)
require.NotEmpty(t, backfilled)
stored, ok := s.GetNodeByID(target)
require.True(t, ok)
require.True(t, stored.IPv6().Valid(), "backfill must assign an IPv6 address")
m := &mapper{state: s, cfg: cfg}
nc := newMockNodeConnection(target)
var self *tailcfg.Node
for _, ch := range change.FilterForNode(target, cs) {
resp, err := generateMapResponse(nc, m, ch)
require.NoError(t, err)
if resp != nil && resp.Node != nil {
self = resp.Node
}
}
require.NotNil(t, self, "the backfilled node must receive its own node")
assert.Contains(t, self.Addresses, netip.PrefixFrom(stored.IPv6().Get(), 128))
}
var errInjectedExpiry = errors.New("injected expiry failure")
// TestFailedExpiryOfPrimaryAnnouncesBackup expires the primary of an HA
// route while the database write fails. The NodeStore already moved the
// route to the backup, so the change returned with the error must still
// tell a client about the new primary, as the successful path does;
// otherwise the client drops the old primary's route with no replacement.
func TestFailedExpiryOfPrimaryAnnouncesBackup(t *testing.T) {
tmp := t.TempDir()
p4 := netip.MustParsePrefix("100.64.0.0/10")
p6 := netip.MustParsePrefix("fd7a:115c:a1e0::/48")
cfg := &types.Config{
Database: types.DatabaseConfig{
Type: types.DatabaseSqlite,
Sqlite: types.SqliteConfig{Path: tmp + "/h.db"},
},
PrefixV4: &p4,
PrefixV6: &p6,
IPAllocation: types.IPAllocationStrategySequential,
BaseDomain: "headscale.test",
Policy: types.PolicyConfig{Mode: types.PolicyModeDB},
DERP: types.DERPConfig{
DERPMap: &tailcfg.DERPMap{
Regions: map[tailcfg.DERPRegionID]*tailcfg.DERPRegion{999: {RegionID: 999}},
},
},
Tuning: types.Tuning{
NodeStoreBatchSize: state.TestBatchSize,
NodeStoreBatchTimeout: state.TestBatchTimeout,
},
}
route := netip.MustParsePrefix("10.77.0.0/24")
database, err := db.NewHeadscaleDatabase(cfg)
require.NoError(t, err)
user := database.CreateUserForTest("u1")
nodes := database.CreateRegisteredNodesForTest(user, 3, "ha")
for _, n := range nodes[:2] {
n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: []netip.Prefix{route}}
n.ApprovedRoutes = []netip.Prefix{route}
require.NoError(t, database.DB.Save(n).Error)
}
require.NoError(t, database.Close())
s, err := state.NewState(cfg)
require.NoError(t, err)
t.Cleanup(func() { _ = s.Close() })
primary, backup, client := nodes[0].ID, nodes[1].ID, nodes[2].ID
for _, id := range []types.NodeID{primary, backup, client} {
s.Connect(id)
}
require.Equal(t, []netip.Prefix{route}, s.GetNodePrimaryRoutes(primary))
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(errInjectedExpiry)
}
}))
t.Cleanup(func() { _ = s.DB().DB.Callback().Update().Remove("fail_node_update") })
past := time.Now().Add(-time.Hour)
_, c, err := s.SetNodeExpiry(primary, &past)
require.ErrorIs(t, err, errInjectedExpiry)
require.Equal(t, []netip.Prefix{route}, s.GetNodePrimaryRoutes(backup),
"the NodeStore moved the route to the backup")
m := &mapper{state: s, cfg: cfg}
nc := newMockNodeConnection(client)
var backupRoutes []netip.Prefix
for _, ch := range change.FilterForNode(client, []change.Change{c}) {
resp, err := generateMapResponse(nc, m, ch)
require.NoError(t, err)
if resp == nil {
continue
}
for _, p := range append(resp.Peers, resp.PeersChanged...) {
if p.ID == backup.NodeID() {
backupRoutes = p.PrimaryRoutes
}
}
}
assert.Contains(t, backupRoutes, route, "the client must learn the backup is primary: %s", c.Type())
}
+6 -3
View File
@@ -993,6 +993,8 @@ func (a *AuthProviderOIDC) handleRegistration(
util.RegisterMethodOIDC, util.RegisterMethodOIDC,
) )
if err != nil { if err != nil {
a.h.Change(nodeChange)
return false, fmt.Errorf("registering node: %w", err) return false, fmt.Errorf("registering node: %w", err)
} }
@@ -1008,13 +1010,14 @@ func (a *AuthProviderOIDC) handleRegistration(
// This works, but might be another good candidate for doing some sort of // This works, but might be another good candidate for doing some sort of
// eventbus. // eventbus.
routesChange, err := a.h.state.AutoApproveRoutes(node) routesChange, err := a.h.state.AutoApproveRoutes(node)
if err != nil {
return false, fmt.Errorf("auto approving routes: %w", err)
}
// Send both changes. Empty changes are ignored by Change(). // Send both changes. Empty changes are ignored by Change().
a.h.Change(nodeChange, routesChange) a.h.Change(nodeChange, routesChange)
if err != nil {
return false, fmt.Errorf("auto approving routes: %w", err)
}
return !nodeChange.IsEmpty(), nil return !nodeChange.IsEmpty(), nil
} }
+35 -4
View File
@@ -37,6 +37,8 @@ type PolicyManager struct {
pol *Policy pol *Policy
users []types.User users []types.User
nodes views.Slice[types.NodeView] nodes views.Slice[types.NodeView]
// nodesByID indexes nodes; see [PolicyManager.cacheableLocked].
nodesByID map[types.NodeID]types.NodeView
filterHash deephash.Sum filterHash deephash.Sum
filter []tailcfg.FilterRule filter []tailcfg.FilterRule
@@ -207,6 +209,7 @@ func NewPolicyManager(b []byte, users []types.User, nodes views.Slice[types.Node
pol: policy, pol: policy,
users: users, users: users,
nodes: nodes, nodes: nodes,
nodesByID: nodeIDViewMap(nodes),
sshPolicyMap: xsync.NewMap[types.NodeID, *tailcfg.SSHPolicy](), sshPolicyMap: xsync.NewMap[types.NodeID, *tailcfg.SSHPolicy](),
filterRulesMap: xsync.NewMap[types.NodeID, []tailcfg.FilterRule](), filterRulesMap: xsync.NewMap[types.NodeID, []tailcfg.FilterRule](),
matchersForNodeMap: xsync.NewMap[types.NodeID, []matcher.Match](), matchersForNodeMap: xsync.NewMap[types.NodeID, []matcher.Match](),
@@ -444,7 +447,9 @@ func (pm *PolicyManager) SSHPolicy(baseURL string, node types.NodeView) (*tailcf
return nil, fmt.Errorf("compiling SSH policy: %w", err) return nil, fmt.Errorf("compiling SSH policy: %w", err)
} }
pm.sshPolicyMap.Store(node.ID(), sshPol) if pm.cacheableLocked(node) {
pm.sshPolicyMap.Store(node.ID(), sshPol)
}
return sshPol, nil return sshPol, nil
} }
@@ -766,7 +771,9 @@ func (pm *PolicyManager) filterForNodeLocked(
} }
reduced := policyutil.ReduceFilterRules(node, unreduced) reduced := policyutil.ReduceFilterRules(node, unreduced)
pm.filterRulesMap.Store(node.ID(), reduced) if pm.cacheableLocked(node) {
pm.filterRulesMap.Store(node.ID(), reduced)
}
return reduced return reduced
} }
@@ -822,7 +829,10 @@ func (pm *PolicyManager) MatchersForNode(node types.NodeView) ([]matcher.Match,
// the stored compiled grants for this specific node. // the stored compiled grants for this specific node.
unreduced := pm.filterRulesForNodeLocked(node) unreduced := pm.filterRulesForNodeLocked(node)
matchers := matcher.MatchesFromFilterRules(unreduced) matchers := matcher.MatchesFromFilterRules(unreduced)
pm.matchersForNodeMap.Store(node.ID(), matchers)
if pm.cacheableLocked(node) {
pm.matchersForNodeMap.Store(node.ID(), matchers)
}
return matchers, nil return matchers, nil
} }
@@ -896,7 +906,9 @@ func (pm *PolicyManager) SetNodes(nodes views.Slice[types.NodeView]) (bool, erro
// For global policies: invalidate only nodes whose properties changed (IPs, routes). // For global policies: invalidate only nodes whose properties changed (IPs, routes).
pm.invalidateNodeCache(nodes) pm.invalidateNodeCache(nodes)
prevNodes, prevByID := pm.nodes, pm.nodesByID
pm.nodes = nodes pm.nodes = nodes
pm.nodesByID = nodeIDViewMap(nodes)
// When policy-affecting node properties change, we must recompile filters because: // When policy-affecting node properties change, we must recompile filters because:
// 1. User/group aliases (like "user1@") resolve to node IPs // 1. User/group aliases (like "user1@") resolve to node IPs
@@ -912,6 +924,12 @@ func (pm *PolicyManager) SetNodes(nodes views.Slice[types.NodeView]) (bool, erro
// Recompile filter with the new node list // Recompile filter with the new node list
needsUpdate, err := pm.updateLocked() needsUpdate, err := pm.updateLocked()
if err != nil { if err != nil {
// Keep the old list so a retry with the same input recompiles
// instead of being treated as unchanged, as in SetUsers. The
// NodeStore writer only logs this error; the writing caller's
// own SetNodes then retries and returns it.
pm.nodes, pm.nodesByID = prevNodes, prevByID
return false, err return false, err
} }
@@ -922,7 +940,9 @@ func (pm *PolicyManager) SetNodes(nodes views.Slice[types.NodeView]) (bool, erro
pm.matchersForNodeMap.Clear() pm.matchersForNodeMap.Clear()
} }
// Always return true when nodes changed, even if filter hash didn't change // 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) // (can happen with autogroup:self or when nodes are added but don't affect rules).
// A SetNodes that moves pm back to an older snapshot counts too; the
// extra PolicyChange it causes is deduplicated downstream.
pm.nodesGen.Add(1) pm.nodesGen.Add(1)
return true, nil return true, nil
@@ -942,6 +962,17 @@ func (pm *PolicyManager) NodesGeneration() uint64 {
return pm.nodesGen.Load() return pm.nodesGen.Load()
} }
// cacheableLocked reports whether a per-node result computed from node may
// be cached under its ID. SetNodes invalidates those caches by diffing its
// own copies of each node, so a result computed from any other view, such
// as a mapper's pre-write snapshot read while the NodeStore writer builds,
// would outlive the invalidation meant to remove it.
func (pm *PolicyManager) cacheableLocked(node types.NodeView) bool {
own, ok := pm.nodesByID[node.ID()]
return ok && !node.HasPolicyChange(own) && !node.HasNetworkChanges(own)
}
// nodeIDViewMap indexes a slice of node views by node ID. On duplicate IDs the // nodeIDViewMap indexes a slice of node views by node ID. On duplicate IDs the
// last view wins, matching the open-coded loops it replaces. // last view wins, matching the open-coded loops it replaces.
func nodeIDViewMap(s views.Slice[types.NodeView]) map[types.NodeID]types.NodeView { func nodeIDViewMap(s views.Slice[types.NodeView]) map[types.NodeID]types.NodeView {
+38
View File
@@ -2899,3 +2899,41 @@ func TestNodesGenerationCountsChangingSetNodes(t *testing.T) {
require.True(t, changed) require.True(t, changed)
require.Equal(t, gen+1, pm.NodesGeneration(), "a changing SetNodes must advance the generation once") require.Equal(t, gen+1, pm.NodesGeneration(), "a changing SetNodes must advance the generation once")
} }
// TestSetNodesRetriesAfterFailedRecompile pins that a SetNodes whose
// recompile fails keeps the previous node list, like SetUsers. Keeping the
// new list would make an identical retry look unchanged, so the recompile
// would never be retried and the caller would never learn the policy moved.
func TestSetNodesRetriesAfterFailedRecompile(t *testing.T) {
users := types.Users{{ID: 1, Name: "user1"}}
nodes := make(types.Nodes, 0, 2)
nodes = append(nodes, node("n1", "100.64.0.1", "fd7a:115c:a1e0::1", users[0]))
nodes[0].ID = 1
pm, err := NewPolicyManager([]byte(`{
"tagOwners": {"tag:a": ["user1@"]},
"acls": [{"action": "accept", "src": ["user1@"], "dst": ["user1@:*"]}]
}`), users, nodes.ViewSlice())
require.NoError(t, err)
added := node("n2", "100.64.0.2", "fd7a:115c:a1e0::2", users[0])
added.ID = 2
grown := append(nodes, added)
// Break tag owner resolution so the recompile fails.
good := pm.pol.TagOwners
missing := Tag("tag:missing")
pm.pol.TagOwners = TagOwners{"tag:a": Owners{&missing}}
_, err = pm.SetNodes(grown.ViewSlice())
require.Error(t, err)
pm.pol.TagOwners = good
gen := pm.NodesGeneration()
changed, err := pm.SetNodes(grown.ViewSlice())
require.NoError(t, err)
require.True(t, changed, "the retry must recompile, not see the failed input as current")
require.Equal(t, gen+1, pm.NodesGeneration())
}
+7 -2
View File
@@ -122,13 +122,13 @@ func (m *mapSession) serve() {
// //
// Process the [tailcfg.MapRequest] to update node state (endpoints, hostinfo, etc.) // Process the [tailcfg.MapRequest] to update node state (endpoints, hostinfo, etc.)
c, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req) c, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req)
m.h.Change(c)
if err != nil { if err != nil {
httpError(m.w, err) httpError(m.w, err)
return return
} }
m.h.Change(c)
// If OmitPeers is true and Stream is false // If OmitPeers is true and Stream is false
// then the server will let clients update their endpoints without // then the server will let clients update their endpoints without
// breaking existing long-polling (Stream == true) connections. // breaking existing long-polling (Stream == true) connections.
@@ -243,6 +243,8 @@ func (m *mapSession) serveLongPoll() {
// the node to be incorrectly removed from AvailableRoutes. // the node to be incorrectly removed from AvailableRoutes.
mapReqChange, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req) mapReqChange, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req)
if err != nil { if err != nil {
m.h.Change(mapReqChange)
m.log.Error().Caller().Err(err).Msg("failed to update node from initial MapRequest") m.log.Error().Caller().Err(err).Msg("failed to update node from initial MapRequest")
// Write an explicit error rather than returning silently: a bare // Write an explicit error rather than returning silently: a bare
// return leaves net/http to send an empty 200, which the client // return leaves net/http to send an empty 200, which the client
@@ -278,6 +280,9 @@ func (m *mapSession) serveLongPoll() {
// time between the node connecting and the batcher being ready. // time between the node connecting and the batcher being ready.
if err := m.h.mapBatcher.AddNode(m.node.ID, m.ch, m.capVer, m.stopFromBatcher); err != nil { //nolint:noinlineerr if err := m.h.mapBatcher.AddNode(m.node.ID, m.ch, m.capVer, m.stopFromBatcher); err != nil { //nolint:noinlineerr
m.log.Error().Caller().Err(err).Msg("failed to add node to batcher") m.log.Error().Caller().Err(err).Msg("failed to add node to batcher")
// The map request already changed state other nodes must see.
m.h.Change(mapReqChange)
// Write an explicit error rather than returning silently: a bare // Write an explicit error rather than returning silently: a bare
// return leaves net/http to send an empty 200, which the client // return leaves net/http to send an empty 200, which the client
// reads as "unexpected EOF" and retries forever (issue #3346). // reads as "unexpected EOF" and retries forever (issue #3346).
+2 -2
View File
@@ -165,11 +165,11 @@ func (h *TestHarness) ChangePolicy(tb testing.TB, policy []byte) {
if changed { if changed {
changes, err := h.Server.State().ReloadPolicy() changes, err := h.Server.State().ReloadPolicy()
h.Server.App.Change(changes...)
if err != nil { if err != nil {
tb.Fatalf("servertest: ReloadPolicy: %v", err) tb.Fatalf("servertest: ReloadPolicy: %v", err)
} }
h.Server.App.Change(changes...)
} }
} }
+420 -18
View File
@@ -15,6 +15,7 @@ import (
"github.com/juanfont/headscale/hscontrol/db" "github.com/juanfont/headscale/hscontrol/db"
policyv2 "github.com/juanfont/headscale/hscontrol/policy/v2" policyv2 "github.com/juanfont/headscale/hscontrol/policy/v2"
"github.com/juanfont/headscale/hscontrol/types" "github.com/juanfont/headscale/hscontrol/types"
"github.com/juanfont/headscale/hscontrol/types/change"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"pgregory.net/rapid" "pgregory.net/rapid"
@@ -1694,9 +1695,7 @@ func nodeStoreWithPolicy(t fatalfer, pol string, users []types.User, nodes types
t.Fatalf("policy: %v", err) t.Fatalf("policy: %v", err)
} }
store := NewNodeStore(nodes, func(ns []types.NodeView) map[types.NodeID][]types.NodeID { store := NewNodeStore(nodes, policyPeersFunc(pm), TestBatchSize, TestBatchTimeout)
return pm.BuildPeerMap(views.SliceOf(ns))
}, TestBatchSize, TestBatchTimeout)
store.Start() store.Start()
return store, pm return store, pm
@@ -1718,12 +1717,21 @@ func syncPolicy(t fatalfer, store *NodeStore, pm *policyv2.PolicyManager) {
// checkAdjacencyMatchesFullBuild compares the NodeStore's current // checkAdjacencyMatchesFullBuild compares the NodeStore's current
// adjacency — however it got there, including the reused-from-previous- // adjacency — however it got there, including the reused-from-previous-
// snapshot path taken for payload-only writes — against a from-scratch // snapshot path taken for payload-only writes — against
// [policyv2.PolicyManager.BuildPeerMap] over the same nodes. Divergence // [policyv2.PolicyManager.BuildPeerMap] from a fresh policy manager over
// means the reuse path served stale adjacency. // the same nodes. The fresh manager keeps the oracle independent of the
func checkAdjacencyMatchesFullBuild(t fatalfer, store *NodeStore, pm *policyv2.PolicyManager) { // 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() 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 { for id := range snap.nodesByID {
got := slices.Sorted(slices.Values(snap.peersByNode[id])) 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 // Check before syncPolicy: the write's own snapshot must
// pre-write nodes, so BuildPeerMap(new nodes) against pm's // already be right, both on the reuse path (updateChanges
// old matchers is exactly what the write's own reuse-vs- // said payload-only) and on the recompute path (the
// recompute decision (updateChanges) should have produced. // peersFunc refreshed pm before building). Checking only
// Checking only after syncPolicy would let a wrong // after syncPolicy would let either mistake hide behind a
// updateChanges classification hide behind the // RebuildPeerMaps.
// RebuildPeerMaps that SetNodes triggers on its own. checkAdjacencyMatchesFullBuild(rt, store, tc.pol, users)
checkAdjacencyMatchesFullBuild(rt, store, pm)
syncPolicy(rt, store, pm) 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()) pm, err := policyv2.NewPolicyManager([]byte(policyGlobal), users, nodes.ViewSlice())
require.NoError(b, err) require.NoError(b, err)
inner := policyPeersFunc(pm)
peersFunc := func(ns []types.NodeView) map[types.NodeID][]types.NodeID { peersFunc := func(ns []types.NodeView) map[types.NodeID][]types.NodeID {
calls.Add(1) calls.Add(1)
return pm.BuildPeerMap(views.SliceOf(ns)) return inner(ns)
} }
store := NewNodeStore(nodes, peersFunc, TestBatchSize, TestBatchTimeout) 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")
})
}
}
+1 -1
View File
@@ -47,7 +47,7 @@ func TestPersistNodeDoesNotClobberConcurrentAdminWrite(t *testing.T) {
"precondition: admin SetNodeTags must have written the tag to the DB") "precondition: admin SetNodeTags must have written the tag to the DB")
// (3) Map-request persists its stale snapshot. // (3) Map-request persists its stale snapshot.
_, _, err = s.persistNodeAndRefreshPolicy(staleView) _, _, err = s.persistNodeAndRefreshPolicy(staleView, s.polMan.NodesGeneration())
require.NoError(t, err) require.NoError(t, err)
// The admin write must survive. // The admin write must survive.
+187 -5
View File
@@ -3,6 +3,7 @@ package state
import ( import (
"errors" "errors"
"net/netip" "net/netip"
"slices"
"sync" "sync"
"testing" "testing"
"time" "time"
@@ -141,7 +142,7 @@ func TestPersistEmptyTags(t *testing.T) {
seeded, ok := s.nodeStore.GetNode(nodeID) seeded, ok := s.nodeStore.GetNode(nodeID)
require.True(t, ok) require.True(t, ok)
_, _, err := s.persistNodeAndRefreshPolicy(seeded) _, _, err := s.persistNodeAndRefreshPolicy(seeded, s.polMan.NodesGeneration())
require.NoError(t, err) require.NoError(t, err)
gotAfterSeed, err := s.DB().GetNodeByID(nodeID) gotAfterSeed, err := s.DB().GetNodeByID(nodeID)
@@ -154,7 +155,7 @@ func TestPersistEmptyTags(t *testing.T) {
}) })
require.True(t, ok) require.True(t, ok)
_, _, err = s.persistNodeAndRefreshPolicy(cleared) _, _, err = s.persistNodeAndRefreshPolicy(cleared, s.polMan.NodesGeneration())
require.NoError(t, err) require.NoError(t, err)
gotAfterClear, err := s.DB().GetNodeByID(nodeID) gotAfterClear, err := s.DB().GetNodeByID(nodeID)
@@ -189,7 +190,7 @@ func TestPersistEmptyEndpoints(t *testing.T) {
seeded, ok := s.nodeStore.GetNode(nodeID) seeded, ok := s.nodeStore.GetNode(nodeID)
require.True(t, ok) require.True(t, ok)
_, _, err := s.persistNodeAndRefreshPolicy(seeded) _, _, err := s.persistNodeAndRefreshPolicy(seeded, s.polMan.NodesGeneration())
require.NoError(t, err) require.NoError(t, err)
gotAfterSeed, err := s.DB().GetNodeByID(nodeID) gotAfterSeed, err := s.DB().GetNodeByID(nodeID)
@@ -202,7 +203,7 @@ func TestPersistEmptyEndpoints(t *testing.T) {
}) })
require.True(t, ok) require.True(t, ok)
_, _, err = s.persistNodeAndRefreshPolicy(cleared) _, _, err = s.persistNodeAndRefreshPolicy(cleared, s.polMan.NodesGeneration())
require.NoError(t, err) require.NoError(t, err)
gotAfterClear, err := s.DB().GetNodeByID(nodeID) gotAfterClear, err := s.DB().GetNodeByID(nodeID)
@@ -724,12 +725,14 @@ func TestPersistNodeAndRefreshPolicyEmptyForPayloadOnlyChange(t *testing.T) {
_, s, nodeID := persistTestSetup(t) _, s, nodeID := persistTestSetup(t)
t.Cleanup(func() { _ = s.Close() }) t.Cleanup(func() { _ = s.Close() })
genBefore := s.polMan.NodesGeneration()
view, ok := s.nodeStore.UpdateNode(nodeID, func(n *types.Node) { view, ok := s.nodeStore.UpdateNode(nodeID, func(n *types.Node) {
n.Hostinfo = &tailcfg.Hostinfo{Hostname: "payload-only"} n.Hostinfo = &tailcfg.Hostinfo{Hostname: "payload-only"}
}) })
require.True(t, ok) require.True(t, ok)
_, c, err := s.persistNodeAndRefreshPolicy(view) _, c, err := s.persistNodeAndRefreshPolicy(view, genBefore)
require.NoError(t, err) require.NoError(t, err)
assert.True(t, c.IsEmpty(), "a payload-only write must not fabricate a change") 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
View File
@@ -268,16 +268,7 @@ func NewState(cfg *types.Config) (*State, error) {
batchTimeout := cmp.Or(cfg.Tuning.NodeStoreBatchTimeout, defaultNodeStoreBatchTimeout) batchTimeout := cmp.Or(cfg.Tuning.NodeStoreBatchTimeout, defaultNodeStoreBatchTimeout)
// [policy.PolicyManager.BuildPeerMap] handles both global and per-node filter complexity. nodeStore := NewNodeStore(nodes, policyPeersFunc(polMan), batchSize, batchTimeout)
// 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.Start() nodeStore.Start()
s := &State{ s := &State{
@@ -302,6 +293,31 @@ func NewState(cfg *types.Config) (*State, error) {
return s, nil 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. // Close gracefully shuts down the [State] instance and releases all resources.
func (s *State) Close() error { func (s *State) Close() error {
s.pings.drain() 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. // 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) { func (s *State) ReloadPolicy() ([]change.Change, error) {
pol, err := hsdb.PolicyBytes(s.db.DB, s.cfg) pol, err := hsdb.PolicyBytes(s.db.DB, s.cfg)
if err != nil { if err != nil {
@@ -369,7 +386,9 @@ func (s *State) ReloadPolicy() ([]change.Change, error) {
// with the current policy. // with the current policy.
rcs, err := s.autoApproveNodes() rcs, err := s.autoApproveNodes()
if err != nil { 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. // 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 // persistNodeAndRefreshPolicy saves the given node state to the database and refreshes the
// policy manager. The exact row written comes from [NodeStore]; see // policy manager. The exact row written comes from [NodeStore]; see
// [State.persistNode]. // [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) fresh, err := s.persistNode(node)
if err != nil { 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 // Check if policy manager needs updating
c, err := s.updatePolicyManagerNodes() c, err := s.updatePolicyManagerNodes(genBefore)
if err != nil { 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 return fresh, c, nil
@@ -584,10 +604,11 @@ func (s *State) SaveNode(node types.NodeView) (types.NodeView, change.Change, er
// Update [NodeStore] first // Update [NodeStore] first
nodePtr := node.AsStruct() nodePtr := node.AsStruct()
genBefore := s.polMan.NodesGeneration()
resultNode := s.nodeStore.PutNode(*nodePtr) resultNode := s.nodeStore.PutNode(*nodePtr)
// Then save to database using the result from [NodeStore.PutNode] // 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. // 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 // publish a non-empty change before handling the error so live sessions are
// still torn down after a committed deletion. // still torn down after a committed deletion.
func (s *State) DeleteNode(node types.NodeView) (change.Change, error) { func (s *State) DeleteNode(node types.NodeView) (change.Change, error) {
genBefore := s.polMan.NodesGeneration()
s.persistMu.Lock() s.persistMu.Lock()
err := s.db.DeleteNode(node.AsStruct()) 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()) c := change.NodeRemoved(node.ID())
// Check if policy manager needs updating after node deletion // Check if policy manager needs updating after node deletion
policyChange, err := s.updatePolicyManagerNodes() policyChange, err := s.updatePolicyManagerNodes(genBefore)
if err != nil { 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() { 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) { func (s *State) SetNodeExpiry(nodeID types.NodeID, expiry *time.Time) (types.NodeView, change.Change, error) {
var onlineChanged bool var onlineChanged bool
genBefore := s.polMan.NodesGeneration()
// Update [NodeStore] before database to ensure consistency. The [NodeStore] update // 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 // 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] // 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 // change will remain, and the change describing it is returned with the error.
// sent to the batcher, preventing inconsistent state propagation.
n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) { n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) {
wasOnline := node.Online() wasOnline := node.Online()
node.Expiry = expiry 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) 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. // Persist expiry change to database directly since persistNodeAndRefreshPolicy omits expiry.
err := s.db.NodeSetExpiry(nodeID, expiry) err := s.db.NodeSetExpiry(nodeID, expiry)
if err != nil { 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. // Update policy manager and generate change notification.
c, err := s.updatePolicyManagerNodes() c, err := s.updatePolicyManagerNodes(genBefore)
if err != nil { 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 // Resolve expiry and online status together from the current snapshot
// when the mapper sends the change, including after a rapid restoration. // when the mapper sends the change, including after a rapid restoration.
c = c.Merge(change.NodeAdded(n.ID())) c = c.Merge(change.NodeAdded(n.ID())).Merge(recompute)
if onlineChanged && s.polMan.NodeNeedsPeerRecompute(n) {
c = c.Merge(change.PolicyChange())
}
return n, c, nil return n, c, nil
} }
@@ -986,6 +1016,8 @@ func (s *State) SetNodeTags(nodeID types.NodeID, tags []string) (types.NodeView,
// Log the operation // Log the operation
logTagOperation(existingNode, validatedTags) logTagOperation(existingNode, validatedTags)
genBefore := s.polMan.NodesGeneration()
// Update [NodeStore] before database to ensure consistency. The [NodeStore] update // 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 // is blocking and will be the source of truth for the batcher. The database update
// must make the exact same change. // 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) 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 { if err != nil {
return nodeView, c, err 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 // because even if the CLI removes an auto-approved route, it will be added
// back automatically. // back automatically.
prevRoutes := s.nodeStore.PrimaryRoutes() prevRoutes := s.nodeStore.PrimaryRoutes()
genBefore := s.polMan.NodesGeneration()
n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) { n, ok := s.nodeStore.UpdateNode(nodeID, func(node *types.Node) {
node.ApprovedRoutes = routes 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 // Persist the node changes to the database
nodeView, c, err := s.persistNodeAndRefreshPolicy(n) nodeView, c, err := s.persistNodeAndRefreshPolicy(n, genBefore)
if err != nil { if err != nil {
return types.NodeView{}, change.Change{}, err return nodeView, c, err
} }
// PolicyChange fans out a fresh netmap whenever the new approved // 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) return types.NodeView{}, change.Change{}, fmt.Errorf("%w: %w", ErrGivenNameInvalid, err)
} }
genBefore := s.polMan.NodesGeneration()
view, err := s.nodeStore.SetGivenName(nodeID, newName) view, err := s.nodeStore.SetGivenName(nodeID, newName)
if err != nil { if err != nil {
switch { 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 { if err != nil {
return nodeView, c, err return nodeView, c, err
} }
@@ -1095,23 +1130,34 @@ func (s *State) RenameNode(nodeID types.NodeID, newName string) (types.NodeView,
return nodeView, c, nil return nodeView, c, nil
} }
// BackfillNodeIPs assigns IP addresses to nodes that don't have them. // BackfillNodeIPs assigns IP addresses to nodes that don't have them. The
func (s *State) BackfillNodeIPs() ([]string, error) { // 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) changes, err := s.db.BackfillNodeIPs(s.ipAlloc)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
var readdressed []types.NodeID
// Refresh [NodeStore] after IP changes to ensure consistency // Refresh [NodeStore] after IP changes to ensure consistency
if len(changes) > 0 { if len(changes) > 0 {
nodes, err := s.db.ListNodes() nodes, err := s.db.ListNodes()
if err != nil { 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 { for _, node := range nodes {
// Preserve online status and NetInfo when refreshing from database // Preserve online status and NetInfo when refreshing from database
existingNode, exists := s.nodeStore.GetNode(node.ID) existingNode, exists := s.nodeStore.GetNode(node.ID)
if !exists || !slices.Equal(existingNode.IPs(), node.IPs()) {
readdressed = append(readdressed, node.ID)
}
if exists && existingNode.Valid() { if exists && existingNode.Valid() {
node.IsOnline = new(existingNode.IsOnline().Get()) 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. // 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). Err(err).
Msg("Failed to persist auto-approved routes") 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") log.Info().EmbedObject(nv).Strs(zf.RoutesApproved, util.PrefixesToString(approved)).Msg("routes approved")
@@ -2295,16 +2356,19 @@ func (s *State) HandleNodeFromAuthPath(
expiry *time.Time, expiry *time.Time,
registrationMethod string, registrationMethod string,
) (types.NodeView, change.Change, error) { ) (types.NodeView, change.Change, error) {
// Read before any NodeStore write below; see updatePolicyManagerNodes.
genBefore := s.polMan.NodesGeneration()
// Get the registration entry from cache // Get the registration entry from cache
regEntry, ok := s.GetAuthCacheEntry(authID) regEntry, ok := s.GetAuthCacheEntry(authID)
if !ok { if !ok {
return types.NodeView{}, change.Change{}, hsdb.ErrNodeNotFoundRegistrationCache return types.NodeView{}, s.policyChangeSince(genBefore), hsdb.ErrNodeNotFoundRegistrationCache
} }
// Get the user // Get the user
user, err := s.db.GetUserByID(userID) user, err := s.db.GetUserByID(userID)
if err != nil { 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() regData := regEntry.RegistrationData()
@@ -2352,7 +2416,7 @@ func (s *State) HandleNodeFromAuthPath(
// present the machine key is in a corrupt/ambiguous state; reject rather // present the machine key is in a corrupt/ambiguous state; reject rather
// than converting an arbitrary node and orphaning the other. // than converting an arbitrary node and orphaning the other.
if existingNodeIsTagged && (nodeExistsForSameUser || existingNodeOwnedByOtherUser) { 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 // Create logger with common fields for all auth operations
@@ -2380,7 +2444,7 @@ func (s *State) HandleNodeFromAuthPath(
finalNode, err = s.applyAuthNodeUpdate(updateParams) finalNode, err = s.applyAuthNodeUpdate(updateParams)
if err != nil { if err != nil {
return types.NodeView{}, change.Change{}, err return types.NodeView{}, s.policyChangeSince(genBefore), err
} }
} else if existingNodeIsTagged { } else if existingNodeIsTagged {
updateParams.ExistingNode = taggedNode updateParams.ExistingNode = taggedNode
@@ -2388,7 +2452,7 @@ func (s *State) HandleNodeFromAuthPath(
finalNode, err = s.applyAuthNodeUpdate(updateParams) finalNode, err = s.applyAuthNodeUpdate(updateParams)
if err != nil { if err != nil {
return types.NodeView{}, change.Change{}, err return types.NodeView{}, s.policyChangeSince(genBefore), err
} }
} else if existingNodeOwnedByOtherUser { } else if existingNodeOwnedByOtherUser {
oldUser := existingNodeOtherUser.User() oldUser := existingNodeOtherUser.User()
@@ -2409,7 +2473,7 @@ func (s *State) HandleNodeFromAuthPath(
expiry, registrationMethod, existingNodeOtherUser, expiry, registrationMethod, existingNodeOtherUser,
) )
if err != nil { if err != nil {
return types.NodeView{}, change.Change{}, err return types.NodeView{}, s.policyChangeSince(genBefore), err
} }
} else { } else {
finalNode, err = s.createNewNodeFromAuth( finalNode, err = s.createNewNodeFromAuth(
@@ -2417,7 +2481,7 @@ func (s *State) HandleNodeFromAuthPath(
expiry, registrationMethod, types.NodeView{}, expiry, registrationMethod, types.NodeView{},
) )
if err != nil { 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 // Update policy managers
usersChange, err := s.updatePolicyManagerUsers() usersChange, err := s.updatePolicyManagerUsers()
if err != nil { 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 { 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() 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. // to a single node rather than racing the find-then-create section.
defer s.lockRegistration(machineKey)() defer s.lockRegistration(machineKey)()
// Read before any NodeStore write below; see updatePolicyManagerNodes.
genBefore := s.polMan.NodesGeneration()
pak, err := s.GetPreAuthKey(regReq.Auth.AuthKey) pak, err := s.GetPreAuthKey(regReq.Auth.AuthKey)
if err != nil { 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. // 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 { 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) existingNodeSameUser, existsSameUser, err := s.findExistingNodeForPAK(machineKey, pak)
if err != nil { if err != nil {
return types.NodeView{}, change.Change{}, err return types.NodeView{}, s.policyChangeSince(genBefore), err
} }
// For existing nodes, skip validation if: // For existing nodes, skip validation if:
@@ -2657,7 +2724,7 @@ func (s *State) HandleNodeFromPreAuthKey(
// New node or NodeKey rotation: require valid auth key. // New node or NodeKey rotation: require valid auth key.
err = pak.Validate() err = pak.Validate()
if err != nil { 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. // NodeStore NodeKey index, denying the victim service.
if existing, ok := s.nodeStore.GetNodeByNodeKey(regReq.NodeKey); ok && if existing, ok := s.nodeStore.GetNodeByNodeKey(regReq.NodeKey); ok &&
existing.MachineKey() != machineKey { 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 // Snapshot the pre-update node so the NodeStore can be rolled back if
@@ -2799,7 +2866,7 @@ func (s *State) HandleNodeFromPreAuthKey(
}) })
if !ok { 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) { _, 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) 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(). log.Trace().
@@ -2928,19 +2995,19 @@ func (s *State) HandleNodeFromPreAuthKey(
ExistingNodeForNetinfo: differentUserNode, ExistingNodeForNetinfo: differentUserNode,
}) })
if err != nil { 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 // Update policy managers
usersChange, err := s.updatePolicyManagerUsers() usersChange, err := s.updatePolicyManagerUsers()
if err != nil { 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 { 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() policyChanged := !usersChange.IsEmpty() || !nodesChange.IsEmpty()
@@ -3006,28 +3073,60 @@ func (s *State) UpdatePolicyManagerUsersForTest() error {
return err return err
} }
// updatePolicyManagerNodes updates the policy manager with current nodes. // updatePolicyManagerNodes refreshes the policy manager with current node
// Returns true if the policy changed and notifications should be sent. // 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 // TODO(kradalby): This is a temporary stepping stone, ultimately we should
// have the list already available so it could go much quicker. Alternatively // have the list already available so it could go much quicker. Alternatively
// the policy manager could have a remove or add list for nodes. // 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(genBefore uint64) (change.Change, error) {
func (s *State) updatePolicyManagerNodes() (change.Change, error) {
nodes := s.ListNodes() nodes := s.ListNodes()
changed, err := s.polMan.SetNodes(nodes) changed, err := s.polMan.SetNodes(nodes)
if err != nil { 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 { if changed {
// Rebuild peer maps because policy-affecting node changes (tags, user, IPs) // The writer refreshes the policy before every relation build, so
// affect ACL visibility. Without this, cached peer relationships use stale data. // a change here means this snapshot raced another writer and moved
// the policy manager away from what adjacency was built with.
s.nodeStore.RebuildPeerMaps() 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. // PingDB checks if the database connection is healthy.
@@ -3071,6 +3170,8 @@ func (s *State) autoApproveNodes() ([]change.Change, error) {
return nil, nil return nil, nil
} }
genBefore := s.polMan.NodesGeneration()
updates := make(map[types.NodeID]UpdateNodeFunc, len(approvedByID)) updates := make(map[types.NodeID]UpdateNodeFunc, len(approvedByID))
for id, approved := range approvedByID { for id, approved := range approvedByID {
updates[id] = func(n *types.Node) { updates[id] = func(n *types.Node) {
@@ -3094,13 +3195,13 @@ func (s *State) autoApproveNodes() ([]change.Change, error) {
_, err := s.persistNode(fresh) _, err := s.persistNode(fresh)
if err != nil { 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 { if err != nil {
return nil, err return []change.Change{c}, err
} }
if c.IsEmpty() { 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 // Snapshot the primary assignment so we can tell whether the
// Hostinfo + auto-approval that follows shifted any prefix. // Hostinfo + auto-approval that follows shifted any prefix.
prevRoutes := s.nodeStore.PrimaryRoutes() prevRoutes := s.nodeStore.PrimaryRoutes()
genBefore := s.polMan.NodesGeneration()
// We need to ensure we update the node as it is in the [NodeStore] at // We need to ensure we update the node as it is in the [NodeStore] at
// the time of the request. // the time of the request.
@@ -3374,16 +3476,16 @@ func (s *State) UpdateNodeFromMapRequest(id types.NodeID, req tailcfg.MapRequest
updatedNode, err = s.persistNode(updatedNode) updatedNode, err = s.persistNode(updatedNode)
if err != nil { 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 // Only refresh the policy manager when something it depends on
// might have moved. Endpoint/key/DERP/LastSeen-only updates do not // might have moved. Endpoint/key/DERP/LastSeen-only updates do not
// affect policy evaluation and are deliberately skipped here. // affect policy evaluation and are deliberately skipped here.
if delta.peerHostinfoChanged || delta.routesChanged { if delta.peerHostinfoChanged || delta.routesChanged {
policyChange, err = s.updatePolicyManagerNodes() policyChange, err = s.updatePolicyManagerNodes(genBefore)
if err != nil { 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)
} }
} }
} }