diff --git a/cmd/tomswall/main.go b/cmd/tomswall/main.go index c617769..c96516d 100644 --- a/cmd/tomswall/main.go +++ b/cmd/tomswall/main.go @@ -33,6 +33,8 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`, root.AddCommand( applyCmd(), + tryCmd(), + confirmCmd(), planCmd(), validateCmd(), statusCmd(), diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go new file mode 100644 index 0000000..543d8d8 --- /dev/null +++ b/cmd/tomswall/try.go @@ -0,0 +1,110 @@ +package main + +import ( + "context" + "fmt" + "os" + "os/signal" + "path/filepath" + "strconv" + "strings" + "syscall" + "time" + + "github.com/spf13/cobra" + + "git.unkin.net/unkin/tomswall/internal/agent" +) + +var tryPIDFile = "/run/tomswall/try.pid" + +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, applies the configuration, then waits +for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is restored +atomically. 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) + + if err := os.MkdirAll(filepath.Dir(tryPIDFile), 0o755); err != nil { + return err + } + if err := os.WriteFile(tryPIDFile, []byte(strconv.Itoa(os.Getpid())), 0o644); err != nil { + return err + } + defer os.Remove(tryPIDFile) + + revert, err := agent.EngineApplier{}.Apply(context.Background(), cfg) + if err != nil { + return fmt.Errorf("applying changes: %w", err) + } + fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout) + if err := confirmOrRevert(confirm, abort, timeout, revert); err != nil { + return err + } + fmt.Println("Confirmed.") + return nil + }, + } + cmd.Flags().DurationVar(&timeout, "timeout", 60*time.Second, "time to wait for confirmation before reverting") + return cmd +} + +// confirmOrRevert waits for confirm; on abort or timeout it runs revert. +func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() error) error { + reason := "not confirmed within " + timeout.String() + select { + case <-confirm: + return nil + case s := <-abort: + reason = "interrupted by " + s.String() + case <-time.After(timeout): + } + if err := revert(); err != nil { + return fmt.Errorf("%s; revert failed: %w", reason, err) + } + 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 { + b, err := os.ReadFile(tryPIDFile) + if os.IsNotExist(err) { + return fmt.Errorf("no pending 'tomswall try'") + } + if err != nil { + return err + } + pid, err := strconv.Atoi(strings.TrimSpace(string(b))) + if err != nil { + return fmt.Errorf("parsing %s: %w", tryPIDFile, err) + } + // A stale pidfile must not signal an unrelated process. + comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)) + if err != nil || strings.TrimSpace(string(comm)) != "tomswall" { + return fmt.Errorf("no pending 'tomswall try' (stale %s)", tryPIDFile) + } + if err := syscall.Kill(pid, syscall.SIGUSR1); err != nil { + return err + } + fmt.Println("Confirmed.") + return nil + }, + } +} diff --git a/cmd/tomswall/try_test.go b/cmd/tomswall/try_test.go new file mode 100644 index 0000000..9652eb5 --- /dev/null +++ b/cmd/tomswall/try_test.go @@ -0,0 +1,51 @@ +package main + +import ( + "errors" + "os" + "syscall" + "testing" + "time" +) + +func TestConfirmOrRevert(t *testing.T) { + tests := []struct { + name string + confirm bool + abort bool + revertErr error + wantRevert bool + wantErr bool + }{ + {name: "confirmed keeps ruleset", confirm: 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}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + confirm := make(chan os.Signal, 1) + abort := make(chan os.Signal, 1) + if tt.confirm { + confirm <- syscall.SIGUSR1 + } + if tt.abort { + abort <- syscall.SIGHUP + } + reverted := false + err := confirmOrRevert(confirm, abort, 20*time.Millisecond, func() error { + reverted = true + return tt.revertErr + }) + if reverted != tt.wantRevert { + t.Errorf("reverted = %v, want %v", reverted, tt.wantRevert) + } + if (err != nil) != tt.wantErr { + t.Errorf("err = %v, wantErr %v", err, tt.wantErr) + } + if tt.revertErr != nil && !errors.Is(err, tt.revertErr) { + t.Errorf("revert error not wrapped: %v", err) + } + }) + } +} diff --git a/internal/agent/agent.go b/internal/agent/agent.go index bf7d68c..03d97d5 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -2,18 +2,21 @@ package agent import ( "context" + "errors" "fmt" "log/slog" + "net/url" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" ) -// Applier applies a translated config to the firewall. Abstracted so the run -// loop is testable without touching the kernel. +// Applier applies a translated config to the firewall and returns a func that +// restores the previous ruleset. Abstracted so the run loop is testable without +// touching the kernel. type Applier interface { - Apply(ctx context.Context, cfg *config.Config) error + Apply(ctx context.Context, cfg *config.Config) (revert func() error, err error) } // Agent runs the pull-apply-report loop for one device. @@ -24,6 +27,10 @@ type Agent struct { Applier Applier // Resolver overrides the DNS resolver (tests); nil derives it per-config. Resolver *Resolver + + // reverted is the last generation rolled back for cutting off the API; it + // is not re-applied until a newer generation is published. + reverted int64 } // Run loops until ctx is cancelled, applying one cycle per Interval (and once @@ -61,16 +68,20 @@ func (a *Agent) RunOnce(ctx context.Context) error { return fmt.Errorf("control plane unreachable and no cached config: %w", err) } // Re-apply last known-good; do not report a generation we didn't fetch. - return a.applyConfig(ctx, cached, false) + return a.applyConfig(ctx, cached, nil) } - if err := a.Cache.Write(raw); err != nil { - slog.Warn("agent: caching config failed", "err", err) + if a.reverted != 0 && rc.Generation == a.reverted { + slog.Warn("agent: skipping reverted generation", "generation", rc.Generation) + return nil } - return a.applyConfig(ctx, rc, true) + return a.applyConfig(ctx, rc, raw) } -func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { +// applyConfig applies rc. A freshly fetched config (raw != nil) is verified by +// reporting status over a new connection through the new ruleset; if the API is +// unreachable the previous ruleset is restored. Only verified configs are cached. +func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error { resolver := a.Resolver if resolver == nil { resolver = NewResolver(rc.Resolver) @@ -81,15 +92,28 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool if err != nil { return fmt.Errorf("translate: %w", err) } - if err := a.Applier.Apply(ctx, cfg); err != nil { + revert, err := a.Applier.Apply(ctx, cfg) + if err != nil { return fmt.Errorf("apply: %w", err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) - if report { + if raw != nil { + a.Client.HTTP.CloseIdleConnections() if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { + var uerr *url.Error + if errors.As(err, &uerr) { + if rerr := revert(); rerr != nil { + return fmt.Errorf("API unreachable after applying generation %d (%v); revert failed: %w", rc.Generation, err, rerr) + } + a.reverted = rc.Generation + return fmt.Errorf("reverted generation %d: API unreachable through new ruleset: %w", rc.Generation, err) + } slog.Warn("agent: reporting status failed", "err", err) } + if err := a.Cache.Write(raw); err != nil { + slog.Warn("agent: caching config failed", "err", err) + } // Report the FIB so the control plane can scope router enforcement. if fib := CollectFIB(ctx); len(fib) > 0 { if err := a.Client.ReportRoutes(ctx, fib); err != nil { @@ -103,18 +127,26 @@ 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. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { +// Apply snapshots the live table, then computes and applies the differential +// change set for cfg. The returned func atomically restores the snapshot. +func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { engine, err := nftables.NewEngine(cfg) if err != nil { - return fmt.Errorf("initializing nftables: %w", err) + return nil, fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return fmt.Errorf("computing changes: %w", err) + return nil, fmt.Errorf("computing changes: %w", err) } if changes.Empty() { - return nil + return func() error { return nil }, nil } - return engine.Apply(changes) + snap, err := engine.Snapshot() + if err != nil { + return nil, fmt.Errorf("snapshotting ruleset: %w", err) + } + if err := engine.Apply(changes); err != nil { + return nil, err + } + return func() error { return engine.Restore(snap) }, nil } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 3019815..cda54f7 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -126,16 +126,17 @@ func TestTranslateRejectsUnknownAction(t *testing.T) { } } -// fakeApplier records applied configs. +// fakeApplier records applied configs and reverts. type fakeApplier struct { count int32 + reverts int32 lastGen int } -func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { +func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { atomic.AddInt32(&f.count, 1) f.lastGen = len(cfg.Rules) - return nil + return func() error { atomic.AddInt32(&f.reverts, 1); return nil }, nil } const renderedYAML = `generation: 7 @@ -238,3 +239,71 @@ func TestRunOnceNoCacheReturnsError(t *testing.T) { t.Fatal("expected error when unreachable and no cache exists") } } + +// statusServer serves renderedYAML and answers status reports with handler. +func statusServer(t *testing.T, status http.HandlerFunc) *httptest.Server { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + _, _ = w.Write([]byte(renderedYAML)) + return + } + status(w, r) + })) + t.Cleanup(srv.Close) + return srv +} + +func TestRunOnceRevertsWhenAPIUnreachableAfterApply(t *testing.T) { + // Dropping the connection mimics a pushed rule that severs the API path. + srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { + conn, _, _ := w.(http.Hijacker).Hijack() + conn.Close() + }) + + applier := &fakeApplier{} + a := &Agent{ + Client: NewClient(srv.URL, "fw-a", "tok"), + Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, + Applier: applier, + } + if err := a.RunOnce(context.Background()); err == nil { + t.Fatal("expected revert error") + } + if applier.reverts != 1 { + t.Errorf("expected 1 revert, got %d", applier.reverts) + } + if cached, _ := a.Cache.Read(); cached != nil { + t.Error("reverted config must not be cached as known-good") + } + + // The same generation is not re-applied on the next cycle. + if err := a.RunOnce(context.Background()); err != nil { + t.Fatalf("second RunOnce: %v", err) + } + if applier.count != 1 { + t.Errorf("reverted generation re-applied: %d applies", applier.count) + } +} + +func TestRunOnceKeepsConfigOnAPIErrorStatus(t *testing.T) { + // An HTTP error still proves the API is reachable — no revert. + srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }) + + applier := &fakeApplier{} + a := &Agent{ + Client: NewClient(srv.URL, "fw-a", "tok"), + Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, + Applier: applier, + } + if err := a.RunOnce(context.Background()); err != nil { + t.Fatalf("RunOnce: %v", err) + } + if applier.reverts != 0 { + t.Errorf("unexpected revert") + } + if cached, _ := a.Cache.Read(); cached == nil { + t.Error("expected config cached") + } +} diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index 82cfa30..faa9860 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -2,6 +2,7 @@ package nftables import ( "fmt" + "sort" "strings" "github.com/google/nftables/expr" @@ -84,6 +85,32 @@ func computeDiff(current, desired *FirewallState) *ChangeSet { return cs } +// 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 +} + func rulesMatch(a, b []ManagedRule) bool { if len(a) != len(b) { return false diff --git a/internal/nftables/diff_test.go b/internal/nftables/diff_test.go new file mode 100644 index 0000000..33c9b2d --- /dev/null +++ b/internal/nftables/diff_test.go @@ -0,0 +1,65 @@ +package nftables + +import ( + "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 !exprsEqual(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 0618a6f..5aba7c1 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -135,31 +135,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 { @@ -168,7 +169,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{ @@ -183,6 +184,36 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { return state, nil } +// Snapshot is the tomswall table as captured live; a nil state means it was absent. +type Snapshot struct { + state *FirewallState +} + +// Snapshot captures the live tomswall table so Restore can roll back to it. +func (e *Engine) Snapshot() (*Snapshot, error) { + t, err := e.findTable() + if err != nil || t == nil { + return &Snapshot{}, err + } + state, err := e.readCurrentState() + if err != nil { + return nil, err + } + return &Snapshot{state: state}, nil +} + +// Restore atomically returns the tomswall table to the snapshot, rule order included. +func (e *Engine) Restore(s *Snapshot) error { + if s.state == nil { + return e.Flush() + } + current, err := e.readCurrentState() + if err != nil { + return err + } + return e.Apply(restoreChangeSet(current, s.state)) +} + func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { return &p }