0b92f1c2f3
apply, flush and purge take the try lock and fail with ErrPending, which now names 'tomswall revert' as the recovery after a failed automatic revert. Each try gets an ID passed to the timer's 'revert --id', so a stale timer cannot revert a newer try. Adds tests for Revert success/failure/stale ID and for restoring a present table at the engine level.
169 lines
4.8 KiB
Go
169 lines
4.8 KiB
Go
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
|
|
}
|