// Package tryapply keeps the state of a pending 'tomswall try' on disk so the // revert survives the try process, backed by a transient systemd timer. package tryapply import ( "crypto/rand" "encoding/hex" "encoding/json" "errors" "fmt" "os" "os/exec" "path/filepath" "strings" "syscall" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" ) // Unit is the transient systemd unit that reverts an unconfirmed try. const Unit = "tomswall-try-revert" var ( // Dir holds the lock and the pending snapshot. Dir = "/var/lib/tomswall" // run executes a systemd command; replaced in tests. run = func(name string, args ...string) error { if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) } return nil } // restore rolls the live table back to a snapshot; replaced in tests. restore = func(s *nftables.Snapshot) error { engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) if err != nil { return err } return engine.Restore(s) } ) // ErrPending means a try awaits confirmation; nothing else may apply meanwhile. var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' to keep it or 'tomswall revert' to restore the previous ruleset (also the recovery if an automatic revert failed)") type pending struct { ID string `json:"id"` PID int `json:"pid"` Snapshot *nftables.Snapshot `json:"snapshot"` } func snapshotPath() string { return filepath.Join(Dir, "try-snapshot.json") } // Acquire takes the exclusive try lock, failing with ErrPending while a try is unconfirmed. func Acquire() (unlock func(), err error) { unlock, err = lock() if err != nil { return nil, err } if _, err := os.Stat(snapshotPath()); err == nil { unlock() return nil, ErrPending } return unlock, nil } func lock() (func(), error) { if err := os.MkdirAll(Dir, 0o755); err != nil { return nil, err } f, err := os.OpenFile(filepath.Join(Dir, "try.lock"), os.O_CREATE|os.O_RDWR, 0o600) if err != nil { return nil, err } if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX); err != nil { f.Close() return nil, fmt.Errorf("locking %s: %w", f.Name(), err) } return func() { f.Close() }, nil } // Arm persists snap and schedules an out-of-process revert after delay, // returning the try ID that scopes later reverts to this try. // The caller must hold the lock from Acquire. func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) { raw := make([]byte, 8) if _, err := rand.Read(raw); err != nil { return "", err } id := hex.EncodeToString(raw) b, err := json.Marshal(pending{ID: id, PID: pid, Snapshot: snap}) if err != nil { return "", err } f, err := os.CreateTemp(Dir, ".try-snapshot-*") if err != nil { return "", err } defer os.Remove(f.Name()) if _, err := f.Write(b); err != nil { f.Close() return "", err } if err := f.Sync(); err != nil { f.Close() return "", err } if err := f.Close(); err != nil { return "", err } if err := os.Rename(f.Name(), snapshotPath()); err != nil { return "", err } exe, err := os.Executable() if err != nil { return "", discardWith(err) } _ = disarm() // a leftover timer from an earlier try would block the unit name if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) } return id, nil } // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. func Discard() error { _ = disarm() if err := os.Remove(snapshotPath()); err != nil && !os.IsNotExist(err) { return err } return nil } func discardWith(err error) error { if derr := Discard(); derr != nil { return fmt.Errorf("%w (discarding snapshot: %v)", err, derr) } return err } func disarm() error { return run("systemctl", "stop", Unit+".timer") } // Confirm keeps the tried ruleset. ok is false when no try was pending, i.e. // it was already reverted; pid is the waiting try process, if any. func Confirm() (pid int, ok bool, err error) { unlock, err := lock() if err != nil { return 0, false, err } defer unlock() p, err := load() if err != nil || p == nil { return 0, false, err } return p.PID, true, Discard() } // Revert restores the pending snapshot. A non-empty id only reverts that try, // so a stale timer cannot undo a newer one. reverted is false when nothing // matching was pending (already confirmed or reverted). A failed restore keeps // the snapshot so 'tomswall revert' can retry. func Revert(id string) (reverted bool, err error) { unlock, err := lock() if err != nil { return false, err } defer unlock() p, err := load() if err != nil || p == nil || (id != "" && p.ID != id) { return false, err } if err := restore(p.Snapshot); err != nil { return false, fmt.Errorf("restoring snapshot: %w", err) } return true, Discard() } func load() (*pending, error) { b, err := os.ReadFile(snapshotPath()) if os.IsNotExist(err) { return nil, nil } if err != nil { return nil, err } var p pending if err := json.Unmarshal(b, &p); err != nil { return nil, fmt.Errorf("parsing %s: %w", snapshotPath(), err) } if p.Snapshot == nil { return nil, fmt.Errorf("%s has no snapshot", snapshotPath()) } return &p, nil }