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 c617769..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 @@ -33,6 +34,9 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`, root.AddCommand( applyCmd(), + tryCmd(), + confirmCmd(), + revertCmd(), planCmd(), validateCmd(), statusCmd(), @@ -86,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) @@ -227,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) @@ -272,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 new file mode 100644 index 0000000..ee9e21e --- /dev/null +++ b/cmd/tomswall/try.go @@ -0,0 +1,168 @@ +package main + +import ( + "fmt" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "github.com/spf13/cobra" + + "git.unkin.net/unkin/tomswall/internal/config" + "git.unkin.net/unkin/tomswall/internal/nftables" + "git.unkin.net/unkin/tomswall/internal/tryapply" +) + +// revertGrace lets the in-process revert win before the systemd fallback fires. +const revertGrace = 30 * time.Second + +func tryCmd() *cobra.Command { + var timeout time.Duration + cmd := &cobra.Command{ + Use: "try", + Short: "Apply configuration and revert unless confirmed within a timeout", + Long: `Try snapshots the live tomswall table to disk, applies the configuration, then +waits for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is +restored atomically. A transient systemd timer restores it too if this process +dies. Confirm from a new session to prove new connections still work.`, + RunE: func(cmd *cobra.Command, args []string) error { + cfg, err := loadConfig() + if err != nil { + return err + } + + // Registered before apply so a confirm or hangup cannot be missed. + signal.Ignore(syscall.SIGPIPE) + confirm := make(chan os.Signal, 1) + signal.Notify(confirm, syscall.SIGUSR1) + abort := make(chan os.Signal, 1) + signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) + + 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, func() (bool, error) { return tryapply.Revert(id) }) + if err != nil { + return err + } + fmt.Println(msg) + return nil + }, + } + cmd.Flags().DurationVar(&timeout, "timeout", 60*time.Second, "time to wait for confirmation before reverting") + return cmd +} + +// 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 "", err + } + defer unlock() + + engine, err := nftables.NewEngine(cfg) + if err != nil { + return "", fmt.Errorf("initializing nftables: %w", err) + } + changes, err := engine.Plan() + if err != nil { + return "", fmt.Errorf("computing changes: %w", err) + } + if changes.Empty() { + fmt.Println("No changes needed — firewall is up to date.") + return "", nil + } + fmt.Println(changes.Summary()) + + snap, err := engine.Snapshot() + if err != nil { + return "", fmt.Errorf("snapshotting ruleset: %w", 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 "", fmt.Errorf("applying changes: %w", err) + } + return id, nil +} + +// confirmOrRevert waits for confirm; on abort or timeout it runs revert. +// A confirm already delivered wins over a simultaneous abort or timeout. +func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() (bool, error)) (string, error) { + reason := "not confirmed within " + timeout.String() + select { + case <-confirm: + return "Confirmed.", nil + case s := <-abort: + reason = "interrupted by " + s.String() + case <-time.After(timeout): + } + select { + case <-confirm: + return "Confirmed.", nil + default: + } + reverted, err := revert() + if err != nil { + return "", fmt.Errorf("%s; revert failed: %w", reason, err) + } + if !reverted { + return "Already resolved by 'tomswall confirm' or 'tomswall revert'.", nil + } + return "", fmt.Errorf("%s: previous ruleset restored", reason) +} + +func confirmCmd() *cobra.Command { + return &cobra.Command{ + Use: "confirm", + Short: "Keep the configuration applied by a pending 'tomswall try'", + RunE: func(cmd *cobra.Command, args []string) error { + pid, ok, err := tryapply.Confirm() + if err != nil { + return err + } + if !ok { + return fmt.Errorf("no pending 'tomswall try': it was already reverted or confirmed") + } + // A recycled PID must not be signalled. + if comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)); err == nil && strings.TrimSpace(string(comm)) == "tomswall" { + _ = syscall.Kill(pid, syscall.SIGUSR1) + } + fmt.Println("Confirmed.") + return nil + }, + } +} + +func revertCmd() *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(id) + if err != nil { + return err + } + if !reverted { + fmt.Println("No pending 'tomswall try'.") + return nil + } + fmt.Println("Previous ruleset restored.") + return nil + }, + } + cmd.Flags().StringVar(&id, "id", "", "only revert the try with this ID (used by the revert timer)") + return cmd +} diff --git a/cmd/tomswall/try_test.go b/cmd/tomswall/try_test.go new file mode 100644 index 0000000..a8f99f8 --- /dev/null +++ b/cmd/tomswall/try_test.go @@ -0,0 +1,63 @@ +package main + +import ( + "errors" + "os" + "syscall" + "testing" + "time" +) + +func TestConfirmOrRevert(t *testing.T) { + tests := []struct { + name string + confirm bool + abort bool + resolved bool + revertErr error + wantRevert bool + wantErr bool + wantResolved bool + }{ + {name: "confirmed keeps ruleset", confirm: true}, + {name: "confirm wins over simultaneous abort", confirm: true, abort: true}, + {name: "timeout reverts", wantRevert: true, wantErr: true}, + {name: "hangup reverts", abort: true, wantRevert: true, wantErr: true}, + {name: "revert failure surfaces", revertErr: errors.New("boom"), wantRevert: true, wantErr: true}, + {name: "resolved elsewhere", resolved: true, wantRevert: true, wantResolved: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for i := 0; i < 20; i++ { // select between ready channels is random + confirm := make(chan os.Signal, 1) + abort := make(chan os.Signal, 1) + if tt.confirm { + confirm <- syscall.SIGUSR1 + } + if tt.abort { + abort <- syscall.SIGHUP + } + called := false + msg, err := confirmOrRevert(confirm, abort, 10*time.Millisecond, func() (bool, error) { + called = true + return !tt.resolved, tt.revertErr + }) + if called != tt.wantRevert { + t.Fatalf("revert called = %v, want %v", called, tt.wantRevert) + } + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if tt.revertErr != nil && !errors.Is(err, tt.revertErr) { + t.Fatalf("revert error not wrapped: %v", err) + } + if tt.wantResolved && msg == "" { + t.Fatal("expected already-resolved message") + } + if !(tt.confirm && tt.abort) { + break + } + } + }) + } +} diff --git a/go.mod b/go.mod index 5d8664e..0f68d47 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.23 require ( github.com/google/nftables v0.2.0 + github.com/mdlayher/netlink v1.7.2 github.com/spf13/cobra v1.8.1 golang.org/x/sys v0.18.0 gopkg.in/yaml.v3 v3.0.1 @@ -13,7 +14,6 @@ require ( github.com/google/go-cmp v0.6.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/josharian/native v1.1.0 // indirect - github.com/mdlayher/netlink v1.7.2 // indirect github.com/mdlayher/socket v0.5.1 // indirect github.com/spf13/pflag v1.0.5 // indirect golang.org/x/net v0.23.0 // indirect diff --git a/internal/agent/agent.go b/internal/agent/agent.go index bf7d68c..868c4e0 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -8,6 +8,7 @@ import ( "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" + "git.unkin.net/unkin/tomswall/internal/tryapply" ) // Applier applies a translated config to the firewall. Abstracted so the run @@ -103,8 +104,14 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} -// Apply computes and applies the differential change set for cfg. +// Apply computes and applies the differential change set for cfg. It refuses +// while a 'tomswall try' awaits confirmation. func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { + 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/internal/nftables/compiler.go b/internal/nftables/compiler.go index cd7964c..416ca06 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -839,7 +839,7 @@ func (c *Compiler) compileMSSClamp(state *FirewallState) { Type: 2, Offset: 2, Len: 2, - Op: 0, + Op: expr.ExthdrOpTcpopt, }, &expr.Cmp{Op: expr.CmpOpGt, Register: 1, Data: mssBytes}, &expr.Immediate{Register: 1, Data: mssBytes}, @@ -848,7 +848,7 @@ func (c *Compiler) compileMSSClamp(state *FirewallState) { Type: 2, Offset: 2, Len: 2, - Op: 1, + Op: expr.ExthdrOpTcpopt, }, ) state.Rules["forward"] = append(state.Rules["forward"], ManagedRule{ diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 8c650c8..623b4c3 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1202,14 +1202,26 @@ func TestCompile_MSSClamp(t *testing.T) { t.Fatalf("Compile() error: %v", err) } - found := false - for _, r := range state.Rules["forward"] { + var rule *ManagedRule + for i, r := range state.Rules["forward"] { if r.Tag == "mss:eth1" { - found = true + rule = &state.Rules["forward"][i] } } - if !found { - t.Error("MSS clamp rule not found in forward chain") + if rule == nil { + t.Fatal("MSS clamp rule not found in forward chain") + } + + mss := []byte{0x05, 0x78} + want := []expr.Any{ + &expr.Exthdr{DestRegister: 1, Type: 2, Offset: 2, Len: 2, Op: expr.ExthdrOpTcpopt}, + &expr.Cmp{Op: expr.CmpOpGt, Register: 1, Data: mss}, + &expr.Immediate{Register: 1, Data: mss}, + &expr.Exthdr{SourceRegister: 1, Type: 2, Offset: 2, Len: 2, Op: expr.ExthdrOpTcpopt}, + } + got := rule.Exprs[len(rule.Exprs)-len(want):] + if !reflect.DeepEqual(got, want) { + t.Errorf("MSS clamp exprs = %#v, want %#v", got, want) } } diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index e4bd988..e0a8167 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -105,3 +105,29 @@ func computeDiff(current, desired *FirewallState) *ChangeSet { func ruleEqual(a, b ManagedRule) bool { return a.Chain == b.Chain && a.Tag == b.Tag && reflect.DeepEqual(a.Exprs, b.Exprs) } + +// restoreChangeSet replaces every managed rule in current with the snapshot's, +// in snapshot order, so a restore cannot reorder rules. +func restoreChangeSet(current, snap *FirewallState) *ChangeSet { + cs := &ChangeSet{} + for _, rules := range current.Rules { + for _, r := range rules { + if r.Tag != "" { + cs.Remove = append(cs.Remove, r) + } + } + } + chains := make([]string, 0, len(snap.Rules)) + for c := range snap.Rules { + chains = append(chains, c) + } + sort.Strings(chains) + for _, c := range chains { + for _, r := range snap.Rules[c] { + if r.Tag != "" { + cs.Add = append(cs.Add, r) + } + } + } + return cs +} diff --git a/internal/nftables/diff_test.go b/internal/nftables/diff_test.go new file mode 100644 index 0000000..f0eff4c --- /dev/null +++ b/internal/nftables/diff_test.go @@ -0,0 +1,66 @@ +package nftables + +import ( + "reflect" + "testing" + + "github.com/google/nftables/expr" +) + +func TestRestoreChangeSet(t *testing.T) { + accept := []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} + drop := []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} + + snap := &FirewallState{Rules: map[string][]ManagedRule{ + "input": { + {Chain: "input", Tag: "ssh", Exprs: accept, Handle: 4}, + {Chain: "input", Tag: "web", Exprs: accept, Handle: 5}, + {Chain: "input", Tag: "", Exprs: drop, Handle: 6}, + }, + "forward": {{Chain: "forward", Tag: "fwd", Exprs: accept, Handle: 7}}, + }} + current := &FirewallState{Rules: map[string][]ManagedRule{ + "input": { + {Chain: "input", Tag: "web", Exprs: accept, Handle: 10}, + {Chain: "input", Tag: "ssh", Exprs: drop, Handle: 11}, + {Chain: "input", Tag: "", Exprs: drop, Handle: 12}, + }, + }} + + cs := restoreChangeSet(current, snap) + + var removed []uint64 + for _, r := range cs.Remove { + removed = append(removed, r.Handle) + } + if len(removed) != 2 || removed[0] != 10 || removed[1] != 11 { + t.Errorf("expected managed handles [10 11] removed, untagged kept; got %v", removed) + } + + var added []string + for _, r := range cs.Add { + added = append(added, r.Tag) + } + want := []string{"fwd", "ssh", "web"} + if len(added) != len(want) { + t.Fatalf("added %v, want %v", added, want) + } + for i := range want { + if added[i] != want[i] { + t.Fatalf("added %v, want %v (snapshot order per chain)", added, want) + } + } + if !reflect.DeepEqual(cs.Add[1].Exprs, accept) { + t.Error("ssh not restored to its snapshot exprs") + } +} + +func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) { + current := &FirewallState{Rules: map[string][]ManagedRule{ + "input": {{Chain: "input", Tag: "x", Handle: 1}}, + }} + cs := restoreChangeSet(current, &FirewallState{Rules: map[string][]ManagedRule{}}) + if len(cs.Remove) != 1 || len(cs.Add) != 0 { + t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add)) + } +} diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 32616cc..77a4afe 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -28,7 +28,8 @@ func (e *Engine) ensureTable() *nftables.Table { }) } -func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain { +// ensureChains declares the base chains; policies overrides their default policy. +func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain { chains := map[string]*nftables.Chain{ "input": { Name: "input", @@ -71,6 +72,9 @@ func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain } for name, chain := range chains { + if p, ok := policies[name]; ok { + chain.Policy = policyPtr(p) + } chains[name] = e.conn.AddChain(chain) } return chains @@ -93,8 +97,12 @@ func (e *Engine) Plan() (*ChangeSet, error) { } func (e *Engine) Apply(changes *ChangeSet) error { + return e.apply(changes, nil) +} + +func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { table := e.ensureTable() - chains := e.ensureChains(table) + chains := e.ensureChains(table, policies) for _, r := range changes.Remove { e.conn.DelRule(&nftables.Rule{ @@ -141,31 +149,32 @@ func (e *Engine) Flush() error { return nil } +func (e *Engine) findTable() (*nftables.Table, error) { + tables, err := e.conn.ListTables() + if err != nil { + return nil, fmt.Errorf("listing tables: %w", err) + } + for _, t := range tables { + if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { + return t, nil + } + } + return nil, nil +} + func (e *Engine) readCurrentState() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } - tables, err := e.conn.ListTables() - if err != nil { - return state, nil - } - - var ourTable *nftables.Table - for _, t := range tables { - if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { - ourTable = t - break - } - } - - if ourTable == nil { - return state, nil + ourTable, err := e.findTable() + if err != nil || ourTable == nil { + return state, err } chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) if err != nil { - return state, nil + return nil, fmt.Errorf("listing chains: %w", err) } for _, chain := range chains { @@ -174,7 +183,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { } rules, err := e.conn.GetRules(ourTable, chain) if err != nil { - continue + return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err) } for _, rule := range rules { state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{ @@ -189,6 +198,72 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { return state, nil } +// Snapshot is the tomswall table as captured live, serialisable so a revert +// survives the process that took it. +type Snapshot struct { + Table string `json:"table"` + Present bool `json:"present"` + Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` + Rules map[string][]SnapshotRule `json:"rules,omitempty"` +} + +// SnapshotRule is a managed rule with its expressions in netlink wire format. +type SnapshotRule struct { + Tag string `json:"tag"` + Exprs [][]byte `json:"exprs"` +} + +// Snapshot captures the live tomswall table so Restore can roll back to it. +func (e *Engine) Snapshot() (*Snapshot, error) { + snap := &Snapshot{Table: e.cfg.Settings.TableName} + t, err := e.findTable() + if err != nil || t == nil { + return snap, err + } + snap.Present = true + + chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) + if err != nil { + return nil, fmt.Errorf("listing chains: %w", err) + } + snap.Policies = make(map[string]nftables.ChainPolicy) + for _, c := range chains { + if c.Table.Name == snap.Table && c.Policy != nil { + snap.Policies[c.Name] = *c.Policy + } + } + + state, err := e.readCurrentState() + if err != nil { + return nil, err + } + snap.Rules, err = encodeState(state) + if err != nil { + return nil, err + } + return snap, nil +} + +// Restore atomically returns the tomswall table to the snapshot: rule order +// and chain policies included, or removed if it was absent. +func (e *Engine) Restore(s *Snapshot) error { + if s.Table != e.cfg.Settings.TableName { + return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) + } + if !s.Present { + return e.Flush() + } + want, err := decodeState(s.Rules) + if err != nil { + return err + } + current, err := e.readCurrentState() + if err != nil { + return err + } + return e.apply(restoreChangeSet(current, want), s.Policies) +} + func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { return &p } diff --git a/internal/nftables/snapshot.go b/internal/nftables/snapshot.go new file mode 100644 index 0000000..bcb9b41 --- /dev/null +++ b/internal/nftables/snapshot.go @@ -0,0 +1,119 @@ +package nftables + +import ( + "encoding/binary" + "fmt" + + "github.com/google/nftables" + "github.com/google/nftables/expr" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +const inet = byte(nftables.TableFamilyINet) + +// exprByName mirrors the expression types google/nftables can parse back from the kernel. +var exprByName = map[string]func() expr.Any{ + "ct": func() expr.Any { return &expr.Ct{} }, + "range": func() expr.Any { return &expr.Range{} }, + "meta": func() expr.Any { return &expr.Meta{} }, + "cmp": func() expr.Any { return &expr.Cmp{} }, + "counter": func() expr.Any { return &expr.Counter{} }, + "objref": func() expr.Any { return &expr.Objref{} }, + "payload": func() expr.Any { return &expr.Payload{} }, + "lookup": func() expr.Any { return &expr.Lookup{} }, + "immediate": func() expr.Any { return &expr.Immediate{} }, + "bitwise": func() expr.Any { return &expr.Bitwise{} }, + "redir": func() expr.Any { return &expr.Redir{} }, + "nat": func() expr.Any { return &expr.NAT{} }, + "limit": func() expr.Any { return &expr.Limit{} }, + "quota": func() expr.Any { return &expr.Quota{} }, + "dynset": func() expr.Any { return &expr.Dynset{} }, + "log": func() expr.Any { return &expr.Log{} }, + "exthdr": func() expr.Any { return &expr.Exthdr{} }, + "connlimit": func() expr.Any { return &expr.Connlimit{} }, + "queue": func() expr.Any { return &expr.Queue{} }, + "flow_offload": func() expr.Any { return &expr.FlowOffload{} }, + "reject": func() expr.Any { return &expr.Reject{} }, + "masq": func() expr.Any { return &expr.Masq{} }, + "hash": func() expr.Any { return &expr.Hash{} }, + "notrack": func() expr.Any { return &expr.Notrack{} }, +} + +func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { + out := make(map[string][]SnapshotRule, len(state.Rules)) + for chain, rules := range state.Rules { + for _, r := range rules { + sr := SnapshotRule{Tag: r.Tag} + for _, e := range r.Exprs { + b, err := expr.Marshal(inet, e) + if err != nil { + return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err) + } + sr.Exprs = append(sr.Exprs, b) + } + out[chain] = append(out[chain], sr) + } + } + return out, nil +} + +func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { + state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))} + for chain, rs := range rules { + for _, sr := range rs { + r := ManagedRule{Chain: chain, Tag: sr.Tag} + for _, b := range sr.Exprs { + e, err := decodeExpr(b) + if err != nil { + return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err) + } + r.Exprs = append(r.Exprs, e) + } + state.Rules[chain] = append(state.Rules[chain], r) + } + } + return state, nil +} + +// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules. +func decodeExpr(b []byte) (expr.Any, error) { + ad, err := netlink.NewAttributeDecoder(b) + if err != nil { + return nil, err + } + ad.ByteOrder = binary.BigEndian + var name string + var data []byte + for ad.Next() { + switch ad.Type() { + case unix.NFTA_EXPR_NAME: + name = ad.String() + case unix.NFTA_EXPR_DATA: + data = ad.Bytes() + } + } + if err := ad.Err(); err != nil { + return nil, err + } + newExpr, ok := exprByName[name] + if !ok { + return nil, fmt.Errorf("unsupported expression %q", name) + } + e := newExpr() + if name == "notrack" { + return e, nil + } + if err := expr.Unmarshal(inet, data, e); err != nil { + return nil, err + } + // A verdict is an immediate into the verdict register with no data. + if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 { + v := &expr.Verdict{} + if err := expr.Unmarshal(inet, data, v); err != nil { + return nil, err + } + return v, nil + } + return e, nil +} diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go new file mode 100644 index 0000000..3098119 --- /dev/null +++ b/internal/nftables/snapshot_test.go @@ -0,0 +1,245 @@ +package nftables + +import ( + "bytes" + "encoding/binary" + "encoding/json" + "reflect" + "testing" + + "github.com/google/nftables" + "github.com/google/nftables/expr" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" + + "git.unkin.net/unkin/tomswall/internal/config" +) + +func TestSnapshotRulesRoundTrip(t *testing.T) { + exprs := []expr.Any{ + &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, + &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}}, + &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, + &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0, 22}}, + &expr.Ct{Register: 1, Key: expr.CtKeySTATE}, + &expr.Notrack{}, + &expr.Verdict{Kind: expr.VerdictAccept}, + } + state := &FirewallState{Rules: map[string][]ManagedRule{ + "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, + }} + + rules, err := encodeState(state) + if err != nil { + t.Fatal(err) + } + b, err := json.Marshal(&Snapshot{Table: "tomswall", Present: true, Rules: rules, + Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}}) + if err != nil { + t.Fatal(err) + } + var snap Snapshot + if err := json.Unmarshal(b, &snap); err != nil { + t.Fatal(err) + } + if snap.Policies["input"] != nftables.ChainPolicyAccept { + t.Errorf("policy lost: %v", snap.Policies) + } + got, err := decodeState(snap.Rules) + if err != nil { + t.Fatal(err) + } + in := got.Rules["input"] + if len(in) != 2 || in[0].Tag != "ssh" || in[1].Tag != "drop" { + t.Fatalf("rules/order lost: %+v", in) + } + if !reflect.DeepEqual(in[0].Exprs, exprs) { + t.Errorf("exprs changed:\n got %#v\nwant %#v", in[0].Exprs, exprs) + } + if _, ok := in[1].Exprs[0].(*expr.Verdict); !ok { + t.Errorf("verdict decoded as %T", in[1].Exprs[0]) + } +} + +func TestEnsureChainsPolicyOverride(t *testing.T) { + e := testEngine(t, nil) + chains := e.ensureChains(e.ensureTable(), map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}) + if *chains["input"].Policy != nftables.ChainPolicyAccept { + t.Error("input policy not overridden") + } + if *chains["forward"].Policy != nftables.ChainPolicyDrop { + t.Error("forward policy should keep its default") + } +} + +func TestSnapshotAndRestoreAbsentTable(t *testing.T) { + tablePresent := false + var sent []netlink.HeaderType + e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { + for _, m := range req { + sent = append(sent, m.Header.Type) + if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent { + data := []byte{inet, 0, 0, 0} + attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) + return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil + } + } + return nil, nil + }) + + snap, err := e.Snapshot() + if err != nil { + t.Fatal(err) + } + if snap.Present || snap.Table != "tomswall" { + t.Fatalf("want absent tomswall snapshot, got %+v", snap) + } + + // The try created the table; restoring the absent snapshot deletes it. + tablePresent = true + sent = nil + if err := e.Restore(snap); err != nil { + t.Fatal(err) + } + deleted := false + for _, ht := range sent { + deleted = deleted || ht == nftType(unix.NFT_MSG_DELTABLE) + } + if !deleted { + t.Errorf("table not deleted; sent %v", sent) + } +} + +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 { + t.Error("expected table mismatch error") + } +} + +func nftType(msg int) netlink.HeaderType { + return netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | msg) +} + +func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, error)) *Engine { + if dial == nil { + dial = func([]netlink.Message) ([]netlink.Message, error) { return nil, nil } + } + conn, err := nftables.New(nftables.WithTestDial(dial)) + if err != nil { + t.Fatal(err) + } + return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn} +} diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go new file mode 100644 index 0000000..c0e8fb4 --- /dev/null +++ b/internal/tryapply/tryapply.go @@ -0,0 +1,200 @@ +// Package tryapply keeps the state of a pending 'tomswall try' on disk so the +// revert survives the try process, backed by a transient systemd timer. +package tryapply + +import ( + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "syscall" + "time" + + "git.unkin.net/unkin/tomswall/internal/config" + "git.unkin.net/unkin/tomswall/internal/nftables" +) + +// Unit is the transient systemd unit that reverts an unconfirmed try. +const Unit = "tomswall-try-revert" + +var ( + // Dir holds the lock and the pending snapshot. + Dir = "/var/lib/tomswall" + // run executes a systemd command; replaced in tests. + run = func(name string, args ...string) error { + if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { + return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) + } + 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' 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"` +} + +func snapshotPath() string { return filepath.Join(Dir, "try-snapshot.json") } + +// Acquire takes the exclusive try lock, failing with ErrPending while a try is unconfirmed. +func Acquire() (unlock func(), err error) { + unlock, err = lock() + if err != nil { + return nil, err + } + if _, err := os.Stat(snapshotPath()); err == nil { + unlock() + return nil, ErrPending + } + return unlock, nil +} + +func lock() (func(), error) { + if err := os.MkdirAll(Dir, 0o755); err != nil { + return nil, err + } + f, err := os.OpenFile(filepath.Join(Dir, "try.lock"), os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, err + } + if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX); err != nil { + f.Close() + return nil, fmt.Errorf("locking %s: %w", f.Name(), err) + } + return func() { f.Close() }, nil +} + +// 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) (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 + } + f, err := os.CreateTemp(Dir, ".try-snapshot-*") + if err != nil { + return "", err + } + defer os.Remove(f.Name()) + if _, err := f.Write(b); err != nil { + f.Close() + return "", err + } + if err := f.Sync(); err != nil { + f.Close() + return "", err + } + if err := f.Close(); err != nil { + return "", err + } + if err := os.Rename(f.Name(), snapshotPath()); err != nil { + return "", err + } + + exe, err := os.Executable() + if err != nil { + 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", "--id", id); err != nil { + return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) + } + return id, nil +} + +// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. +func Discard() error { + _ = disarm() + if err := os.Remove(snapshotPath()); err != nil && !os.IsNotExist(err) { + return err + } + return nil +} + +func discardWith(err error) error { + if derr := Discard(); derr != nil { + return fmt.Errorf("%w (discarding snapshot: %v)", err, derr) + } + return err +} + +func disarm() error { + return run("systemctl", "stop", Unit+".timer") +} + +// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e. +// it was already reverted; pid is the waiting try process, if any. +func Confirm() (pid int, ok bool, err error) { + unlock, err := lock() + if err != nil { + return 0, false, err + } + defer unlock() + p, err := load() + if err != nil || p == nil { + return 0, false, err + } + return p.PID, true, Discard() +} + +// 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 || (id != "" && p.ID != id) { + return false, err + } + if err := restore(p.Snapshot); err != nil { + return false, fmt.Errorf("restoring snapshot: %w", err) + } + return true, Discard() +} + +func load() (*pending, error) { + b, err := os.ReadFile(snapshotPath()) + if os.IsNotExist(err) { + return nil, nil + } + if err != nil { + return nil, err + } + var p pending + if err := json.Unmarshal(b, &p); err != nil { + return nil, fmt.Errorf("parsing %s: %w", snapshotPath(), err) + } + if p.Snapshot == nil { + return nil, fmt.Errorf("%s has no snapshot", snapshotPath()) + } + return &p, nil +} diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go new file mode 100644 index 0000000..5b00f2c --- /dev/null +++ b/internal/tryapply/tryapply_test.go @@ -0,0 +1,211 @@ +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) + } +}