Add try/confirm safe-apply and agent auto-revert
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-03 20:45:57 +10:00
parent 410109515e
commit cc12c4a43a
8 changed files with 423 additions and 36 deletions
+2
View File
@@ -33,6 +33,8 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`,
root.AddCommand(
applyCmd(),
tryCmd(),
confirmCmd(),
planCmd(),
validateCmd(),
statusCmd(),
+110
View File
@@ -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
},
}
}
+51
View File
@@ -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)
}
})
}
}
+48 -16
View File
@@ -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
}
+72 -3
View File
@@ -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")
}
}
+27
View File
@@ -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
+65
View File
@@ -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))
}
}
+48 -17
View File
@@ -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
}