Persist try snapshot and arm a systemd revert timer
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-03 20:52:42 +10:00
parent 7f9c010e1a
commit 6ac03e1012
9 changed files with 757 additions and 76 deletions
+1
View File
@@ -35,6 +35,7 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`,
applyCmd(),
tryCmd(),
confirmCmd(),
revertCmd(),
planCmd(),
validateCmd(),
statusCmd(),
+91 -39
View File
@@ -1,31 +1,32 @@
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"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
var tryPIDFile = "/run/tomswall/try.pid"
// 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, 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.`,
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 {
@@ -39,23 +40,16 @@ atomically. 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)
if err := os.MkdirAll(filepath.Dir(tryPIDFile), 0o755); err != nil {
applied, err := tryApply(cfg, timeout+revertGrace)
if err != nil || !applied {
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 {
msg, err := confirmOrRevert(confirm, abort, timeout, tryapply.Revert)
if err != nil {
return err
}
fmt.Println("Confirmed.")
fmt.Println(msg)
return nil
},
}
@@ -63,20 +57,67 @@ atomically. Confirm from a new session to prove new connections still work.`,
return cmd
}
func tryApply(cfg *config.Config, fallback time.Duration) (bool, error) {
unlock, err := tryapply.Acquire()
if err != nil {
return false, err
}
defer unlock()
engine, err := nftables.NewEngine(cfg)
if err != nil {
return false, fmt.Errorf("initializing nftables: %w", err)
}
changes, err := engine.Plan()
if err != nil {
return false, fmt.Errorf("computing changes: %w", err)
}
if changes.Empty() {
fmt.Println("No changes needed — firewall is up to date.")
return false, nil
}
fmt.Println(changes.Summary())
snap, err := engine.Snapshot()
if err != nil {
return false, fmt.Errorf("snapshotting ruleset: %w", err)
}
if err := tryapply.Arm(snap, os.Getpid(), fallback); err != nil {
return false, 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 true, nil
}
// 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 {
// 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 nil
return "Confirmed.", 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)
select {
case <-confirm:
return "Confirmed.", nil
default:
}
return fmt.Errorf("%s: previous ruleset restored", reason)
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 {
@@ -84,27 +125,38 @@ func confirmCmd() *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'")
}
pid, ok, err := tryapply.Confirm()
if err != nil {
return err
}
pid, err := strconv.Atoi(strings.TrimSpace(string(b)))
if err != nil {
return fmt.Errorf("parsing %s: %w", tryPIDFile, err)
if !ok {
return fmt.Errorf("no pending 'tomswall try': it was already reverted or confirmed")
}
// 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
// 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 {
return &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()
if err != nil {
return err
}
if !reverted {
fmt.Println("No pending 'tomswall try'.")
return nil
}
fmt.Println("Previous ruleset restored.")
return nil
},
}
}
+39 -27
View File
@@ -10,41 +10,53 @@ import (
func TestConfirmOrRevert(t *testing.T) {
tests := []struct {
name string
confirm bool
abort bool
revertErr error
wantRevert bool
wantErr bool
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) {
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)
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
}
}
})
}