From 0b92f1c2f3d9360f8a12c0b4fab0ae696d0c9353 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:56:52 +1000 Subject: [PATCH] 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. --- cmd/tomswall/guard_test.go | 30 ++++++++ cmd/tomswall/main.go | 21 ++++++ cmd/tomswall/try.go | 36 +++++---- internal/nftables/snapshot_test.go | 114 +++++++++++++++++++++++++++++ internal/tryapply/tryapply.go | 61 +++++++++------ internal/tryapply/tryapply_test.go | 90 +++++++++++++++++++++-- 6 files changed, 306 insertions(+), 46 deletions(-) create mode 100644 cmd/tomswall/guard_test.go diff --git a/cmd/tomswall/guard_test.go b/cmd/tomswall/guard_test.go new file mode 100644 index 0000000..12459a3 --- /dev/null +++ b/cmd/tomswall/guard_test.go @@ -0,0 +1,30 @@ +package main + +import ( + "errors" + "os" + "path/filepath" + "testing" + + "github.com/spf13/cobra" + + "git.unkin.net/unkin/tomswall/internal/tryapply" +) + +func TestMutatingCommandsRefuseWhileTryPending(t *testing.T) { + orig := tryapply.Dir + tryapply.Dir = t.TempDir() + t.Cleanup(func() { tryapply.Dir = orig }) + if err := os.WriteFile(filepath.Join(tryapply.Dir, "try-snapshot.json"), []byte("{}"), 0o600); err != nil { + t.Fatal(err) + } + configPath = "../../tomswall.example.yaml" + + for _, cmd := range []*cobra.Command{applyCmd(), flushCmd(), purgeCmd()} { + t.Run(cmd.Use, func(t *testing.T) { + if err := cmd.RunE(cmd, nil); !errors.Is(err, tryapply.ErrPending) { + t.Errorf("got %v, want ErrPending", err) + } + }) + } +} diff --git a/cmd/tomswall/main.go b/cmd/tomswall/main.go index cc4563a..bd886e1 100644 --- a/cmd/tomswall/main.go +++ b/cmd/tomswall/main.go @@ -12,6 +12,7 @@ import ( "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/shorewall" + "git.unkin.net/unkin/tomswall/internal/tryapply" ) var configPath string @@ -89,6 +90,12 @@ The firewall is never torn down — existing connections are preserved.`, return err } + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) @@ -230,6 +237,14 @@ func purgeCmd() *cobra.Command { return err } + if !dryRun { + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + } + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) @@ -275,6 +290,12 @@ func flushCmd() *cobra.Command { return err } + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go index 0af4a92..ee9e21e 100644 --- a/cmd/tomswall/try.go +++ b/cmd/tomswall/try.go @@ -40,12 +40,12 @@ dies. Confirm from a new session to prove new connections still work.`, abort := make(chan os.Signal, 1) signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) - applied, err := tryApply(cfg, timeout+revertGrace) - if err != nil || !applied { + id, err := tryApply(cfg, timeout+revertGrace) + if err != nil || id == "" { return err } fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout) - msg, err := confirmOrRevert(confirm, abort, timeout, tryapply.Revert) + msg, err := confirmOrRevert(confirm, abort, timeout, func() (bool, error) { return tryapply.Revert(id) }) if err != nil { return err } @@ -57,41 +57,44 @@ dies. Confirm from a new session to prove new connections still work.`, return cmd } -func tryApply(cfg *config.Config, fallback time.Duration) (bool, error) { +// tryApply applies cfg under a pending try and returns its ID, or "" when +// there was nothing to change. +func tryApply(cfg *config.Config, fallback time.Duration) (string, error) { unlock, err := tryapply.Acquire() if err != nil { - return false, err + return "", err } defer unlock() engine, err := nftables.NewEngine(cfg) if err != nil { - return false, fmt.Errorf("initializing nftables: %w", err) + return "", fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return false, fmt.Errorf("computing changes: %w", err) + return "", fmt.Errorf("computing changes: %w", err) } if changes.Empty() { fmt.Println("No changes needed — firewall is up to date.") - return false, nil + return "", nil } fmt.Println(changes.Summary()) snap, err := engine.Snapshot() if err != nil { - return false, fmt.Errorf("snapshotting ruleset: %w", err) + return "", fmt.Errorf("snapshotting ruleset: %w", err) } - if err := tryapply.Arm(snap, os.Getpid(), fallback); err != nil { - return false, err + id, err := tryapply.Arm(snap, os.Getpid(), fallback) + if err != nil { + return "", err } if err := engine.Apply(changes); err != nil { if derr := tryapply.Discard(); derr != nil { err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) } - return false, fmt.Errorf("applying changes: %w", err) + return "", fmt.Errorf("applying changes: %w", err) } - return true, nil + return id, nil } // confirmOrRevert waits for confirm; on abort or timeout it runs revert. @@ -143,11 +146,12 @@ func confirmCmd() *cobra.Command { } func revertCmd() *cobra.Command { - return &cobra.Command{ + var id string + cmd := &cobra.Command{ Use: "revert", Short: "Restore the ruleset saved by a pending 'tomswall try'", RunE: func(cmd *cobra.Command, args []string) error { - reverted, err := tryapply.Revert() + reverted, err := tryapply.Revert(id) if err != nil { return err } @@ -159,4 +163,6 @@ func revertCmd() *cobra.Command { return nil }, } + cmd.Flags().StringVar(&id, "id", "", "only revert the try with this ID (used by the revert timer)") + return cmd } diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go index 8eab0a7..3098119 100644 --- a/internal/nftables/snapshot_test.go +++ b/internal/nftables/snapshot_test.go @@ -1,6 +1,8 @@ package nftables import ( + "bytes" + "encoding/binary" "encoding/json" "reflect" "testing" @@ -108,6 +110,118 @@ func TestSnapshotAndRestoreAbsentTable(t *testing.T) { } } +func TestRestorePresentTable(t *testing.T) { + want := []SnapshotRule{} + for _, r := range []ManagedRule{ + {Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}}, + {Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}, + } { + enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}) + if err != nil { + t.Fatal(err) + } + want = append(want, enc["input"]...) + } + snap := &Snapshot{Table: "tomswall", Present: true, + Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}, + Rules: map[string][]SnapshotRule{"input": want}} + + // Live state: the tried ruleset left one managed rule (handle 7) in input. + attrs := func(a ...netlink.Attribute) []byte { + b, err := netlink.MarshalAttributes(a) + if err != nil { + t.Fatal(err) + } + return append([]byte{inet, 0, 0, 0}, b...) + } + handle := make([]byte, 8) + binary.BigEndian.PutUint64(handle, 7) + var batch []netlink.Message + e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { + if len(req) == 0 { + return nil, nil + } + reply := func(msg int, data []byte) ([]netlink.Message, error) { + return []netlink.Message{{Header: netlink.Header{Type: nftType(msg), Sequence: req[0].Header.Sequence}, Data: data}}, nil + } + switch req[0].Header.Type { + case nftType(unix.NFT_MSG_GETTABLE): + return reply(unix.NFT_MSG_NEWTABLE, attrs(netlink.Attribute{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")})) + case nftType(unix.NFT_MSG_GETCHAIN): + return reply(unix.NFT_MSG_NEWCHAIN, attrs( + netlink.Attribute{Type: unix.NFTA_CHAIN_TABLE, Data: []byte("tomswall\x00")}, + netlink.Attribute{Type: unix.NFTA_CHAIN_NAME, Data: []byte("input\x00")})) + case nftType(unix.NFT_MSG_GETRULE): + return reply(unix.NFT_MSG_NEWRULE, attrs( + netlink.Attribute{Type: unix.NFTA_RULE_TABLE, Data: []byte("tomswall\x00")}, + netlink.Attribute{Type: unix.NFTA_RULE_CHAIN, Data: []byte("input\x00")}, + netlink.Attribute{Type: unix.NFTA_RULE_HANDLE, Data: handle}, + netlink.Attribute{Type: unix.NFTA_RULE_USERDATA, Data: []byte("tried")})) + } + batch = append(batch, req...) + return nil, nil + }) + + if err := e.Restore(snap); err != nil { + t.Fatal(err) + } + + var deleted []uint64 + var added []SnapshotRule + policy := map[string]uint32{} + for _, m := range batch { + ad, err := netlink.NewAttributeDecoder(m.Data[4:]) + if err != nil { + t.Fatal(err) + } + ad.ByteOrder = binary.BigEndian + var name string + var r SnapshotRule + var h uint64 + var pol *uint32 + for ad.Next() { + switch { + case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_NAME: + name = ad.String() + case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_POLICY: + v := ad.Uint32() + pol = &v + case m.Header.Type == nftType(unix.NFT_MSG_DELRULE) && ad.Type() == unix.NFTA_RULE_HANDLE: + h = ad.Uint64() + case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_USERDATA: + r.Tag = string(ad.Bytes()) + case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_EXPRESSIONS: + ad.Nested(func(nad *netlink.AttributeDecoder) error { + for nad.Next() { + r.Exprs = append(r.Exprs, bytes.Clone(nad.Bytes())) + } + return nil + }) + } + } + switch m.Header.Type { + case nftType(unix.NFT_MSG_NEWCHAIN): + if pol != nil { + policy[name] = *pol + } + case nftType(unix.NFT_MSG_DELRULE): + deleted = append(deleted, h) + case nftType(unix.NFT_MSG_NEWRULE): + added = append(added, r) + } + } + + if !reflect.DeepEqual(deleted, []uint64{7}) { + t.Errorf("deleted handles %v, want [7]", deleted) + } + if !reflect.DeepEqual(added, want) { + t.Errorf("restored rules differ from snapshot:\n got %+v\nwant %+v", added, want) + } + if policy["input"] != uint32(nftables.ChainPolicyAccept) || policy["forward"] != uint32(nftables.ChainPolicyDrop) { + t.Errorf("chain policies %v: input must be restored to accept, forward keep drop", policy) + } +} + func TestRestoreRejectsOtherTable(t *testing.T) { e := testEngine(t, nil) if err := e.Restore(&Snapshot{Table: "other"}); err == nil { diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go index e8b5f30..c0e8fb4 100644 --- a/internal/tryapply/tryapply.go +++ b/internal/tryapply/tryapply.go @@ -3,6 +3,8 @@ package tryapply import ( + "crypto/rand" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -30,12 +32,21 @@ var ( } return nil } + // restore rolls the live table back to a snapshot; replaced in tests. + restore = func(s *nftables.Snapshot) error { + engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) + if err != nil { + return err + } + return engine.Restore(s) + } ) // ErrPending means a try awaits confirmation; nothing else may apply meanwhile. -var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' or 'tomswall revert'") +var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' to keep it or 'tomswall revert' to restore the previous ruleset (also the recovery if an automatic revert failed)") type pending struct { + ID string `json:"id"` PID int `json:"pid"` Snapshot *nftables.Snapshot `json:"snapshot"` } @@ -70,43 +81,49 @@ func lock() (func(), error) { return func() { f.Close() }, nil } -// Arm persists snap and schedules an out-of-process revert after delay. +// Arm persists snap and schedules an out-of-process revert after delay, +// returning the try ID that scopes later reverts to this try. // The caller must hold the lock from Acquire. -func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) error { - b, err := json.Marshal(pending{PID: pid, Snapshot: snap}) +func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) { + raw := make([]byte, 8) + if _, err := rand.Read(raw); err != nil { + return "", err + } + id := hex.EncodeToString(raw) + b, err := json.Marshal(pending{ID: id, PID: pid, Snapshot: snap}) if err != nil { - return err + return "", err } f, err := os.CreateTemp(Dir, ".try-snapshot-*") if err != nil { - return err + return "", err } defer os.Remove(f.Name()) if _, err := f.Write(b); err != nil { f.Close() - return err + return "", err } if err := f.Sync(); err != nil { f.Close() - return err + return "", err } if err := f.Close(); err != nil { - return err + return "", err } if err := os.Rename(f.Name(), snapshotPath()); err != nil { - return err + return "", err } exe, err := os.Executable() if err != nil { - return discardWith(err) + return "", discardWith(err) } _ = disarm() // a leftover timer from an earlier try would block the unit name if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, - fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert"); err != nil { - return discardWith(fmt.Errorf("arming revert timer: %w", err)) + fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { + return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) } - return nil + return id, nil } // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. @@ -144,23 +161,21 @@ func Confirm() (pid int, ok bool, err error) { return p.PID, true, Discard() } -// Revert restores the pending snapshot. reverted is false when nothing was -// pending (already confirmed or reverted). -func Revert() (reverted bool, err error) { +// Revert restores the pending snapshot. A non-empty id only reverts that try, +// so a stale timer cannot undo a newer one. reverted is false when nothing +// matching was pending (already confirmed or reverted). A failed restore keeps +// the snapshot so 'tomswall revert' can retry. +func Revert(id string) (reverted bool, err error) { unlock, err := lock() if err != nil { return false, err } defer unlock() p, err := load() - if err != nil || p == nil { + if err != nil || p == nil || (id != "" && p.ID != id) { return false, err } - engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: p.Snapshot.Table}}) - if err != nil { - return false, err - } - if err := engine.Restore(p.Snapshot); err != nil { + if err := restore(p.Snapshot); err != nil { return false, fmt.Errorf("restoring snapshot: %w", err) } return true, Discard() diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go index ab72aa1..5b00f2c 100644 --- a/internal/tryapply/tryapply_test.go +++ b/internal/tryapply/tryapply_test.go @@ -24,23 +24,25 @@ func setup(t *testing.T) *[]string { return &cmds } -func arm(t *testing.T, snap *nftables.Snapshot) { +func arm(t *testing.T, snap *nftables.Snapshot) string { t.Helper() unlock, err := Acquire() if err != nil { t.Fatal(err) } defer unlock() - if err := Arm(snap, 4242, 90*time.Second); err != nil { + 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}}}}}} - arm(t, snap) + id := arm(t, snap) info, err := os.Stat(snapshotPath()) if err != nil { @@ -53,12 +55,12 @@ func TestArmPersistsSnapshotAndTimer(t *testing.T) { if err != nil { t.Fatal(err) } - if p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) { + 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") { + !strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert --id "+id) { t.Errorf("unexpected arm command %q", last) } } @@ -78,9 +80,13 @@ func TestAbsentTableSnapshotRoundTrip(t *testing.T) { func TestAcquireRefusesWhilePending(t *testing.T) { setup(t) arm(t, &nftables.Snapshot{Table: "tomswall"}) - if _, err := Acquire(); !errors.Is(err, ErrPending) { + _, 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) { @@ -96,7 +102,7 @@ func TestArmFailureDiscardsSnapshot(t *testing.T) { t.Fatal(err) } defer unlock() - if err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil { + 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) { @@ -130,8 +136,76 @@ func TestConfirmAfterRevertFails(t *testing.T) { if err != nil || ok { t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err) } - reverted, err := Revert() + 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) + } +}