Persist try snapshot and arm a systemd revert timer
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-03 20:52:42 +10:00
parent 7f9c010e1a
commit 6ac03e1012
9 changed files with 757 additions and 76 deletions
+53 -9
View File
@@ -28,7 +28,8 @@ func (e *Engine) ensureTable() *nftables.Table {
})
}
func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain {
// ensureChains declares the base chains; policies overrides their default policy.
func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain {
chains := map[string]*nftables.Chain{
"input": {
Name: "input",
@@ -71,6 +72,9 @@ func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain
}
for name, chain := range chains {
if p, ok := policies[name]; ok {
chain.Policy = policyPtr(p)
}
chains[name] = e.conn.AddChain(chain)
}
return chains
@@ -93,8 +97,12 @@ func (e *Engine) Plan() (*ChangeSet, error) {
}
func (e *Engine) Apply(changes *ChangeSet) error {
return e.apply(changes, nil)
}
func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error {
table := e.ensureTable()
chains := e.ensureChains(table)
chains := e.ensureChains(table, policies)
for _, r := range changes.Remove {
e.conn.DelRule(&nftables.Rule{
@@ -184,34 +192,70 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
return state, nil
}
// Snapshot is the tomswall table as captured live; a nil state means it was absent.
// Snapshot is the tomswall table as captured live, serialisable so a revert
// survives the process that took it.
type Snapshot struct {
state *FirewallState
Table string `json:"table"`
Present bool `json:"present"`
Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"`
Rules map[string][]SnapshotRule `json:"rules,omitempty"`
}
// SnapshotRule is a managed rule with its expressions in netlink wire format.
type SnapshotRule struct {
Tag string `json:"tag"`
Exprs [][]byte `json:"exprs"`
}
// Snapshot captures the live tomswall table so Restore can roll back to it.
func (e *Engine) Snapshot() (*Snapshot, error) {
snap := &Snapshot{Table: e.cfg.Settings.TableName}
t, err := e.findTable()
if err != nil || t == nil {
return &Snapshot{}, err
return snap, err
}
snap.Present = true
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
if err != nil {
return nil, fmt.Errorf("listing chains: %w", err)
}
snap.Policies = make(map[string]nftables.ChainPolicy)
for _, c := range chains {
if c.Table.Name == snap.Table && c.Policy != nil {
snap.Policies[c.Name] = *c.Policy
}
}
state, err := e.readCurrentState()
if err != nil {
return nil, err
}
return &Snapshot{state: state}, nil
snap.Rules, err = encodeState(state)
if err != nil {
return nil, err
}
return snap, nil
}
// Restore atomically returns the tomswall table to the snapshot, rule order included.
// Restore atomically returns the tomswall table to the snapshot: rule order
// and chain policies included, or removed if it was absent.
func (e *Engine) Restore(s *Snapshot) error {
if s.state == nil {
if s.Table != e.cfg.Settings.TableName {
return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName)
}
if !s.Present {
return e.Flush()
}
want, err := decodeState(s.Rules)
if err != nil {
return err
}
current, err := e.readCurrentState()
if err != nil {
return err
}
return e.Apply(restoreChangeSet(current, s.state))
return e.apply(restoreChangeSet(current, want), s.Policies)
}
func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy {
+119
View File
@@ -0,0 +1,119 @@
package nftables
import (
"encoding/binary"
"fmt"
"github.com/google/nftables"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
const inet = byte(nftables.TableFamilyINet)
// exprByName mirrors the expression types google/nftables can parse back from the kernel.
var exprByName = map[string]func() expr.Any{
"ct": func() expr.Any { return &expr.Ct{} },
"range": func() expr.Any { return &expr.Range{} },
"meta": func() expr.Any { return &expr.Meta{} },
"cmp": func() expr.Any { return &expr.Cmp{} },
"counter": func() expr.Any { return &expr.Counter{} },
"objref": func() expr.Any { return &expr.Objref{} },
"payload": func() expr.Any { return &expr.Payload{} },
"lookup": func() expr.Any { return &expr.Lookup{} },
"immediate": func() expr.Any { return &expr.Immediate{} },
"bitwise": func() expr.Any { return &expr.Bitwise{} },
"redir": func() expr.Any { return &expr.Redir{} },
"nat": func() expr.Any { return &expr.NAT{} },
"limit": func() expr.Any { return &expr.Limit{} },
"quota": func() expr.Any { return &expr.Quota{} },
"dynset": func() expr.Any { return &expr.Dynset{} },
"log": func() expr.Any { return &expr.Log{} },
"exthdr": func() expr.Any { return &expr.Exthdr{} },
"connlimit": func() expr.Any { return &expr.Connlimit{} },
"queue": func() expr.Any { return &expr.Queue{} },
"flow_offload": func() expr.Any { return &expr.FlowOffload{} },
"reject": func() expr.Any { return &expr.Reject{} },
"masq": func() expr.Any { return &expr.Masq{} },
"hash": func() expr.Any { return &expr.Hash{} },
"notrack": func() expr.Any { return &expr.Notrack{} },
}
func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) {
out := make(map[string][]SnapshotRule, len(state.Rules))
for chain, rules := range state.Rules {
for _, r := range rules {
sr := SnapshotRule{Tag: r.Tag}
for _, e := range r.Exprs {
b, err := expr.Marshal(inet, e)
if err != nil {
return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err)
}
sr.Exprs = append(sr.Exprs, b)
}
out[chain] = append(out[chain], sr)
}
}
return out, nil
}
func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) {
state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))}
for chain, rs := range rules {
for _, sr := range rs {
r := ManagedRule{Chain: chain, Tag: sr.Tag}
for _, b := range sr.Exprs {
e, err := decodeExpr(b)
if err != nil {
return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err)
}
r.Exprs = append(r.Exprs, e)
}
state.Rules[chain] = append(state.Rules[chain], r)
}
}
return state, nil
}
// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules.
func decodeExpr(b []byte) (expr.Any, error) {
ad, err := netlink.NewAttributeDecoder(b)
if err != nil {
return nil, err
}
ad.ByteOrder = binary.BigEndian
var name string
var data []byte
for ad.Next() {
switch ad.Type() {
case unix.NFTA_EXPR_NAME:
name = ad.String()
case unix.NFTA_EXPR_DATA:
data = ad.Bytes()
}
}
if err := ad.Err(); err != nil {
return nil, err
}
newExpr, ok := exprByName[name]
if !ok {
return nil, fmt.Errorf("unsupported expression %q", name)
}
e := newExpr()
if name == "notrack" {
return e, nil
}
if err := expr.Unmarshal(inet, data, e); err != nil {
return nil, err
}
// A verdict is an immediate into the verdict register with no data.
if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 {
v := &expr.Verdict{}
if err := expr.Unmarshal(inet, data, v); err != nil {
return nil, err
}
return v, nil
}
return e, nil
}
+131
View File
@@ -0,0 +1,131 @@
package nftables
import (
"encoding/json"
"reflect"
"testing"
"github.com/google/nftables"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
func TestSnapshotRulesRoundTrip(t *testing.T) {
exprs := []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0, 22}},
&expr.Ct{Register: 1, Key: expr.CtKeySTATE},
&expr.Notrack{},
&expr.Verdict{Kind: expr.VerdictAccept},
}
state := &FirewallState{Rules: map[string][]ManagedRule{
"input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}},
}}
rules, err := encodeState(state)
if err != nil {
t.Fatal(err)
}
b, err := json.Marshal(&Snapshot{Table: "tomswall", Present: true, Rules: rules,
Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}})
if err != nil {
t.Fatal(err)
}
var snap Snapshot
if err := json.Unmarshal(b, &snap); err != nil {
t.Fatal(err)
}
if snap.Policies["input"] != nftables.ChainPolicyAccept {
t.Errorf("policy lost: %v", snap.Policies)
}
got, err := decodeState(snap.Rules)
if err != nil {
t.Fatal(err)
}
in := got.Rules["input"]
if len(in) != 2 || in[0].Tag != "ssh" || in[1].Tag != "drop" {
t.Fatalf("rules/order lost: %+v", in)
}
if !reflect.DeepEqual(in[0].Exprs, exprs) {
t.Errorf("exprs changed:\n got %#v\nwant %#v", in[0].Exprs, exprs)
}
if _, ok := in[1].Exprs[0].(*expr.Verdict); !ok {
t.Errorf("verdict decoded as %T", in[1].Exprs[0])
}
}
func TestEnsureChainsPolicyOverride(t *testing.T) {
e := testEngine(t, nil)
chains := e.ensureChains(e.ensureTable(), map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept})
if *chains["input"].Policy != nftables.ChainPolicyAccept {
t.Error("input policy not overridden")
}
if *chains["forward"].Policy != nftables.ChainPolicyDrop {
t.Error("forward policy should keep its default")
}
}
func TestSnapshotAndRestoreAbsentTable(t *testing.T) {
tablePresent := false
var sent []netlink.HeaderType
e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) {
for _, m := range req {
sent = append(sent, m.Header.Type)
if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent {
data := []byte{inet, 0, 0, 0}
attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}})
return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil
}
}
return nil, nil
})
snap, err := e.Snapshot()
if err != nil {
t.Fatal(err)
}
if snap.Present || snap.Table != "tomswall" {
t.Fatalf("want absent tomswall snapshot, got %+v", snap)
}
// The try created the table; restoring the absent snapshot deletes it.
tablePresent = true
sent = nil
if err := e.Restore(snap); err != nil {
t.Fatal(err)
}
deleted := false
for _, ht := range sent {
deleted = deleted || ht == nftType(unix.NFT_MSG_DELTABLE)
}
if !deleted {
t.Errorf("table not deleted; sent %v", sent)
}
}
func TestRestoreRejectsOtherTable(t *testing.T) {
e := testEngine(t, nil)
if err := e.Restore(&Snapshot{Table: "other"}); err == nil {
t.Error("expected table mismatch error")
}
}
func nftType(msg int) netlink.HeaderType {
return netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | msg)
}
func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, error)) *Engine {
if dial == nil {
dial = func([]netlink.Message) ([]netlink.Message, error) { return nil, nil }
}
conn, err := nftables.New(nftables.WithTestDial(dial))
if err != nil {
t.Fatal(err)
}
return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn}
}
+185
View File
@@ -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
}
+137
View File
@@ -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)
}
}