Persist try snapshot and arm a systemd revert timer
This commit is contained in:
@@ -0,0 +1,185 @@
|
||||
// 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
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
package tryapply
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/nftables"
|
||||
)
|
||||
|
||||
func setup(t *testing.T) *[]string {
|
||||
t.Helper()
|
||||
Dir = t.TempDir()
|
||||
var cmds []string
|
||||
orig := run
|
||||
run = func(name string, args ...string) error {
|
||||
cmds = append(cmds, name+" "+strings.Join(args, " "))
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() { run = orig })
|
||||
return &cmds
|
||||
}
|
||||
|
||||
func arm(t *testing.T, snap *nftables.Snapshot) {
|
||||
t.Helper()
|
||||
unlock, err := Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer unlock()
|
||||
if err := Arm(snap, 4242, 90*time.Second); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArmPersistsSnapshotAndTimer(t *testing.T) {
|
||||
cmds := setup(t)
|
||||
snap := &nftables.Snapshot{Table: "tomswall", Present: true,
|
||||
Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1, 2, 3}}}}}}
|
||||
arm(t, snap)
|
||||
|
||||
info, err := os.Stat(snapshotPath())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Errorf("snapshot mode %v, want 0600", info.Mode().Perm())
|
||||
}
|
||||
p, err := load()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) {
|
||||
t.Errorf("round trip mismatch: %+v", p)
|
||||
}
|
||||
last := (*cmds)[len(*cmds)-1]
|
||||
if !strings.HasPrefix(last, "systemd-run ") || !strings.Contains(last, "--unit "+Unit) ||
|
||||
!strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert") {
|
||||
t.Errorf("unexpected arm command %q", last)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAbsentTableSnapshotRoundTrip(t *testing.T) {
|
||||
setup(t)
|
||||
arm(t, &nftables.Snapshot{Table: "tomswall"})
|
||||
p, err := load()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if p.Snapshot.Present || p.Snapshot.Table != "tomswall" {
|
||||
t.Errorf("absent table not preserved: %+v", p.Snapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAcquireRefusesWhilePending(t *testing.T) {
|
||||
setup(t)
|
||||
arm(t, &nftables.Snapshot{Table: "tomswall"})
|
||||
if _, err := Acquire(); !errors.Is(err, ErrPending) {
|
||||
t.Fatalf("second try: got %v, want ErrPending", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArmFailureDiscardsSnapshot(t *testing.T) {
|
||||
setup(t)
|
||||
run = func(name string, args ...string) error {
|
||||
if name == "systemd-run" {
|
||||
return errors.New("no systemd")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
unlock, err := Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer unlock()
|
||||
if err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil {
|
||||
t.Fatal("expected arm error")
|
||||
}
|
||||
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
|
||||
t.Error("snapshot left behind without a revert timer")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfirmPendingDisarms(t *testing.T) {
|
||||
cmds := setup(t)
|
||||
arm(t, &nftables.Snapshot{Table: "tomswall"})
|
||||
pid, ok, err := Confirm()
|
||||
if err != nil || !ok || pid != 4242 {
|
||||
t.Fatalf("Confirm = %d, %v, %v", pid, ok, err)
|
||||
}
|
||||
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
|
||||
t.Error("snapshot not removed")
|
||||
}
|
||||
if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" {
|
||||
t.Errorf("timer not stopped, last command %q", last)
|
||||
}
|
||||
unlock, err := Acquire()
|
||||
if err != nil {
|
||||
t.Fatalf("new try refused after confirm: %v", err)
|
||||
}
|
||||
unlock()
|
||||
}
|
||||
|
||||
func TestConfirmAfterRevertFails(t *testing.T) {
|
||||
setup(t)
|
||||
_, ok, err := Confirm()
|
||||
if err != nil || ok {
|
||||
t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err)
|
||||
}
|
||||
reverted, err := Revert()
|
||||
if err != nil || reverted {
|
||||
t.Fatalf("Revert with nothing pending = %v, %v", reverted, err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user