mirror of
https://github.com/juanfont/headscale.git
synced 2026-09-17 22:12:03 +09:00
policy,types: skip recompile when the user list is unchanged
SetUsers now also reports whether user-derived peer adjacency moved. Updates #3417
This commit is contained in:
@@ -25,7 +25,10 @@ type PolicyManager interface {
|
|||||||
// from the current policy, avoiding trust of client-provided URL params.
|
// from the current policy, avoiding trust of client-provided URL params.
|
||||||
SSHCheckParams(srcNodeID, dstNodeID types.NodeID) (time.Duration, bool)
|
SSHCheckParams(srcNodeID, dstNodeID types.NodeID) (time.Duration, bool)
|
||||||
SetPolicy(pol []byte) (bool, error)
|
SetPolicy(pol []byte) (bool, error)
|
||||||
SetUsers(users []types.User) (bool, error)
|
// SetUsers replaces the user list. policyChanged reports whether clients
|
||||||
|
// need a policy refresh; peerMapChanged reports whether user-derived peer
|
||||||
|
// adjacency may have changed. Both are false when the list is unchanged.
|
||||||
|
SetUsers(users []types.User) (policyChanged, peerMapChanged bool, err error)
|
||||||
SetNodes(nodes views.Slice[types.NodeView]) (bool, error)
|
SetNodes(nodes views.Slice[types.NodeView]) (bool, error)
|
||||||
// NodeCanHaveTag reports whether the given node can have the given tag.
|
// NodeCanHaveTag reports whether the given node can have the given tag.
|
||||||
NodeCanHaveTag(node types.NodeView, tag string) bool
|
NodeCanHaveTag(node types.NodeView, tag string) bool
|
||||||
|
|||||||
@@ -803,34 +803,57 @@ func (pm *PolicyManager) MatchersForNode(node types.NodeView) ([]matcher.Match,
|
|||||||
return matchers, nil
|
return matchers, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetUsers updates the users in the policy manager and updates the filter rules.
|
// SetUsers replaces the user list and recompiles when it changed. Both results
|
||||||
func (pm *PolicyManager) SetUsers(users []types.User) (bool, error) {
|
// are false for an unchanged list, so callers can skip the peer-map rebuild
|
||||||
|
// and the client refresh that a user change would otherwise require.
|
||||||
|
func (pm *PolicyManager) SetUsers(users []types.User) (bool, bool, error) {
|
||||||
if pm == nil {
|
if pm == nil {
|
||||||
return false, nil
|
return false, false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
pm.mu.Lock()
|
pm.mu.Lock()
|
||||||
defer pm.mu.Unlock()
|
defer pm.mu.Unlock()
|
||||||
|
|
||||||
|
if equalUsers(pm.users, users) {
|
||||||
|
return false, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
prev := pm.users
|
||||||
pm.users = users
|
pm.users = users
|
||||||
|
|
||||||
// Clear SSH policy map when users change to force SSH policy recomputation
|
// SSH policies resolve users by name, so they are recomputed on any
|
||||||
// This ensures that if SSH policy compilation previously failed due to missing users,
|
// user change.
|
||||||
// it will be retried with the new user list
|
|
||||||
pm.sshPolicyMap.Clear()
|
pm.sshPolicyMap.Clear()
|
||||||
|
|
||||||
changed, err := pm.updateLocked()
|
policyChanged, err := pm.updateLocked()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
// Keep the old list so a retry with the same input recompiles
|
||||||
|
// instead of being treated as unchanged.
|
||||||
|
pm.users = prev
|
||||||
|
|
||||||
|
return false, false, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// If SSH policies exist, force a policy change when users are updated
|
// SSH rules embed user identity, so a user change needs a client refresh
|
||||||
// This ensures nodes get updated SSH policies even if other policy hashes didn't change
|
// even when the filter hash did not move.
|
||||||
if pm.pol != nil && len(pm.pol.SSHs) > 0 {
|
if pm.pol != nil && len(pm.pol.SSHs) > 0 {
|
||||||
return true, nil
|
policyChanged = true
|
||||||
}
|
}
|
||||||
|
|
||||||
return changed, nil
|
return policyChanged, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// equalUsers compares user lists ignoring order and every field the policy
|
||||||
|
// does not read, so a row touch such as an OIDC login is not a change.
|
||||||
|
func equalUsers(a, b []types.User) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
byID := func(l, r types.User) int { return cmp.Compare(l.ID, r.ID) }
|
||||||
|
a, b = slices.SortedFunc(slices.Values(a), byID), slices.SortedFunc(slices.Values(b), byID)
|
||||||
|
|
||||||
|
return slices.EqualFunc(a, b, func(l, r types.User) bool { return l.PolicyEqual(&r) })
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetNodes updates the nodes in the policy manager and updates the filter rules.
|
// SetNodes updates the nodes in the policy manager and updates the filter rules.
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
"github.com/juanfont/headscale/hscontrol/policy/matcher"
|
"github.com/juanfont/headscale/hscontrol/policy/matcher"
|
||||||
@@ -355,6 +356,83 @@ func TestSSHCheckParamsUnhydratedUserNoPanic(t *testing.T) {
|
|||||||
}, "SSHCheckParams must not panic when a non-tagged node has an unhydrated User")
|
}, "SSHCheckParams must not panic when a non-tagged node has an unhydrated User")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSetUsers(t *testing.T) {
|
||||||
|
const allowAll = `{"acls":[{"action":"accept","src":["*"],"dst":["*:*"]}]}`
|
||||||
|
|
||||||
|
const sshCheck = `{
|
||||||
|
"ssh": [
|
||||||
|
{
|
||||||
|
"action": "check",
|
||||||
|
"src": ["user1@headscale.net"],
|
||||||
|
"dst": ["autogroup:self"],
|
||||||
|
"users": ["root"]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}`
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
policy string
|
||||||
|
mutate func(*types.User)
|
||||||
|
|
||||||
|
wantPolicyChanged bool
|
||||||
|
wantPeerMapChanged bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "identical users without ssh",
|
||||||
|
policy: allowAll,
|
||||||
|
mutate: func(*types.User) {},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "identical users with ssh",
|
||||||
|
policy: sshCheck,
|
||||||
|
mutate: func(*types.User) {},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "timestamp bump only",
|
||||||
|
policy: sshCheck,
|
||||||
|
mutate: func(u *types.User) { u.UpdatedAt = u.UpdatedAt.Add(time.Hour) },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "display name change without ssh",
|
||||||
|
policy: allowAll,
|
||||||
|
mutate: func(u *types.User) { u.DisplayName = "Renamed" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "email change without ssh",
|
||||||
|
policy: allowAll,
|
||||||
|
mutate: func(u *types.User) { u.Email = "other@headscale.net" },
|
||||||
|
|
||||||
|
wantPeerMapChanged: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "rename with ssh",
|
||||||
|
policy: sshCheck,
|
||||||
|
mutate: func(u *types.User) { u.Name = "renamed" },
|
||||||
|
|
||||||
|
wantPolicyChanged: true,
|
||||||
|
wantPeerMapChanged: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
users := types.Users{{ID: 1, Name: "user1", Email: "user1@headscale.net"}}
|
||||||
|
|
||||||
|
pm, err := NewPolicyManager([]byte(tt.policy), users, types.Nodes{}.ViewSlice())
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
updated := slices.Clone(users)
|
||||||
|
tt.mutate(&updated[0])
|
||||||
|
|
||||||
|
policyChanged, peerMapChanged, err := pm.SetUsers(updated)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, tt.wantPolicyChanged, policyChanged, "policyChanged")
|
||||||
|
require.Equal(t, tt.wantPeerMapChanged, peerMapChanged, "peerMapChanged")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestInvalidateGlobalPolicyCache tests the cache invalidation logic for global policies.
|
// TestInvalidateGlobalPolicyCache tests the cache invalidation logic for global policies.
|
||||||
func TestInvalidateGlobalPolicyCache(t *testing.T) {
|
func TestInvalidateGlobalPolicyCache(t *testing.T) {
|
||||||
mustIPPtr := func(s string) *netip.Addr {
|
mustIPPtr := func(s string) *netip.Addr {
|
||||||
|
|||||||
@@ -647,3 +647,24 @@ func TestConcurrentPreAuthKeyRegistrationSameMachineKey(t *testing.T) {
|
|||||||
require.Equal(t, 1, s.ListNodes().Len(),
|
require.Equal(t, 1, s.ListNodes().Len(),
|
||||||
"concurrent registrations of one machine key must yield a single node")
|
"concurrent registrations of one machine key must yield a single node")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestUpdatePolicyManagerUsersUnchangedKeepsSnapshot ensures re-sending the
|
||||||
|
// same user list does not rebuild peer adjacency, while a real user change
|
||||||
|
// does.
|
||||||
|
func TestUpdatePolicyManagerUsersUnchangedKeepsSnapshot(t *testing.T) {
|
||||||
|
_, s, _ := persistTestSetup(t)
|
||||||
|
t.Cleanup(func() { _ = s.Close() })
|
||||||
|
|
||||||
|
require.NoError(t, s.UpdatePolicyManagerUsersForTest())
|
||||||
|
|
||||||
|
before := s.nodeStore.data.Load()
|
||||||
|
|
||||||
|
require.NoError(t, s.UpdatePolicyManagerUsersForTest())
|
||||||
|
require.Same(t, before, s.nodeStore.data.Load(),
|
||||||
|
"unchanged users must not rebuild the peer map")
|
||||||
|
|
||||||
|
_, _, err := s.CreateUser(types.User{Name: "second"})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotSame(t, before, s.nodeStore.data.Load(),
|
||||||
|
"a user change must rebuild the peer map")
|
||||||
|
}
|
||||||
|
|||||||
@@ -2878,12 +2878,9 @@ func reauthChange(node types.NodeView, isRelogin, policyChanged bool) change.Cha
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// updatePolicyManagerUsers updates the policy manager with current users.
|
// updatePolicyManagerUsers pushes the current user list into the policy
|
||||||
// Returns true if the policy changed and notifications should be sent.
|
// manager, rebuilds peer adjacency when user identity changed, and returns
|
||||||
// TODO(kradalby): This is a temporary stepping stone, ultimately we should
|
// a PolicyChange when clients need a refresh.
|
||||||
// have the list already available so it could go much quicker. Alternatively
|
|
||||||
// the policy manager could have a remove or add list for users.
|
|
||||||
// updatePolicyManagerUsers refreshes the policy manager with current user data.
|
|
||||||
func (s *State) updatePolicyManagerUsers() (change.Change, error) {
|
func (s *State) updatePolicyManagerUsers() (change.Change, error) {
|
||||||
users, err := s.ListAllUsers()
|
users, err := s.ListAllUsers()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -2892,13 +2889,20 @@ func (s *State) updatePolicyManagerUsers() (change.Change, error) {
|
|||||||
|
|
||||||
log.Debug().Caller().Int("user.count", len(users)).Msg("policy manager user update initiated because user list modification detected")
|
log.Debug().Caller().Int("user.count", len(users)).Msg("policy manager user update initiated because user list modification detected")
|
||||||
|
|
||||||
changed, err := s.polMan.SetUsers(users)
|
changed, peerMapChanged, err := s.polMan.SetUsers(users)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return change.Change{}, fmt.Errorf("updating policy manager users: %w", err)
|
return change.Change{}, fmt.Errorf("updating policy manager users: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug().Caller().Bool("policy.changed", changed).Msg("policy manager user update completed because SetUsers operation finished")
|
log.Debug().Caller().Bool("policy.changed", changed).Msg("policy manager user update completed because SetUsers operation finished")
|
||||||
|
|
||||||
|
if peerMapChanged {
|
||||||
|
// User-driven matcher state changed: rebuild candidate adjacency
|
||||||
|
// so peer visibility reflects the new policy. Without this, the
|
||||||
|
// cached peersByNode stays stale until the next node write.
|
||||||
|
s.nodeStore.RebuildPeerMaps()
|
||||||
|
}
|
||||||
|
|
||||||
if changed {
|
if changed {
|
||||||
return change.PolicyChange(), nil
|
return change.PolicyChange(), nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -105,6 +105,16 @@ func (u *User) StringID() string {
|
|||||||
return strconv.FormatUint(uint64(u.ID), 10)
|
return strconv.FormatUint(uint64(u.ID), 10)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PolicyEqual reports whether the policy would resolve both users the same
|
||||||
|
// way: the same row, and the same name, email, and provider identity that
|
||||||
|
// user aliases match on.
|
||||||
|
func (u *User) PolicyEqual(o *User) bool {
|
||||||
|
return u.ID == o.ID &&
|
||||||
|
u.Name == o.Name &&
|
||||||
|
u.Email == o.Email &&
|
||||||
|
u.ProviderIdentifier == o.ProviderIdentifier
|
||||||
|
}
|
||||||
|
|
||||||
// TypedID returns a pointer to the user's ID as a [UserID] type.
|
// TypedID returns a pointer to the user's ID as a [UserID] type.
|
||||||
// This is a convenience method to avoid ugly casting like ptr.To(types.UserID(user.ID)).
|
// This is a convenience method to avoid ugly casting like ptr.To(types.UserID(user.ID)).
|
||||||
func (u *User) TypedID() *UserID {
|
func (u *User) TypedID() *UserID {
|
||||||
|
|||||||
Reference in New Issue
Block a user