Files
tomswall/internal/tryapply/tryapply_test.go
T
unkin-agent 0b92f1c2f3
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Refuse mutating commands during a try and scope reverts to the try ID
apply, flush and purge take the try lock and fail with ErrPending, which now
names 'tomswall revert' as the recovery after a failed automatic revert. Each
try gets an ID passed to the timer's 'revert --id', so a stale timer cannot
revert a newer try. Adds tests for Revert success/failure/stale ID and for
restoring a present table at the engine level.
2026-10-03 20:56:52 +10:00

212 lines
5.4 KiB
Go

package tryapply
import (
"errors"
"os"
"reflect"
"strings"
"testing"
"time"
"git.unkin.net/unkin/tomswall/internal/nftables"
)
func setup(t *testing.T) *[]string {
t.Helper()
Dir = t.TempDir()
var cmds []string
orig := run
run = func(name string, args ...string) error {
cmds = append(cmds, name+" "+strings.Join(args, " "))
return nil
}
t.Cleanup(func() { run = orig })
return &cmds
}
func arm(t *testing.T, snap *nftables.Snapshot) string {
t.Helper()
unlock, err := Acquire()
if err != nil {
t.Fatal(err)
}
defer unlock()
id, err := Arm(snap, 4242, 90*time.Second)
if err != nil {
t.Fatal(err)
}
return id
}
func TestArmPersistsSnapshotAndTimer(t *testing.T) {
cmds := setup(t)
snap := &nftables.Snapshot{Table: "tomswall", Present: true,
Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1, 2, 3}}}}}}
id := arm(t, snap)
info, err := os.Stat(snapshotPath())
if err != nil {
t.Fatal(err)
}
if info.Mode().Perm() != 0o600 {
t.Errorf("snapshot mode %v, want 0600", info.Mode().Perm())
}
p, err := load()
if err != nil {
t.Fatal(err)
}
if p.ID != id || id == "" || p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) {
t.Errorf("round trip mismatch: %+v", p)
}
last := (*cmds)[len(*cmds)-1]
if !strings.HasPrefix(last, "systemd-run ") || !strings.Contains(last, "--unit "+Unit) ||
!strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert --id "+id) {
t.Errorf("unexpected arm command %q", last)
}
}
func TestAbsentTableSnapshotRoundTrip(t *testing.T) {
setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
p, err := load()
if err != nil {
t.Fatal(err)
}
if p.Snapshot.Present || p.Snapshot.Table != "tomswall" {
t.Errorf("absent table not preserved: %+v", p.Snapshot)
}
}
func TestAcquireRefusesWhilePending(t *testing.T) {
setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
_, err := Acquire()
if !errors.Is(err, ErrPending) {
t.Fatalf("second try: got %v, want ErrPending", err)
}
if !strings.Contains(err.Error(), "tomswall revert") {
t.Errorf("ErrPending does not name the recovery: %v", err)
}
}
func TestArmFailureDiscardsSnapshot(t *testing.T) {
setup(t)
run = func(name string, args ...string) error {
if name == "systemd-run" {
return errors.New("no systemd")
}
return nil
}
unlock, err := Acquire()
if err != nil {
t.Fatal(err)
}
defer unlock()
if _, err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil {
t.Fatal("expected arm error")
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot left behind without a revert timer")
}
}
func TestConfirmPendingDisarms(t *testing.T) {
cmds := setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
pid, ok, err := Confirm()
if err != nil || !ok || pid != 4242 {
t.Fatalf("Confirm = %d, %v, %v", pid, ok, err)
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot not removed")
}
if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" {
t.Errorf("timer not stopped, last command %q", last)
}
unlock, err := Acquire()
if err != nil {
t.Fatalf("new try refused after confirm: %v", err)
}
unlock()
}
func TestConfirmAfterRevertFails(t *testing.T) {
setup(t)
_, ok, err := Confirm()
if err != nil || ok {
t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err)
}
reverted, err := Revert("")
if err != nil || reverted {
t.Fatalf("Revert with nothing pending = %v, %v", reverted, err)
}
}
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
t.Helper()
var got []*nftables.Snapshot
orig := restore
restore = func(s *nftables.Snapshot) error {
got = append(got, s)
return err
}
t.Cleanup(func() { restore = orig })
return &got
}
func TestRevertPendingRestoresAndDisarms(t *testing.T) {
cmds := setup(t)
restored := stubRestore(t, nil)
snap := &nftables.Snapshot{Table: "tomswall", Present: true,
Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1}}}}}}
id := arm(t, snap)
reverted, err := Revert(id)
if err != nil || !reverted {
t.Fatalf("Revert = %v, %v", reverted, err)
}
if len(*restored) != 1 || !reflect.DeepEqual((*restored)[0], snap) {
t.Errorf("restored %+v, want the armed snapshot", *restored)
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot not removed")
}
if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" {
t.Errorf("timer not stopped, last command %q", last)
}
}
func TestRevertFailureKeepsSnapshot(t *testing.T) {
setup(t)
boom := errors.New("netlink down")
stubRestore(t, boom)
id := arm(t, &nftables.Snapshot{Table: "tomswall"})
if _, err := Revert(id); !errors.Is(err, boom) {
t.Fatalf("Revert error = %v, want %v", err, boom)
}
if _, err := os.Stat(snapshotPath()); err != nil {
t.Fatalf("snapshot gone after failed revert: %v", err)
}
if _, err := Acquire(); !errors.Is(err, ErrPending) {
t.Errorf("failed revert must keep the try pending, got %v", err)
}
}
func TestRevertStaleIDIgnored(t *testing.T) {
setup(t)
restored := stubRestore(t, nil)
arm(t, &nftables.Snapshot{Table: "tomswall"})
reverted, err := Revert("stale-try")
if err != nil || reverted {
t.Fatalf("stale Revert = %v, %v; want no-op", reverted, err)
}
if len(*restored) != 0 {
t.Error("stale timer restored a newer try's snapshot")
}
if _, err := os.Stat(snapshotPath()); err != nil {
t.Errorf("newer try's snapshot removed: %v", err)
}
}