From bcc71f9b572965947972a71a33dd6d4cfdeab4bd Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Wed, 7 Oct 2026 14:29:56 +0000 Subject: [PATCH] integration: reuse images and remove redundant scenario setup --- integration/auth_key_test.go | 529 ++++++++++++++++----------------- integration/cli_policy_test.go | 140 +++++---- integration/scenario.go | 46 ++- 3 files changed, 369 insertions(+), 346 deletions(-) diff --git a/integration/auth_key_test.go b/integration/auth_key_test.go index 82db8b819..983aea5ba 100644 --- a/integration/auth_key_test.go +++ b/integration/auth_key_test.go @@ -1,7 +1,6 @@ package integration import ( - "fmt" "net/netip" "slices" "testing" @@ -22,196 +21,182 @@ import ( func TestAuthKeyLogoutAndReloginSameUser(t *testing.T) { IntegrationSkip(t) - for _, https := range []bool{true, false} { - t.Run(fmt.Sprintf("with-https-%t", https), func(t *testing.T) { - spec := ScenarioSpec{ - NodesPerUser: len(MustTestVersions), - Users: []string{"user1", "user2"}, - } - - scenario, err := NewScenario(spec) - - require.NoError(t, err) - defer scenario.ShutdownAssertNoPanics(t) - - opts := []hsic.Option{ - hsic.WithTestName("authkey-relogsame"), - } - - err = scenario.CreateHeadscaleEnv([]tsic.Option{}, opts...) - requireNoErrHeadscaleEnv(t, err) - - allClients, err := scenario.ListTailscaleClients() - requireNoErrListClients(t, err) - - allIps, err := scenario.ListTailscaleClientsIPs() - requireNoErrListClientIPs(t, err) - - err = scenario.WaitForTailscaleSync() - requireNoErrSync(t, err) - - headscale, err := scenario.Headscale() - requireNoErrGetHeadscale(t, err) - - expectedNodes := collectExpectedNodeIDs(t, allClients) - requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected", integrationutil.ScaledTimeout(120*time.Second)) - - // Validate that all nodes have [tailcfg.NetInfo] and DERP servers before logout - requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP before logout", 3*time.Minute) - - // assertClientsState(t, allClients) - - clientIPs := make(map[TailscaleClient][]netip.Addr) - - for _, client := range allClients { - ips, err := client.IPs() - if err != nil { - t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) - } - - clientIPs[client] = ips - } - - var ( - listNodes []*clientv1.Node - nodeCountBeforeLogout int - ) - - assert.EventuallyWithT(t, func(c *assert.CollectT) { - var err error - - listNodes, err = headscale.ListNodes() - assert.NoError(c, err) - assert.Len(c, listNodes, len(allClients)) - - for _, node := range listNodes { - assertLastSeenSetWithCollect(c, node) - } - }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for expected node list before logout") - - nodeCountBeforeLogout = len(listNodes) - t.Logf("node count before logout: %d", nodeCountBeforeLogout) - - for _, client := range allClients { - err := client.Logout() - if err != nil { - t.Fatalf("failed to logout client %s: %s", client.Hostname(), err) - } - } - - err = scenario.WaitForTailscaleLogout() - requireNoErrLogout(t, err) - - // After taking down all nodes, verify all systems show nodes offline - requireAllClientsOnline(t, headscale, expectedNodes, false, "all nodes should have logged out", integrationutil.ScaledTimeout(120*time.Second)) - - t.Logf("all clients logged out") - - t.Logf("Validating node persistence after logout at %s", time.Now().Format(TimestampFormat)) - assert.EventuallyWithT(t, func(ct *assert.CollectT) { - var err error - - listNodes, err = headscale.ListNodes() - assert.NoError(ct, err, "Failed to list nodes after logout") - assert.Len(ct, listNodes, nodeCountBeforeLogout, "Node count should match before logout count - expected %d nodes, got %d", nodeCountBeforeLogout, len(listNodes)) - }, integrationutil.StatusReadyTimeout, 2*time.Second, "validating node persistence after logout (nodes should remain in database)") - - for _, node := range listNodes { - assertLastSeenSet(t, node) - } - - // if the server is not running with HTTPS, we have to wait a bit before - // reconnection as the newest Tailscale client has a measure that will only - // reconnect over HTTPS if they saw a noise connection previously. - // https://github.com/tailscale/tailscale/commit/1eaad7d3deb0815e8932e913ca1a862afa34db38 - // https://github.com/juanfont/headscale/issues/2164 - if !https { - //nolint:forbidigo // Intentional delay: Tailscale client requires 5 min wait before reconnecting over non-HTTPS - time.Sleep(5 * time.Minute) - } - - userMap, err := headscale.MapUsers() - require.NoError(t, err) - - for _, userName := range spec.Users { - key, err := scenario.CreatePreAuthKey(mustParseID(userMap[userName].Id), true, false) - if err != nil { - t.Fatalf("failed to create pre-auth key for user %s: %s", userName, err) - } - - err = scenario.RunTailscaleUp(userName, headscale.GetEndpoint(), key.Key) - if err != nil { - t.Fatalf("failed to run tailscale up for user %s: %s", userName, err) - } - } - - t.Logf("Validating node persistence after relogin at %s", time.Now().Format(TimestampFormat)) - assert.EventuallyWithT(t, func(ct *assert.CollectT) { - var err error - - listNodes, err = headscale.ListNodes() - assert.NoError(ct, err, "Failed to list nodes after relogin") - assert.Len(ct, listNodes, nodeCountBeforeLogout, "Node count should remain unchanged after relogin - expected %d nodes, got %d", nodeCountBeforeLogout, len(listNodes)) - }, integrationutil.HAConvergeTimeout, 2*time.Second, "validating node count stability after same-user auth key relogin") - - for _, node := range listNodes { - assertLastSeenSet(t, node) - } - - requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected to batcher", integrationutil.ScaledTimeout(120*time.Second)) - - // Wait for Tailscale sync before validating [tailcfg.NetInfo] to ensure proper state propagation - err = scenario.WaitForTailscaleSync() - requireNoErrSync(t, err) - - // Validate that all nodes have [tailcfg.NetInfo] and DERP servers after reconnection - requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP after reconnection", 3*time.Minute) - - err = scenario.WaitForTailscaleSync() - requireNoErrSync(t, err) - - allAddrs := lo.Map(allIps, func(x netip.Addr, index int) string { - return x.String() - }) - - assertPingAll(t, allClients, allAddrs) - - for _, client := range allClients { - ips, err := client.IPs() - if err != nil { - t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) - } - - // lets check if the IPs are the same - if len(ips) != len(clientIPs[client]) { - t.Fatalf("IPs changed for client %s", client.Hostname()) - } - - for _, ip := range ips { - if !slices.Contains(clientIPs[client], ip) { - t.Fatalf( - "IPs changed for client %s. Used to be %v now %v", - client.Hostname(), - clientIPs[client], - ips, - ) - } - } - } - - assert.EventuallyWithT(t, func(c *assert.CollectT) { - var err error - - listNodes, err = headscale.ListNodes() - assert.NoError(c, err) - assert.Len(c, listNodes, nodeCountBeforeLogout) - - for _, node := range listNodes { - assertLastSeenSetWithCollect(c, node) - } - }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for node list after relogin") - }) + spec := ScenarioSpec{ + NodesPerUser: len(MustTestVersions), + Users: []string{"user1", "user2"}, } + + scenario, err := NewScenario(spec) + + require.NoError(t, err) + defer scenario.ShutdownAssertNoPanics(t) + + opts := []hsic.Option{ + hsic.WithTestName("authkey-relogsame"), + } + + err = scenario.CreateHeadscaleEnv([]tsic.Option{}, opts...) + requireNoErrHeadscaleEnv(t, err) + + allClients, err := scenario.ListTailscaleClients() + requireNoErrListClients(t, err) + + allIps, err := scenario.ListTailscaleClientsIPs() + requireNoErrListClientIPs(t, err) + + err = scenario.WaitForTailscaleSync() + requireNoErrSync(t, err) + + headscale, err := scenario.Headscale() + requireNoErrGetHeadscale(t, err) + + expectedNodes := collectExpectedNodeIDs(t, allClients) + requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected", integrationutil.ScaledTimeout(120*time.Second)) + + // Validate that all nodes have [tailcfg.NetInfo] and DERP servers before logout + requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP before logout", 3*time.Minute) + + // assertClientsState(t, allClients) + + clientIPs := make(map[TailscaleClient][]netip.Addr) + + for _, client := range allClients { + ips, err := client.IPs() + if err != nil { + t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) + } + + clientIPs[client] = ips + } + + var ( + listNodes []*clientv1.Node + nodeCountBeforeLogout int + ) + + assert.EventuallyWithT(t, func(c *assert.CollectT) { + var err error + + listNodes, err = headscale.ListNodes() + assert.NoError(c, err) + assert.Len(c, listNodes, len(allClients)) + + for _, node := range listNodes { + assertLastSeenSetWithCollect(c, node) + } + }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for expected node list before logout") + + nodeCountBeforeLogout = len(listNodes) + t.Logf("node count before logout: %d", nodeCountBeforeLogout) + + for _, client := range allClients { + err := client.Logout() + if err != nil { + t.Fatalf("failed to logout client %s: %s", client.Hostname(), err) + } + } + + err = scenario.WaitForTailscaleLogout() + requireNoErrLogout(t, err) + + // After taking down all nodes, verify all systems show nodes offline + requireAllClientsOnline(t, headscale, expectedNodes, false, "all nodes should have logged out", integrationutil.ScaledTimeout(120*time.Second)) + + t.Logf("all clients logged out") + + t.Logf("Validating node persistence after logout at %s", time.Now().Format(TimestampFormat)) + assert.EventuallyWithT(t, func(ct *assert.CollectT) { + var err error + + listNodes, err = headscale.ListNodes() + assert.NoError(ct, err, "Failed to list nodes after logout") + assert.Len(ct, listNodes, nodeCountBeforeLogout, "Node count should match before logout count - expected %d nodes, got %d", nodeCountBeforeLogout, len(listNodes)) + }, integrationutil.StatusReadyTimeout, 2*time.Second, "validating node persistence after logout (nodes should remain in database)") + + for _, node := range listNodes { + assertLastSeenSet(t, node) + } + + userMap, err := headscale.MapUsers() + require.NoError(t, err) + + for _, userName := range spec.Users { + key, err := scenario.CreatePreAuthKey(mustParseID(userMap[userName].Id), true, false) + if err != nil { + t.Fatalf("failed to create pre-auth key for user %s: %s", userName, err) + } + + err = scenario.RunTailscaleUp(userName, headscale.GetEndpoint(), key.Key) + if err != nil { + t.Fatalf("failed to run tailscale up for user %s: %s", userName, err) + } + } + + t.Logf("Validating node persistence after relogin at %s", time.Now().Format(TimestampFormat)) + assert.EventuallyWithT(t, func(ct *assert.CollectT) { + var err error + + listNodes, err = headscale.ListNodes() + assert.NoError(ct, err, "Failed to list nodes after relogin") + assert.Len(ct, listNodes, nodeCountBeforeLogout, "Node count should remain unchanged after relogin - expected %d nodes, got %d", nodeCountBeforeLogout, len(listNodes)) + }, integrationutil.HAConvergeTimeout, 2*time.Second, "validating node count stability after same-user auth key relogin") + + for _, node := range listNodes { + assertLastSeenSet(t, node) + } + + requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected to batcher", integrationutil.ScaledTimeout(120*time.Second)) + + // Wait for Tailscale sync before validating [tailcfg.NetInfo] to ensure proper state propagation + err = scenario.WaitForTailscaleSync() + requireNoErrSync(t, err) + + // Validate that all nodes have [tailcfg.NetInfo] and DERP servers after reconnection + requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP after reconnection", 3*time.Minute) + + err = scenario.WaitForTailscaleSync() + requireNoErrSync(t, err) + + allAddrs := lo.Map(allIps, func(x netip.Addr, index int) string { + return x.String() + }) + + assertPingAll(t, allClients, allAddrs) + + for _, client := range allClients { + ips, err := client.IPs() + if err != nil { + t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) + } + + // lets check if the IPs are the same + if len(ips) != len(clientIPs[client]) { + t.Fatalf("IPs changed for client %s", client.Hostname()) + } + + for _, ip := range ips { + if !slices.Contains(clientIPs[client], ip) { + t.Fatalf( + "IPs changed for client %s. Used to be %v now %v", + client.Hostname(), + clientIPs[client], + ips, + ) + } + } + } + + assert.EventuallyWithT(t, func(c *assert.CollectT) { + var err error + + listNodes, err = headscale.ListNodes() + assert.NoError(c, err) + assert.Len(c, listNodes, nodeCountBeforeLogout) + + for _, node := range listNodes { + assertLastSeenSetWithCollect(c, node) + } + }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for node list after relogin") } // This test will first log in two sets of nodes to two sets of users, then @@ -352,120 +337,106 @@ func TestAuthKeyLogoutAndReloginNewUser(t *testing.T) { func TestAuthKeyLogoutAndReloginSameUserExpiredKey(t *testing.T) { IntegrationSkip(t) - for _, https := range []bool{true, false} { - t.Run(fmt.Sprintf("with-https-%t", https), func(t *testing.T) { - spec := ScenarioSpec{ - NodesPerUser: len(MustTestVersions), - Users: []string{"user1", "user2"}, - } + spec := ScenarioSpec{ + NodesPerUser: len(MustTestVersions), + Users: []string{"user1", "user2"}, + } - scenario, err := NewScenario(spec) + scenario, err := NewScenario(spec) - require.NoError(t, err) - defer scenario.ShutdownAssertNoPanics(t) + require.NoError(t, err) + defer scenario.ShutdownAssertNoPanics(t) - opts := []hsic.Option{ - hsic.WithTestName("authkey-rlogexpired"), - } + opts := []hsic.Option{ + hsic.WithTestName("authkey-rlogexpired"), + } - err = scenario.CreateHeadscaleEnv([]tsic.Option{}, opts...) - requireNoErrHeadscaleEnv(t, err) + err = scenario.CreateHeadscaleEnv([]tsic.Option{}, opts...) + requireNoErrHeadscaleEnv(t, err) - allClients, err := scenario.ListTailscaleClients() - requireNoErrListClients(t, err) + allClients, err := scenario.ListTailscaleClients() + requireNoErrListClients(t, err) - err = scenario.WaitForTailscaleSync() - requireNoErrSync(t, err) + err = scenario.WaitForTailscaleSync() + requireNoErrSync(t, err) - // assertClientsState(t, allClients) + // assertClientsState(t, allClients) - clientIPs := make(map[TailscaleClient][]netip.Addr) + clientIPs := make(map[TailscaleClient][]netip.Addr) - for _, client := range allClients { - ips, err := client.IPs() - if err != nil { - t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) - } + for _, client := range allClients { + ips, err := client.IPs() + if err != nil { + t.Fatalf("failed to get IPs for client %s: %s", client.Hostname(), err) + } - clientIPs[client] = ips - } + clientIPs[client] = ips + } - headscale, err := scenario.Headscale() - requireNoErrGetHeadscale(t, err) + headscale, err := scenario.Headscale() + requireNoErrGetHeadscale(t, err) - // Collect expected node IDs for validation - expectedNodes := collectExpectedNodeIDs(t, allClients) + // Collect expected node IDs for validation + expectedNodes := collectExpectedNodeIDs(t, allClients) - // Validate initial connection state - requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected after initial login", integrationutil.ScaledTimeout(120*time.Second)) - requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP after initial login", 3*time.Minute) + // Validate initial connection state + requireAllClientsOnline(t, headscale, expectedNodes, true, "all clients should be connected after initial login", integrationutil.ScaledTimeout(120*time.Second)) + requireAllClientsNetInfoAndDERP(t, headscale, expectedNodes, "all clients should have NetInfo and DERP after initial login", 3*time.Minute) - var ( - listNodes []*clientv1.Node - nodeCountBeforeLogout int - ) + var ( + listNodes []*clientv1.Node + nodeCountBeforeLogout int + ) - assert.EventuallyWithT(t, func(c *assert.CollectT) { - var err error + assert.EventuallyWithT(t, func(c *assert.CollectT) { + var err error - listNodes, err = headscale.ListNodes() - assert.NoError(c, err) - assert.Len(c, listNodes, len(allClients)) - }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for expected node list before logout") + listNodes, err = headscale.ListNodes() + assert.NoError(c, err) + assert.Len(c, listNodes, len(allClients)) + }, integrationutil.ScaledTimeout(10*time.Second), integrationutil.FastPoll, "Waiting for expected node list before logout") - nodeCountBeforeLogout = len(listNodes) - t.Logf("node count before logout: %d", nodeCountBeforeLogout) + nodeCountBeforeLogout = len(listNodes) + t.Logf("node count before logout: %d", nodeCountBeforeLogout) - for _, client := range allClients { - err := client.Logout() - if err != nil { - t.Fatalf("failed to logout client %s: %s", client.Hostname(), err) - } - } + for _, client := range allClients { + err := client.Logout() + if err != nil { + t.Fatalf("failed to logout client %s: %s", client.Hostname(), err) + } + } - err = scenario.WaitForTailscaleLogout() - requireNoErrLogout(t, err) + err = scenario.WaitForTailscaleLogout() + requireNoErrLogout(t, err) - // Validate that all nodes are offline after logout - requireAllClientsOnline(t, headscale, expectedNodes, false, "all nodes should be offline after logout", integrationutil.ScaledTimeout(120*time.Second)) + // Validate that all nodes are offline after logout + requireAllClientsOnline(t, headscale, expectedNodes, false, "all nodes should be offline after logout", integrationutil.ScaledTimeout(120*time.Second)) - t.Logf("all clients logged out") + t.Logf("all clients logged out") - // if the server is not running with HTTPS, we have to wait a bit before - // reconnection as the newest Tailscale client has a measure that will only - // reconnect over HTTPS if they saw a noise connection previously. - // https://github.com/tailscale/tailscale/commit/1eaad7d3deb0815e8932e913ca1a862afa34db38 - // https://github.com/juanfont/headscale/issues/2164 - if !https { - //nolint:forbidigo // Intentional delay: Tailscale client requires 5 min wait before reconnecting over non-HTTPS - time.Sleep(5 * time.Minute) - } + userMap, err := headscale.MapUsers() + require.NoError(t, err) - userMap, err := headscale.MapUsers() - require.NoError(t, err) + for _, userName := range spec.Users { + key, err := scenario.CreatePreAuthKey(mustParseID(userMap[userName].Id), true, false) + if err != nil { + t.Fatalf("failed to create pre-auth key for user %s: %s", userName, err) + } - for _, userName := range spec.Users { - key, err := scenario.CreatePreAuthKey(mustParseID(userMap[userName].Id), true, false) - if err != nil { - t.Fatalf("failed to create pre-auth key for user %s: %s", userName, err) - } + // Expire the key so it can't be used + _, err = headscale.Execute( + []string{ + "headscale", + "preauthkeys", + "expire", + "--id", + key.Id, + }) + require.NoError(t, err) + require.NoError(t, err) - // Expire the key so it can't be used - _, err = headscale.Execute( - []string{ - "headscale", - "preauthkeys", - "expire", - "--id", - key.Id, - }) - require.NoError(t, err) - require.NoError(t, err) - - err = scenario.RunTailscaleUp(userName, headscale.GetEndpoint(), key.Key) - assert.ErrorContains(t, err, "authkey expired") - } - }) + err = scenario.RunTailscaleUp(userName, headscale.GetEndpoint(), key.Key) + assert.ErrorContains(t, err, "authkey expired") } } diff --git a/integration/cli_policy_test.go b/integration/cli_policy_test.go index 20e003ed3..13f20dad1 100644 --- a/integration/cli_policy_test.go +++ b/integration/cli_policy_test.go @@ -25,9 +25,10 @@ import ( // - bypass: no-bypass talks to the server over gRPC; bypass opens the // database directly. // -// Each row spins up its own scenario because policy_mode is fixed at boot -// via `HEADSCALE_POLICY_MODE`. The two users + two nodes give the tests -// block real `user@` aliases to resolve against. +// Each policy mode shares one scenario across its sequential checks because +// policy_mode is fixed at boot. The two users + two nodes give the tests +// block real `user@` aliases to resolve against. Check inputs are separate +// from the live policy, which must remain unchanged after every check. func TestPolicyCheckCommand(t *testing.T) { IntegrationSkip(t) @@ -80,44 +81,56 @@ func TestPolicyCheckCommand(t *testing.T) { type row struct { name string - policyMode string fixture fixture bypass bool wantErr string wantStdout string } - modes := []string{"file", "database"} //nolint:goconst // axis labels match HEADSCALE_POLICY_MODE values + modes := []types.PolicyMode{types.PolicyModeFile, types.PolicyModeDB} bypasses := []bool{false, true} - rows := make([]row, 0, len(modes)*len(fixtures)*len(bypasses)) + rows := make([]row, 0, len(fixtures)*len(bypasses)) - for _, mode := range modes { - for _, f := range fixtures { - for _, bypass := range bypasses { - suffix := "no-bypass" - if bypass { - suffix = "bypass" - } - - r := row{ - name: mode + "-" + f.name + "-" + suffix, - policyMode: mode, - fixture: f, - bypass: bypass, - wantStdout: "Policy is valid", - } - if f.name == "acl-plus-failing-tests" { - r.wantErr = "test(s) failed" - r.wantStdout = "" - } - - rows = append(rows, r) + for _, f := range fixtures { + for _, bypass := range bypasses { + suffix := "no-bypass" + if bypass { + suffix = "bypass" } + + r := row{ + name: f.name + "-" + suffix, + fixture: f, + bypass: bypass, + wantStdout: "Policy is valid", + } + if f.name == "acl-plus-failing-tests" { + r.wantErr = "test(s) failed" + r.wantStdout = "" + } + + rows = append(rows, r) } } - for _, tt := range rows { - t.Run(tt.name, func(t *testing.T) { + // Use a live policy distinct from every check fixture so an accidental + // policy update is observable, including for the ACL-only fixture. + livePolicy := policyv2.Policy{ + ACLs: []policyv2.ACL{ + { + Action: policyv2.ActionAccept, + Sources: []policyv2.Alias{wildcard()}, + Destinations: []policyv2.AliasWithPorts{ + aliasWithPorts(wildcard(), tailcfg.PortRangeAny), + }, + }, + }, + } + livePolicyBytes, err := json.Marshal(livePolicy) + require.NoError(t, err) + + for _, mode := range modes { + t.Run(string(mode), func(t *testing.T) { spec := ScenarioSpec{ NodesPerUser: 1, Users: []string{"user1", "user2"}, //nolint:goconst // matches usernamep("user1@")/("user2@") above @@ -126,45 +139,62 @@ func TestPolicyCheckCommand(t *testing.T) { scenario, err := NewScenario(spec) require.NoError(t, err) - defer scenario.ShutdownAssertNoPanics(t) + t.Cleanup(func() { scenario.ShutdownAssertNoPanics(t) }) err = scenario.CreateHeadscaleEnv( []tsic.Option{}, hsic.WithTestName("cli-policycheck"), - hsic.WithConfigEnv(map[string]string{ - "HEADSCALE_POLICY_MODE": tt.policyMode, //nolint:goconst // env var name from hscontrol/types/config.go - }), + hsic.WithPolicyMode(mode), + hsic.WithACLPolicy(&livePolicy), ) require.NoError(t, err) headscale, err := scenario.Headscale() require.NoError(t, err) - pBytes, err := json.Marshal(tt.fixture.policy) - require.NoError(t, err) + require.EventuallyWithT(t, func(c *assert.CollectT) { + stdout, err := headscale.Execute([]string{"headscale", "policy", "get"}) + assert.NoError(c, err) + assert.JSONEq(c, string(livePolicyBytes), stdout) + }, integrationutil.StatusReadyTimeout, integrationutil.FastPoll, "live policy should be loaded before checks") - policyFilePath := "/etc/headscale/policy.json" //nolint:goconst // standard headscale policy path - err = headscale.WriteFile(policyFilePath, pBytes) - require.NoError(t, err) + for _, tt := range rows { + t.Run(tt.name, func(t *testing.T) { + pBytes, err := json.Marshal(tt.fixture.policy) + require.NoError(t, err) - cmd := []string{"headscale", "policy", "check", "-f", policyFilePath} //nolint:goconst // CLI invocation - if tt.bypass { - // --force suppresses the "is the server running?" - // confirmation prompt so the command can run - // non-interactively under the test harness. - cmd = append(cmd, "--bypass-server-and-access-database-directly", "--force") + policyFilePath := "/etc/headscale/policy-check-" + tt.name + ".json" + err = headscale.WriteFile(policyFilePath, pBytes) + require.NoError(t, err) + + t.Cleanup(func() { + assert.EventuallyWithT(t, func(c *assert.CollectT) { + stdout, err := headscale.Execute([]string{"headscale", "policy", "get"}) + assert.NoError(c, err) + assert.JSONEq(c, string(livePolicyBytes), stdout) + }, integrationutil.StatusReadyTimeout, integrationutil.FastPoll, "policy check must not alter the live policy") + }) + + cmd := []string{"headscale", "policy", "check", "-f", policyFilePath} //nolint:goconst // CLI invocation + if tt.bypass { + // --force suppresses the "is the server running?" + // confirmation prompt so the command can run + // non-interactively under the test harness. + cmd = append(cmd, "--bypass-server-and-access-database-directly", "--force") + } + + stdout, err := headscale.Execute(cmd) + + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + + return + } + + require.NoError(t, err) + require.Contains(t, stdout, tt.wantStdout) + }) } - - stdout, err := headscale.Execute(cmd) - - if tt.wantErr != "" { - require.ErrorContains(t, err, tt.wantErr) - - return - } - - require.NoError(t, err) - require.Contains(t, stdout, tt.wantStdout) }) } } diff --git a/integration/scenario.go b/integration/scenario.go index f6dd3da53..110225765 100644 --- a/integration/scenario.go +++ b/integration/scenario.go @@ -52,11 +52,12 @@ const ( var usePostgresForTest = envknob.Bool("HEADSCALE_INTEGRATION_POSTGRES") var ( - errNoHeadscaleAvailable = errors.New("no headscale available") - errNoUserAvailable = errors.New("no user available") - errNoClientFound = errors.New("client not found") - errInvalidMockOIDCImage = errors.New("invalid HEADSCALE_INTEGRATION_HEADSCALE_IMAGE format, expected repository:tag") - errMockOIDCImageRequiredInCI = errors.New("HEADSCALE_INTEGRATION_HEADSCALE_IMAGE must be set for mock OIDC in CI") + errNoHeadscaleAvailable = errors.New("no headscale available") + errNoUserAvailable = errors.New("no user available") + errNoClientFound = errors.New("client not found") + errInvalidHeadscaleImageFormat = errors.New("invalid HEADSCALE_INTEGRATION_HEADSCALE_IMAGE format, expected repository:tag") + errInvalidMockOIDCImage = errors.New("invalid HEADSCALE_INTEGRATION_HEADSCALE_IMAGE format, expected repository:tag") + errMockOIDCImageRequiredInCI = errors.New("HEADSCALE_INTEGRATION_HEADSCALE_IMAGE must be set for mock OIDC in CI") // AllVersions represents a list of Tailscale versions the suite // uses to test compatibility with the [ControlServer]. @@ -686,8 +687,8 @@ func (s *Scenario) CreateTailscaleNodesInUser( s.mu.Lock() - opts = append( - opts, + clientOpts := append( + slices.Clone(opts), tsic.WithCACert(cert), tsic.WithHeadscaleName(hostname), tsic.WithExtraHosts(extraHosts), @@ -700,7 +701,7 @@ func (s *Scenario) CreateTailscaleNodesInUser( tsClient, err := tsic.New( s.pool, version, - opts..., + clientOpts..., ) s.mu.Unlock() @@ -1757,11 +1758,32 @@ func Webservice(s *Scenario, networkName string) (*dockertest.Resource, error) { ContextDir: dockerContextPath, } - web, err := s.pool.BuildAndRunWithBuildOptions( - webBOpts, - webOpts, - dockertestutil.DockerRestartPolicy, + var ( + web *dockertest.Resource + err error ) + + if prebuiltImage := os.Getenv("HEADSCALE_INTEGRATION_HEADSCALE_IMAGE"); prebuiltImage != "" { + repo, tag, ok := strings.Cut(prebuiltImage, ":") + if !ok { + return nil, errInvalidHeadscaleImageFormat + } + + webOpts.Repository = repo + webOpts.Tag = tag + + web, err = s.pool.RunWithOptions( + webOpts, + dockertestutil.DockerRestartPolicy, + ) + } else { + web, err = s.pool.BuildAndRunWithBuildOptions( + webBOpts, + webOpts, + dockertestutil.DockerRestartPolicy, + ) + } + if err != nil { return nil, err }