0b92f1c2f3
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.
212 lines
5.4 KiB
Go
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)
|
|
}
|
|
}
|