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)
}
})
}
}