138 lines
3.3 KiB
Go
138 lines
3.3 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) {
|
|
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)
|
|
}
|
|
}
|