mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-05 22:30:07 +09:00
policy/v2: drop per-node filter caches when users change
autogroup:self sources resolve users by name outside the filter hash, so SetUsers left stale self rules cached and reported no change.
This commit is contained in:
@@ -855,9 +855,12 @@ func (pm *PolicyManager) SetUsers(users []types.User) (bool, bool, error) {
|
||||
prev := pm.users
|
||||
pm.users = users
|
||||
|
||||
// SSH policies resolve users by name, so they are recomputed on any
|
||||
// user change.
|
||||
// SSH policies and autogroup:self sources resolve users by name, and
|
||||
// the self sources are outside the filter hash, so updateLocked can
|
||||
// report no change while per-node results moved.
|
||||
pm.sshPolicyMap.Clear()
|
||||
pm.filterRulesMap.Clear()
|
||||
pm.matchersForNodeMap.Clear()
|
||||
|
||||
policyChanged, err := pm.updateLocked()
|
||||
if err != nil {
|
||||
@@ -868,9 +871,9 @@ func (pm *PolicyManager) SetUsers(users []types.User) (bool, bool, error) {
|
||||
return false, false, err
|
||||
}
|
||||
|
||||
// SSH rules embed user identity, so a user change needs a client refresh
|
||||
// even when the filter hash did not move.
|
||||
if pm.pol != nil && len(pm.pol.SSHs) > 0 {
|
||||
// SSH rules and per-node filters embed user identity outside the filter
|
||||
// hash, so a user change needs a client refresh even when it did not move.
|
||||
if pm.needsPerNodeFilter || (pm.pol != nil && len(pm.pol.SSHs) > 0) {
|
||||
policyChanged = true
|
||||
}
|
||||
|
||||
|
||||
@@ -472,6 +472,85 @@ func TestSetUsers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func matcherStrings(ms []matcher.Match) []string {
|
||||
out := make([]string, 0, len(ms))
|
||||
for i := range ms {
|
||||
out = append(out, ms[i].DebugString())
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// TestSetUsersDropsStaleSelfFilters pins that a user change which moves
|
||||
// autogroup:self sources, without touching the global filter, still drops
|
||||
// the cached per-node filters and matchers built from the old sources and
|
||||
// reports a policy change so clients receive their new filter.
|
||||
func TestSetUsersDropsStaleSelfFilters(t *testing.T) {
|
||||
pol := `{
|
||||
"groups": {"group:a": ["u1@", "u3@"]},
|
||||
"acls": [{"action": "accept", "src": ["group:a"], "dst": ["autogroup:self:*"]}]}`
|
||||
|
||||
// ID is set by assignment so this builds where types.User embeds
|
||||
// gorm.Model and promoted-field literals are not allowed.
|
||||
u1, u3, x3 := types.User{Name: "u1"}, types.User{Name: "u3"}, types.User{Name: "x3"}
|
||||
u1.ID, u3.ID, x3.ID = 1, 3, 3
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
before types.Users
|
||||
after types.Users
|
||||
}{
|
||||
{name: "user-added", before: types.Users{u1}, after: types.Users{u1, u3}},
|
||||
{name: "user-renamed", before: types.Users{u1, x3}, after: types.Users{u1, u3}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
nodes := types.Nodes{
|
||||
node("u1-a", "100.64.0.1", "fd7a:115c:a1e0::1", u1),
|
||||
node("u3-a", "100.64.0.3", "fd7a:115c:a1e0::3", u3),
|
||||
node("u3-b", "100.64.0.4", "fd7a:115c:a1e0::4", u3),
|
||||
}
|
||||
for i, n := range nodes {
|
||||
n.ID = types.NodeID(i + 1) //nolint:gosec // safe conversion in test
|
||||
}
|
||||
|
||||
pm, err := NewPolicyManager([]byte(pol), tt.before, nodes.ViewSlice())
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, n := range nodes {
|
||||
_, err := pm.FilterForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
_, err = pm.MatchersForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
policyChanged, _, err := pm.SetUsers(tt.after)
|
||||
require.NoError(t, err)
|
||||
require.True(t, policyChanged, "moved self sources must reach clients")
|
||||
|
||||
fresh, err := NewPolicyManager([]byte(pol), tt.after, nodes.ViewSlice())
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, n := range nodes {
|
||||
got, err := pm.FilterForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
|
||||
want, err := fresh.FilterForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff(want, got), "node %d FilterForNode (-fresh +cached)", n.ID)
|
||||
|
||||
gotM, err := pm.MatchersForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
|
||||
wantM, err := fresh.MatchersForNode(n.View())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, matcherStrings(wantM), matcherStrings(gotM), "node %d MatchersForNode", n.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvalidateGlobalPolicyCache tests the cache invalidation logic for global policies.
|
||||
func TestInvalidateGlobalPolicyCache(t *testing.T) {
|
||||
mustIPPtr := func(s string) *netip.Addr {
|
||||
|
||||
Reference in New Issue
Block a user