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.
201 lines
5.5 KiB
Go
201 lines
5.5 KiB
Go
// 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
|
|
}
|