diff --git a/hscontrol/policy/v2/policy_test.go b/hscontrol/policy/v2/policy_test.go index 67603f650..e40594677 100644 --- a/hscontrol/policy/v2/policy_test.go +++ b/hscontrol/policy/v2/policy_test.go @@ -12,6 +12,7 @@ import ( "github.com/juanfont/headscale/hscontrol/types" "github.com/puzpuzpuz/xsync/v4" "github.com/stretchr/testify/require" + "pgregory.net/rapid" "tailscale.com/net/tsaddr" "tailscale.com/tailcfg" ) @@ -3017,6 +3018,190 @@ func TestSetNodesRetriesAfterFailedRecompile(t *testing.T) { require.Equal(t, gen+1, pm.NodesGeneration()) } +// TestSetNodesCachedResultsMatchFresh drives random node writes through +// one PolicyManager, reading every node between writes so its caches +// fill, and checks each read against a PolicyManager built fresh from the +// same nodes. A cache entry that a write should have dropped shows up as +// a difference. +func TestSetNodesCachedResultsMatchFresh(t *testing.T) { + users := types.Users{{ID: 1, Name: "u1"}, {ID: 2, Name: "u2"}, {ID: 3, Name: "u3"}} + + policies := []struct { + name string + policy string + }{ + {name: "global", policy: `{ + "groups": {"group:a": ["u1@", "u2@"]}, + "tagOwners": {"tag:srv": ["u1@"], "tag:router": ["u1@"], "tag:other": ["u1@"]}, + "acls": [ + {"action": "accept", "src": ["group:a"], "dst": ["tag:srv:*"]}, + {"action": "accept", "src": ["u3@"], "dst": ["10.33.0.0/16:*", "u1@:22"]} + ], + "ssh": [{"action": "accept", "src": ["group:a"], "dst": ["tag:srv"], "users": ["root"]}]}`}, + {name: "autogroup-self", policy: `{ + "groups": {"group:a": ["u1@"]}, + "tagOwners": {"tag:srv": ["u1@"], "tag:router": ["u1@"], "tag:other": ["u1@"]}, + "acls": [ + {"action": "accept", "src": ["autogroup:member"], "dst": ["autogroup:self:*"]}, + {"action": "accept", "src": ["group:a"], "dst": ["tag:srv:*"]} + ], + "ssh": [ + {"action": "accept", "src": ["autogroup:member"], "dst": ["autogroup:self"], "users": ["autogroup:nonroot"]}, + {"action": "accept", "src": ["group:a"], "dst": ["tag:srv"], "users": ["root"]} + ]}`}, + {name: "via", policy: `{ + "tagOwners": {"tag:srv": ["u1@"], "tag:router": ["u1@"], "tag:other": ["u1@"]}, + "grants": [ + {"src": ["u2@"], "dst": ["10.33.0.0/16"], "ip": ["*"], "via": ["tag:router"]}, + {"src": ["u3@", "tag:srv"], "dst": ["autogroup:internet"], "ip": ["*"], "via": ["tag:router"]}, + {"src": ["u1@"], "dst": ["tag:srv"], "ip": ["*"]} + ]}`}, + {name: "autogroup-self-and-via", policy: `{ + "groups": {"group:a": ["u1@", "u2@"]}, + "tagOwners": {"tag:srv": ["u1@"], "tag:router": ["u1@"], "tag:other": ["u1@"]}, + "grants": [ + {"src": ["group:a"], "dst": ["autogroup:self"], "ip": ["*"]}, + {"src": ["u2@", "tag:srv"], "dst": ["10.33.0.0/16"], "ip": ["*"], "via": ["tag:router"]}, + {"src": ["u3@"], "dst": ["tag:srv"], "ip": ["*"]} + ]}`}, + } + + subnets := []netip.Prefix{ + netip.MustParsePrefix("10.33.0.0/24"), + netip.MustParsePrefix("10.33.1.0/24"), + } + exits := []netip.Prefix{tsaddr.AllIPv4(), tsaddr.AllIPv6()} + + for _, pc := range policies { + t.Run(pc.name, func(t *testing.T) { + rapid.Check(t, func(rt *rapid.T) { + nodes := make(types.Nodes, 0, 8) + nextIP := 1 + + newNode := func(id types.NodeID, u types.User) *types.Node { + n := node(fmt.Sprintf("n%d", id), fmt.Sprintf("100.64.0.%d", nextIP), fmt.Sprintf("fd7a:115c:a1e0::%d", nextIP), u) + n.ID = id + nextIP++ + + return n + } + + for i := range 6 { + nodes = append(nodes, newNode(types.NodeID(i+1), users[i%len(users)])) //nolint:gosec // safe conversion in test + } + + nextID := types.NodeID(len(nodes) + 1) //nolint:gosec // safe conversion in test + + pm, err := NewPolicyManager([]byte(pc.policy), users, nodes.ViewSlice()) + if err != nil { + rt.Fatalf("new policy manager: %v", err) + } + + check := func() { + fresh, err := NewPolicyManager([]byte(pc.policy), users, nodes.ViewSlice()) + if err != nil { + rt.Fatalf("fresh policy manager: %v", err) + } + + for _, n := range nodes { + nv := n.View() + + got, _ := pm.FilterForNode(nv) + + want, _ := fresh.FilterForNode(nv) + if diff := cmp.Diff(want, got); diff != "" { + rt.Fatalf("node %d FilterForNode (-fresh +cached):\n%s", n.ID, diff) + } + + gotM, _ := pm.MatchersForNode(nv) + + wantM, _ := fresh.MatchersForNode(nv) + if diff := cmp.Diff(matcherStrings(wantM), matcherStrings(gotM)); diff != "" { + rt.Fatalf("node %d MatchersForNode (-fresh +cached):\n%s", n.ID, diff) + } + + gotS, err := pm.SSHPolicy("", nv) + if err != nil { + rt.Fatalf("node %d SSHPolicy: %v", n.ID, err) + } + + wantS, _ := fresh.SSHPolicy("", nv) + if diff := cmp.Diff(wantS, gotS); diff != "" { + rt.Fatalf("node %d SSHPolicy (-fresh +cached):\n%s", n.ID, diff) + } + } + } + + check() + + steps := rapid.IntRange(1, 25).Draw(rt, "steps") + for range steps { + // Several writes per SetNodes, as a NodeStore batch applies them. + writes := rapid.IntRange(1, 3).Draw(rt, "writes") + for range writes { + i := rapid.IntRange(0, len(nodes)-1).Draw(rt, "node") + // Mutate a copy: pm still holds views of the old node. + n := nodes[i].View().AsStruct() + + switch rapid.IntRange(0, 9).Draw(rt, "op") { + case 0: // tag + n.Tags = []string{rapid.SampledFrom([]string{"tag:srv", "tag:router", "tag:other"}).Draw(rt, "tag")} + n.UserID, n.User = nil, nil + case 1: // (re)assign user, untagging + u := rapid.SampledFrom(users).Draw(rt, "user") + n.Tags = nil + n.UserID, n.User = new(u.ID), new(u) + case 2: // new IPs + ip4 := netip.MustParseAddr(fmt.Sprintf("100.64.1.%d", nextIP)) + ip6 := netip.MustParseAddr(fmt.Sprintf("fd7a:115c:a1e0::1:%d", nextIP)) + n.IPv4, n.IPv6 = &ip4, &ip6 + nextIP++ + case 3: // announce and approve a subnet + p := rapid.SampledFrom(subnets).Draw(rt, "subnet") + n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: []netip.Prefix{p}} + n.ApprovedRoutes = []netip.Prefix{p} + case 4: // drop approvals + n.ApprovedRoutes = nil + case 5: // exit node + n.Hostinfo = &tailcfg.Hostinfo{RoutableIPs: exits} + n.ApprovedRoutes = exits + case 6: // add a node + nodes = append(nodes, newNode(nextID, rapid.SampledFrom(users).Draw(rt, "user"))) + nextID++ + + n = nil + case 7: // remove the node + if len(nodes) > 1 { + nodes = slices.Delete(slices.Clone(nodes), i, i+1) + } + + n = nil + case 8: // owner association not loaded + if !n.IsTagged() { + n.User = nil + } + case 9: // payload only + n.Hostname += "x" + } + + if n != nil { + nodes = slices.Clone(nodes) + nodes[i] = n + } + } + + _, err := pm.SetNodes(nodes.ViewSlice()) + if err != nil { + rt.Fatalf("SetNodes: %v", err) + } + + check() + } + }) + }) + } +} + // TestFailedSetUsersKeepsCompiledPolicy renames alice to charlie while // nodeAttrs still names alice, so the recompile resolves charlie's grant // and then fails. A failed Set* must leave every compiled result as it was: