From b8146aa7cdcc175e169bf3d0335a1e6b75a35453 Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Wed, 7 Oct 2026 09:36:25 +0000 Subject: [PATCH] tests: preserve SSH regressions after combining policy fixes Updates #3517 --- hscontrol/noise_test.go | 9 +++++---- hscontrol/policy/v2/policy_test.go | 2 +- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/hscontrol/noise_test.go b/hscontrol/noise_test.go index dcaf40969..91ad75603 100644 --- a/hscontrol/noise_test.go +++ b/hscontrol/noise_test.go @@ -914,7 +914,7 @@ func newSSHVerdictFixture(t *testing.T) *sshVerdictFixture { }`, user.Name+"@")) require.NoError(t, err) - period, checkFound := app.state.SSHCheckParams(ids[0], ids[1]) + period, checkFound := app.state.SSHCheckParams(ids[0], ids[1], sshTestLocalUser) require.True(t, checkFound, "test setup: pair must be subject to a check") require.Zero(t, period, "test setup: checkPeriod must be always") @@ -1053,12 +1053,13 @@ func TestSSHActionFollowUp_ConsumedVerdictNotReplayed(t *testing.T) { _, err := f.ns.headscale.state.SetPolicy([]byte(`{}`)) require.NoError(t, err) - _, checkFound := f.ns.headscale.state.SSHCheckParams(f.src, f.dst) + _, checkFound := f.ns.headscale.state.SSHCheckParams(f.src, f.dst, sshTestLocalUser) require.False(t, checkFound, "test setup: pair must no longer be subject to a check") rec := f.followUp(t, authID) - assert.Equal(t, http.StatusBadRequest, rec.Code, - "replay without a check must be refused, body=%s", rec.Body.String()) + action := sshActionFromRecorder(t, rec) + assert.True(t, action.Reject, "replay without a check must be refused") + assert.False(t, action.Accept, "replay without a check must never approve access") }) } } diff --git a/hscontrol/policy/v2/policy_test.go b/hscontrol/policy/v2/policy_test.go index c5a427ebf..4420a7520 100644 --- a/hscontrol/policy/v2/policy_test.go +++ b/hscontrol/policy/v2/policy_test.go @@ -2713,7 +2713,7 @@ func TestUnregisteredUsersAreNoOp(t *testing.T) { continue } - period, ok := pm.SSHCheckParams(n.ID, p.ID) + period, ok := pm.SSHCheckParams(n.ID, p.ID, "alice") got[n.Hostname+"->"+p.Hostname+" via"] = pm.ViaRoutesForPeer(nv, p.View()) got[n.Hostname+"->"+p.Hostname+" check"] = checkParams{period, ok} }