From fe689e99edf7e9161b8d8bf5a625efce3b3ab0ff Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sun, 4 Oct 2026 15:34:08 +1100 Subject: [PATCH] Raise netlink buffers for large batches; restore snapshot when try apply fails --- cmd/tomswall/try.go | 6 +-- internal/nftables/engine.go | 34 ++++++++++++++++- internal/nftables/engine_test.go | 59 ++++++++++++++++++++++++++++++ internal/tryapply/tryapply.go | 24 ++++++++++-- internal/tryapply/tryapply_test.go | 38 +++++++++++++++++++ 5 files changed, 154 insertions(+), 7 deletions(-) create mode 100644 internal/nftables/engine_test.go diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go index ee9e21e..7c3e6a7 100644 --- a/cmd/tomswall/try.go +++ b/cmd/tomswall/try.go @@ -89,10 +89,10 @@ func tryApply(cfg *config.Config, fallback time.Duration) (string, error) { return "", err } if err := engine.Apply(changes); err != nil { - if derr := tryapply.Discard(); derr != nil { - err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) + if aerr := tryapply.Abort(); aerr != nil { + return "", fmt.Errorf("applying changes: %w; %v; the revert timer restores the previous ruleset within %s", err, aerr, fallback) } - return "", fmt.Errorf("applying changes: %w", err) + return "", fmt.Errorf("applying changes: %w: previous ruleset restored", err) } return id, nil } diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index c6781c6..ed4ccb5 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -5,6 +5,8 @@ import ( "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" ) @@ -15,13 +17,43 @@ type Engine struct { } func NewEngine(cfg *config.Config) (*Engine, error) { - conn, err := nftables.New() + conn, err := nftables.New(nftables.WithSockOptions(largeBuffers)) if err != nil { return nil, fmt.Errorf("connecting to nftables: %w", err) } return &Engine{cfg: cfg, conn: conn}, nil } +// batchBufSize bounds one batch: the kernel rejects a batch larger than the +// send buffer (EMSGSIZE) and drops ACKs beyond the receive buffer (ENOBUFS) +// after committing it. +// ponytail: fixed cap of tens of thousands of rules; size per batch if exceeded. +const batchBufSize = 64 << 20 + +// largeBuffers raises both socket buffers, ignoring rmem_max/wmem_max when +// CAP_NET_ADMIN allows it and falling back to the capped sizes otherwise. +func largeBuffers(c *netlink.Conn) error { + rc, err := c.SyscallConn() + if err != nil { + return err + } + var serr error + err = rc.Control(func(fd uintptr) { + for _, o := range [][2]int{{unix.SO_SNDBUFFORCE, unix.SO_SNDBUF}, {unix.SO_RCVBUFFORCE, unix.SO_RCVBUF}} { + if unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[0], batchBufSize) == nil { + continue + } + if serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[1], batchBufSize); serr != nil { + return + } + } + }) + if err != nil { + return err + } + return serr +} + func (e *Engine) ensureTable() *nftables.Table { return e.conn.AddTable(&nftables.Table{ Family: nftables.TableFamilyINet, diff --git a/internal/nftables/engine_test.go b/internal/nftables/engine_test.go new file mode 100644 index 0000000..a19c2c7 --- /dev/null +++ b/internal/nftables/engine_test.go @@ -0,0 +1,59 @@ +package nftables + +import ( + "os" + "strconv" + "strings" + "testing" + + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +func TestLargeBuffersRaisesSocketBuffers(t *testing.T) { + c, err := netlink.Dial(unix.NETLINK_NETFILTER, nil) + if err != nil { + t.Skipf("netlink unavailable: %v", err) + } + defer c.Close() + if err := largeBuffers(c); err != nil { + t.Fatal(err) + } + // Without CAP_NET_ADMIN the kernel caps at the sysctl max; it doubles either way. + for opt, sysctl := range map[int]string{unix.SO_RCVBUF: "rmem_max", unix.SO_SNDBUF: "wmem_max"} { + want := 2 * min(batchBufSize, procInt(t, "/proc/sys/net/core/"+sysctl)) + if got := sockBuf(t, c, opt); got < want { + t.Errorf("%s-bounded buffer = %d, want >= %d", sysctl, got, want) + } + } +} + +func procInt(t *testing.T, path string) int { + t.Helper() + b, err := os.ReadFile(path) + if err != nil { + t.Skipf("reading %s: %v", path, err) + } + v, err := strconv.Atoi(strings.TrimSpace(string(b))) + if err != nil { + t.Fatal(err) + } + return v +} + +func sockBuf(t *testing.T, c *netlink.Conn, opt int) int { + t.Helper() + rc, err := c.SyscallConn() + if err != nil { + t.Fatal(err) + } + var v int + var serr error + if err := rc.Control(func(fd uintptr) { v, serr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt) }); err != nil { + t.Fatal(err) + } + if serr != nil { + t.Fatal(serr) + } + return v +} diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go index c0e8fb4..a68e024 100644 --- a/internal/tryapply/tryapply.go +++ b/internal/tryapply/tryapply.go @@ -175,10 +175,28 @@ func Revert(id string) (reverted bool, err error) { 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, restorePending(p) +} + +// Abort restores the pending snapshot after a failed apply, which may have +// committed partially. A failed restore keeps the snapshot and timer so the +// timer still reverts. The caller must hold the lock. +func Abort() error { + p, err := load() + if err != nil { + return err } - return true, Discard() + if p == nil { + return errors.New("no pending try to abort") + } + return restorePending(p) +} + +func restorePending(p *pending) error { + if err := restore(p.Snapshot); err != nil { + return fmt.Errorf("restoring snapshot: %w", err) + } + return Discard() } func load() (*pending, error) { diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go index 5b00f2c..e82dc83 100644 --- a/internal/tryapply/tryapply_test.go +++ b/internal/tryapply/tryapply_test.go @@ -209,3 +209,41 @@ func TestRevertStaleIDIgnored(t *testing.T) { t.Errorf("newer try's snapshot removed: %v", err) } } + +func TestAbortRestoresAndDisarms(t *testing.T) { + cmds := setup(t) + restored := stubRestore(t, nil) + snap := &nftables.Snapshot{Table: "tomswall", Present: true} + arm(t, snap) + + if err := Abort(); err != nil { + t.Fatal(err) + } + if len(*restored) != 1 || !reflect.DeepEqual((*restored)[0], snap) { + t.Errorf("restored %+v, want the armed snapshot", *restored) + } + 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) + } +} + +func TestAbortFailureKeepsSnapshotAndTimer(t *testing.T) { + cmds := setup(t) + boom := errors.New("netlink down") + stubRestore(t, boom) + arm(t, &nftables.Snapshot{Table: "tomswall"}) + armed := len(*cmds) + + if err := Abort(); !errors.Is(err, boom) { + t.Fatalf("Abort error = %v, want %v", err, boom) + } + if _, err := os.Stat(snapshotPath()); err != nil { + t.Fatalf("snapshot gone after failed abort: %v", err) + } + if len(*cmds) != armed { + t.Errorf("revert timer touched after failed abort: %v", (*cmds)[armed:]) + } +}