diff --git a/hscontrol/api/v1/auth.go b/hscontrol/api/v1/auth.go index e6ff980f7..ca63c3cca 100644 --- a/hscontrol/api/v1/auth.go +++ b/hscontrol/api/v1/auth.go @@ -91,16 +91,18 @@ func registerAuth(api huma.API, b Backend) { util.RegisterMethodCLI, ) if err != nil { + b.Change(nodeChange) + return nil, mapError("registering node", err) } routeChange, err := b.State.AutoApproveRoutes(node) + b.Change(nodeChange, routeChange) + if err != nil { return nil, huma.Error500InternalServerError("auto approving routes", err) } - b.Change(nodeChange, routeChange) - out := &authRegisterOutput{} out.Body.Node = nodeFromView(node) diff --git a/hscontrol/api/v1/nodes.go b/hscontrol/api/v1/nodes.go index 73ea5f9ae..c082ffa0c 100644 --- a/hscontrol/api/v1/nodes.go +++ b/hscontrol/api/v1/nodes.go @@ -309,12 +309,12 @@ func registerNodeWriteOps(api huma.API, b Backend) { switch { case disableExpiry: node, nodeChange, expErr := b.State.SetNodeExpiry(nodeID, nil) + b.Change(nodeChange) + if expErr != nil { return nil, mapError("expiring node", expErr) } - b.Change(nodeChange) - out := &nodeOutput{} out.Body.Node = nodeFromView(node) @@ -324,12 +324,12 @@ func registerNodeWriteOps(api huma.API, b Backend) { } node, nodeChange, err := b.State.SetNodeExpiry(nodeID, &expiry) + b.Change(nodeChange) + if err != nil { return nil, mapError("expiring node", err) } - b.Change(nodeChange) - out := &nodeOutput{} 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) + b.Change(nodeChange) + if err != nil { return nil, mapError("renaming node", err) } - b.Change(nodeChange) - out := &nodeOutput{} 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) + b.Change(nodeChange) + if err != nil { return nil, huma.Error400BadRequest("setting tags", err) } - b.Change(nodeChange) - out := &nodeOutput{} out.Body.Node = nodeFromView(node) @@ -444,12 +444,12 @@ func registerNodeAdminOps(api huma.API, b Backend) { newApproved = slices.Compact(newApproved) node, nodeChange, err := b.State.SetApprovedRoutes(nodeID, newApproved) + b.Change(nodeChange) + if err != nil { return nil, mapError("setting approved routes", err) } - b.Change(nodeChange) - out := &nodeOutput{} out.Body.Node = nodeFromView(node) // SubnetRoutes here excludes exit routes, unlike the list handler. @@ -485,17 +485,20 @@ func registerNodeAdminOps(api huma.API, b Backend) { util.RegisterMethodCLI, ) if err != nil { + b.Change(nodeChange) + return nil, mapError("registering node", err) } 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. b.Change(nodeChange, routeChange) + if err != nil { + return nil, huma.Error500InternalServerError("auto approving routes", err) + } + out := &nodeOutput{} out.Body.Node = nodeFromView(node) @@ -514,7 +517,9 @@ func registerNodeAdminOps(api huma.API, b Backend) { 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 { return nil, huma.Error500InternalServerError("backfilling node IPs", err) } diff --git a/hscontrol/api/v1/policy.go b/hscontrol/api/v1/policy.go index 452148113..bfae2d139 100644 --- a/hscontrol/api/v1/policy.go +++ b/hscontrol/api/v1/policy.go @@ -143,14 +143,12 @@ func registerPolicy(api huma.API, b Backend) { // Reload even when content is unchanged: routes manually disabled before // may now qualify for auto-approval, so they must be re-evaluated. cs, err := b.State.ReloadPolicy() + b.Change(cs...) + if err != nil { return nil, huma.Error500InternalServerError("reloading policy", err) } - if len(cs) > 0 { - b.Change(cs...) - } - out := &setPolicyOutput{} out.Body.Policy = updated.Data out.Body.UpdatedAt = updated.UpdatedAt diff --git a/hscontrol/api/v2/acl.go b/hscontrol/api/v2/acl.go index ff6917933..3d6329fe6 100644 --- a/hscontrol/api/v2/acl.go +++ b/hscontrol/api/v2/acl.go @@ -134,14 +134,12 @@ func registerACL(api huma.API, b Backend) { } cs, err := b.State.ReloadPolicy() + b.Change(cs...) + if err != nil { return nil, huma.Error500InternalServerError("reloading policy", err) } - if len(cs) > 0 { - b.Change(cs...) - } - return streamPolicy([]byte(updated.Data), aclContentType(in.Accept)), nil }) } diff --git a/hscontrol/api/v2/devices.go b/hscontrol/api/v2/devices.go index efef676b6..8eeb3d258 100644 --- a/hscontrol/api/v2/devices.go +++ b/hscontrol/api/v2/devices.go @@ -232,12 +232,12 @@ func registerDevices(api huma.API, b Backend) { } _, nodeChange, err := b.State.RenameNode(node.ID(), in.Body.Name) + b.Change(nodeChange) + if err != nil { return nil, mapError("renaming device", err) } - b.Change(nodeChange) - return &emptyOutput{}, nil }) @@ -278,12 +278,12 @@ func registerDevices(api huma.API, b Backend) { } _, nodeChange, err := b.State.SetNodeTags(node.ID(), in.Body.Tags) + b.Change(nodeChange) + if err != nil { return nil, mapError("setting device tags", err) } - b.Change(nodeChange) - return &emptyOutput{}, nil }) @@ -311,12 +311,12 @@ func registerDevices(api huma.API, b Backend) { } _, nodeChange, err := b.State.SetNodeExpiry(node.ID(), nil) + b.Change(nodeChange) + if err != nil { return nil, mapError("setting device key expiry", err) } - b.Change(nodeChange) - return &emptyOutput{}, nil }) @@ -341,12 +341,12 @@ func registerDevices(api huma.API, b Backend) { } updated, nodeChange, err := b.State.SetApprovedRoutes(node.ID(), approved) + b.Change(nodeChange) + if err != nil { return nil, mapError("setting device routes", err) } - b.Change(nodeChange) - return &deviceRoutesOutput{Body: routesFromView(updated)}, nil }) diff --git a/hscontrol/app.go b/hscontrol/app.go index a44ac5af7..de8338cd1 100644 --- a/hscontrol/app.go +++ b/hscontrol/app.go @@ -862,13 +862,13 @@ func (h *Headscale) Serve() error { } changes, err := h.state.ReloadPolicy() + h.Change(changes...) + if err != nil { log.Error().Err(err).Msgf("reloading policy") continue } - h.Change(changes...) - default: info := func(msg string) { log.Info().Msg(msg) } diff --git a/hscontrol/auth.go b/hscontrol/auth.go index 387f88062..69985e2bc 100644 --- a/hscontrol/auth.go +++ b/hscontrol/auth.go @@ -263,12 +263,12 @@ func (h *Headscale) handleLogout( } updatedNode, c, err := h.state.SetNodeExpiry(node.ID(), &expiry) + h.Change(c) + if err != nil { return nil, fmt.Errorf("setting node expiry: %w", err) } - h.Change(c) - return nodeToRegisterResponse(updatedNode), nil } @@ -420,6 +420,8 @@ func (h *Headscale) handleRegisterWithAuthKey( machineKey, ) if err != nil { + h.Change(changed) + if errors.Is(err, gorm.ErrRecordNotFound) { 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? // now since we dont update the node/pol here anymore 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(). h.Change(changed, routesChange) + if err != nil { + return nil, fmt.Errorf("auto approving routes: %w", err) + } + resp := &tailcfg.RegisterResponse{ MachineAuthorized: true, NodeKeyExpired: node.IsExpired(), diff --git a/hscontrol/mapper/mapper_test.go b/hscontrol/mapper/mapper_test.go index 6425e59fd..9ab4125bd 100644 --- a/hscontrol/mapper/mapper_test.go +++ b/hscontrol/mapper/mapper_test.go @@ -1,6 +1,7 @@ package mapper import ( + "errors" "fmt" "net/netip" "strings" @@ -15,6 +16,7 @@ import ( "github.com/juanfont/headscale/hscontrol/types/change" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" "tailscale.com/tailcfg" "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()) +} diff --git a/hscontrol/oidc.go b/hscontrol/oidc.go index e36678495..898b1842d 100644 --- a/hscontrol/oidc.go +++ b/hscontrol/oidc.go @@ -993,6 +993,8 @@ func (a *AuthProviderOIDC) handleRegistration( util.RegisterMethodOIDC, ) if err != nil { + a.h.Change(nodeChange) + 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 // eventbus. 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(). a.h.Change(nodeChange, routesChange) + if err != nil { + return false, fmt.Errorf("auto approving routes: %w", err) + } + return !nodeChange.IsEmpty(), nil } diff --git a/hscontrol/policy/v2/policy.go b/hscontrol/policy/v2/policy.go index 8ca9819fe..230f02863 100644 --- a/hscontrol/policy/v2/policy.go +++ b/hscontrol/policy/v2/policy.go @@ -37,6 +37,8 @@ type PolicyManager struct { pol *Policy users []types.User nodes views.Slice[types.NodeView] + // nodesByID indexes nodes; see [PolicyManager.cacheableLocked]. + nodesByID map[types.NodeID]types.NodeView filterHash deephash.Sum filter []tailcfg.FilterRule @@ -207,6 +209,7 @@ func NewPolicyManager(b []byte, users []types.User, nodes views.Slice[types.Node pol: policy, users: users, nodes: nodes, + nodesByID: nodeIDViewMap(nodes), sshPolicyMap: xsync.NewMap[types.NodeID, *tailcfg.SSHPolicy](), filterRulesMap: xsync.NewMap[types.NodeID, []tailcfg.FilterRule](), 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) } - pm.sshPolicyMap.Store(node.ID(), sshPol) + if pm.cacheableLocked(node) { + pm.sshPolicyMap.Store(node.ID(), sshPol) + } return sshPol, nil } @@ -766,7 +771,9 @@ func (pm *PolicyManager) filterForNodeLocked( } reduced := policyutil.ReduceFilterRules(node, unreduced) - pm.filterRulesMap.Store(node.ID(), reduced) + if pm.cacheableLocked(node) { + pm.filterRulesMap.Store(node.ID(), reduced) + } return reduced } @@ -822,7 +829,10 @@ func (pm *PolicyManager) MatchersForNode(node types.NodeView) ([]matcher.Match, // the stored compiled grants for this specific node. unreduced := pm.filterRulesForNodeLocked(node) matchers := matcher.MatchesFromFilterRules(unreduced) - pm.matchersForNodeMap.Store(node.ID(), matchers) + + if pm.cacheableLocked(node) { + pm.matchersForNodeMap.Store(node.ID(), matchers) + } 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). pm.invalidateNodeCache(nodes) + prevNodes, prevByID := pm.nodes, pm.nodesByID pm.nodes = nodes + pm.nodesByID = nodeIDViewMap(nodes) // When policy-affecting node properties change, we must recompile filters because: // 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 needsUpdate, err := pm.updateLocked() 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 } @@ -922,7 +940,9 @@ func (pm *PolicyManager) SetNodes(nodes views.Slice[types.NodeView]) (bool, erro pm.matchersForNodeMap.Clear() } // 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) return true, nil @@ -942,6 +962,17 @@ func (pm *PolicyManager) NodesGeneration() uint64 { 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 // last view wins, matching the open-coded loops it replaces. func nodeIDViewMap(s views.Slice[types.NodeView]) map[types.NodeID]types.NodeView { diff --git a/hscontrol/policy/v2/policy_test.go b/hscontrol/policy/v2/policy_test.go index 87b9ab3df..579af5830 100644 --- a/hscontrol/policy/v2/policy_test.go +++ b/hscontrol/policy/v2/policy_test.go @@ -2899,3 +2899,41 @@ func TestNodesGenerationCountsChangingSetNodes(t *testing.T) { require.True(t, changed) 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()) +} diff --git a/hscontrol/poll.go b/hscontrol/poll.go index 689669d90..58c820e3d 100644 --- a/hscontrol/poll.go +++ b/hscontrol/poll.go @@ -122,13 +122,13 @@ func (m *mapSession) serve() { // // Process the [tailcfg.MapRequest] to update node state (endpoints, hostinfo, etc.) c, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req) + m.h.Change(c) + if err != nil { httpError(m.w, err) return } - m.h.Change(c) - // If OmitPeers is true and Stream is false // then the server will let clients update their endpoints without // breaking existing long-polling (Stream == true) connections. @@ -243,6 +243,8 @@ func (m *mapSession) serveLongPoll() { // the node to be incorrectly removed from AvailableRoutes. mapReqChange, err := m.h.state.UpdateNodeFromMapRequest(m.node.ID, m.req) if err != nil { + m.h.Change(mapReqChange) + m.log.Error().Caller().Err(err).Msg("failed to update node from initial MapRequest") // Write an explicit error rather than returning silently: a bare // 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. 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") + // The map request already changed state other nodes must see. + m.h.Change(mapReqChange) + // Write an explicit error rather than returning silently: a bare // return leaves net/http to send an empty 200, which the client // reads as "unexpected EOF" and retries forever (issue #3346). diff --git a/hscontrol/servertest/harness.go b/hscontrol/servertest/harness.go index 7461b1cb7..0bf458643 100644 --- a/hscontrol/servertest/harness.go +++ b/hscontrol/servertest/harness.go @@ -165,11 +165,11 @@ func (h *TestHarness) ChangePolicy(tb testing.TB, policy []byte) { if changed { changes, err := h.Server.State().ReloadPolicy() + h.Server.App.Change(changes...) + if err != nil { tb.Fatalf("servertest: ReloadPolicy: %v", err) } - - h.Server.App.Change(changes...) } } diff --git a/hscontrol/state/node_store_test.go b/hscontrol/state/node_store_test.go index 053ab88ea..a62a8148f 100644 --- a/hscontrol/state/node_store_test.go +++ b/hscontrol/state/node_store_test.go @@ -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") + }) + } +} diff --git a/hscontrol/state/persist_clobber_test.go b/hscontrol/state/persist_clobber_test.go index 860f578de..aa00f3b19 100644 --- a/hscontrol/state/persist_clobber_test.go +++ b/hscontrol/state/persist_clobber_test.go @@ -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. diff --git a/hscontrol/state/persist_test.go b/hscontrol/state/persist_test.go index 233eb15da..4a082fe3e 100644 --- a/hscontrol/state/persist_test.go +++ b/hscontrol/state/persist_test.go @@ -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()) + }) + } +} diff --git a/hscontrol/state/state.go b/hscontrol/state/state.go index edbaf415f..b02ce6f3a 100644 --- a/hscontrol/state/state.go +++ b/hscontrol/state/state.go @@ -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) } } }