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) } }