// 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 ( "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 } ) // ErrPending means a try awaits confirmation; nothing else may apply meanwhile. var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' or 'tomswall revert'") type pending struct { 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. // The caller must hold the lock from Acquire. func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) error { b, err := json.Marshal(pending{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"); err != nil { return discardWith(fmt.Errorf("arming revert timer: %w", err)) } return 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. reverted is false when nothing was // pending (already confirmed or reverted). func Revert() (reverted bool, err error) { unlock, err := lock() if err != nil { return false, err } defer unlock() p, err := load() if err != nil || p == nil { return false, err } engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: p.Snapshot.Table}}) if err != nil { return false, err } if err := engine.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 }