mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-06 06:40:06 +09:00
state: refresh policy nodes inside the peer map build
One peer build per tag/user/IP/route write; callers detect policy moves via NodesGeneration. Per-node caches only store results for the node pm holds, so a mapper reading mid-build cannot pin a stale filter.
This commit is contained in:
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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(),
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
@@ -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).
|
||||||
|
|||||||
@@ -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...)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user