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) { t.Helper() unlock, err := Acquire() if err != nil { t.Fatal(err) } defer unlock() if err := Arm(snap, 4242, 90*time.Second); err != nil { t.Fatal(err) } } 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}}}}}} 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.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") { 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"}) if _, err := Acquire(); !errors.Is(err, ErrPending) { t.Fatalf("second try: got %v, want ErrPending", 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) } }