diff --git a/Makefile b/Makefile index 24c71459..d84bc6fa 100644 --- a/Makefile +++ b/Makefile @@ -161,6 +161,10 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11 bin-pkcs11: CGO_ENABLED = 1 bin-pkcs11: bin +# Build with the pprof debug server (serves on :6060). See startPprofServer. +debug: BUILD_ARGS += -tags debug +debug: bin + bin: go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH} go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert @@ -280,5 +284,5 @@ smoke-vagrant/%: bin-docker build/%/nebula cd .github/workflows/smoke/ && ./smoke-vagrant.sh $* .FORCE: -.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/% +.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/% .DEFAULT_GOAL := bin diff --git a/connection_state.go b/connection_state.go index 7d1f0091..3018048b 100644 --- a/connection_state.go +++ b/connection_state.go @@ -14,7 +14,7 @@ import ( ) const ( - ReplayWindow = 1024 + ReplayWindow = 8192 // RehandshakeAfterMessages rolls keys inside the AES-GCM data-volume margin (~2^-36 advantage at 64KB frames). RehandshakeAfterMessages = uint64(1) << 34 @@ -26,6 +26,14 @@ const ( // RehandshakeAfterMessages must stay below RejectAfterMessages so tunnels roll before the hard send stop. const _ = RejectAfterMessages - RehandshakeAfterMessages +// sessionEpoch hands out a receiver-local ordinal to every ConnectionState at creation. The RX +// staging sort (overlay/batch) orders packets by (epoch, message counter). A re-handshake never +// rekeys an existing tunnel; it brings up a new hostinfo and ConnectionState with a counter space +// starting near zero, while the old tunnel keeps decrypting until torn down. During that cutover +// one flush batch can hold packets from both tunnels, and the epoch keeps the old tunnel's +// packets sorted first. +var sessionEpoch atomic.Uint64 + type ConnectionState struct { eKey noiseutil.CipherState dKey noiseutil.CipherState @@ -36,6 +44,8 @@ type ConnectionState struct { window *Bits decryptLock sync.Mutex writeLock sync.Mutex + // epoch is this session's sessionEpoch ordinal. Immutable after creation. + epoch uint64 } // newConnectionStateFromResult builds a fully-populated ConnectionState from a @@ -55,6 +65,7 @@ func newConnectionStateFromResult(r *handshake.Result) (*ConnectionState, error) eKey: noiseutil.NewCipherState(r.EKey, r.Cipher), dKey: noiseutil.NewCipherState(r.DKey, r.Cipher), window: NewBits(ReplayWindow), + epoch: sessionEpoch.Add(1), } ci.messageCounter.Add(r.MessageIndex) for i := uint64(1); i <= r.MessageIndex; i++ { @@ -85,8 +96,7 @@ func (cs *ConnectionState) Curve() cert.Curve { return cs.myCert.Curve() } -func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) { - var err error +func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) { cs.decryptLock.Lock() result := cs.window.Check(l, messageCounter) cs.decryptLock.Unlock() @@ -94,7 +104,7 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out [] return nil, ErrAlreadySeen } - out, err = cs.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb) + out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb) if err != nil { return nil, err } @@ -108,7 +118,6 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out [] return out, nil } -// VerifyRelay verifies AEAD protected (but not encrypted) relay frames. packet must be length-checked by the caller. func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error { cs.decryptLock.Lock() result := cs.window.Check(l, messageCounter) @@ -117,6 +126,11 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa return ErrAlreadySeen } + // The entire body is sent as AD, not encrypted. + // The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value. + // The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's + // otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice + // which will gracefully fail in the DecryptDanger call. signedPayload := packet[:len(packet)-cs.dKey.Overhead()] signatureValue := packet[len(packet)-cs.dKey.Overhead():] _, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb) @@ -130,6 +144,5 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa if !result { return ErrAlreadySeen } - return nil } diff --git a/control.go b/control.go index 7df5a09e..6d50005c 100644 --- a/control.go +++ b/control.go @@ -115,7 +115,7 @@ func (c *Control) Start() error { c.lighthouseStart() } - c.f.triggerShutdown = c.Stop + c.f.triggerShutdown = func() { go c.Stop() } // Start reading packets. c.f.run() diff --git a/control_lifecycle_test.go b/control_lifecycle_test.go index 0b5d106d..a9f5323f 100644 --- a/control_lifecycle_test.go +++ b/control_lifecycle_test.go @@ -11,6 +11,8 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/batch" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/test" "github.com/slackhq/nebula/udp" @@ -30,9 +32,9 @@ func newFakeDevice() *fakeDevice { // Read blocks until Close like a real tun with no traffic, then reports EOF // the same way a closed device does -func (d *fakeDevice) Read(p []byte) (int, error) { +func (d *fakeDevice) Read() ([]tio.Packet, error) { <-d.closedCh - return 0, io.EOF + return nil, io.EOF } func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil } @@ -49,10 +51,8 @@ func (d *fakeDevice) Activate() error { return nil } func (d *fakeDevice) Networks() []netip.Prefix { return nil } func (d *fakeDevice) Name() string { return "fake" } func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil } -func (d *fakeDevice) SupportsMultiqueue() bool { return false } -func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, errors.New("unsupported") -} + +func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil } // newReadyControl hand-builds the minimum Control that Main would have // produced right before Start, including the construction token NewInterface @@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) { inside: dev, outside: conn, writers: []udp.Conn{conn}, - readers: make([]io.ReadWriteCloser, 1), + batchers: make([]*batch.MultiCoalescer, 1), routines: 1, hostMap: newHostMap(l), lightHouse: lh, @@ -109,7 +109,8 @@ func TestControl_StopBeforeStart(t *testing.T) { require.NoError(t, c.Wait()) // A stopped control can never be started - require.ErrorIs(t, c.Start(), ErrAlreadyStopped) + err := c.Start() + require.ErrorIs(t, err, ErrAlreadyStopped) // A second Stop is a harmless no-op c.Stop() @@ -143,19 +144,29 @@ type fakeConn struct { rebinds int } -func (c *fakeConn) Rebind() error { c.rebinds++; return nil } -func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil } -func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil } -func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil } -func (c *fakeConn) ReloadConfig(_ *config.C) {} -func (c *fakeConn) SupportsMultipleReaders() bool { return true } -func (c *fakeConn) Close() error { c.closed = true; return nil } +func (c *fakeConn) Rebind() error { c.rebinds++; return nil } +func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil } +func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil } +func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil } +func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) { + return len(bufs), nil +} +func (c *fakeConn) ReloadConfig(_ *config.C) {} +func (c *fakeConn) SupportsMultipleReaders() bool { return true } +func (c *fakeConn) Close() error { c.closed = true; return nil } type multiqueueDevice struct { *fakeDevice } -func (d *multiqueueDevice) SupportsMultiqueue() bool { return true } +// Queues claims multiqueue support but fails to open the second queue, +// exercising the activation error path. +func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) { + if n > 1 { + return nil, errors.New("second queue failed to open") + } + return d.fakeDevice.Queues(n) +} func TestControl_StartMultiqueueFailureReleases(t *testing.T) { dev := &multiqueueDevice{fakeDevice: newFakeDevice()} @@ -166,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) { inside: dev, outside: conn, writers: []udp.Conn{conn}, - readers: make([]io.ReadWriteCloser, 2), + batchers: make([]*batch.MultiCoalescer, 2), routines: 2, l: test.NewLogger(), } @@ -181,7 +192,8 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) { } // The second reader fails to open, everything must be released - require.Error(t, c.Start()) + err := c.Start() + require.Error(t, err) assert.Equal(t, StateStopped, c.State()) assert.True(t, dev.closed, "the tun device should have been closed") assert.True(t, conn.closed, "the udp socket should have been closed") @@ -251,15 +263,18 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) { // panic and Wait must observe the final state require.NoError(t, c.Wait()) assert.Equal(t, StateStopped, c.State()) - require.ErrorIs(t, c.Start(), ErrAlreadyStopped) + err := c.Start() + require.ErrorIs(t, err, ErrAlreadyStopped) } func TestControl_StartStopLifecycle(t *testing.T) { c, dev, conn := newReadyControl(t) - require.NoError(t, c.Start()) + err := c.Start() + require.NoError(t, err) assert.Equal(t, StateStarted, c.State()) - require.ErrorIs(t, c.Start(), ErrAlreadyStarted) + err = c.Start() + require.ErrorIs(t, err, ErrAlreadyStarted) // Stop must unpark the reader blocked in the device and release everything c.Stop() @@ -270,7 +285,8 @@ func TestControl_StartStopLifecycle(t *testing.T) { // The reader drained off a closed device, that is not a fatal error require.NoError(t, c.Wait()) - require.ErrorIs(t, c.Start(), ErrAlreadyStopped) + err = c.Start() + require.ErrorIs(t, err, ErrAlreadyStopped) } func TestControl_RebindIsGatedByState(t *testing.T) { @@ -280,7 +296,8 @@ func TestControl_RebindIsGatedByState(t *testing.T) { c.RebindUDPServer() assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op") - require.NoError(t, c.Start()) + err := c.Start() + require.NoError(t, err) c.RebindUDPServer() assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn") diff --git a/cpupick/cpupick.go b/cpupick/cpupick.go new file mode 100644 index 00000000..49f791e8 --- /dev/null +++ b/cpupick/cpupick.go @@ -0,0 +1,187 @@ +// Package cpupick chooses which CPUs the tun reader threads pin to when the +// operator has not chosen for us (tun.cpu_affinity). The stock spread — +// allowed[i] for routine i — has two failure modes this package exists to fix: +// +// - every co-located nebula starts its spread at allowed[0], so N instances +// on one box stack their readers onto the same cores, and allowed[0] is +// usually CPU 0, the core housekeeping and default IRQ affinity already +// favor; +// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact +// cores) low IDs are not necessarily fast cores, and pinning an encrypt +// thread to an efficiency core caps that queue's throughput. +// +// Default instead returns a preference-ordered pin list: the allowed set +// filtered to performance cores (when the platform distinguishes them and +// enough remain for every routine), confined to a single NUMA node and spread +// across distinct physical cores when the topology permits, CPU 0's physical +// core demoted to last resort, and the order rotated by a stable per-instance +// key so co-located instances spread instead of stacking. +package cpupick + +import ( + "log/slog" + + "github.com/slackhq/nebula/util" +) + +// topology is the slice of machine layout arrange consults: the NUMA node +// and the physical core behind each candidate CPU, plus which core CPU 0 +// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's +// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate). +// Probed from sysfs on Linux; flatTopology stands in when the platform can't +// say, which turns every topology rule into a no-op rather than a wrong +// answer. +type topology struct { + nodeOf map[int]int + coreOf map[int]int + zeroCore int +} + +// flatTopology places every CPU on node 0 and on a physical core of its own. +func flatTopology(cpus []int) topology { + t := topology{ + nodeOf: make(map[int]int, len(cpus)), + coreOf: make(map[int]int, len(cpus)), + zeroCore: -1, + } + for i, c := range cpus { + t.nodeOf[c] = 0 + t.coreOf[c] = i + if c == 0 { + t.zeroCore = i + } + } + return t +} + +// Default computes the pin order for `routines` tun readers. key is any +// stable per-instance value; the bound UDP port is ideal — distinct across +// co-located instances, stable across restarts so benchmark runs stay +// comparable. Returns nil when there is nothing useful to say (no affinity +// support on this platform, lookup failure); callers keep their existing +// fallback spread. +func Default(routines int, key uint64, l *slog.Logger) []int { + allowed, err := util.AllowedCPUs() + if err != nil || len(allowed) == 0 { + return nil + } + perf, signal := perfCPUs(allowed) + cands := pickCandidates(allowed, perf, routines) + if len(cands) == 0 { + return nil + } + if len(perf) < routines { + signal = "" + } + cpus := arrange(cands, readTopology(cands), routines, splitmix64(key)) + if l != nil { + l.Info("chose default pin CPUs for tun readers", + "cpus", cpus[:min(routines, len(cpus))], + "perfSignal", signal) + } + return cpus +} + +// pickCandidates applies the enough-for-everyone guard: a perf filter that +// leaves fewer candidates than routines is discarded — giving every reader +// its own (possibly slow) core beats stacking two readers on a fast one. +func pickCandidates(allowed, perf []int, routines int) []int { + if len(perf) < routines { + return allowed + } + return perf +} + +// arrange turns the candidate set into the final pin order: +// +// 1. NUMA: when at least one node holds enough candidates for every +// routine, confine to one such node, chosen by the instance hash. The +// readers share hostmap and cipher state, so splitting one instance +// across nodes taxes every packet — and co-located instances that hash +// to different nodes stop competing entirely. When no node is big +// enough, span nodes rather than stack readers. +// 2. Rotate the preferred candidates by the hash so instances spread. +// 3. SMT: emit one thread per physical core before any of their siblings — +// two encrypt threads on one core split its execution units. Siblings +// still follow for the routines > cores case. +// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ +// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the +// sibling precedes CPU 0 itself, which only catches the bleed-through. +// +// The rotation happens before the SMT pass so each instance's one-per-core +// walk also starts at a different core, and CPU 0's core is excluded from +// the rotation so no hash value can put it back at the front. +func arrange(cands []int, topo topology, routines int, h uint64) []int { + byNode := map[int][]int{} + var nodes []int + for _, c := range cands { + n := topo.nodeOf[c] + if _, ok := byNode[n]; !ok { + nodes = append(nodes, n) + } + byNode[n] = append(byNode[n], c) + } + var eligible []int + for _, n := range nodes { + if len(byNode[n]) >= routines { + eligible = append(eligible, n) + } + } + if len(eligible) > 0 { + cands = byNode[eligible[int(h%uint64(len(eligible)))]] + } + + // Split off CPU 0's core: its siblings tail the list, CPU 0 tails them. + preferred := make([]int, 0, len(cands)) + var zeroTail []int + hasZero := false + for _, c := range cands { + switch { + case c == 0: + hasZero = true + case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore: + zeroTail = append(zeroTail, c) + default: + preferred = append(preferred, c) + } + } + if hasZero { + zeroTail = append(zeroTail, 0) + } + if len(preferred) == 0 { + return zeroTail // CPU 0's core is all we have + } + + // The node pick consumed the low hash bits; rotate by the high ones so + // the two choices stay independent. + off := int((h >> 32) % uint64(len(preferred))) + rot := make([]int, 0, len(preferred)) + rot = append(rot, preferred[off:]...) + rot = append(rot, preferred[:off]...) + + seenCore := make(map[int]bool, len(rot)) + out := make([]int, 0, len(cands)) + var siblings []int + for _, c := range rot { + g := topo.coreOf[c] + if seenCore[g] { + siblings = append(siblings, c) + continue + } + seenCore[g] = true + out = append(out, c) + } + out = append(out, siblings...) + out = append(out, zeroTail...) + return out +} + +// splitmix64 decorrelates instance keys before the selection modulos: ports +// on one box often share spacing (4242/4243, or round steps like +1000) that +// raw key%len arithmetic would fold onto the same offset. +func splitmix64(x uint64) uint64 { + x += 0x9e3779b97f4a7c15 + x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9 + x = (x ^ (x >> 27)) * 0x94d049bb133111eb + return x ^ (x >> 31) +} diff --git a/cpupick/cpupick_test.go b/cpupick/cpupick_test.go new file mode 100644 index 00000000..bd81c1c9 --- /dev/null +++ b/cpupick/cpupick_test.go @@ -0,0 +1,171 @@ +package cpupick + +import ( + "slices" + "testing" +) + +// pairTopo builds a topology where consecutive candidate pairs are SMT +// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ... +// All CPUs land on node 0. +func pairTopo(cpus []int) topology { + t := topology{ + nodeOf: make(map[int]int, len(cpus)), + coreOf: make(map[int]int, len(cpus)), + zeroCore: -1, + } + for i, c := range cpus { + t.nodeOf[c] = 0 + t.coreOf[c] = i / 2 + if c == 0 { + t.zeroCore = i / 2 + } + } + return t +} + +func TestArrangeDemotesZeroForEveryKey(t *testing.T) { + candidates := []int{0, 1, 2, 3, 4, 5, 6, 7} + for key := range uint64(64) { + got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key)) + if len(got) != len(candidates) { + t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates)) + } + if got[0] == 0 { + t.Errorf("key %d: CPU 0 at the front: %v", key, got) + } + if got[len(got)-1] != 0 { + t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got) + } + sorted := slices.Clone(got) + slices.Sort(sorted) + if !slices.Equal(sorted, candidates) { + t.Errorf("key %d: not a permutation: %v", key, got) + } + } +} + +func TestArrangeDemotesZeroSiblings(t *testing.T) { + // Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 — + // must tail the list, sibling ahead of 0 itself. + candidates := []int{0, 1, 2, 3, 4, 5, 6, 7} + for key := range uint64(64) { + got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key)) + n := len(got) + if got[n-1] != 0 || got[n-2] != 1 { + t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got) + } + } +} + +func TestArrangeZeroSiblingWithoutZero(t *testing.T) { + // CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still + // tails the list when the topology knows which core CPU 0 lives on. + candidates := []int{1, 2, 3, 4, 5} + topo := pairTopo([]int{0, 1, 2, 3, 4, 5}) + got := arrange(candidates, topo, 2, splitmix64(7)) + if got[len(got)-1] != 1 { + t.Errorf("CPU 0's sibling not demoted: %v", got) + } +} + +func TestArrangeRotatesByKey(t *testing.T) { + candidates := []int{1, 2, 3, 4, 5, 6, 7, 8} + seen := map[int]bool{} + for key := range uint64(64) { + seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true + } + // 64 hashed keys over 8 slots must hit more than one starting CPU, or + // co-located instances would all stack again. + if len(seen) < 2 { + t.Errorf("rotation never varied across keys: %v", seen) + } +} + +func TestArrangeStableForSameKey(t *testing.T) { + candidates := []int{0, 2, 4, 6} + topo := flatTopology(candidates) + a := arrange(candidates, topo, 2, splitmix64(4242)) + b := arrange(candidates, topo, 2, splitmix64(4242)) + if !slices.Equal(a, b) { + t.Errorf("same key ordered differently: %v vs %v", a, b) + } +} + +func TestArrangeZeroOnly(t *testing.T) { + if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) { + t.Errorf("sole CPU 0 must survive: %v", got) + } +} + +func TestArrangeSMTSiblingsLast(t *testing.T) { + // Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four + // distinct physical cores before any sibling repeats. + candidates := []int{1, 2, 3, 4, 5, 6, 7, 8} + topo := pairTopo(candidates) + for key := range uint64(16) { + got := arrange(candidates, topo, 4, splitmix64(key)) + seen := map[int]bool{} + for _, c := range got[:4] { + g := topo.coreOf[c] + if seen[g] { + t.Fatalf("key %d: sibling before all cores covered: %v", key, got) + } + seen[g] = true + } + } +} + +func TestArrangeNUMAConfinesToOneNode(t *testing.T) { + // Two nodes of four; both fit routines=3, so the result must sit + // entirely inside one of them, and the hash must pick both across keys. + candidates := []int{1, 2, 3, 4, 10, 11, 12, 13} + topo := flatTopology(candidates) + for _, c := range []int{10, 11, 12, 13} { + topo.nodeOf[c] = 1 + } + nodesSeen := map[int]bool{} + for key := range uint64(32) { + got := arrange(candidates, topo, 3, splitmix64(key)) + if len(got) != 4 { + t.Fatalf("key %d: not confined to one node: %v", key, got) + } + n := topo.nodeOf[got[0]] + for _, c := range got { + if topo.nodeOf[c] != n { + t.Fatalf("key %d: spans nodes: %v", key, got) + } + } + nodesSeen[n] = true + } + if len(nodesSeen) != 2 { + t.Errorf("hash never spread instances across nodes: %v", nodesSeen) + } +} + +func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) { + candidates := []int{1, 2, 3, 4, 10, 11, 12, 13} + topo := flatTopology(candidates) + for _, c := range []int{10, 11, 12, 13} { + topo.nodeOf[c] = 1 + } + got := arrange(candidates, topo, 6, splitmix64(1)) + if len(got) != len(candidates) { + t.Errorf("undersized nodes must span, got %v", got) + } +} + +func TestPickCandidates(t *testing.T) { + allowed := []int{0, 1, 2, 3, 4, 5, 6, 7} + perf := []int{4, 5} + + // Enough perf cores for every routine: only they are used. + if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) { + t.Errorf("perf filter not applied: %v", got) + } + // Perf filter too small for the routine count: discarded, everyone + // gets their own core from the full allowed set. + if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) { + t.Errorf("undersized perf filter not discarded: %v", got) + } +} diff --git a/cpupick/perf_linux.go b/cpupick/perf_linux.go new file mode 100644 index 00000000..f5bf76b5 --- /dev/null +++ b/cpupick/perf_linux.go @@ -0,0 +1,154 @@ +//go:build linux + +package cpupick + +import ( + "fmt" + "os" + "path/filepath" + "strconv" + "strings" +) + +// capacityKeepPct is the cpu_capacity admission threshold, relative to the +// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big +// core's 1024 while mid cores sit at ~75%+, so half of max separates little +// from the rest without splitting prime from mid on three-tier parts. +const capacityKeepPct = 50 + +// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core +// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE +// cores, and AMD compact cores all sit >= 20% below their siblings' max. +const freqKeepPct = 85 + +// perfCPUs partitions allowed into the subset that are "performance" cores, +// consulting (in order of authority): +// +// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on +// arm/arm64/riscv; the scheduler's own view of big vs LITTLE. +// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present +// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright. +// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact +// cores, which neither of the above covers. +// +// Returns allowed unchanged (signal "") when nothing distinguishes the +// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable. +func perfCPUs(allowed []int) ([]int, string) { + return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed) +} + +func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) { + if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok { + return cpus, "cpu_capacity" + } + if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok { + return cpus, "intel_core_pmu" + } + if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok { + return cpus, "max_freq" + } + return allowed, "" +} + +// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least +// keepPct percent of the maximum across allowed. Inconclusive (ok=false) +// when any CPU is missing the file or when every value is equal. +func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) { + vals := make([]int, len(allowed)) + minV, maxV := 0, 0 + for i, cpu := range allowed { + v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file)) + if err != nil { + return nil, false + } + vals[i] = v + if i == 0 || v < minV { + minV = v + } + if v > maxV { + maxV = v + } + } + if minV == maxV { + return nil, false // homogeneous by this signal; try the next one + } + keep := make([]int, 0, len(allowed)) + for i, cpu := range allowed { + if vals[i]*100 >= maxV*keepPct { + keep = append(keep, cpu) + } + } + return keep, true +} + +// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU +// mask. Inconclusive when the file is absent (non-hybrid x86, other arches) +// or no allowed CPU is in the mask (the process was deliberately confined +// to E-cores; nothing useful to prefer within that). +func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) { + b, err := os.ReadFile(maskPath) + if err != nil { + return nil, false + } + set, err := parseCPUList(strings.TrimSpace(string(b))) + if err != nil || len(set) == 0 { + return nil, false + } + pcore := make(map[int]bool, len(set)) + for _, c := range set { + pcore[c] = true + } + keep := make([]int, 0, len(allowed)) + for _, cpu := range allowed { + if pcore[cpu] { + keep = append(keep, cpu) + } + } + if len(keep) == 0 { + return nil, false + } + return keep, true +} + +// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into +// individual CPU IDs. Empty input yields an empty list. +func parseCPUList(s string) ([]int, error) { + if s == "" { + return nil, nil + } + var out []int + for part := range strings.SplitSeq(s, ",") { + part = strings.TrimSpace(part) + if part == "" { + continue + } + lo, hi, isRange := strings.Cut(part, "-") + a, err := strconv.Atoi(lo) + if err != nil { + return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err) + } + if !isRange { + out = append(out, a) + continue + } + b, err := strconv.Atoi(hi) + if err != nil { + return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err) + } + if b < a || b-a > 8192 { + return nil, fmt.Errorf("bad cpulist range %q", part) + } + for v := a; v <= b; v++ { + out = append(out, v) + } + } + return out, nil +} + +func readIntFile(path string) (int, error) { + b, err := os.ReadFile(path) + if err != nil { + return 0, err + } + return strconv.Atoi(strings.TrimSpace(string(b))) +} diff --git a/cpupick/perf_linux_test.go b/cpupick/perf_linux_test.go new file mode 100644 index 00000000..00aee422 --- /dev/null +++ b/cpupick/perf_linux_test.go @@ -0,0 +1,163 @@ +//go:build linux + +package cpupick + +import ( + "fmt" + "os" + "path/filepath" + "slices" + "testing" +) + +// fakeSysfs builds a cpuDir tree with the given per-CPU file values. +// A nil map for a file means "file absent on every CPU". +func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string { + t.Helper() + dir := t.TempDir() + write := func(cpu int, rel string, v int) { + p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil { + t.Fatal(err) + } + } + for cpu, v := range capacity { + write(cpu, "cpu_capacity", v) + } + for cpu, v := range maxFreq { + write(cpu, "cpufreq/cpuinfo_max_freq", v) + } + return dir +} + +func writeCoreMask(t *testing.T, mask string) string { + t.Helper() + p := filepath.Join(t.TempDir(), "cpus") + if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil { + t.Fatal(err) + } + return p +} + +func TestPerfCPUsBigLittleCapacity(t *testing.T) { + // 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM. + dir := fakeSysfs(t, map[int]int{ + 0: 1024, 1: 1024, 2: 1024, 3: 1024, + 4: 290, 5: 290, 6: 290, 7: 290, + }, nil) + got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7}) + if signal != "cpu_capacity" { + t.Fatalf("signal = %q", signal) + } + if !slices.Equal(got, []int{0, 1, 2, 3}) { + t.Errorf("got %v", got) + } +} + +func TestPerfCPUsThreeTierKeepsMid(t *testing.T) { + // prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid. + dir := fakeSysfs(t, map[int]int{ + 0: 280, 1: 280, 2: 280, 3: 280, + 4: 780, 5: 780, 6: 780, + 7: 1024, + }, nil) + got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7}) + if !slices.Equal(got, []int{4, 5, 6, 7}) { + t.Errorf("got %v", got) + } +} + +func TestPerfCPUsIntelHybridMask(t *testing.T) { + // No cpu_capacity on x86; the P-core PMU mask decides. + dir := fakeSysfs(t, nil, nil) + mask := writeCoreMask(t, "0-7") + got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11}) + if signal != "intel_core_pmu" { + t.Fatalf("signal = %q", signal) + } + if !slices.Equal(got, []int{0, 1, 2, 3}) { + t.Errorf("got %v", got) + } +} + +func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) { + // Confined to E-cores only: the mask can't help, and equal freqs below + // mean nothing else distinguishes them either -> allowed unchanged. + dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000}) + mask := writeCoreMask(t, "0-7") + got, signal := perfCPUsFrom(dir, mask, []int{8, 9}) + if signal != "" || !slices.Equal(got, []int{8, 9}) { + t.Errorf("got %v signal %q", got, signal) + } +} + +func TestPerfCPUsMaxFreqCompactCores(t *testing.T) { + // AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz. + dir := fakeSysfs(t, nil, map[int]int{ + 0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000, + }) + got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3}) + if signal != "max_freq" { + t.Fatalf("signal = %q", signal) + } + if !slices.Equal(got, []int{0, 1}) { + t.Errorf("got %v", got) + } +} + +func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) { + // Turbo Boost Max favored cores run a few percent hot; they must not + // shrink the candidate set to one or two cores. + dir := fakeSysfs(t, nil, map[int]int{ + 0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000, + }) + got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3}) + if !slices.Equal(got, []int{0, 1, 2, 3}) { + t.Errorf("favored-core skew filtered CPUs: %v", got) + } +} + +func TestPerfCPUsHomogeneousInconclusive(t *testing.T) { + dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000}) + got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1}) + if signal != "" || !slices.Equal(got, []int{0, 1}) { + t.Errorf("got %v signal %q", got, signal) + } +} + +func TestPerfCPUsNoSysfs(t *testing.T) { + dir := t.TempDir() + got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2}) + if signal != "" || !slices.Equal(got, []int{0, 1, 2}) { + t.Errorf("got %v signal %q", got, signal) + } +} + +func TestParseCPUList(t *testing.T) { + cases := []struct { + in string + want []int + wantErr bool + }{ + {"0-3", []int{0, 1, 2, 3}, false}, + {"0-1,16-17", []int{0, 1, 16, 17}, false}, + {"5", []int{5}, false}, + {"", nil, false}, + {"3-1", nil, true}, + {"a-b", nil, true}, + {"1,x", nil, true}, + } + for _, c := range cases { + got, err := parseCPUList(c.in) + if (err != nil) != c.wantErr { + t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr) + continue + } + if !c.wantErr && !slices.Equal(got, c.want) { + t.Errorf("%q: got %v want %v", c.in, got, c.want) + } + } +} diff --git a/cpupick/perf_other.go b/cpupick/perf_other.go new file mode 100644 index 00000000..9be94583 --- /dev/null +++ b/cpupick/perf_other.go @@ -0,0 +1,10 @@ +//go:build !linux + +package cpupick + +// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction". +// Default already returns nil off-Linux (util.AllowedCPUs has no answer +// there), so this exists to keep the package compiling everywhere. +func perfCPUs(allowed []int) ([]int, string) { + return allowed, "" +} diff --git a/cpupick/topo_linux.go b/cpupick/topo_linux.go new file mode 100644 index 00000000..f72ed950 --- /dev/null +++ b/cpupick/topo_linux.go @@ -0,0 +1,118 @@ +//go:build linux + +package cpupick + +import ( + "fmt" + "os" + "path/filepath" + "strconv" + "strings" +) + +// readTopology probes the NUMA node and physical-core layout of cpus from +// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown +// node becomes node 0, an unknown core becomes a core of its own — either +// way the corresponding arrange rule becomes a no-op instead of a wrong +// answer. +func readTopology(cpus []int) topology { + return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus) +} + +func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology { + coreOf, zeroCore := coreGroups(cpuDir, cpus) + return topology{ + nodeOf: numaNodes(nodeDir, cpus), + coreOf: coreOf, + zeroCore: zeroCore, + } +} + +// numaNodes maps each cpu to its NUMA node via +// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node +// dirs at all: VMs, non-NUMA kernels) land on node 0. +func numaNodes(nodeDir string, cpus []int) map[int]int { + out := make(map[int]int, len(cpus)) + for _, c := range cpus { + out[c] = 0 + } + entries, err := os.ReadDir(nodeDir) + if err != nil { + return out + } + want := make(map[int]bool, len(cpus)) + for _, c := range cpus { + want[c] = true + } + for _, e := range entries { + id, ok := strings.CutPrefix(e.Name(), "node") + if !ok { + continue + } + n, err := strconv.Atoi(id) + if err != nil { + continue // has_cpu, possible, ... share the prefix + } + b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist")) + if err != nil { + continue + } + list, err := parseCPUList(strings.TrimSpace(string(b))) + if err != nil { + continue + } + for _, c := range list { + if want[c] { + out[c] = n + } + } + } + return out +} + +// coreGroups maps each cpu to a dense physical-core id derived from its +// (physical_package_id, core_id) pair — core_id alone repeats across +// sockets. CPUs whose topology files are unreadable get a core of their own. +// The second return is the group id of the core CPU 0 lives on, or -1 when +// that can't be determined; CPU 0's own files are consulted even when 0 is +// not a candidate, so its SMT siblings are recognized under cpusets that +// exclude CPU 0 itself. +func coreGroups(cpuDir string, cpus []int) (map[int]int, int) { + type pkgCore struct{ pkg, core int } + pairOf := func(cpu int) (pkgCore, bool) { + topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology") + pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id")) + core, err2 := readIntFile(filepath.Join(topoDir, "core_id")) + if err1 != nil || err2 != nil { + return pkgCore{}, false + } + return pkgCore{pkg, core}, true + } + + ids := map[pkgCore]int{} + out := make(map[int]int, len(cpus)) + next := 0 + for _, cpu := range cpus { + k, ok := pairOf(cpu) + if !ok { + out[cpu] = next + next++ + continue + } + id, ok := ids[k] + if !ok { + id = next + next++ + ids[k] = id + } + out[cpu] = id + } + + zeroCore := -1 + if k, ok := pairOf(0); ok { + if id, ok := ids[k]; ok { + zeroCore = id + } + } + return out, zeroCore +} diff --git a/cpupick/topo_linux_test.go b/cpupick/topo_linux_test.go new file mode 100644 index 00000000..2b914328 --- /dev/null +++ b/cpupick/topo_linux_test.go @@ -0,0 +1,111 @@ +//go:build linux + +package cpupick + +import ( + "fmt" + "os" + "path/filepath" + "testing" +) + +// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist +// string; cores maps cpu -> (package, core) pair. +func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) { + t.Helper() + base := t.TempDir() + nodeDir := filepath.Join(base, "node") + cpuDir := filepath.Join(base, "cpu") + for n, list := range nodes { + d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n)) + if err := os.MkdirAll(d, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil { + t.Fatal(err) + } + } + for cpu, pc := range cores { + d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology") + if err := os.MkdirAll(d, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil { + t.Fatal(err) + } + } + return nodeDir, cpuDir +} + +func TestReadTopology(t *testing.T) { + // Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1. + // core_id repeats across packages on purpose: the pair must disambiguate. + nodeDir, cpuDir := fakeTopoSysfs(t, + map[int]string{0: "0-1,4-5", 1: "2-3,6-7"}, + map[int][2]int{ + 0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}, + 2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1}, + }) + cpus := []int{0, 1, 2, 3, 4, 5, 6, 7} + topo := readTopologyFrom(nodeDir, cpuDir, cpus) + + for _, c := range []int{0, 1, 4, 5} { + if topo.nodeOf[c] != 0 { + t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c]) + } + } + for _, c := range []int{2, 3, 6, 7} { + if topo.nodeOf[c] != 1 { + t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c]) + } + } + pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}} + for _, p := range pairs { + if topo.coreOf[p[0]] != topo.coreOf[p[1]] { + t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]]) + } + } + if topo.coreOf[0] == topo.coreOf[2] { + t.Error("cross-package cores with equal core_id must not merge") + } + if topo.zeroCore != topo.coreOf[0] { + t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0]) + } +} + +func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) { + // CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is: + // zeroCore must still identify their shared core. + nodeDir, cpuDir := fakeTopoSysfs(t, + map[int]string{0: "0-7"}, + map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}}) + topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5}) + if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore { + t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4]) + } + if topo.coreOf[1] == topo.zeroCore { + t.Error("cpu 1 wrongly grouped with CPU 0's core") + } +} + +func TestReadTopologyMissingSysfs(t *testing.T) { + base := t.TempDir() + cpus := []int{0, 1, 2} + topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus) + seen := map[int]bool{} + for _, c := range cpus { + if topo.nodeOf[c] != 0 { + t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c]) + } + if seen[topo.coreOf[c]] { + t.Errorf("cpu %d shares a fallback core group", c) + } + seen[topo.coreOf[c]] = true + } + if topo.zeroCore != -1 { + t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore) + } +} diff --git a/cpupick/topo_other.go b/cpupick/topo_other.go new file mode 100644 index 00000000..8003ce4a --- /dev/null +++ b/cpupick/topo_other.go @@ -0,0 +1,10 @@ +//go:build !linux + +package cpupick + +// readTopology has no sysfs to consult off Linux; the flat stand-in makes +// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux +// (util.AllowedCPUs has no answer there) — this keeps the package compiling. +func readTopology(cpus []int) topology { + return flatTopology(cpus) +} diff --git a/e2e/helpers_test.go b/e2e/helpers_test.go index b555fbc4..1691aeab 100644 --- a/e2e/helpers_test.go +++ b/e2e/helpers_test.go @@ -4,15 +4,13 @@ package e2e import ( - "io" + "log/slog" "net/netip" "os" "strings" "testing" "time" - "log/slog" - "dario.cat/mergo" "github.com/google/gopacket" "github.com/google/gopacket/layers" @@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr { func NewTestLogger() *slog.Logger { v := os.Getenv("TEST_LOGS") if v == "" { - return slog.New(slog.NewTextHandler(io.Discard, nil)) + return slog.New(slog.DiscardHandler) } level := slog.LevelInfo diff --git a/examples/config.yml b/examples/config.yml index d5409bce..fc0282a9 100644 --- a/examples/config.yml +++ b/examples/config.yml @@ -131,6 +131,9 @@ listen: port: 4242 # Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg) # default is 64, does not support reload + # Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced + # superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the + # default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts. #batch: 64 # Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel # Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default) @@ -169,6 +172,8 @@ listen: # allowing for more precise routing decisions based on the packet tags. Default is 0 meaning no mark is set. # This setting is reloadable. #so_mark: 0 + # the udp_offloads setting controls if Nebula will attempt to enable GSO and GRO for its UDP socket(s). Linux only, not reloadable. + # udp_offloads: true # Routines is the number of thread pairs to run that consume from the tun and UDP queues. # Currently, this defaults to 1 which means we have 1 tun queue reader and 1 @@ -262,6 +267,33 @@ tun: # Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic mtu: 1300 + # the use_offloads setting controls if Nebula will attempt to enable GSO and GRO for the tun device. Linux only, not reloadable. + #use_offloads: true + + # Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's + # batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire + # instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1. + #pin_threads: true + + # pin_threads_key helps the CPU-auto-selector shuffle which CPUs are chosen for pinning. + # Valid options are "pid" or "port". Use "port" if you want Nebula to choose the same cores every time, which is nice for benchmarking. + # Linux only, not reloadable. + #pin_threads_key: "pid" + + # Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine + # (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra + # entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets; + # a non-integer or not-allowed entry disables the override, leaving the default pin selection described below. + # Only meaningful while pin_threads is true. Not reloadable. + # When unset (or rejected), the default spread prefers performance cores on heterogeneous CPUs (ARM big.LITTLE, + # Intel P/E hybrids, AMD compact cores), keeps all readers on one NUMA node and on distinct physical cores when the + # topology allows (SMT siblings last), leaves CPU 0's physical core as a last resort, and rotates its starting + # point per instance (keyed by the bound UDP port) so co-located nebulas don't stack their readers onto the + # same cores. + #cpu_affinity: + # - 2 + # - 4 + # Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here routes: #- mtu: 8800 diff --git a/firewall/cache.go b/firewall/cache.go index ba4b9732..3e34e6ea 100644 --- a/firewall/cache.go +++ b/firewall/cache.go @@ -5,6 +5,8 @@ import ( "log/slog" "sync/atomic" "time" + + "github.com/slackhq/nebula/logging" ) // ConntrackCache is used as a local routine cache to know if a given flow @@ -56,8 +58,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache { if tick := c.cacheTick.Load(); tick != c.cacheV { c.cacheV = tick if ll := len(c.cache); ll > 0 { - if c.l.Enabled(context.Background(), slog.LevelDebug) { - c.l.Debug("resetting conntrack cache", "len", ll) + if c.l.Enabled(context.Background(), logging.LevelTrace) { + c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll) } c.cache = make(ConntrackCache, ll) } diff --git a/firewall/cache_test.go b/firewall/cache_test.go index ab807984..3baf2326 100644 --- a/firewall/cache_test.go +++ b/firewall/cache_test.go @@ -6,6 +6,7 @@ import ( "strings" "testing" + "github.com/slackhq/nebula/logging" "github.com/slackhq/nebula/test" "github.com/stretchr/testify/assert" ) @@ -30,27 +31,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) { buf := &bytes.Buffer{} - l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug) + l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace) c := newFixedTicker(t, l, 3) c.Get() - assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String()) + assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String()) } func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) { buf := &bytes.Buffer{} - l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug) + l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace) c := newFixedTicker(t, l, 2) c.Get() - assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String())) + assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String())) } -func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) { +func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) { buf := &bytes.Buffer{} - l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo) + l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug) c := newFixedTicker(t, l, 5) c.Get() @@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) { func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) { buf := &bytes.Buffer{} - l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug) + l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace) c := newFixedTicker(t, l, 0) c.Get() diff --git a/firewall/packet.go b/firewall/packet.go index 2cbfb5ea..8e2999e5 100644 --- a/firewall/packet.go +++ b/firewall/packet.go @@ -65,3 +65,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) { "Fragment": fp.Fragment, }) } + +// ParsedPacket is a Packet plus the parse byproducts the RX path reuses +type ParsedPacket struct { + Packet + IPHdrLen int + // FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6. + // Distinct from Packet.Fragment, which is true only for NON-FIRST fragments + FragAny bool +} diff --git a/handshake_manager.go b/handshake_manager.go index f3e801d3..fa9ae154 100644 --- a/handshake_manager.go +++ b/handshake_manager.go @@ -987,6 +987,9 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn nb := make([]byte, 12, 12) out := make([]byte, mtu) for _, cp := range hh.packetStore { + // TODO: use a SendBatch here. Each callback lands in + // sendNoMetrics -> WriteTo: one syscall per cached packet, + // where one sendmmsg could flush the whole store. cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out) } f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore))) @@ -1098,7 +1101,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos // We received a valid handshake on this relay, so make sure the relay // state reflects that, in case it had been marked Disestablished. via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established) - f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false) + f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0) f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...) } } diff --git a/handshake_manager_test.go b/handshake_manager_test.go index 5f8383e4..03483d87 100644 --- a/handshake_manager_test.go +++ b/handshake_manager_test.go @@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes return } -func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) { +func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) { return } diff --git a/header/header.go b/header/header.go index b973141f..d0848c55 100644 --- a/header/header.go +++ b/header/header.go @@ -190,13 +190,18 @@ func SubTypeName(t MessageType, s MessageSubType) string { } func IsValidSubType(t MessageType, s MessageSubType) bool { - if n, ok := subTypeMap[t]; ok { - if _, ok := (*n)[s]; ok { - return true - } + switch t { + case Message: + return s == MessageNone || s == MessageRelay + case Handshake: + return s == HandshakeIXPSK0 + case Test: + return s == TestReply || s == TestRequest + case Control, CloseTunnel, RecvError, LightHouse: + return s == 0 + default: + return false } - - return false } // NewHeader turns bytes into a header diff --git a/header/header_test.go b/header/header_test.go index a7e53742..63c5b1a5 100644 --- a/header/header_test.go +++ b/header/header_test.go @@ -102,6 +102,57 @@ func TestTypeMap(t *testing.T) { }, subTypeMap) } +// mapIsValidSubType is the pre-refactor, map-driven definition of a valid +// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps +// the original behavior around so we can prove the switch is equivalent to it. +func mapIsValidSubType(t MessageType, s MessageSubType) bool { + if n, ok := subTypeMap[t]; ok { + if _, ok := (*n)[s]; ok { + return true + } + } + return false +} + +func TestIsValidSubType(t *testing.T) { + // Explicit intent table: documents exactly which subtypes are valid so the + // test stays meaningful even if both the switch and subTypeMap change. + assert.True(t, IsValidSubType(Message, MessageNone)) + assert.True(t, IsValidSubType(Message, MessageRelay)) + assert.False(t, IsValidSubType(Message, 2)) + + assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0)) + // HandshakeXXPSK0 is defined but not a wire-valid subtype. + assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0)) + + assert.True(t, IsValidSubType(Test, TestRequest)) + assert.True(t, IsValidSubType(Test, TestReply)) + assert.False(t, IsValidSubType(Test, 2)) + + // These types only ever carry subtype 0. + for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} { + assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt) + assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt) + } + + // Unknown/unassigned types are never valid. + assert.False(t, IsValidSubType(99, 0)) + + // Exhaustive proof of equivalence with the original map-driven logic across + // the entire (type, subtype) input space. + for ti := 0; ti <= 0xff; ti++ { + for si := 0; si <= 0xff; si++ { + mt, mst := MessageType(ti), MessageSubType(si) + assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst), + "IsValidSubType(%d, %d) diverged from map-driven definition", ti, si) + } + } + + // H method must delegate to the package function. + assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType()) + assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType()) +} + func TestHeader_String(t *testing.T) { assert.Equal( t, diff --git a/hostmap.go b/hostmap.go index 45515fc3..fb518552 100644 --- a/hostmap.go +++ b/hostmap.go @@ -543,6 +543,17 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool { return final } +func (hm *HostMap) QueryIndexCached(index uint32, cache map[uint32]*HostInfo) *HostInfo { + if out, ok := cache[index]; ok { + return out + } + out := hm.QueryIndex(index) + if out != nil { + cache[index] = out + } + return out +} + func (hm *HostMap) QueryIndex(index uint32) *HostInfo { hm.RLock() if h, ok := hm.Indexes[index]; ok { diff --git a/inside.go b/inside.go index c85afc2f..4f42fe84 100644 --- a/inside.go +++ b/inside.go @@ -2,6 +2,8 @@ package nebula import ( "context" + "fmt" + "io" "log/slog" "net/netip" @@ -9,10 +11,24 @@ import ( "github.com/slackhq/nebula/header" "github.com/slackhq/nebula/iputil" "github.com/slackhq/nebula/noiseutil" + "github.com/slackhq/nebula/overlay/batch" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" ) -func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) { + // borrowed: pkt.Bytes is owned by the originating tio.Queue and is + // only valid until the next Read on that queue. Every consumer below + // (parse, self-forward, handshake cache, sendInsideMessage) reads it + // synchronously; do not retain pkt outside this call. If a future + // caller needs to keep the packet, use pkt.Clone() to detach it from + // the borrow. + // + // pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO + // superpacket. In both cases the L3+L4 headers at the start describe + // the same 5-tuple every segment will share, so a single newPacket / + // firewall check covers the whole superpacket. + packet := pkt.Bytes err := newPacket(packet, false, fwPacket) if err != nil { if f.l.Enabled(context.Background(), slog.LevelDebug) { @@ -37,7 +53,14 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet // routes packets from the Nebula addr to the Nebula addr through the Nebula // TUN device. if immediatelyForwardToSelf { - _, err := f.readers[q].Write(packet) + // Write copies into the kernel queue synchronously, so seg's lifetime ends at return. + // A self-forwarded superpacket would be re-handed to the + // kernel as one giant blob; segment first so the loopback + // path sees one IP datagram per Write. + err := tio.SegmentSuperpacket(pkt, func(seg []byte) error { + _, werr := f.queues[q].Write(seg) + return werr + }) if err != nil { f.l.Error("Failed to forward to tun", "error", err) } @@ -52,12 +75,24 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet return } - hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) { - hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics) + hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) { + // borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt + // bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket), + // so retaining segments past the loop is safe. + err := tio.SegmentSuperpacket(pkt, func(seg []byte) error { + hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics) + return nil + }) + if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) { + f.l.Debug("Failed to segment superpacket for handshake cache", + "error", err, + "vpnAddr", fwPacket.RemoteAddr, + ) + } }) if hostinfo == nil { - f.rejectInside(packet, out, q) + f.rejectInside(packet, rejectBuf, q) if f.l.Enabled(context.Background(), slog.LevelDebug) { f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks", "vpnAddr", fwPacket.RemoteAddr, @@ -71,12 +106,11 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet return } - dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache) if dropReason == nil { - f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q) - + f.sendInsideMessage(hostinfo, pkt, nb, sendBatch) } else { - f.rejectInside(packet, out, q) + f.rejectInside(packet, rejectBuf, q) if f.l.Enabled(context.Background(), slog.LevelDebug) { hostinfo.logger(f.l).Debug("dropping outbound packet", "fwPacket", fwPacket, @@ -86,6 +120,125 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet } } +func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte { + if noiseutil.EncryptLockNeeded { + ci.writeLock.Lock() + } + c := ci.messageCounter.Add(1) + + out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c) + + out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb) + if noiseutil.EncryptLockNeeded { + ci.writeLock.Unlock() + } + if encErr != nil { + hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet", + "error", encErr, + "udpAddr", hostinfo.GetRemote(), + "counter", c, + ) + // Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here. + return nil + } + + return out +} + +// sendInsideMessage encrypts a firewall-approved inside packet (or every +// segment of a TSO/USO superpacket) into the caller's batch slot for +// later sendmmsg flush. Segmentation is fused with encryption here so the +// kernel-supplied superpacket bytes never get written into a separate +// scratch arena: SegmentSuperpacket builds each segment's plaintext in +// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot. +func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) { + ci := hostinfo.ConnectionState + if ci.eKey == nil { + return + } + + // One traffic-out mark covers every segment of the superpacket; doing it + // per segment in sendInsideEncrypt paid an atomic store up to ~45 extra + // times per TSO packet, inside writeLock under boring crypto. + f.connectionManager.Out(hostinfo) + + remote := hostinfo.GetRemote() + if hostinfo.lastRebindCount != f.rebindCount { + //NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is + // finally used again. This tunnel would eventually be torn down and recreated if this action didn't help. + f.lightHouse.QueryServer(hostinfo.vpnAddrs[0]) + hostinfo.lastRebindCount = f.rebindCount + if f.l.Enabled(context.Background(), slog.LevelDebug) { + hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter", + "vpnAddrs", hostinfo.vpnAddrs, + ) + } + } + + if !remote.IsValid() { //the relay path + //first, find our relay hostinfo: + var relayHostInfo *HostInfo + var relay *Relay + var err error + for _, relayIP := range hostinfo.relayState.CopyRelayIps() { + relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP) + if err != nil { + hostinfo.relayState.DeleteRelay(relayIP) + hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo", + "relay", relayIP, + "error", err, + ) + continue + } + break + } + if relayHostInfo == nil || relay == nil { + //failure already logged + return + } + + err = tio.SegmentSuperpacket(pkt, func(seg []byte) error { + //relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag + scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16) + + innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb) + if innerPacket == nil { + return nil + } + + //now we need to do a relay-encrypt: + toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true) + if err != nil { + //already logged + return nil + } + + sendBatch.Commit(toSend, relayHostInfo.GetRemote()) + return nil + }) + if err != nil { + hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err) + } + return + } + + err := tio.SegmentSuperpacket(pkt, func(seg []byte) error { + // header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + scratch := sendBatch.Reserve(header.Len + len(seg) + 16) + + out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb) + if out == nil { + return nil + } + + sendBatch.Commit(out, remote) + return nil + }) + if err != nil { + hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err) + } +} + func (f *Interface) rejectInside(packet []byte, out []byte, q int) { if !f.firewall.OutboundSendReject { return @@ -96,33 +249,36 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) { return } - _, err := f.readers[q].Write(out) + _, err := f.queues[q].Write(out) if err != nil { f.l.Error("Failed to write to tun", "error", err) } } -func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) { +func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) { if !f.firewall.InboundSendReject { return } - out = iputil.CreateRejectPacket(packet, out) + // split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything + // we can't re-use packet, if we need to send an icmp reject, it won't be long enough. + half := len(rejectBuf) / 2 + encryptBuf := rejectBuf[0:0:half] //the first half of rejectBuf's capacity, len set to 0 + buildBuf := rejectBuf[half:] + + out := iputil.CreateRejectPacket(packet, buildBuf) if len(out) == 0 { return } if len(out) > iputil.MaxRejectPacketSize { if f.l.Enabled(context.Background(), slog.LevelInfo) { - f.l.Info("rejectOutside: packet too big, not sending", - "packet", packet, - "outPacket", out, - ) + f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out) } return } - f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q) + f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, q) } // Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established @@ -216,7 +372,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac } func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) { - fp := &firewall.Packet{} + fp := &firewall.ParsedPacket{} err := newPacket(p, false, fp) if err != nil { f.l.Warn("error while parsing outgoing packet for firewall check", "error", err) @@ -224,7 +380,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp } // check if packet is in outbound fw rules - dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil) + dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil) if dropReason != nil { if f.l.Enabled(context.Background(), slog.LevelDebug) { f.l.Debug("dropping cached packet", @@ -283,21 +439,13 @@ func (f *Interface) dropExhausted(hostinfo *HostInfo, c uint64, msg string) { } } -// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done -// to the payload for the ultimate target host, making this a useful method for sending -// handshake messages to peers through relay tunnels. -// via is the HostInfo through which the message is relayed. -// ad is the plaintext data to authenticate, but not encrypt -// nb is a buffer used to store the nonce value, re-used for performance reasons. -// out is a buffer used to store the result of the Encrypt operation -// q indicates which writer to use to send the packet. -func (f *Interface) SendVia(via *HostInfo, +func (f *Interface) prepareSendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, -) { +) ([]byte, error) { if noiseutil.EncryptLockNeeded { // NOTE: for goboring AESGCMTLS we need to lock because of the nonce check via.ConnectionState.writeLock.Lock() @@ -308,7 +456,7 @@ func (f *Interface) SendVia(via *HostInfo, via.ConnectionState.writeLock.Unlock() } f.dropExhausted(via, c, "Dropping outbound relay packets, tunnel message counter is exhausted") - return + return nil, fmt.Errorf("tunnel message counter is exhausted") } out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c) @@ -326,7 +474,7 @@ func (f *Interface) SendVia(via *HostInfo, "headerLen", len(out), "cipherOverhead", via.ConnectionState.eKey.Overhead(), ) - return + return nil, io.ErrShortBuffer } // The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload. @@ -346,13 +494,31 @@ func (f *Interface) SendVia(via *HostInfo, } if err != nil { via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err) + return nil, err + } + f.connectionManager.RelayUsed(relay.LocalIndex) + return out, nil +} + +// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done +// to the payload for the ultimate target host, making this a useful method for sending +// handshake messages to peers through relay tunnels. +// via is the HostInfo through which the message is relayed. +// ad is the plaintext data to authenticate, but not encrypt +// nb is a buffer used to store the nonce value, re-used for performance reasons. +// out is a buffer used to store the result of the Encrypt operation +// q indicates which writer to use to send the packet. +func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) { + toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy) + if err != nil { + // already logged by prepareSendVia return } - err = f.writers[0].WriteTo(out, via.GetRemote()) + + err = f.writers[q].WriteTo(toSend, via.GetRemote()) if err != nil { via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err) } - f.connectionManager.RelayUsed(relay.LocalIndex) } func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) { @@ -445,7 +611,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType ) continue } - f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true) + f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q) break } } diff --git a/interface.go b/interface.go index c44f38b3..972b60fa 100644 --- a/interface.go +++ b/interface.go @@ -4,9 +4,9 @@ import ( "context" "errors" "fmt" - "io" "log/slog" "net/netip" + "runtime" "slices" "sync" "sync/atomic" @@ -14,12 +14,15 @@ import ( "github.com/gaissmai/bart" "github.com/rcrowley/go-metrics" + "github.com/slackhq/nebula/util" "github.com/slackhq/nebula/cert" "github.com/slackhq/nebula/config" "github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/header" "github.com/slackhq/nebula/overlay" + "github.com/slackhq/nebula/overlay/batch" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/udp" ) @@ -49,7 +52,19 @@ type InterfaceConfig struct { reQueryWait time.Duration ConntrackCacheTimeout time.Duration - l *slog.Logger + + // CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine + // should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] — + // shorter lists than `routines` cycle. Empty list keeps the default + // pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true. + CpuAffinity []int + // PinThreads controls whether each TUN reader OS thread is pinned to a + // single CPU (via tun.pin_threads, default true). Pinning keeps each + // goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow + // packets stay ordered on the wire. + PinThreads bool + + l *slog.Logger } type Interface struct { @@ -73,7 +88,16 @@ type Interface struct { routines int disconnectInvalid atomic.Bool closed atomic.Bool - relayManager *relayManager + // cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine + // should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)]. + // Empty falls back to the default pin-to-(allowed CPU) behavior. + // Only consulted when pinThreads is true. + cpuAffinity []int + // pinThreads controls whether listenIn pins each TUN reader OS thread to + // a CPU at all (tun.pin_threads, default true). When false, threads are + // left free to migrate as on stock nebula. + pinThreads bool + relayManager *relayManager tryPromoteEvery atomic.Uint32 reQueryEvery atomic.Uint32 @@ -90,8 +114,14 @@ type Interface struct { ctx context.Context writers []udp.Conn - readers []io.ReadWriteCloser - wg sync.WaitGroup + queues []tio.Queue + // batchers is one per tun queue, wrapping queues[i]. readOutsidePackets + // commits plaintext into the batcher; the plaintext is decrypted + // in place inside the UDP receive buffers, so listenOut must call Flush + // at the end of each UDP recvmmsg batch, before those buffers are + // reused (every udp.Conn ListenOut guarantees that ordering). + batchers []*batch.MultiCoalescer + wg sync.WaitGroup // fatalErr holds the first unexpected reader error that caused shutdown. // nil means "no fatal error" (yet) @@ -102,18 +132,13 @@ type Interface struct { metricHandshakes metrics.Histogram messageMetrics *MessageMetrics cachedPacketMetrics *cachedPacketMetrics + metricTxDropped metrics.Counter l *slog.Logger } type EncWriter interface { - SendVia(via *HostInfo, - relay *Relay, - ad, - nb, - out []byte, - nocopy bool, - ) + SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) Handshake(vpnAddr netip.Addr) @@ -172,6 +197,10 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) { return nil, errors.New("no connection manager") } + if c.routines <= 1 { + c.PinThreads = false //pinning is not useful unless there's more than one tun reader + } + cs := c.pki.getCertState() ifce := &Interface{ ctx: ctx, @@ -189,7 +218,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) { routines: c.routines, version: c.version, writers: make([]udp.Conn, c.routines), - readers: make([]io.ReadWriteCloser, c.routines), + batchers: make([]*batch.MultiCoalescer, c.routines), myVpnNetworks: cs.myVpnNetworks, myVpnNetworksTable: cs.myVpnNetworksTable, myVpnAddrs: cs.myVpnAddrs, @@ -198,8 +227,11 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) { relayManager: c.relayManager, connectionManager: c.connectionManager, conntrackCacheTimeout: c.ConntrackCacheTimeout, + cpuAffinity: c.CpuAffinity, + pinThreads: c.PinThreads, metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)), + metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil), messageMetrics: c.MessageMetrics, cachedPacketMetrics: &cachedPacketMetrics{ sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil), @@ -240,25 +272,36 @@ func (f *Interface) activate() error { "boringcrypto", boringEnabled(), ) - if f.routines > 1 { - if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() { - f.routines = 1 - f.l.Warn("routines is not supported on this platform, falling back to a single routine") - } + if f.routines > 1 && !f.outside.SupportsMultipleReaders() { + f.routines = 1 + f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine") } + // Prepare the tun queues. A device that can't open that many hands back + // fewer (a single queue on platforms without multiqueue support) and we + // size the reader routines to what we actually got. + queues, err := f.inside.Queues(f.routines) + if err != nil { + return err + } + if len(queues) < f.routines { + // TODO: this clamp is only safe because it is unreachable when the + // udp side has multiple readers (linux Queues opens exactly n or + // errors; every other platform already clamped routines to 1 above). + // If a platform ever returns fewer queues than routines with + // SO_REUSEPORT sockets already bound, the surplus sockets get no + // listenOut and the kernel blackholes every flow it hashes to them — + // fail loudly or close the extra sockets instead. + f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines", + "requested", f.routines, "opened", len(queues)) + f.routines = len(queues) + } + f.queues = queues + metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines)) - // Prepare n tun queues - var reader io.ReadWriteCloser = f.inside - for i := 0; i < f.routines; i++ { - if i > 0 { - reader, err = f.inside.NewMultiQueueReader() - if err != nil { - return err - } - } - f.readers[i] = reader + for i := range f.queues { + f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l) } // On error the caller owns the cleanup, Control.Start cancels the service context @@ -281,7 +324,7 @@ func (f *Interface) run() { // Launch n queues to read packets from tun dev for i := 0; i < f.routines; i++ { f.wg.Go(func() { - f.listenIn(f.readers[i], i) + f.listenIn(f.queues[i], i) }) } @@ -306,6 +349,31 @@ func (f *Interface) onFatal(err error) { } } +type rxContext struct { + q int + scratch []byte + // nb is a re-usable nonce buffer for decrypt calls to use + nb []byte + h *header.H + fwPacket *firewall.ParsedPacket + hostmapCache map[uint32]*HostInfo + lhh *LightHouseHandler + ctCache *firewall.ConntrackCacheTicker +} + +func newRxContext(f *Interface, q int) *rxContext { + return &rxContext{ + q: q, + scratch: make([]byte, mtu), + nb: make([]byte, 12, 12), + h: &header.H{}, + fwPacket: &firewall.ParsedPacket{}, + hostmapCache: map[uint32]*HostInfo{}, + lhh: f.lightHouse.NewRequestHandler(), + ctCache: firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout), + } +} + func (f *Interface) listenOut(i int) { var li udp.Conn if i > 0 { @@ -314,16 +382,20 @@ func (f *Interface) listenOut(i int) { li = f.outside } - ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) - lhh := f.lightHouse.NewRequestHandler() - plaintext := make([]byte, udp.MTU) - h := &header.H{} - fwPacket := &firewall.Packet{} - nb := make([]byte, 12, 12) + rxc := newRxContext(f, i) - err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) { - f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get()) - }) + listener := func(fromUdpAddr netip.AddrPort, payload []byte) { + f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc) + } + + flusher := func() { + if err := f.batchers[i].Flush(); err != nil { + f.l.Error("Failed to flush tun coalescer", "error", err) + } + clear(rxc.hostmapCache) + } + + err := li.ListenOut(listener, flusher) // An error after teardown began is shutdown noise, the closed flag covers resources // Close releases itself and the cancelled ctx covers ones torn down by their owners @@ -336,16 +408,42 @@ func (f *Interface) listenOut(i int) { f.l.Debug("underlay reader is done", "reader", i) } -func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { - packet := make([]byte, mtu) - out := make([]byte, mtu) - fwPacket := &firewall.Packet{} +func (f *Interface) pinThisThread(i int) { + var cpu int + if n := len(f.cpuAffinity); n > 0 { + // Explicit tun.cpu_affinity list wins; parseCpuAffinity already + // validated the entries against the allowed CPU set. + cpu = f.cpuAffinity[i%n] + } else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 { + // Default: spread queues across the CPUs we're actually allowed to + // run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so + // i % NumCPU would pick unrunnable IDs and every pin would fail. + cpu = allowed[i%len(allowed)] + } else { + cpu = i % runtime.NumCPU() + } + if err := util.PinThreadToCPU(cpu); err != nil { + f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err) + } +} + +func (f *Interface) listenIn(queue tio.Queue, i int) { + // Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the + // same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false. + if f.pinThreads { + f.pinThisThread(i) + } + + rejectBuf := make([]byte, mtu) + arenaSize := batch.SendBatchCap * (udp.MTU + 32) + sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize) + fwPacket := &firewall.ParsedPacket{} nb := make([]byte, 12, 12) conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) for { - n, err := reader.Read(packet) + pkts, err := queue.Read() if err != nil { // Same shutdown noise handling as listenOut if !f.closed.Load() && f.ctx.Err() == nil { @@ -355,12 +453,35 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { break } - f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get()) + for _, pkt := range pkts { + f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get()) + // Flush incrementally once a full sendmmsg batch has + // accumulated so the first packets of a deep read drain + // hit the wire while the rest are still being encrypted. + if sb.Len() >= batch.SendBatchCap { + f.flushSendBatch(sb, i) + } + } + f.flushSendBatch(sb, i) } f.l.Debug("overlay reader is done", "reader", i) } +// flushSendBatch drains sb to the underlay and accounts for anything it could not deliver. A shortfall means +// specific destinations were undeliverable (a stale remote, a reject rule), which the backend logs per peer at +// debug; here it is only a counter, so one unreachable peer cannot spam a log line per batch. +func (f *Interface) flushSendBatch(sb *batch.SendBatch, q int) { + queued := sb.Len() + written, err := sb.Flush() + if err != nil { + f.l.Error("Failed to write outgoing batch", "error", err, "writer", q) + } + if dropped := queued - written; dropped > 0 { + f.metricTxDropped.Inc(int64(dropped)) + } +} + func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) { c.RegisterReloadCallback(f.reloadFirewall) c.RegisterReloadCallback(f.reloadSendRecvError) diff --git a/iputil/packet.go b/iputil/packet.go index f91988e2..9307f2b4 100644 --- a/iputil/packet.go +++ b/iputil/packet.go @@ -204,7 +204,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte { } func ipv6CreateRejectPacket(packet []byte, out []byte) []byte { - proto, offset, isFragment, err := IPv6FindUpperProtocol(packet) + proto, offset, isFragment, _, err := IPv6FindUpperProtocol(packet) if err != nil || isFragment { return nil } @@ -346,36 +346,38 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte { // protocol and offset points at the fragment header, there is no transport header to locate. Returns // ErrIPv6CouldNotFindPayload if packet is smaller than an ipv6 header or the chain is truncated before a // terminal protocol is reached. -func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool, err error) { +func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool, anyFragment bool, err error) { + const maxIPv6ExtHeaders = 8 if len(packet) < ipv6.HeaderLen { - return 0, 0, false, ErrIPv6CouldNotFindPayload + return 0, 0, false, false, ErrIPv6CouldNotFindPayload } nextHeader = packet[6] offset = ipv6.HeaderLen - for { + for range maxIPv6ExtHeaders { switch nextHeader { case 0, 43, 60: // Hop-by-Hop, Routing, Destination if len(packet) < offset+2 { - return nextHeader, offset, isFragment, ErrIPv6CouldNotFindPayload + return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload } nextHeader = packet[offset] offset += (int(packet[offset+1]) + 1) << 3 case 44: // Fragment if len(packet) < offset+8 { - return nextHeader, offset, isFragment, ErrIPv6CouldNotFindPayload + return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload } + anyFragment = true // Non-first fragments carry no transport header, report the fragmented protocol and stop if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 { - return packet[offset], offset, true, nil + return packet[offset], offset, true, anyFragment, nil } nextHeader = packet[offset] offset += 8 case 51: // AH if len(packet) < offset+2 { - return nextHeader, offset, isFragment, ErrIPv6CouldNotFindPayload + return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload } nextHeader = packet[offset] offset += (int(packet[offset+1]) + 2) << 2 @@ -384,11 +386,12 @@ func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm // A prior extension header can declare a length that advances offset past the packet. The terminal // protocol's header isn't actually here, so treat the chain as truncated rather than classifying it. if offset > len(packet) { - return nextHeader, offset, isFragment, ErrIPv6CouldNotFindPayload + return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload } - return nextHeader, offset, isFragment, nil + return nextHeader, offset, isFragment, anyFragment, nil } } + return nextHeader, offset, isFragment, anyFragment, nil } func CreateICMPEchoResponse(packet, out []byte) []byte { diff --git a/iputil/packet_test.go b/iputil/packet_test.go index 00de5382..3baaef21 100644 --- a/iputil/packet_test.go +++ b/iputil/packet_test.go @@ -1,6 +1,7 @@ package iputil import ( + "bytes" "encoding/binary" "net" "testing" @@ -180,6 +181,46 @@ func Test_CreateRejectPacket_NoICMPError(t *testing.T) { } } +// Test_CreateRejectPacket_RespectsCap ensures it is impossible for +// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes. +func Test_CreateRejectPacket_RespectsCap(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + + // Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet + // plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes + // than the inner packet length. + inner := makeIPv6Packet(src, dst, 17, make([]byte, 20)) + + // The ciphertext scratch reused as the reject buffer is the received + // datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only + // 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes. + const nebulaOverhead = 32 + segLen := len(inner) + nebulaOverhead + + // Shared backing row laid out as [segment][neighbor's 16-byte Nebula header]. + const neighborHdr = 16 + sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr) + + // Uncapped: the slice's capacity reaches into the neighbor, reproducing + // the overrun that silently drops the neighbor packet. + backing := make([]byte, segLen+neighborHdr) + copy(backing[segLen:], sentinel) + reject := CreateRejectPacket(inner, backing[:segLen]) + assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built") + assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr], + "without the cap the oversized reject overruns into the neighbor segment") + + // Capped (the fix): cap==len, so the builder cannot exceed the segment. The + // reject does not fit, so it is refused rather than corrupting the neighbor. + backing = make([]byte, segLen+neighborHdr) + copy(backing[segLen:], sentinel) + reject = CreateRejectPacket(inner, backing[:segLen:segLen]) + assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused") + assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr], + "capped segment must leave the neighbor untouched") +} + func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte { b := make([]byte, ipv6.HeaderLen+len(payload)) b[0] = ipv6.Version << 4 @@ -496,26 +537,27 @@ func Test_IPv6FindUpperProtocol(t *testing.T) { wantProto uint8 wantOffset int wantFragment bool + wantAnyFrag bool wantErr error }{ - {"plain udp", 17, transport, 17, ipv6.HeaderLen, false, nil}, - {"hop-by-hop then tcp", 0, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, nil}, - {"routing then tcp", 43, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, nil}, - {"destination then udp", 60, append(extToUDP, transport...), 17, ipv6.HeaderLen + 8, false, nil}, - {"hop-by-hop, routing, then tcp", 0, append(append(extToRouting, extToTCP...), transport...), 6, ipv6.HeaderLen + 16, false, nil}, - {"ah then udp", 51, append(ahToUDP, transport...), 17, ipv6.HeaderLen + 8, false, nil}, - {"first fragment walks to transport", 44, append(firstFragToUDP, transport...), 17, ipv6.HeaderLen + 8, false, nil}, - {"non-first fragment stops", 44, append(nonFirstFrag, transport...), 17, ipv6.HeaderLen, true, nil}, - {"unknown protocol is terminal", 132, transport, 132, ipv6.HeaderLen, false, nil}, // SCTP - {"truncated extension header", 0, nil, 0, ipv6.HeaderLen, false, ErrIPv6CouldNotFindPayload}, + {"plain udp", 17, transport, 17, ipv6.HeaderLen, false, false, nil}, + {"hop-by-hop then tcp", 0, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, false, nil}, + {"routing then tcp", 43, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, false, nil}, + {"destination then udp", 60, append(extToUDP, transport...), 17, ipv6.HeaderLen + 8, false, false, nil}, + {"hop-by-hop, routing, then tcp", 0, append(append(extToRouting, extToTCP...), transport...), 6, ipv6.HeaderLen + 16, false, false, nil}, + {"ah then udp", 51, append(ahToUDP, transport...), 17, ipv6.HeaderLen + 8, false, false, nil}, + {"first fragment walks to transport", 44, append(firstFragToUDP, transport...), 17, ipv6.HeaderLen + 8, false, true, nil}, + {"non-first fragment stops", 44, append(nonFirstFrag, transport...), 17, ipv6.HeaderLen, true, true, nil}, + {"unknown protocol is terminal", 132, transport, 132, ipv6.HeaderLen, false, false, nil}, // SCTP + {"truncated extension header", 0, nil, 0, ipv6.HeaderLen, false, false, ErrIPv6CouldNotFindPayload}, // Destination Options with a declared length (255+1)*8 = 2048 that runs past the 48 byte buffer, next = SCTP - {"extension length past buffer", 60, []byte{132, 255, 0, 0, 0, 0, 0, 0}, 132, ipv6.HeaderLen + 2048, false, ErrIPv6CouldNotFindPayload}, + {"extension length past buffer", 60, []byte{132, 255, 0, 0, 0, 0, 0, 0}, 132, ipv6.HeaderLen + 2048, false, false, ErrIPv6CouldNotFindPayload}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { packet := makeIPv6Packet(src, dst, tt.nextHeader, tt.payload) - proto, offset, isFragment, err := IPv6FindUpperProtocol(packet) + proto, offset, isFragment, anyFragment, err := IPv6FindUpperProtocol(packet) if tt.wantErr != nil { assert.ErrorIs(t, err, tt.wantErr) return @@ -524,12 +566,13 @@ func Test_IPv6FindUpperProtocol(t *testing.T) { assert.Equal(t, tt.wantProto, proto) assert.Equal(t, tt.wantOffset, offset) assert.Equal(t, tt.wantFragment, isFragment) + assert.Equal(t, tt.wantAnyFrag, anyFragment) }) } // A packet smaller than an ipv6 header must error rather than panic reading byte 6 t.Run("shorter than ipv6 header", func(t *testing.T) { - _, _, _, err := IPv6FindUpperProtocol(make([]byte, 6)) + _, _, _, _, err := IPv6FindUpperProtocol(make([]byte, 6)) assert.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) }) } diff --git a/lighthouse_test.go b/lighthouse_test.go index 81c883ff..7a81e5d2 100644 --- a/lighthouse_test.go +++ b/lighthouse_test.go @@ -498,7 +498,7 @@ type testEncWriter struct { protocolVersion cert.Version } -func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) { +func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) { } func (tw *testEncWriter) Handshake(vpnIp netip.Addr) { } diff --git a/main.go b/main.go index da2776f1..1c243332 100644 --- a/main.go +++ b/main.go @@ -6,11 +6,14 @@ import ( "log/slog" "net" "net/netip" + "os" "runtime/debug" + "slices" "strings" "time" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/cpupick" "github.com/slackhq/nebula/noiseutil" "github.com/slackhq/nebula/overlay" "github.com/slackhq/nebula/sshd" @@ -40,6 +43,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev buildVersion = moduleVersion() } + // Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise. + startPprofServer(ctx, l) + // Print the config if in test, the exit comes later if configTest { b, err := yaml.Marshal(c.Settings) @@ -170,8 +176,21 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev } for i := 0; i < routines; i++ { - l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port))) - udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64)) + listen := netip.AddrPortFrom(listenHost, uint16(port)) + l.Info("listening", "addr", listen) + batchSize := c.GetInt("listen.batch", 64) + if batchSize < 1 { + oldBatch := batchSize + batchSize = 1 + l.Warn("listen.batch size is invalid", "provided", oldBatch, "overridden to", batchSize) + } + udpSettings := udp.Settings{ + Listen: listen, + Multi: routines > 1, + Batch: batchSize, + Offloads: c.GetBool("listen.udp_offloads", true), + } + udpServer, err := udp.NewListener(l, udpSettings) if err != nil { return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err) } @@ -220,6 +239,33 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev l.Warn("Failed to start DNS responder", "error", err) } + pinThreads := c.GetBool("tun.pin_threads", true) + cpuAffinity := parseCpuAffinity(c, l, routines) + if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest { + // The operator didn't choose pin CPUs, so pick a default set that + // prefers performance cores and doesn't stack co-located instances + // onto allowed[0]. The bound UDP port keys the per-instance spread: + // distinct across instances sharing a box, stable across restarts. + // A nil result keeps listenIn's stock allowed[i] fallback. + key := uint64(os.Getpid()) + pinKeyStr := strings.ToLower(c.GetString("tun.pin_threads_key", "")) + switch pinKeyStr { + case "": + l.Debug("tun.pin_threads_key is empty, using PID") + case "pid": + l.Debug("tun.pin_threads_key is PID") + case "port": + l.Info("tun.pin_threads_key is port number") + default: + l.Warn("tun.pin_threads_key is invalid, using PID") + } + + if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 { + key = uint64(ap.Port()) + } + cpuAffinity = cpupick.Default(routines, key, l) + } + ifConfig := &InterfaceConfig{ HostMap: hostMap, Inside: tun, @@ -241,6 +287,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev relayManager: NewRelayManager(ctx, l, hostMap, c), punchy: punchy, ConntrackCacheTimeout: conntrackCacheTimeout, + CpuAffinity: cpuAffinity, + PinThreads: pinThreads, l: l, } @@ -295,6 +343,70 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev }, nil } +// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of +// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil +// (listenIn falls back to spreading queues across the allowed CPU set). +// Length mismatch with `routines` is a warning, not an error: shorter lists +// are modulo-cycled across queues, longer lists' tail is ignored. Invalid +// entries (non-integer, or a CPU ID we're not allowed to run on) are also a +// warning and disable the override entirely so we don't silently pin to the +// wrong CPU. Entries are validated against the process's current affinity +// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or +// taskset the runnable IDs are frequently not that contiguous range, and +// pinning to an unrunnable ID always fails. If the allowed set can't be +// determined we fall back to a plain non-negative check. +func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int { + raw := c.Get("tun.cpu_affinity") + if raw == nil { + return nil + } + rv, ok := raw.([]any) + if !ok { + l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw) + return nil + } + // allowed is the set of CPU IDs we're actually permitted to run on. A nil + // slice (unsupported platform or lookup error) means "can't tell", so we + // only apply the weaker non-negative check in that case. + allowed, err := util.AllowedCPUs() + if err != nil { + l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err) + allowed = nil + } + cpus := make([]int, 0, len(rv)) + for i, e := range rv { + var cpu int + switch v := e.(type) { + case int: + cpu = v + case int64: + cpu = int(v) + case float64: + cpu = int(v) + default: + l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity", + "index", i, "value", e) + return nil + } + if cpu < 0 { + l.Warn("tun.cpu_affinity entry out of range; ignoring affinity", + "index", i, "cpu", cpu) + return nil + } + if len(allowed) > 0 && !slices.Contains(allowed, cpu) { + l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity", + "index", i, "cpu", cpu, "allowed", allowed) + return nil + } + cpus = append(cpus, cpu) + } + if len(cpus) != routines { + l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list", + "affinity_len", len(cpus), "routines", routines) + } + return cpus +} + func moduleVersion() string { info, ok := debug.ReadBuildInfo() if !ok { diff --git a/main_test.go b/main_test.go new file mode 100644 index 00000000..bbaef347 --- /dev/null +++ b/main_test.go @@ -0,0 +1,51 @@ +package nebula + +import ( + "testing" + + "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/test" + "github.com/slackhq/nebula/util" + "github.com/stretchr/testify/assert" +) + +func TestParseCpuAffinity(t *testing.T) { + l := test.NewLogger() + + // newConfig returns a config.C with tun.cpu_affinity set to v. A nil v + // leaves the key unset. + newConfig := func(v any) *config.C { + c := config.NewC(l) + if v != nil { + c.Settings["tun"] = map[string]any{"cpu_affinity": v} + } + return c + } + + // unset -> nil (listenIn falls back to spreading across the allowed set) + assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1)) + + // Pick a CPU we're actually allowed to run on so a valid list survives + // validation regardless of the host's affinity mask. + allowed, _ := util.AllowedCPUs() + validCPU := 0 + if len(allowed) > 0 { + validCPU = allowed[0] + } + + // valid list -> parsed through unchanged + assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2)) + + // a negative entry is out of range on every platform -> disables the override + assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2)) + + // a non-integer entry -> disables the override + assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2)) + + // a CPU id outside the allowed set -> disables the override. Only assertable + // where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond + // any representable CPU id so it can never be in the mask. + if len(allowed) > 0 { + assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1)) + } +} diff --git a/noiseutil/cipher_state_test.go b/noiseutil/cipher_state_test.go index 01cb959f..cb7b2703 100644 --- a/noiseutil/cipher_state_test.go +++ b/noiseutil/cipher_state_test.go @@ -183,3 +183,48 @@ func TestCipherStateNilSafety(t *testing.T) { assert.Empty(t, out) assert.Equal(t, 0, cc.Overhead()) } + +func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) { + enc, dec := buildCipherStates(t, CipherAESGCM) + inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec)) +} + +func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) { + enc, dec := buildCipherStates(t, noise.CipherChaChaPoly) + inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec)) +} + +func inPlaceDecrypt(t *testing.T, enc, dec CipherState) { + t.Helper() + const hdrLen = 16 + plaintext := []byte("in-place decrypt should replace the ciphertext bytes") + nb := make([]byte, 12) + + // packet = [16-byte header | ciphertext+tag], like a nebula Message. + packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead()) + for i := range packet { + packet[i] = byte(i) + } + packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb) + require.NoError(t, err) + + // Simulate a GRO row: [packet | next segment]. A failed auth on packet + // may zero packet's plaintext region but must not touch the header, the + // tag, or the neighboring segment. + neighbor := []byte("next coalesced segment, must stay intact") + row := append(append([]byte(nil), packet...), neighbor...) + tampered := row[:len(packet)] + tampered[hdrLen] ^= 0x01 + _, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb) + require.Error(t, err) + assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header") + assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():], + "failed auth must not touch the tag") + assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment") + + out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb) + require.NoError(t, err) + assert.Equal(t, plaintext, out) + // The plaintext must be IN the packet buffer, not a fresh allocation. + assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer") +} diff --git a/outside.go b/outside.go index b135110e..3b1d18bd 100644 --- a/outside.go +++ b/outside.go @@ -14,6 +14,7 @@ import ( "github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/header" "github.com/slackhq/nebula/iputil" + "github.com/slackhq/nebula/overlay/batch" "golang.org/x/net/ipv4" ) @@ -23,7 +24,11 @@ const ( var ErrOutOfWindow = errors.New("out of window packet") -func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { +// readOutsidePackets processes one received underlay packet. +// Message payloads are decrypted IN PLACE, so packet must stay untouched +// by the caller until the batcher for queue q has been flushed +func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxContext) { + h := rxc.h err := h.Parse(packet) if err != nil { // Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors @@ -91,7 +96,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, if isMessageRelay { hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex) } else { - hostinfo = f.hostMap.QueryIndex(h.RemoteIndex) + hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache) } // At this point we should have a valid existing tunnel, verify and send @@ -114,17 +119,18 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, // All remaining packets are encrypted if isMessageRelay { // Relay packets are special, this branch should always early-return - if err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, nb); err != nil { + err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb) + if err != nil { if f.l.Enabled(context.Background(), slog.LevelDebug) { hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h) } return } - f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache) + f.handleOutsideRelayPacket(hostinfo, via, packet, rxc) return } - out, err = hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, out, packet, nb) + out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb) if err != nil { if f.l.Enabled(context.Background(), slog.LevelDebug) { hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h) @@ -140,7 +146,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, case header.Message: switch h.Subtype { case header.MessageNone: - f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache) + f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc) default: hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h) return @@ -148,15 +154,23 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, case header.LightHouse: //TODO: assert via is not relayed - lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f) + rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f) case header.Test: switch h.Subtype { case header.TestReply: // No-op, useful for the Roaming and connectionManager side-effects above case header.TestRequest: - //recycle the input packet ciphertext as our output buffer - f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, nb, packet) + const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const + const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead + if maxOverhead+len(out) > len(rxc.scratch) { + // A reply that cannot fit in scratch is dropped no matter the log level. + if f.l.Enabled(context.Background(), slog.LevelDebug) { + hostinfo.logger(f.l).Debug("dropping oversized test request", "payloadLen", len(out), "from", via) + } + return + } + f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, rxc.nb, rxc.scratch[:0]) default: hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h) return @@ -174,7 +188,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, } } -func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) { + h := rxc.h // Successfully validated the thing. Get rid of the Relay header and the AEAD tag signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()] // Pull the Roaming parts up here, and return in all call paths. @@ -187,9 +202,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, if !ok { // The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing // its internal mapping. This should never happen. - hostinfo.logger(f.l).Error("HostInfo missing remote relay index", - "relayRemoteIndex", h.RemoteIndex, - ) + hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex) return } @@ -203,7 +216,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, relay: relay, IsRelayed: true, } - f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache) + f.readOutsidePackets(via, signedPayload, rxc) case ForwardingType: // Find the target HostInfo relay object targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr) @@ -222,8 +235,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, case ForwardingType: // Forward this packet through the relay tunnel, rebuilding it in place. // Encode overwrites the old outer header, and the new AEAD tag lands where the old one was - fwdBuf := packet[:0:len(packet)] // Cap to len(packet) to protect memory from a larger parent buffer - f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true) + fwdBuf := packet[:0] + //todo it would potentially be nice to batch these + f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q) case TerminalType: hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal") return @@ -303,7 +317,11 @@ var ( ) // newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers -func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { +func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error { + // fp is reused across packets; reset the parse byproducts so an early-error return cannot + // leak the previous packet's offsets. + fp.IPHdrLen = 0 + fp.FragAny = false if len(data) < 1 { return ErrPacketTooShort } @@ -318,7 +336,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { return ErrUnknownIPVersion } -func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error { dataLen := len(data) if dataLen < ipv6.HeaderLen { return ErrIPv6PacketTooShort @@ -335,13 +353,15 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { // Walk the extension header chain to the upper layer protocol. iputil.IPv6FindUpperProtocol is the single // source of truth for which headers are extension headers, so this stays in lockstep with the reject path // and cannot drift into misreading an unknown protocol (SCTP, GRE, etc.) as a forged transport. - proto, offset, isFragment, err := iputil.IPv6FindUpperProtocol(data) + proto, offset, isFragment, anyFragment, err := iputil.IPv6FindUpperProtocol(data) if err != nil { return ErrIPv6PacketTooShort } fp.Protocol = proto fp.Fragment = isFragment + fp.FragAny = anyFragment + fp.IPHdrLen = offset if isFragment { // Non-first fragments carry no transport header, so we have no ports to read fp.RemotePort = 0 @@ -387,7 +407,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { return nil } -func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error { // Do we at least have an ipv4 header worth of data? if len(data) < ipv4.HeaderLen { return ErrIPv4PacketTooShort @@ -404,6 +424,10 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { // Check if this is the second or further fragment of a fragmented packet. flagsfrags := binary.BigEndian.Uint16(data[6:8]) fp.Fragment = (flagsfrags & 0x1FFF) != 0 + // Any fragmentation at all (MF or offset): first fragments have readable ports for the + // firewall but must never be coalesced. + fp.FragAny = (flagsfrags & 0x3fff) != 0 + fp.IPHdrLen = ihl // Firewall handles protocol checks fp.Protocol = data[9] @@ -447,31 +471,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { return nil } -func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) { - err := newPacket(out, true, fwPacket) +func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) { + err := newPacket(out, true, rxc.fwPacket) if err != nil { - hostinfo.logger(f.l).Warn("Error while validating inbound packet", - "error", err, - "packet", out, - ) + hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out) return } - dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get()) if dropReason != nil { - // NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore - // This gives us a buffer to build the reject packet in - f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q) + f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q) if f.l.Enabled(context.Background(), slog.LevelDebug) { - hostinfo.logger(f.l).Debug("dropping inbound packet", - "fwPacket", fwPacket, - "reason", dropReason, - ) + hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason) } return } - _, err = f.readers[q].Write(out) + err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket) if err != nil { f.l.Error("Failed to write to tun", "error", err) } diff --git a/outside_test.go b/outside_test.go index ebed4564..230944cf 100644 --- a/outside_test.go +++ b/outside_test.go @@ -18,7 +18,7 @@ import ( ) func Test_newPacket(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} // length fails err := newPacket([]byte{}, true, p) @@ -97,7 +97,7 @@ func Test_newPacket(t *testing.T) { } func Test_newPacket_v6(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} // invalid ipv6 ip := layers.IPv6{ @@ -362,7 +362,7 @@ func Test_newPacket_v6(t *testing.T) { } func Test_newPacket_ipv6Fragment(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} ip := &layers.IPv6{ Version: 6, @@ -542,7 +542,7 @@ func BenchmarkParseV6(b *testing.B) { secondFrag = append(secondFrag, fragHeader...) secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...) - fp := &firewall.Packet{} + fp := &firewall.ParsedPacket{} b.Run("Normal", func(b *testing.B) { for i := 0; i < b.N; i++ { @@ -666,7 +666,7 @@ func serializeAH(ah *layers.IPSecAH) []byte { // host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land // on the same offset the host does. func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} const ( hdrLen = 40 // IPv6 header @@ -697,7 +697,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { // advances the walk past the end of the packet. The upper layer protocol's header isn't actually present, // so parseV6 must drop the packet rather than classify it as the terminal protocol with no ports. func Test_newPacket_v6ExtHeaderPastBuffer(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} pkt := make([]byte, 48) pkt[0] = 0x60 @@ -716,7 +716,7 @@ func Test_newPacket_v6ExtHeaderPastBuffer(t *testing.T) { // would trust while the host delivered the real SCTP datagram. The fix fails closed: the packet // is classified as its true protocol with no ports, so it only matches an `any` rule. func Test_newPacket_v6ExtHeaderConfusion(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} pkt := make([]byte, 52) pkt[0] = 0x60 // version 6 @@ -755,3 +755,88 @@ func Test_newPacket_v6ExtHeaderConfusion(t *testing.T) { assert.Equal(t, uint16(0), p.LocalPort) assert.False(t, p.Fragment) } + +// Test_newPacket_parsedFields pins the ParsedPacket byproducts the RX +// batcher consumes: IPHdrLen (the true L4 offset) and FragAny (any fragment +// shape at all — unlike Packet.Fragment, which is port-oriented and true +// only for non-first fragments). +func Test_newPacket_parsedFields(t *testing.T) { + p := &firewall.ParsedPacket{} + + // Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape. + v4 := make([]byte, 28) + v4[0] = 0x45 + v4[9] = firewall.ProtoTCP + binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only + require.NoError(t, newPacket(v4, true, p)) + assert.Equal(t, 20, p.IPHdrLen) + assert.False(t, p.FragAny) + assert.False(t, p.Fragment) + + // IPv4 first fragment (MF set, offset 0): the firewall can read ports + // (Fragment false) but the coalescer must not touch it (FragAny true). + ff := make([]byte, 28) + ff[0] = 0x45 + ff[9] = firewall.ProtoUDP + binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0 + require.NoError(t, newPacket(ff, true, p)) + assert.False(t, p.Fragment) + assert.True(t, p.FragAny) + assert.Equal(t, 20, p.IPHdrLen) + + // IPv4 non-first fragment (nonzero offset): both flags set. + nf := make([]byte, 28) + nf[0] = 0x45 + nf[9] = firewall.ProtoUDP + binary.BigEndian.PutUint16(nf[6:8], 0x00b9) + require.NoError(t, newPacket(nf, true, p)) + assert.True(t, p.Fragment) + assert.True(t, p.FragAny) + + // IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset. + opts := make([]byte, 32) + opts[0] = 0x46 + opts[9] = firewall.ProtoTCP + binary.BigEndian.PutUint16(opts[6:8], 0x4000) + require.NoError(t, newPacket(opts, true, p)) + assert.Equal(t, 24, p.IPHdrLen) + assert.False(t, p.FragAny) + + // Plain IPv6 TCP: L4 at 40. + v6 := make([]byte, 60) + v6[0] = 0x60 + v6[6] = firewall.ProtoTCP + require.NoError(t, newPacket(v6, true, p)) + assert.Equal(t, 40, p.IPHdrLen) + assert.False(t, p.FragAny) + + // IPv6 hop-by-hop then TCP: IPHdrLen lands past the extension header. + hbh := make([]byte, 60) + hbh[0] = 0x60 + hbh[6] = 0 // hop-by-hop + hbh[40] = firewall.ProtoTCP + hbh[41] = 0 // HdrExtLen 0 -> 8-byte header + require.NoError(t, newPacket(hbh, true, p)) + assert.Equal(t, 48, p.IPHdrLen) + assert.False(t, p.FragAny) + + // IPv6 first fragment: terminal proto resolved, FragAny set, Fragment not. + f6 := make([]byte, 60) + f6[0] = 0x60 + f6[6] = 44 // fragment extension header + f6[40] = firewall.ProtoUDP + require.NoError(t, newPacket(f6, true, p)) + assert.True(t, p.FragAny) + assert.False(t, p.Fragment) + assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) + + // IPv6 non-first fragment: both set, walk stops at the fragment header. + f6n := make([]byte, 60) + f6n[0] = 0x60 + f6n[6] = 44 + f6n[40] = firewall.ProtoUDP + binary.BigEndian.PutUint16(f6n[42:44], 0x0008) + require.NoError(t, newPacket(f6n, true, p)) + assert.True(t, p.Fragment) + assert.True(t, p.FragAny) +} diff --git a/overlay/batch/checksum_seed_test.go b/overlay/batch/checksum_seed_test.go new file mode 100644 index 00000000..dfa19784 --- /dev/null +++ b/overlay/batch/checksum_seed_test.go @@ -0,0 +1,187 @@ +package batch + +import ( + "encoding/binary" + "math/rand" + "testing" +) + +// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4 +// checksum field is pre-loaded with the folded (not inverted) pseudo-header +// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed +// produces packets every receiver silently drops, with nothing failing on +// our side — so these tests check the helpers against an independent +// RFC 1071 reference built from explicit pseudo-header bytes, never against +// the production checksum code. + +// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded) +// into a wide one's-complement accumulator. +func refSum(b []byte) uint64 { + var s uint64 + for i := 0; i+1 < len(b); i += 2 { + s += uint64(b[i])<<8 | uint64(b[i+1]) + } + if len(b)%2 == 1 { + s += uint64(b[len(b)-1]) << 8 + } + return s +} + +// refFold folds a wide one's-complement accumulator to 16 bits. +func refFold(s uint64) uint16 { + for s>>16 != 0 { + s = s&0xffff + s>>16 + } + return uint16(s) +} + +func TestFoldOnceNoInvertEdgeCases(t *testing.T) { + cases := []uint32{ + 0, 1, 0xffff, + 0x10000, // single carry + 0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff + 0xffff0000, // high half only + 0xfffeffff, // fold yields 0x1fffd: needs a second fold + 0xffffffff, // worst case + 0x00010001, // simple two-word + } + for _, c := range cases { + want := refFold(uint64(c)) + if got := foldOnceNoInvert(c); got != want { + t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want) + } + // Folding a folded value must be a no-op. + if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) { + t.Errorf("foldOnceNoInvert not idempotent at %#x", c) + } + } +} + +func TestPseudoSumIPv4MatchesReference(t *testing.T) { + cases := []struct { + name string + src, dst [4]byte + proto byte + l4Len int + }{ + {"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20}, + {"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0}, + {"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535}, + {"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535}, + {"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + // RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2). + ph := make([]byte, 12) + copy(ph[0:4], c.src[:]) + copy(ph[4:8], c.dst[:]) + ph[9] = c.proto + binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len)) + want := refFold(refSum(ph)) + + got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len)) + if got != want { + t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want) + } + }) + } +} + +func TestPseudoSumIPv6MatchesReference(t *testing.T) { + ones := func(b byte) (a [16]byte) { + for i := range a { + a[i] = b + } + return + } + cases := []struct { + name string + src, dst [16]byte + proto byte + l4Len int + }{ + {"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20}, + {"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0}, + {"max-u16-len", ones(0xff), ones(0xfe), 6, 65535}, + {"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + // RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1). + ph := make([]byte, 40) + copy(ph[0:16], c.src[:]) + copy(ph[16:32], c.dst[:]) + binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len)) + ph[39] = c.proto + want := refFold(refSum(ph)) + + got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len)) + if got != want { + t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want) + } + }) + } +} + +func TestIPv4HdrChecksumMatchesReference(t *testing.T) { + rng := rand.New(rand.NewSource(0x1791)) + for _, hdrLen := range []int{20, 24, 40, 60} { + for trial := 0; trial < 200; trial++ { + hdr := make([]byte, hdrLen) + rng.Read(hdr) + hdr[0] = 0x40 | byte(hdrLen/4) + hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires + + want := ^refFold(refSum(hdr)) + got := ipv4HdrChecksum(hdr) + if got != want { + t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want) + } + + // Receiver-side property: with the checksum stored, the full + // header must sum to all-ones. + binary.BigEndian.PutUint16(hdr[10:12], got) + if v := refFold(refSum(hdr)); v != 0xffff { + t.Fatalf("stored checksum does not validate: full-header fold = %#x", v) + } + } + } +} + +// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers +// exist for: seed the TCP checksum field with fold(pseudoSum), do what the +// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4 +// bytes including the seed, then invert, then store), and verify the result +// the way a receiver does (pseudo-header + L4 must sum to all-ones). +func TestChecksumSeedReceiverAcceptance(t *testing.T) { + rng := rand.New(rand.NewSource(0x1826)) + for trial := 0; trial < 200; trial++ { + src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))} + dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))} + payLen := rng.Intn(1500) + l4 := make([]byte, 20+payLen) + rng.Read(l4) + + // Seed exactly as flushSlot does. + seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4))) + binary.BigEndian.PutUint16(l4[16:18], seed) + + // Kernel NEEDS_CSUM completion: sum the L4 region (seed included, + // which is equivalent to summing with the field zeroed and folding + // the seed in), invert, store. + final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:])) + binary.BigEndian.PutUint16(l4[16:18], final) + + // Receiver validation. + ph := make([]byte, 12) + copy(ph[0:4], src[:]) + copy(ph[4:8], dst[:]) + ph[9] = 6 + binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4))) + if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff { + t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)", + trial, v, seed, final, payLen) + } + } +} diff --git a/overlay/batch/coalesce_core.go b/overlay/batch/coalesce_core.go new file mode 100644 index 00000000..801ab43c --- /dev/null +++ b/overlay/batch/coalesce_core.go @@ -0,0 +1,169 @@ +package batch + +import ( + "bytes" + "encoding/binary" +) + +// SortKey identifies a packet's position in its sender's transmission order. +type SortKey struct { + // Epoch is a receiver-local ordinal for the tunnel (ConnectionState) that decrypted the packet: + // a re-handshake replaces the tunnel outright and the replacement's epoch is higher, + // so the old tunnel's packets sort first during the cutover overlap. + Epoch uint64 + // Counter is the packet's AEAD message counter within that tunnel. + Counter uint64 +} + +// flowKey identifies a transport flow by {src, dst, sport, dport, family}. +// Comparable, so map lookups and linear scans over the slot list stay tight. +// Shared by the TCP and UDP coalescers; each coalescer keeps its own +// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias. +type flowKey struct { + src, dst [16]byte + sport, dport uint16 + isV6 bool +} + +// initialSlots is the starting capacity of the slot pool. +// One flow per packet is the worst case, so this matches a typical carrier-side recvmmsg batch on the UDP socket. +const initialSlots = 64 + +// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol +// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is +// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable +// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at byte 40. +// +// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must +// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt +// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these +// functions are too big to inline, and returning structs by value put five 64-byte copies on the +// per-packet path. +func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) { + if len(pkt) < 20 { + return nil, false + } + switch pkt[0] >> 4 { + case 4: + if ipHdrLen != 20 { + return nil, false + } + return fk.parseIPv4Prologue(pkt) + case 6: + if ipHdrLen != 40 || len(pkt) < 40 { + return nil, false + } + return fk.parseIPv6Prologue(pkt) + } + return nil, false +} + +// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified +// len(pkt) >= 20 and the version. +func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) { + ihl := int(pkt[0]&0x0f) * 4 + if ihl != 20 { + return nil, false + } + // Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept + // as defense in depth, since a fragment folded into a superpacket would corrupt reassembly. + if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 { + return nil, false + } + totalLen := int(binary.BigEndian.Uint16(pkt[2:4])) + if totalLen > len(pkt) || totalLen < ihl { + return nil, false + } + fk.isV6 = false + copy(fk.src[:4], pkt[12:16]) + copy(fk.dst[:4], pkt[16:20]) + return pkt[:totalLen], true +} + +// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version, +// and that the L4 header sits at byte 40. +func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) { + payloadLen := int(binary.BigEndian.Uint16(pkt[4:6])) + if 40+payloadLen > len(pkt) { + return nil, false + } + fk.isV6 = true + copy(fk.src[:], pkt[8:24]) + copy(fk.dst[:], pkt[24:40]) + return pkt[:40+payloadLen], true +} + +// ipHeadersMatch compares the IP portion of two packet header prefixes for +// byte-for-byte equality on every field that must be identical across coalesced segments. +// Size/IPID/IPCsum are masked out. +// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO: +// segments with differing ECN codepoints must not coalesce, +// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable. +// +// The transport (L4) portion of the header is checked separately by the per-protocol matcher. +func ipHeadersMatch(a, b []byte, isV6 bool) bool { + if isV6 { + // IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match), + // [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len. + return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40]) + } + // IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match), + // [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst. + // Skip [2:4] total len, [4:6] id, [10:12] csum. + return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20]) +} + +// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6). +const ipv4FlagDF = 0x40 + +// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at +// nextHdr may join a chain whose seed header is seedHdr as segment index seg +// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as +// seed_id+n, so coalescing is only transparent when that re-stamp is either +// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or +// reproduces the original IDs exactly (DF clear + IDs already sequential — +// the same admission rule kernel GRO applies). Without this, a DF=0 sender +// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs +// rewritten into ranges that collide across superpackets, corrupting +// reassembly if the packets are fragmented after the TUN write. +// +// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6 +// is inside its compared range), so checking the seed's copy suffices. +func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool { + if seedHdr[6]&ipv4FlagDF != 0 { + return true + } + expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg) + return binary.BigEndian.Uint16(nextHdr[4:6]) == expect +} + +// Arena is an injectable byte-slab that hands out non-overlapping borrowed +// slices via Reserve and releases them in bulk via Reset. +type Arena struct { + buf []byte +} + +// NewArena returns an Arena with a pre-allocated backing of the given capacity. +func NewArena(capacity int) *Arena { + return &Arena{buf: make([]byte, 0, capacity)} +} + +// Reserve hands out a non-overlapping sz-byte slice from the arena. +// If the request doesn't fit the current backing, a fresh, larger backing is allocated. +// Already-borrowed slices reference the old backing and remain valid until Reset. +func (a *Arena) Reserve(sz int) []byte { + if len(a.buf)+sz > cap(a.buf) { + newCap := max(cap(a.buf)*2, sz) + a.buf = make([]byte, 0, newCap) + } + start := len(a.buf) + a.buf = a.buf[:start+sz] + return a.buf[start : start+sz : start+sz] +} + +// Reset releases every slice handed out since the last Reset. +// Callers must not use any previously-borrowed slice after this returns. +// The underlying backing array is retained so subsequent Reserves don't re-allocate. +func (a *Arena) Reset() { + a.buf = a.buf[:0] +} diff --git a/overlay/batch/dispatch_bench_test.go b/overlay/batch/dispatch_bench_test.go new file mode 100644 index 00000000..b523be22 --- /dev/null +++ b/overlay/batch/dispatch_bench_test.go @@ -0,0 +1,112 @@ +package batch + +import ( + "testing" + + "github.com/slackhq/nebula/test" +) + +// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks +// bypass staging and the sort entirely. +func stagePackets(pkts [][]byte) []stagedPacket { + staged := make([]stagedPacket, len(pkts)) + for i, p := range pkts { + pp := testPP(p) + staged[i] = stagedPacket{ + pkt: p, + key: SortKey{Epoch: 1, Counter: uint64(i + 1)}, + proto: pp.Protocol, + fragAny: pp.FragAny, + ipHdrLen: uint16(pp.IPHdrLen), + } + } + return staged +} + +func flushLanes(b *testing.B, m *MultiCoalescer) { + b.Helper() + if m.tcp != nil { + if err := m.tcp.Flush(); err != nil { + b.Fatal(err) + } + } + if m.udp != nil { + if err := m.udp.Flush(); err != nil { + b.Fatal(err) + } + } + if err := m.pt.Flush(); err != nil { + b.Fatal(err) + } +} + +// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the +// batcher, which is where the production profile concentrates. +func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) { + b.Helper() + m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger()) + staged := stagePackets(pkts) + b.ReportAllocs() + b.SetBytes(int64(len(pkts[0]))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := m.dispatch(staged[i%len(staged)]); err != nil { + b.Fatal(err) + } + if (i+1)%batchSize == 0 { + flushLanes(b, m) + } + } + b.StopTimer() + flushLanes(b, m) +} + +// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends. +func BenchmarkDispatchSingleFlow(b *testing.B) { + runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs) +} + +// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the +// lastSlot cache on every packet. +func BenchmarkDispatchInterleaved16(b *testing.B) { + pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200) + runDispatchBench(b, pkts, len(pkts)) +} + +// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a +// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction). +func BenchmarkDispatchAckHeavy(b *testing.B) { + pay := make([]byte, 1200) + var pkts [][]byte + seq := uint32(1000) + for range tcpCoalesceMaxSegs / 2 { + pkts = append(pkts, buildTCPv4(seq, tcpAck, pay)) + seq += uint32(len(pay)) + pkts = append(pkts, buildTCPv4(seq, tcpAck, nil)) + } + runDispatchBench(b, pkts, len(pkts)) +} + +// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape. +func BenchmarkDispatchUDPFlow(b *testing.B) { + pay := make([]byte, 1200) + pkts := make([][]byte, udpCoalesceMaxSegs) + for i := range pkts { + pkts[i] = buildUDPv4(2000, 443, pay) + } + runDispatchBench(b, pkts, len(pkts)) +} + +// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes +// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including +// the parsedTCP-to-slot field transfer) can cost. +func BenchmarkDispatchSeedHeavy(b *testing.B) { + pay := make([]byte, 1200) + pkts := make([][]byte, tcpCoalesceMaxSegs) + seq := uint32(1000) + for i := range pkts { + pkts[i] = buildTCPv4(seq, tcpAckPsh, pay) + seq += uint32(len(pay)) + } + runDispatchBench(b, pkts, len(pkts)) +} diff --git a/overlay/batch/lane_entry_test.go b/overlay/batch/lane_entry_test.go new file mode 100644 index 00000000..5aeb9237 --- /dev/null +++ b/overlay/batch/lane_entry_test.go @@ -0,0 +1,76 @@ +package batch + +//TODO refactor this away +// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind +// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At +// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top +// of a local parse, so tests and benches can drive one lane with nothing but a packet. + +// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP, +// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4 +// offset; fk must be zero on entry and is filled in place. +func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) { + if len(pkt) < 20 { + return nil, 0, false + } + switch pkt[0] >> 4 { + case 4: + if pkt[9] != wantProto { + return nil, 0, false + } + trimmed, ok := fk.parseIPv4Prologue(pkt) + return trimmed, 20, ok + case 6: + if len(pkt) < 40 { + return nil, 0, false + } + if pkt[6] != wantProto { + return nil, 0, false + } + trimmed, ok := fk.parseIPv6Prologue(pkt) + return trimmed, 40, ok + } + return nil, 0, false +} + +// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for +// coalescing or not. Returns false for non-TCP or malformed input. +func (p *parsedTCP) parseBase(pkt []byte) bool { + trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP) + if !ok { + return false + } + return p.parseTail(trimmed, ipHdrLen) +} + +// parseBase extracts the flow key and IP/UDP offsets for a UDP packet. +func (p *parsedUDP) parseBase(pkt []byte) bool { + trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP) + if !ok { + return false + } + return p.parseTail(trimmed, ipHdrLen) +} + +// Commit borrows pkt. The caller must keep pkt valid until the next Flush. +func (c *TCPCoalescer) Commit(pkt []byte) error { + var info parsedTCP + if !info.parseBase(pkt) { + // Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it. + c.sealAllOpen() + c.addVerbatim(pkt) + return nil + } + return c.commitParsed(pkt, &info) +} + +// Commit borrows pkt. The caller must keep pkt valid until the next Flush. +func (c *UDPCoalescer) Commit(pkt []byte) error { + var info parsedUDP + if !info.parseBase(pkt) { + c.sealAllOpen() + c.addVerbatim(pkt) + return nil + } + return c.commitParsed(pkt, &info) +} diff --git a/overlay/batch/multi_coalesce.go b/overlay/batch/multi_coalesce.go new file mode 100644 index 00000000..c8457e2d --- /dev/null +++ b/overlay/batch/multi_coalesce.go @@ -0,0 +1,133 @@ +package batch + +import ( + "cmp" + "errors" + "io" + "log/slog" + "slices" + + "github.com/slackhq/nebula/firewall" +) + +// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush, +// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol. +// +// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in +// transmission order, builds slots in that order, and emits them in creation order. Wire reorder +// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the +// lanes carry no reorder-repair machinery. +// +// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP +// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late +// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow +// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing +// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling +// to the later-flushed pt lane. +// +// Cross-lane order (TCP vs UDP vs everything else) is not preserved. +type MultiCoalescer struct { + tcp *TCPCoalescer + udp *UDPCoalescer + pt *Passthrough + + // staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps + // each pkt alive until Flush returns. + staged []stagedPacket +} + +// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by +// value: pp is reused by the caller per packet and must not be retained past Commit. +type stagedPacket struct { + pkt []byte + key SortKey + proto byte + fragAny bool + ipHdrLen uint16 +} + +// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The +// staging sort applies even when no GSO lane is available: passthrough-only platforms still get +// transmission-order repair. +func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer { + m := &MultiCoalescer{ + pt: NewPassthrough(w), + staged: make([]stagedPacket, 0, initialSlots), + } + m.tcp = NewTCPCoalescer(w, l) + m.udp = NewUDPCoalescer(w) + return m +} + +// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in +// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed: +// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a +// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only +// for this call, so the fields dispatch needs are copied here. +func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error { + m.staged = append(m.staged, stagedPacket{ + pkt: pkt, + key: key, + proto: pp.Protocol, + fragAny: pp.FragAny, + ipHdrLen: uint16(pp.IPHdrLen), + }) + return nil +} + +// compareStaged orders staged packets by (epoch, counter) +func compareStaged(a, b stagedPacket) int { + if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 { + return c + } + return cmp.Compare(a.key.Counter, b.key.Counter) +} + +// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim +// passthrough when the lane has no GSO support. +func (m *MultiCoalescer) dispatch(sp stagedPacket) error { + switch sp.proto { + case ipProtoTCP: + if m.tcp != nil { + return m.tcp.commitStaged(sp) + } + case ipProtoUDP: + if m.udp != nil { + return m.udp.commitStaged(sp) + } + } + return m.pt.enqueue(sp.pkt) +} + +// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane. +// Drains everything and returns the joined errors; one bad packet does not hold up the rest. +// After Flush returns, committed payload slices may be recycled. +func (m *MultiCoalescer) Flush() error { + // Arrival order is already almost sorted (reorder is the exception), which pdqsort detects + // and handles in near-linear time. + slices.SortFunc(m.staged, compareStaged) + + var errs []error + for _, sp := range m.staged { + if err := m.dispatch(sp); err != nil { + errs = append(errs, err) + } + } + clear(m.staged) // drop borrowed pkt refs + m.staged = m.staged[:0] + + if m.tcp != nil { + if err := m.tcp.Flush(); err != nil { + errs = append(errs, err) + } + } + if m.udp != nil { + if err := m.udp.Flush(); err != nil { + errs = append(errs, err) + } + } + if err := m.pt.Flush(); err != nil { + errs = append(errs, err) + } + return errors.Join(errs...) +} diff --git a/overlay/batch/multi_coalesce_test.go b/overlay/batch/multi_coalesce_test.go new file mode 100644 index 00000000..c40b739a --- /dev/null +++ b/overlay/batch/multi_coalesce_test.go @@ -0,0 +1,437 @@ +package batch + +import ( + "bytes" + "encoding/binary" + "io" + "testing" + + "github.com/slackhq/nebula/firewall" + "github.com/slackhq/nebula/test" +) + +// keySeq hands out SortKeys with ascending counters in a fixed epoch, for +// tests where commit order IS transmission order. +type keySeq struct { + epoch, counter uint64 +} + +func (k *keySeq) next() SortKey { + k.counter++ + return SortKey{Epoch: k.epoch, Counter: k.counter} +} + +// newTestMultiCoalescer builds a batcher over w. +func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer { + tb.Helper() + return NewMultiCoalescer(w, test.NewLogger()) +} + +// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right +// lane: TCP and UDP get coalesced when their lanes are enabled, anything +// else (ICMP here) falls through to plain Write. +func TestMultiCoalescerRoutesByProto(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + k := &keySeq{epoch: 1} + + tcpPay := make([]byte, 1200) + udpPay := make([]byte, 1200) + icmp := make([]byte, 28) + icmp[0] = 0x45 + icmp[2] = 0 + icmp[3] = 28 + icmp[9] = 1 + + if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + // 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes. + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites)) + } + if len(w.writes) != 1 { + t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes)) + } +} + +// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort +// property: packets committed out of counter order (wire reorder inside one +// flush batch) are replayed into the lanes in transmission order, so the +// reorder never fragments the coalesce chain — one superpacket, in seq +// order, exactly as if the wire had never reordered. The retransmit shape +// falls out of the same key: a retransmit carries a lower seq but a HIGHER +// counter (it was encrypted later), so it emits after the data it trails. +func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + pay := make([]byte, 1200) + + // Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3). + // Arrival order: 3400, 1000, 2200. + if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 || len(w.writes) != 0 { + t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } + g := w.gsoWrites[0] + if len(g.pays) != 3 { + t.Fatalf("segs=%d want 3", len(g.pays)) + } + const ipHdrLen = 20 + if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 { + t.Errorf("seed seq=%d want 1000", seedSeq) + } + + // Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3). + w.writes, w.gsoWrites, w.order = nil, nil, nil + if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 2 { + t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites)) + } + first := binary.BigEndian.Uint32(w.writes[0][24:28]) + second := binary.BigEndian.Uint32(w.writes[1][24:28]) + if first != 4600 || second != 1000 { + t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second) + } +} + +// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows; +// the staging sort must repair each flow into one superpacket without any +// cross-flow contamination. +func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + pay := make([]byte, 1200) + + // Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4). + // Arrival: A.1300, B.1700, A.100, B.500. + if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + for i, g := range w.gsoWrites { + if len(g.pays) != 2 { + t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays)) + } + const ipHdrLen = 20 + seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]) + sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2]) + switch sport { + case 1000: + if seedSeq != 100 { + t.Errorf("flow A seed seq=%d want 100", seedSeq) + } + case 3000: + if seedSeq != 500 { + t.Errorf("flow B seed seq=%d want 500", seedSeq) + } + default: + t.Errorf("unexpected sport %d", sport) + } + } +} + +// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces +// the tunnel, and the replacement's counter space starts near zero — raw +// counter order would emit the new tunnel's packets first while the old +// tunnel's backlog is still arriving. The epoch key must dominate: +// everything from the old tunnel emits before anything from the new one. +func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + pay := make([]byte, 1200) + + // New session's first data arrives before the old session's last data. + if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + // Same flow, contiguous seq, identical headers: after the epoch sort the + // two segments append into one superpacket seeded by the OLD session's + // packet. + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + const ipHdrLen = 20 + if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 { + t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq) + } +} + +// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO +// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP +// packets still reach the kernel via verbatim rather than being lost. +func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true, noUSO: true} + m := newTestMultiCoalescer(t, w) + k := &keySeq{epoch: 1} + if m.udp != nil { + t.Fatal("UDP lane must not come up without USO") + } + + if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 0 { + t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites)) + } + if len(w.writes) != 2 { + t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes)) + } +} + +// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload +// anything. Both lane constructors refuse, so every packet rides the +// verbatim lane — but the staging sort still applies, so emission follows +// transmission order even without GSO. +func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: false} + m := newTestMultiCoalescer(t, w) + if m.tcp != nil || m.udp != nil { + t.Fatal("no lane may come up without offloads") + } + pkts := [][]byte{ + buildTCPv4(1000, tcpAck, make([]byte, 1200)), + buildUDPv4(1000, 53, make([]byte, 800)), + buildTCPv4(2200, tcpAck, make([]byte, 1200)), + } + // Committed in reverse transmission order; keys carry the truth. + for i := len(pkts) - 1; i >= 0; i-- { + if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil { + t.Fatal(err) + } + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 0 { + t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites)) + } + if len(w.writes) != len(pkts) { + t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes)) + } + // One lane for everything means the sorted order survives end to end. + for i, want := range pkts { + if !bytes.Equal(w.writes[i], want) { + t.Errorf("write %d out of order or corrupt", i) + } + } +} + +// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a +// single fragment header (NH=44) naming UDP as the terminal protocol — +// a first fragment (offset 0, MF set) carrying the UDP header and a +// partial payload. +func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte { + const ipHdrLen = 40 + const fragHdrLen = 8 + const udpHdrLen = 8 + total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload) + pkt := make([]byte, total) + + pkt[0] = 0x60 + binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen)) + pkt[6] = 44 // fragment extension header + pkt[7] = 64 + pkt[8] = 0xfe + pkt[9] = 0x80 + pkt[23] = 1 + pkt[24] = 0xfe + pkt[25] = 0x80 + pkt[39] = 2 + + pkt[40] = ipProtoUDP // fragment's next header + binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set + binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification + + binary.BigEndian.PutUint16(pkt[48:50], sport) + binary.BigEndian.PutUint16(pkt[50:52], dport) + binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload))) + copy(pkt[56:], payload) + return pkt +} + +// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header +// routing: a fragment whose chain terminates in UDP must ride the UDP lane +// as an in-lane verbatim — emitted ahead of later same-flow datagrams — +// not the verbatim lane, which flushes after every coalescer lane and +// would reorder it behind data that arrived after it. +func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + k := &keySeq{epoch: 1} + + if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 { + t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes)) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites)) + } + // Transmission order was fragment-then-data; same-lane routing must keep it. + if w.order[0] != "write" { + t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order) + } +} + +// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram +// (fragment) seals every open UDP chain, so datagrams from before and after +// it land in separate superpackets and the fragment holds its transmission- +// order position between them. +func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + m := newTestMultiCoalescer(t, w) + k := &keySeq{epoch: 1} + + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites)) + } + if len(w.writes) != 1 { + t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes)) + } + want := []string{"gso", "write", "gso"} + if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] { + t.Fatalf("emission order = %v, want %v", w.order, want) + } +} + +// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case. +func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true, noTSO: true} + m := newTestMultiCoalescer(t, w) + k := &keySeq{epoch: 1} + if m.tcp != nil { + t.Fatal("TCP lane must not come up without TSO") + } + + pay := make([]byte, 1200) + if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil { + t.Fatal(err) + } + if err := m.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 0 { + t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites)) + } + if len(w.writes) != 2 { + t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes)) + } +} + +// testPP derives the ParsedPacket newPacket would produce for the packet +// shapes the tests build: plain v4/v6, v4 with options or fragment bits set, +// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything +// unrecognizable stays zero (proto 0 routes to the passthrough lane). +func testPP(pkt []byte) *firewall.ParsedPacket { + pp := &firewall.ParsedPacket{} + if len(pkt) < 20 { + return pp + } + switch pkt[0] >> 4 { + case 4: + pp.Protocol = pkt[9] + pp.IPHdrLen = int(pkt[0]&0x0f) * 4 + pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 + case 6: + pp.Protocol = pkt[6] + pp.IPHdrLen = 40 + if pp.Protocol == 44 { // fragment extension header + pp.Protocol = pkt[40] + pp.IPHdrLen = 48 + pp.FragAny = true + } + } + return pp +} diff --git a/overlay/batch/passthrough.go b/overlay/batch/passthrough.go new file mode 100644 index 00000000..3d3fc4de --- /dev/null +++ b/overlay/batch/passthrough.go @@ -0,0 +1,38 @@ +package batch + +import ( + "io" +) + +// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the +// order enqueued. +type Passthrough struct { + out io.Writer + slots [][]byte +} + +func NewPassthrough(w io.Writer) *Passthrough { + return &Passthrough{ + out: w, + slots: make([][]byte, 0, 128), + } +} + +// enqueue accepts one packet, already sorted into transmission order by dispatch. +func (p *Passthrough) enqueue(pkt []byte) error { + p.slots = append(p.slots, pkt) + return nil +} + +func (p *Passthrough) Flush() error { + var firstErr error + for _, s := range p.slots { + _, err := p.out.Write(s) + if err != nil && firstErr == nil { + firstErr = err + } + } + clear(p.slots) + p.slots = p.slots[:0] + return firstErr +} diff --git a/overlay/batch/tcp_coalesce.go b/overlay/batch/tcp_coalesce.go new file mode 100644 index 00000000..e6f78021 --- /dev/null +++ b/overlay/batch/tcp_coalesce.go @@ -0,0 +1,472 @@ +package batch + +import ( + "bytes" + "encoding/binary" + "io" + "log/slog" + + "github.com/slackhq/nebula/overlay/tio" +) + +// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out. +const ipProtoTCP = 6 + +// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's +// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway. +const tcpCoalesceBufSize = 65535 + +// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single +// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency. +const tcpCoalesceMaxSegs = 64 + +// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single +// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a +// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the +// caller's plaintext buffers; the caller must keep them alive until Flush. +type coalesceSlot struct { + verbatim bool + // rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce + // slots. A slot that never grows past one segment is emitted from rawPkt so its original + // (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it. + // A multi-segment slot's superpacket header is rawPkt's, patched in place at flush. + rawPkt []byte + + fk flowKey + hdrLen int + ipHdrLen int + isV6 bool + gsoSize int + numSeg int + totalPay int + nextSeq uint32 + payIovs [][]byte +} + +// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and +// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender +// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted +// in creation order, so emission reproduces transmission order except for the pure-ACK case in +// commitParsed. Owns no locks; one coalescer per TUN write queue. +type TCPCoalescer struct { + w tio.GSOWriter + + // slots is the ordered event queue. Flush walks it once and emits each + // entry as either a WriteGSO (coalesced) or a w.Write (verbatim). + slots []*coalesceSlot + // openSlots maps a flow key to its open slot so new segments can extend an in-progress + // superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a + // non-admissible packet for the flow, or in Flush. + openSlots map[flowKey]*coalesceSlot + // lastSlot caches the most recently touched open slot. Bulk traffic + // arrives in same-flow runs (single-flow steady state, or GRO bursts + // under multi-flow), so comparing the incoming key against the cached + // slot's own fk lets the hot path skip the map lookup (and the aeshash + // of a 38-byte key) for the length of each run. + // Kept in lockstep with openSlots: nil whenever the slot it pointed + // at is removed. + lastSlot *coalesceSlot + pool []*coalesceSlot // free list for reuse + l *slog.Logger +} + +// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes. +func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer { + gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP) + if !ok { + return nil + } + return &TCPCoalescer{ + w: gw, + slots: make([]*coalesceSlot, 0, initialSlots), + openSlots: make(map[flowKey]*coalesceSlot, initialSlots), + pool: make([]*coalesceSlot, 0, initialSlots), + l: l, + } +} + +// parsedTCP holds the fields extracted from a single parse so later steps +// (admission, slot lookup, canAppend) don't re-walk the header. +type parsedTCP struct { + fk flowKey + ipHdrLen int + hdrLen int + payLen int + seq uint32 + flags byte +} + +// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is +// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on +// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input +// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers). +func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool { + trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen) + if !ok { + return false + } + return p.parseTail(trimmed, ipHdrLen) +} + +// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet; +// fk's addresses are already filled. +func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool { + if len(pkt) < ipHdrLen+20 { + return false + } + tcpOff := int(pkt[ipHdrLen+12]>>4) * 4 + if tcpOff < 20 || tcpOff > 60 { + return false + } + if len(pkt) < ipHdrLen+tcpOff { + return false + } + p.ipHdrLen = ipHdrLen + p.hdrLen = ipHdrLen + tcpOff + p.payLen = len(pkt) - p.hdrLen + p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2]) + p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4]) + p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8]) + p.flags = pkt[ipHdrLen+13] + return true +} + +// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named; +// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed. +const ( + tcpFlagPsh = 0x08 + tcpFlagAck = 0x10 + tcpFlagEce = 0x40 +) + +// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is +// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet. +func (c *TCPCoalescer) sealAllOpen() { + clear(c.openSlots) + c.lastSlot = nil +} + +// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips +// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues). +func (c *TCPCoalescer) sealFlow(fk flowKey) { + if len(c.openSlots) == 0 { + return + } + if last := c.lastSlot; last != nil && last.fk == fk { + c.lastSlot = nil + } + delete(c.openSlots, fk) +} + +// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot +// coalesce (any fragmentation, unparseable header) seals every open chain +// and rides the lane as an in-lane verbatim, still in transmission order. +func (c *TCPCoalescer) commitStaged(sp stagedPacket) error { + if sp.fragAny { + c.sealAllOpen() + c.addVerbatim(sp.pkt) + return nil + } + var info parsedTCP + if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) { + c.sealAllOpen() + c.addVerbatim(sp.pkt) + return nil + } + return c.commitParsed(sp.pkt, &info) +} + +// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a +// valid parse so the header is not re-walked here. +func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error { + // Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a + // one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN + // reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN. + if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 { + // SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later + // in-flow packets cannot extend it and emit ahead of this verbatim. + c.sealFlow(info.fk) + c.addVerbatim(pkt) + return nil + } + if info.payLen == 0 { + // Pure ACK: no ordering obligation toward the flow's data. Delivering it after + // later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing + // keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching + // kernel GRO. This is the only place emission deviates from transmission order. + c.addVerbatim(pkt) + return nil + } + + // Cached-slot fast path. Arrival isn't per-packet interleaved even with + // many flows: wire-side GRO delivers runs of same-flow packets + // (deliverSegments splits a superdatagram into up to 64), so the cache + // hits for the length of each run and a miss costs one fk compare + // before the map lookup carries the weight. + var open *coalesceSlot + if last := c.lastSlot; last != nil && last.fk == info.fk { + open = last + } else { + open = c.openSlots[info.fk] + } + if open != nil { + if c.canAppend(open, pkt, info) { + if c.appendPayload(open, pkt, info) { + // Chain closed (PSH or short segment): stop extending it. + c.sealFlow(info.fk) + } else { + c.lastSlot = open + } + return nil + } + // Can't extend (seq gap from upstream loss, header change, or a full + // chain): evict it from openSlots and fall through to seed a fresh slot. + c.sealFlow(info.fk) + } + c.seed(pkt, info) + return nil +} + +func (c *TCPCoalescer) Flush() error { + var first error + for _, s := range c.slots { + var err error + if s.verbatim || s.numSeg == 1 { + // A slot that never grew is byte-identical to its seed packet; ship the original so + // its valid checksum rides the DATA_VALID path instead of a kernel software csum. + // rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is + // pristine here. + _, err = c.w.Write(s.rawPkt) + } else { + err = c.flushSlot(s) + } + if err != nil && first == nil { + first = err + } + c.release(s) + } + clear(c.slots) + c.slots = c.slots[:0] + clear(c.openSlots) + c.lastSlot = nil + + return first +} + +func (c *TCPCoalescer) addVerbatim(pkt []byte) { + s := c.take() + s.verbatim = true + s.rawPkt = pkt + c.slots = append(c.slots, s) +} + +func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) { + if info.hdrLen+info.payLen > tcpCoalesceBufSize { + // Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can + // be open here (commitParsed evicts before seeding), so sealFlow is defense in depth + // against a stale cache entry absorbing later data. + c.sealFlow(info.fk) + c.addVerbatim(pkt) + return + } + s := c.take() + s.verbatim = false + // rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is + // the superpacket header flushSlot patches in place. + s.rawPkt = pkt + s.hdrLen = info.hdrLen + s.ipHdrLen = info.ipHdrLen + s.isV6 = info.fk.isV6 + s.fk = info.fk + s.gsoSize = info.payLen + s.numSeg = 1 + s.totalPay = info.payLen + s.nextSeq = info.seq + uint32(info.payLen) + s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen]) + c.slots = append(c.slots, s) + if info.flags&tcpFlagPsh == 0 { + c.openSlots[info.fk] = s + c.lastSlot = s + } else { + // PSH on the seed closes the chain immediately; it is never registered as open. + // Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds). + c.sealFlow(info.fk) + } +} + +// canAppend reports whether info's packet extends the slot's seed: same header shape and stable +// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the +// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine: +// the only pre-flush mutation is the PSH propagate, which also closes the chain. +func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool { + if info.hdrLen != s.hdrLen { + return false + } + if info.seq != s.nextSeq { + return false + } + if s.numSeg >= tcpCoalesceMaxSegs { + return false + } + if info.payLen > s.gsoSize { + return false + } + if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize { + return false + } + // ECE state must be stable across a burst. + // Receivers expect the flag set on every segment of a CE-echoing window or none. + seedFlags := s.rawPkt[s.ipHdrLen+13] + if (seedFlags^info.flags)&tcpFlagEce != 0 { + return false + } + if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) { + return false + } + if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) { + return false + } + return true +} + +// appendPayload folds info's packet into s and reports whether the chain is now closed: the +// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH. +// The caller must deregister a closed slot from openSlots. +func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool { + s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen]) + s.numSeg++ + s.totalPay += info.payLen + s.nextSeq = info.seq + uint32(info.payLen) + if info.flags&tcpFlagPsh != 0 { + // Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating + // rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header. + s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh + } + return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 +} + +func (c *TCPCoalescer) take() *coalesceSlot { + if n := len(c.pool); n > 0 { + s := c.pool[n-1] + c.pool[n-1] = nil + c.pool = c.pool[:n-1] + return s + } + return &coalesceSlot{} +} + +func (c *TCPCoalescer) release(s *coalesceSlot) { + clear(s.payIovs) + *s = coalesceSlot{payIovs: s.payIovs[:0]} + c.pool = append(c.pool, s) +} + +// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header +// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after, +// so nothing re-reads the patched header. Does not remove the slot from c.slots. +func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error { + total := s.hdrLen + s.totalPay + l4Len := total - s.ipHdrLen + hdr := s.rawPkt[:s.hdrLen] + + if s.isV6 { + binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len)) + } else { + binary.BigEndian.PutUint16(hdr[2:4], uint16(total)) + hdr[10] = 0 + hdr[11] = 0 + binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen])) + } + + var psum uint32 + if s.isV6 { + psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len) + } else { + psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len) + } + tcsum := s.ipHdrLen + 16 + binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum)) + + return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP) +} + +// headersMatch compares two IP+TCP header prefixes for byte-for-byte +// equality on every field that must be identical across coalesced +// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out. +func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool { + if len(a) != len(b) { + return false + } + if !ipHeadersMatch(a, b, isV6) { + return false + } + // TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window, + // [18:tcpHdrLen] options (incl. urgent). + tcp := ipHdrLen + if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) { + return false + } + if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) { + return false + } + if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) { + return false + } + if !bytes.Equal(a[tcp+18:], b[tcp+18:]) { + return false + } + return true +} + +// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must +// already have its checksum field zeroed) and returns the folded/inverted +// 16-bit value to store. +func ipv4HdrChecksum(hdr []byte) uint16 { + var sum uint32 + for i := 0; i+1 < len(hdr); i += 2 { + sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2])) + } + if len(hdr)%2 == 1 { + sum += uint32(hdr[len(hdr)-1]) << 8 + } + for sum>>16 != 0 { + sum = (sum & 0xffff) + (sum >> 16) + } + return ^uint16(sum) +} + +// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum +// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator +// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer +// reuses these helpers. +func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 { + var sum uint32 + sum += uint32(binary.BigEndian.Uint16(src[0:2])) + sum += uint32(binary.BigEndian.Uint16(src[2:4])) + sum += uint32(binary.BigEndian.Uint16(dst[0:2])) + sum += uint32(binary.BigEndian.Uint16(dst[2:4])) + sum += uint32(proto) + sum += uint32(l4Len) + return sum +} + +func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 { + var sum uint32 + for i := 0; i < 16; i += 2 { + sum += uint32(binary.BigEndian.Uint16(src[i : i+2])) + sum += uint32(binary.BigEndian.Uint16(dst[i : i+2])) + } + sum += uint32(l4Len >> 16) + sum += uint32(l4Len & 0xffff) + sum += uint32(proto) + return sum +} + +// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement). +// This is what virtio NEEDS_CSUM wants in the L4 checksum field +func foldOnceNoInvert(sum uint32) uint16 { + for sum>>16 != 0 { + sum = (sum & 0xffff) + (sum >> 16) + } + return uint16(sum) +} diff --git a/overlay/batch/tcp_coalesce_bench_test.go b/overlay/batch/tcp_coalesce_bench_test.go new file mode 100644 index 00000000..ab88138b --- /dev/null +++ b/overlay/batch/tcp_coalesce_bench_test.go @@ -0,0 +1,214 @@ +package batch + +import ( + "encoding/binary" + "testing" + + "github.com/slackhq/nebula/firewall" + "github.com/slackhq/nebula/overlay/tio" + "github.com/slackhq/nebula/test" +) + +// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards +// everything but satisfies the interface the coalescer detects. +type nopTunWriter struct{} + +func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil } +func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error { + return nil +} +func (nopTunWriter) Capabilities() tio.Capabilities { + return tio.Capabilities{TSO: true, USO: true} +} + +// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments +// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are +// contiguous so every packet is coalesceable onto the previous one. +func buildTCPv4BulkFlow(n, payloadLen int) [][]byte { + pkts := make([][]byte, n) + pay := make([]byte, payloadLen) + seq := uint32(1000) + for i := range n { + pkts[i] = buildTCPv4(seq, tcpAck, pay) + seq += uint32(payloadLen) + } + return pkts +} + +// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow +// seq continuity but round-robin across flows — worst case for any +// "last-slot" cache. +func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte { + pay := make([]byte, payloadLen) + seqs := make([]uint32, nFlows) + for i := range seqs { + seqs[i] = uint32(1000 + i*1000000) + } + pkts := make([][]byte, 0, nFlows*perFlow) + for range perFlow { + for f := range nFlows { + sport := uint16(10000 + f) + pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay)) + seqs[f] += uint32(payloadLen) + } + } + return pkts +} + +// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in +// runs of runLen per flow — the arrival pattern wire-side GRO actually +// produces (deliverSegments splits each superdatagram into up to 64 +// same-flow packets back to back). Contrast with buildTCPv4Interleaved's +// per-packet round-robin, the adversarial worst case for a last-slot cache. +func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte { + pay := make([]byte, payloadLen) + seqs := make([]uint32, nFlows) + for i := range seqs { + seqs[i] = uint32(1000 + i*1000000) + } + pkts := make([][]byte, 0, nFlows*perFlow) + for done := 0; done < perFlow; done += runLen { + for f := range nFlows { + sport := uint16(10000 + f) + for range runLen { + pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay)) + seqs[f] += uint32(payloadLen) + } + } + } + return pkts +} + +// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim +// branch in Commit. +func buildICMPv4() []byte { + pkt := make([]byte, 28) + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], 28) + pkt[9] = 1 // ICMP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + return pkt +} + +// runCommitBench drives Commit over pkts batchSize at a time, flushing +// between batches, and reports per-packet cost. +func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) { + b.Helper() + c := newTestTCPCoalescer(b, nopTunWriter{}) + b.ReportAllocs() + b.SetBytes(int64(len(pkts[0]))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + pkt := pkts[i%len(pkts)] + if err := c.Commit(pkt); err != nil { + b.Fatal(err) + } + if (i+1)%batchSize == 0 { + if err := c.Flush(); err != nil { + b.Fatal(err) + } + } + } + // Drain any trailing partial batch so slot state doesn't leak across runs. + _ = c.Flush() +} + +// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow, +// contiguous seq, 1200-byte payloads. Every packet past the seed should +// append onto the open slot. This is the case we most care about. +func BenchmarkCommitSingleFlow(b *testing.B) { + pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200) + runCommitBench(b, pkts, tcpCoalesceMaxSegs) +} + +// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined. +// A single-entry fast-path cache will miss on every packet; an N-way +// cache or map lookup carries the weight. +func BenchmarkCommitInterleaved4(b *testing.B) { + pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200) + runCommitBench(b, pkts, len(pkts)) +} + +// BenchmarkCommitInterleaved16 stresses the map at higher flow counts. +func BenchmarkCommitInterleaved16(b *testing.B) { + pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200) + runCommitBench(b, pkts, len(pkts)) +} + +// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in +// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot +// cache hits for the length of each run; the per-packet round-robin +// benches above are its worst case. +func BenchmarkCommitRunInterleaved4(b *testing.B) { + pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200) + runCommitBench(b, pkts, len(pkts)) +} + +// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase +// bails early and addVerbatim is the only work. +func BenchmarkCommitPassthrough(b *testing.B) { + pkt := buildICMPv4() + pkts := make([][]byte, 64) + for i := range pkts { + pkts[i] = pkt + } + runCommitBench(b, pkts, 64) +} + +// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow. +// Each packet takes the "TCP but not admissible" branch which does a +// map delete + verbatim. Measures the seal-without-slot cost. +func BenchmarkCommitNonCoalesceableTCP(b *testing.B) { + pay := make([]byte, 0) + pkts := make([][]byte, 64) + for i := range pkts { + pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay) + } + runCommitBench(b, pkts, 64) +} + +// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so +// it includes the staging sort's already-sorted fast path plus the +// dispatch-time parse — the full steady-state cost of the batcher. The +// ParsedPackets are precomputed: in production they fall out of the +// firewall's newPacket, which this bench does not model. +func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) { + b.Helper() + m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger()) + pps := make([]*firewall.ParsedPacket, len(pkts)) + for i, p := range pkts { + pps[i] = testPP(p) + } + b.ReportAllocs() + b.SetBytes(int64(len(pkts[0]))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + j := i % len(pkts) + if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil { + b.Fatal(err) + } + if (i+1)%batchSize == 0 { + if err := m.Flush(); err != nil { + b.Fatal(err) + } + } + } + _ = m.Flush() +} + +// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of +// BenchmarkCommitSingleFlow — same workload but routed through the +// dispatcher. The delta vs the single-lane bench measures dispatcher +// overhead. +func BenchmarkMultiCommitSingleFlow(b *testing.B) { + pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200) + runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs) +} + +// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4 +// through the dispatcher. +func BenchmarkMultiCommitInterleaved4(b *testing.B) { + pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200) + runMultiCommitBench(b, pkts, len(pkts)) +} diff --git a/overlay/batch/tcp_coalesce_test.go b/overlay/batch/tcp_coalesce_test.go new file mode 100644 index 00000000..9c6c4791 --- /dev/null +++ b/overlay/batch/tcp_coalesce_test.go @@ -0,0 +1,1287 @@ +package batch + +import ( + "bytes" + "encoding/binary" + "io" + "testing" + + "github.com/slackhq/nebula/overlay/tio" + "github.com/slackhq/nebula/test" +) + +// fakeTunWriter records plain Writes and WriteGSO calls without touching a +// real TUN fd. WriteGSO records the IP header, transport header, and +// borrowed payload fragments separately so tests can inspect each. +// noTSO / noUSO withhold one offload from an otherwise GSO-capable writer, so +// tests can build the half-capable queues real kernels hand us (USO needs a +// newer kernel than TSO). +type fakeTunWriter struct { + gsoEnabled bool + noTSO bool + noUSO bool + writes [][]byte + gsoWrites []fakeGSOWrite + // order records the interleaving of Write ("write") and WriteGSO ("gso") + // calls for tests that assert cross-call emission order. + order []string +} + +// fakeGSOWrite captures one WriteGSO call. hdr is the concatenation of the +// IP and transport headers (in that order), gsoSize / isV6 / csumStart are +// derived from the call so existing assertions keep working unchanged. +type fakeGSOWrite struct { + hdr []byte + pays [][]byte + gsoSize uint16 + isV6 bool + csumStart uint16 +} + +// total returns hdrLen + sum of pay lens. +func (g fakeGSOWrite) total() int { + n := len(g.hdr) + for _, p := range g.pays { + n += len(p) + } + return n +} + +// payLen sums the pays. +func (g fakeGSOWrite) payLen() int { + var n int + for _, p := range g.pays { + n += len(p) + } + return n +} + +func (w *fakeTunWriter) Write(p []byte) (int, error) { + buf := make([]byte, len(p)) + copy(buf, p) + w.writes = append(w.writes, buf) + w.order = append(w.order, "write") + return len(p), nil +} + +func (w *fakeTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error { + hcopy := make([]byte, len(hdr)+len(transportHdr)) + copy(hcopy, hdr) + copy(hcopy[len(hdr):], transportHdr) + paysCopy := make([][]byte, len(pays)) + for i, p := range pays { + pc := make([]byte, len(p)) + copy(pc, p) + paysCopy[i] = pc + } + var gsoSize uint16 + if len(pays) > 1 { + gsoSize = uint16(len(pays[0])) + } + isV6 := len(hdr) > 0 && hdr[0]>>4 == 6 + w.gsoWrites = append(w.gsoWrites, fakeGSOWrite{ + hdr: hcopy, + pays: paysCopy, + gsoSize: gsoSize, + isV6: isV6, + csumStart: uint16(len(hdr)), + }) + w.order = append(w.order, "gso") + return nil +} + +func (w *fakeTunWriter) Capabilities() tio.Capabilities { + return tio.Capabilities{TSO: w.gsoEnabled && !w.noTSO, USO: w.gsoEnabled && !w.noUSO} +} + +// buildTCPv4 constructs a minimal IPv4+TCP packet with the given payload, +// seq, and flags. Assumes no IP options and a 20-byte TCP header. +func buildTCPv4(seq uint32, flags byte, payload []byte) []byte { + return buildTCPv4Ports(1000, 2000, seq, flags, payload) +} + +// buildTCPv4Ports is buildTCPv4 with caller-specified ports so tests can +// build distinct flows. +func buildTCPv4Ports(sport, dport uint16, seq uint32, flags byte, payload []byte) []byte { + const ipHdrLen = 20 + const tcpHdrLen = 20 + total := ipHdrLen + tcpHdrLen + len(payload) + pkt := make([]byte, total) + + pkt[0] = 0x45 + pkt[1] = 0x00 + binary.BigEndian.PutUint16(pkt[2:4], uint16(total)) + binary.BigEndian.PutUint16(pkt[4:6], 0) + binary.BigEndian.PutUint16(pkt[6:8], 0x4000) + pkt[8] = 64 + pkt[9] = ipProtoTCP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + + binary.BigEndian.PutUint16(pkt[20:22], sport) + binary.BigEndian.PutUint16(pkt[22:24], dport) + binary.BigEndian.PutUint32(pkt[24:28], seq) + binary.BigEndian.PutUint32(pkt[28:32], 12345) + pkt[32] = 0x50 + pkt[33] = flags + binary.BigEndian.PutUint16(pkt[34:36], 0xffff) + + copy(pkt[40:], payload) + return pkt +} + +const ( + tcpAck = 0x10 + tcpPsh = 0x08 + tcpSyn = 0x02 + tcpFin = 0x01 + tcpAckPsh = tcpAck | tcpPsh +) + +// setIPv4ID stamps an IPv4 ID and DF state onto a builder packet. The +// builders default to DF=1/ID=0 (an atomic datagram); the ID-admission +// tests use this to fabricate non-atomic (DF=0) senders. +func setIPv4ID(pkt []byte, id uint16, df bool) { + binary.BigEndian.PutUint16(pkt[4:6], id) + var flags uint16 + if df { + flags = 0x4000 + } + binary.BigEndian.PutUint16(pkt[6:8], flags) +} + +// newTestTCPCoalescer builds a coalescer over w and fails the test if w can't +// do TSO. Every test but TestNewTCPCoalescerRefusesWhenGSOUnavailable wants the +// GSO path, and the constructor now hands back a nil coalescer otherwise. +func newTestTCPCoalescer(tb testing.TB, w io.Writer) *TCPCoalescer { + tb.Helper() + c := NewTCPCoalescer(w, test.NewLogger()) + if c == nil { + tb.Fatal("NewTCPCoalescer: writer does not support TSO") + } + return c +} + +// TestNewTCPCoalescerRefusesWhenGSOUnavailable pins the constructor +// precondition: no TSO, no coalescer. There's no degraded mode — the caller +// (MultiCoalescer) sends TCP down the verbatim lane instead. +func TestNewTCPCoalescerRefusesWhenGSOUnavailable(t *testing.T) { + if c := NewTCPCoalescer(&fakeTunWriter{gsoEnabled: false}, test.NewLogger()); c != nil { + t.Fatalf("want nil for a non-TSO writer, got %v", c) + } + // A writer that isn't a GSOWriter at all is refused the same way. + if c := NewTCPCoalescer(&plainOnlyWriter{}, test.NewLogger()); c != nil { + t.Fatalf("want nil for a plain writer, got %v", c) + } +} + +// plainOnlyWriter is an io.Writer with no GSO support at all — the +// single-packet Queue shape. +type plainOnlyWriter struct{ writes int } + +func (w *plainOnlyWriter) Write(p []byte) (int, error) { + w.writes++ + return len(p), nil +} + +func TestCoalescerNonTCPPassthrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pkt := make([]byte, 28) + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], 28) + pkt[9] = 1 + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("ICMP should pass through unchanged") + } +} + +func TestCoalescerSeedThenFlushAlone(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pkt := buildTCPv4(1000, tcpAck, make([]byte, 1000)) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if len(w.writes) != 0 || len(w.gsoWrites) != 0 { + t.Fatalf("unexpected output before flush") + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // A slot that never grew past one segment flushes as a plain Write of + // the original packet bytes: the original (already valid) checksum + // ships via the DATA_VALID path, so the kernel does no csum work. + // WriteGSO is reserved for slots that actually coalesced (>=2 segs). + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if !bytes.Equal(w.writes[0], pkt) { + t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt)) + } +} + +// TestCoalescerPureAckDoesNotSealRun pins the pure-ACK fast path: a bare +// acknowledgment (zero payload, nothing beyond ACK|PSH|ECE) rides its lane +// as a verbatim WITHOUT sealing the flow's open slot, so an inbound data +// run on a bidirectional connection keeps coalescing across the peer ACKs +// interleaved into it. The ACK is emitted after the superpacket (stale ACKs +// are ignored by receivers, so the reorder is harmless by design). +func TestCoalescerPureAckDoesNotSealRun(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + ack := buildTCPv4(2200, tcpAck, nil) + if err := c.Commit(ack); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 || len(w.writes) != 1 { + t.Fatalf("ACK sealed the run: writes=%d gso=%d, want 1 gso (2 pays) + 1 plain", len(w.writes), len(w.gsoWrites)) + } + if got := len(w.gsoWrites[0].pays); got != 2 { + t.Errorf("pay count=%d want 2 (data kept coalescing across the ACK)", got) + } + if !bytes.Equal(w.writes[0], ack) { + t.Errorf("plain write is not the ACK packet: got %d bytes want %d", len(w.writes[0]), len(ack)) + } + if got, want := w.order, []string{"gso", "write"}; !stringSliceEq(got, want) { + t.Errorf("flush order=%v want %v (slot order: data run seeded first)", got, want) + } +} + +// TestCoalescerFinStillSealsRun is the guard rail for the pure-ACK fast +// path: control flags (here FIN|ACK, zero payload) must keep sealing the +// open slot so data never reorders across a flow-state transition. +func TestCoalescerFinStillSealsRun(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpFin|tcpAck, nil)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // FIN evicts the open slot; the third packet seeds a fresh one. All + // three stay single-segment, so all three emit as plain writes in + // arrival order — any gso write would mean data coalesced across FIN. + if len(w.writes) != 3 || len(w.gsoWrites) != 0 { + t.Fatalf("FIN must seal the run: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestCoalescerCoalescesAdjacentACKs(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + g := w.gsoWrites[0] + if g.gsoSize != 1200 { + t.Errorf("gsoSize=%d want 1200", g.gsoSize) + } + if len(g.hdr) != 40 { + t.Errorf("hdrLen=%d want 40", len(g.hdr)) + } + if g.csumStart != 20 { + t.Errorf("csumStart=%d want 20", g.csumStart) + } + if len(g.pays) != 3 { + t.Errorf("pay count=%d want 3", len(g.pays)) + } + if g.total() != 40+3*1200 { + t.Errorf("superpacket len=%d want %d", g.total(), 40+3*1200) + } + if tot := binary.BigEndian.Uint16(g.hdr[2:4]); int(tot) != g.total() { + t.Errorf("ip total_length=%d want %d", tot, g.total()) + } +} + +func TestCoalescerRejectsSeqGap(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(3000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Each packet stays a single-segment slot and flushes as its own plain + // write of the original bytes. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("seq gap: want 2 plain writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestCoalescerRejectsFlagMismatch(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // SYN|ACK is non-admissible. Must flush the matching flow's slot — + // single-segment, so a plain write of the original bytes — and then + // plain-write the SYN packet itself. + syn := buildTCPv4(2200, tcpSyn|tcpAck, pay) + if err := c.Commit(syn); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("flag mismatch: want 2 plain writes (flushed seed + SYN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if !bytes.Equal(w.writes[1], syn) { + t.Errorf("second plain write should be the SYN packet, got %d bytes", len(w.writes[1])) + } +} + +func TestCoalescerRejectsFIN(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + fin := buildTCPv4(1000, tcpAck|tcpFin, []byte("x")) + if err := c.Commit(fin); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // FIN isn't admissible — verbatim as plain, no slot, no gso. + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("FIN should be verbatim, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestCoalescerShortLastSegmentClosesChain(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + full := make([]byte, 1200) + half := make([]byte, 500) + if err := c.Commit(buildTCPv4(1000, tcpAck, full)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, half)); err != nil { + t.Fatal(err) + } + // Chain now closed; next packet seeds a new slot on the same flow + // after flushing the old one. + if err := c.Commit(buildTCPv4(2700, tcpAck, full)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Expect one gso write for the first two packets coalesced, then the + // third — still single-segment — flushed as a plain write of the + // original packet. + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write got %d", len(w.gsoWrites)) + } + if len(w.writes) != 1 { + t.Fatalf("want 1 plain write got %d", len(w.writes)) + } + if w.gsoWrites[0].gsoSize != 1200 { + t.Errorf("gsoSize=%d want 1200", w.gsoWrites[0].gsoSize) + } + if got, want := w.gsoWrites[0].total(), 40+1200+500; got != want { + t.Errorf("super len=%d want %d", got, want) + } + if got, want := len(w.writes[0]), 40+1200; got != want { + t.Errorf("plain write len=%d want %d", got, want) + } +} + +func TestCoalescerPSHFinalizesChain(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // First two coalesce into one gso write; the third seeds a fresh slot + // that stays single-segment and flushes as a plain write. + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write got %d", len(w.gsoWrites)) + } + if len(w.writes) != 1 { + t.Fatalf("want 1 plain write got %d", len(w.writes)) + } +} + +// TestCoalescerPropagatesPSHFromAppended ensures that when an appended +// segment carries PSH (or is short, sealing the chain), the PSH bit ends +// up in the emitted superpacket's TCP flags. The kernel TSO path keeps +// PSH only on the last segment iff the input header has it set; if the +// coalescer drops it the sender's push signal never reaches the receiver. +func TestCoalescerPropagatesPSHFromAppended(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + // Seed has no PSH; second segment carries PSH and seals the chain. + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write got %d", len(w.gsoWrites)) + } + g := w.gsoWrites[0] + const ipHdrLen = 20 + flags := g.hdr[ipHdrLen+13] + if flags&tcpPsh == 0 { + t.Fatalf("PSH lost from coalesced superpacket: flags=0x%02x", flags) + } + if flags&tcpAck == 0 { + t.Fatalf("ACK missing from coalesced superpacket: flags=0x%02x", flags) + } +} + +func TestCoalescerRejectsDifferentFlow(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + p1 := buildTCPv4(1000, tcpAck, pay) + p2 := buildTCPv4(2200, tcpAck, pay) + binary.BigEndian.PutUint16(p2[20:22], 9999) + if err := c.Commit(p1); err != nil { + t.Fatal(err) + } + if err := c.Commit(p2); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Two independent flows, each stays single-segment and flushes as its + // own plain write of the original bytes. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("diff flow: want 2 plain writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestCoalescerRejectsIPOptions(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 500) + pkt := buildTCPv4(1000, tcpAck, pay) + // Bump IHL to 6 to simulate 4 bytes of IP options. Don't actually add + // bytes — parser should bail before it matters. + pkt[0] = 0x46 + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Non-admissible parse → verbatim as plain. + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("IP options should verbatim, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestCoalescerCapBySegments(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 512) + seq := uint32(1000) + for i := 0; i < tcpCoalesceMaxSegs+5; i++ { + if err := c.Commit(buildTCPv4(seq, tcpAck, pay)); err != nil { + t.Fatal(err) + } + seq += uint32(len(pay)) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + for _, g := range w.gsoWrites { + segs := len(g.pays) + if segs > tcpCoalesceMaxSegs { + t.Fatalf("super exceeded seg cap: %d > %d", segs, tcpCoalesceMaxSegs) + } + } +} + +// TestCoalescerMultipleFlowsInSameBatch proves two interleaved bulk TCP +// flows coalesce independently in a single Flush. +func TestCoalescerMultipleFlowsInSameBatch(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + // Flow A: sport 1000. Flow B: sport 3000. + if err := c.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(1000, 2000, 2500, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(3000, 2000, 2900, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites)) + } + if len(w.writes) != 0 { + t.Fatalf("want no plain writes, got %d", len(w.writes)) + } + // Each superpacket should carry 3 segments. + for i, g := range w.gsoWrites { + if len(g.pays) != 3 { + t.Errorf("gso[%d]: segs=%d want 3", i, len(g.pays)) + } + if g.gsoSize != 1200 { + t.Errorf("gso[%d]: gsoSize=%d want 1200", i, g.gsoSize) + } + } + // Verify each superpacket carries the source port it was seeded with. + seenSports := map[uint16]bool{} + for _, g := range w.gsoWrites { + sp := binary.BigEndian.Uint16(g.hdr[20:22]) + seenSports[sp] = true + } + if !seenSports[1000] || !seenSports[3000] { + t.Errorf("expected superpackets for sports 1000 and 3000, got %v", seenSports) + } +} + +// TestCoalescerPreservesArrivalOrder confirms that with verbatim and +// coalesced events both queued, Flush emits them in Add order rather than +// writing verbatim packets synchronously. +func TestCoalescerPreservesArrivalOrder(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + // Sequence: coalesceable TCP, ICMP (verbatim), coalesceable TCP on + // a different flow. Both TCP slots stay single-segment, so all three + // emit as plain writes; the packet order (X, ICMP, Y) is asserted by + // byte content since the kinds no longer distinguish them. + pay := make([]byte, 1200) + tcpX := buildTCPv4Ports(1000, 2000, 100, tcpAck, pay) + if err := c.Commit(tcpX); err != nil { + t.Fatal(err) + } + icmp := make([]byte, 28) + icmp[0] = 0x45 + binary.BigEndian.PutUint16(icmp[2:4], 28) + icmp[9] = 1 + copy(icmp[12:16], []byte{10, 0, 0, 1}) + copy(icmp[16:20], []byte{10, 0, 0, 3}) + if err := c.Commit(icmp); err != nil { + t.Fatal(err) + } + tcpY := buildTCPv4Ports(3000, 2000, 500, tcpAck, pay) + if err := c.Commit(tcpY); err != nil { + t.Fatal(err) + } + // Nothing should have hit the writer synchronously. + if len(w.order) != 0 { + t.Fatalf("Add emitted events synchronously: %v", w.order) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if got, want := w.order, []string{"write", "write", "write"}; !stringSliceEq(got, want) { + t.Fatalf("flush order=%v want %v", got, want) + } + for i, want := range [][]byte{tcpX, icmp, tcpY} { + if !bytes.Equal(w.writes[i], want) { + t.Fatalf("write %d out of arrival order: got %d bytes, want %d bytes", i, len(w.writes[i]), len(want)) + } + } +} + +func stringSliceEq(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +// TestCoalescerInterleavedFlowsPreserveOrdering checks that a non-admissible +// packet (SYN) mid-flow only flushes its own flow, not others. +func TestCoalescerInterleavedFlowsPreserveOrdering(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + // Flow A two segments. + if err := c.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Flow B two segments. + if err := c.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Flow A SYN (non-admissible) — must flush only flow A's slot. + syn := buildTCPv4Ports(1000, 2000, 9999, tcpSyn|tcpAck, pay) + if err := c.Commit(syn); err != nil { + t.Fatal(err) + } + // Flow B continues — should still be coalesced with its seed. + if err := c.Commit(buildTCPv4Ports(3000, 2000, 2900, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + + // Expected: + // - 1 gso for flow A (first 2 segments) + // - 1 plain for flow A SYN + // - 1 gso for flow B (3 segments) + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes, got %d", len(w.gsoWrites)) + } + if len(w.writes) != 1 { + t.Fatalf("want 1 plain write (SYN), got %d", len(w.writes)) + } + // Find the 3-segment gso (flow B) and the 2-segment gso (flow A). + var segCounts []int + for _, g := range w.gsoWrites { + segCounts = append(segCounts, len(g.pays)) + } + if !(segCounts[0] == 2 && segCounts[1] == 3) && !(segCounts[0] == 3 && segCounts[1] == 2) { + t.Errorf("unexpected segment counts: %v (want 2 and 3)", segCounts) + } +} + +// ECN test helpers and constants. + +const ( + tcpEce = 0x40 + tcpCwr = 0x80 + + // 2-bit IP-level ECN codepoints (lower 2 bits of IPv4 ToS / IPv6 TC). + ecnNotECT = 0x00 + ecnECT1 = 0x01 + ecnECT0 = 0x02 + ecnCE = 0x03 +) + +// buildTCPv4WithToS is buildTCPv4 with caller-specified IPv4 ToS so tests can +// drive DSCP and ECN bits. +func buildTCPv4WithToS(tos byte, seq uint32, flags byte, payload []byte) []byte { + pkt := buildTCPv4(seq, flags, payload) + pkt[1] = tos + return pkt +} + +// buildTCPv6 mirrors buildTCPv4 for IPv6. tcLow is the low 4 bits of Traffic +// Class, which carries the ECN codepoint (mask 0x03) and the bottom 2 DSCP +// bits — enough to drive the ECN paths under test. +func buildTCPv6(tcLow byte, seq uint32, flags byte, payload []byte) []byte { + const ipHdrLen = 40 + const tcpHdrLen = 20 + pkt := make([]byte, ipHdrLen+tcpHdrLen+len(payload)) + + pkt[0] = 0x60 // version=6, TC[7:4]=0 + pkt[1] = (tcLow & 0x0f) << 4 // TC[3:0] in high nibble; flow=0 + binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpHdrLen+len(payload))) + pkt[6] = ipProtoTCP + pkt[7] = 64 + copy(pkt[8:24], []byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}) + copy(pkt[24:40], []byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}) + + binary.BigEndian.PutUint16(pkt[40:42], 1000) + binary.BigEndian.PutUint16(pkt[42:44], 2000) + binary.BigEndian.PutUint32(pkt[44:48], seq) + binary.BigEndian.PutUint32(pkt[48:52], 12345) + pkt[52] = 0x50 + pkt[53] = flags + binary.BigEndian.PutUint16(pkt[54:56], 0xffff) + + copy(pkt[60:], payload) + return pkt +} + +// TestCoalescerCoalescesEceFlow confirms that ECN-Echo-marked ACKs (an +// ECN-aware flow under congestion) keep getting coalesced into a TSO +// superpacket instead of falling out to verbatim, and that the seed +// retains ECE on the wire. +func TestCoalescerCoalescesEceFlow(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + flags := byte(tcpAck | tcpEce) + if err := c.Commit(buildTCPv4(1000, flags, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, flags, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + g := w.gsoWrites[0] + if len(g.pays) != 2 { + t.Errorf("pay count=%d want 2", len(g.pays)) + } + if seedFlags := g.hdr[20+13]; seedFlags&tcpEce == 0 { + t.Errorf("seed flags=0x%02x want ECE preserved", seedFlags) + } +} + +// TestCoalescerCwrSealsFlow confirms that a CWR-bearing segment in the +// middle of a flow goes to verbatim and seals the open slot, so a later +// in-flow segment seeds a new slot rather than extending the prior burst. +func TestCoalescerCwrSealsFlow(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck|tcpCwr, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // All three emissions are plain writes: the seed before CWR and the + // fresh seed after both stay single-segment, and the CWR packet itself + // is verbatim. Order: seed, CWR, reseed. + if len(w.writes) != 3 || len(w.gsoWrites) != 0 { + t.Fatalf("want 3 plain writes (seed, CWR, reseed), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if flags := w.writes[1][20+13]; flags&tcpCwr == 0 { + t.Errorf("middle write flags=0x%02x want CWR (verbatim in arrival order)", flags) + } +} + +// TestCoalescerEceMismatchReseeds confirms that toggling ECE mid-flow does +// not silently merge — receivers expect ECE either set on every segment of +// a CE-echoing window or none. +func TestCoalescerEceMismatchReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck|tcpEce, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Each seed stays single-segment and flushes as its own plain write. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if flags := w.writes[0][20+13]; flags&tcpEce == 0 { + t.Errorf("first write lost ECE: flags=0x%02x", flags) + } + if flags := w.writes[1][20+13]; flags&tcpEce != 0 { + t.Errorf("second write gained ECE: flags=0x%02x", flags) + } +} + +// TestCoalescerDifferingECNReseeds confirms that segments with differing IP +// ECN codepoints do NOT coalesce: headersMatch compares the full ToS byte, +// matching kernel GRO. Two ECT(0) segments merge into a superpacket; a CE +// stamp mid-run seals the ECT(0) chain and reseeds, and the trailing ECT(0) +// reseeds again — those reseeds stay single-segment and ship as plain +// writes of the original packets, each keeping its own codepoint. ORing +// the marks (the old buggy behavior) would have fabricated a false CE +// across the whole burst. +func TestCoalescerDifferingECNReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4WithToS(ecnECT0, 2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Router along the path stamped CE on this one. + if err := c.Commit(buildTCPv4WithToS(ecnCE, 3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4WithToS(ecnECT0, 4600, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // gso: the two ECT(0) segments merged; then plain CE alone; then plain + // trailing ECT(0) alone. Emitted in seq order. + if len(w.gsoWrites) != 1 || len(w.writes) != 2 { + t.Fatalf("want 1 gso (ECT0 pair) + 2 plain (ECN split), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } + if got, want := w.order, []string{"gso", "write", "write"}; !stringSliceEq(got, want) { + t.Fatalf("emission order=%v want %v", got, want) + } + g := w.gsoWrites[0] + if len(g.pays) != 2 { + t.Errorf("gso pay count=%d want 2", len(g.pays)) + } + if got := g.hdr[1] & 0x03; got != ecnECT0 { + t.Errorf("gso ECN=0x%02x want 0x%02x", got, ecnECT0) + } + wantECN := []byte{ecnCE, ecnECT0} + for i, wnt := range wantECN { + if got := w.writes[i][1] & 0x03; got != wnt { + t.Errorf("plain %d ECN=0x%02x want 0x%02x", i, got, wnt) + } + } +} + +// TestCoalescerECT0ThenECT1NoCE is the core regression for the ECN merge +// bug: ORing ECT(0)=0b10 with ECT(1)=0b01 fabricates CE=0b11. The two +// segments must land in separate emissions — both stay single-segment, so +// each ships as a plain write of its original bytes, preserving its own +// codepoint — and neither may end up CE-marked. +func TestCoalescerECT0ThenECT1NoCE(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4WithToS(ecnECT1, 2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("want 2 separate plain writes (ECT0 vs ECT1), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + wantECN := []byte{ecnECT0, ecnECT1} + for i, p := range w.writes { + if got := p[1] & 0x03; got != wantECN[i] { + t.Errorf("write %d ECN=0x%02x want 0x%02x", i, got, wantECN[i]) + } + if got := p[1] & 0x03; got == ecnCE { + t.Errorf("write %d fabricated CE from ECT merge", i) + } + } +} + +// TestCoalescerDscpMismatchReseeds confirms that a DSCP difference (same +// ECN) still splits — headersMatch compares the full ToS byte, so the upper +// six DSCP bits must match too. +func TestCoalescerDscpMismatchReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + // Same ECN (Not-ECT), different DSCP (0x10 vs 0x20 in upper 6 bits). + tosA := byte(0x10<<2) | ecnNotECT + tosB := byte(0x20<<2) | ecnNotECT + if err := c.Commit(buildTCPv4WithToS(tosA, 1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4WithToS(tosB, 2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Both seeds stay single-segment → two plain writes, no gso. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +// TestCoalescerIPv6CoalescesEceFlow is the IPv6 analogue of +// TestCoalescerCoalescesEceFlow. +func TestCoalescerIPv6CoalescesEceFlow(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + flags := byte(tcpAck | tcpEce) + if err := c.Commit(buildTCPv6(0, 1000, flags, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv6(0, 2200, flags, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites)) + } + g := w.gsoWrites[0] + if seedFlags := g.hdr[40+13]; seedFlags&tcpEce == 0 { + t.Errorf("seed flags=0x%02x want ECE preserved", seedFlags) + } +} + +// TestCoalescerPSHKeepsChainBoundary verifies that a PSH-sealed chain is +// not extended by a later seq-contiguous segment — PSH placement is part of +// the wire signal and growing the superpacket past it would shift the +// receiver's push boundary by an arbitrary number of segments. +func TestCoalescerPSHKeepsChainBoundary(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + // Seq 1000 (no PSH) + 2200 (PSH) → seal one slot with PSH set. + // Seq 3400 is contiguous to the sealed chain's nextSeq; without the + // seal check it would append in. + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // The PSH-sealed pair is a real superpacket; the fresh seed stays + // single-segment and flushes as a plain write. + if len(w.gsoWrites) != 1 || len(w.writes) != 1 { + t.Fatalf("want 1 gso (PSH-sealed pair) + 1 plain (fresh seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } +} + +// TestCoalescerSynSealsFlowChain confirms a non-admissible in-flow packet +// (SYN+ACK here) seals its flow's open chain and holds its emission +// position: data committed after it seeds a fresh slot and emits after it, +// never extending a chain created before it. +func TestCoalescerSynSealsFlowChain(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Discontiguous seq: evicts the 1000 slot and seeds its own. + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Non-coalesceable packet (SYN+ACK) seals the flow's open slot and + // becomes a verbatim slot in c.slots. + if err := c.Commit(buildTCPv4(9999, tcpSyn|tcpAck, pay)); err != nil { + t.Fatal(err) + } + // Post-SYN data: must emit after the SYN, in its own slot. + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // All four packets emit as plain writes in creation order: 1000 and + // 3400 are separate single-segment slots, the SYN is verbatim, and the + // post-SYN 2200 is a fresh single-segment slot after it. + if len(w.writes) != 4 || len(w.gsoWrites) != 0 { + t.Fatalf("want 4 plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + wantSeqs := []uint32{1000, 3400, 9999, 2200} + for i, want := range wantSeqs { + if seq := binary.BigEndian.Uint32(w.writes[i][24:28]); seq != want { + t.Errorf("write %d seq=%d want %d", i, seq, want) + } + } +} + +// TestCoalescerIPv6DifferingECNReseeds is the IPv6 analogue of +// TestCoalescerDifferingECNReseeds. ECN bits live in TC[1:0] = byte 1 mask +// 0x30, so ipHeadersMatch (comparing byte 1 fully) still splits them. +func TestCoalescerIPv6DifferingECNReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + // tcLow is the low 4 bits of TC; ECN occupies the bottom 2 of those. + if err := c.Commit(buildTCPv6(ecnECT0, 1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv6(ecnECT0, 2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv6(ecnCE, 3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv6(ecnECT0, 4600, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Like the v4 test: the ECT(0) pair merges into one superpacket; the CE + // and trailing ECT(0) reseeds stay single-segment and ship as plain + // writes, in seq order. + if len(w.gsoWrites) != 1 || len(w.writes) != 2 { + t.Fatalf("want 1 gso (ECT0 pair) + 2 plain (ECN split), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } + // Byte 1 high nibble holds TC[3:0]; ECN is the low 2 bits of that nibble, + // which appears in byte 1 mask 0x30 (>>4 to read the codepoint value). + g := w.gsoWrites[0] + if len(g.pays) != 2 { + t.Errorf("gso pay count=%d want 2", len(g.pays)) + } + if got := (g.hdr[1] >> 4) & 0x03; got != ecnECT0 { + t.Errorf("gso v6 ECN=0x%02x want 0x%02x", got, ecnECT0) + } + wantECN := []byte{ecnCE, ecnECT0} + for i, wnt := range wantECN { + if got := (w.writes[i][1] >> 4) & 0x03; got != wnt { + t.Errorf("plain %d v6 ECN=0x%02x want 0x%02x", i, got, wnt) + } + } +} + +// TestCoalescerNonAtomicSequentialIDsCoalesce: with DF clear, coalescing +// is allowed when the IPv4 IDs already run seed+1 per segment — kernel +// TSO's re-stamp then reproduces the originals exactly (the kernel GRO +// admission rule). +func TestCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + seq := uint32(1000) + for i := range 3 { + pkt := buildTCPv4(seq, tcpAck, pay) + setIPv4ID(pkt, uint16(700+i), false) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + seq += uint32(len(pay)) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 3 { + t.Fatalf("sequential-ID DF=0 chain must coalesce: gso=%d", len(w.gsoWrites)) + } + if id := binary.BigEndian.Uint16(w.gsoWrites[0].hdr[4:6]); id != 700 { + t.Errorf("superpacket seed ID=%d want 700", id) + } +} + +// TestCoalescerNonAtomicIDGapDoesNotCoalesce: with DF clear and an ID jump +// mid-flow, neither the append path nor the flush-time merge may combine +// the segments — TSO would re-stamp seed+n and rewrite the second +// packet's ID, which is meaningful on non-atomic datagrams. Each stays a +// single-segment slot and flushes as a plain write with its original ID. +func TestCoalescerNonAtomicIDGapDoesNotCoalesce(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + p1 := buildTCPv4(1000, tcpAck, pay) + setIPv4ID(p1, 700, false) + p2 := buildTCPv4(1000+uint32(len(pay)), tcpAck, pay) + setIPv4ID(p2, 900, false) + + if err := c.Commit(p1); err != nil { + t.Fatal(err) + } + if err := c.Commit(p2); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("ID gap on DF=0 must not coalesce (append or merge): writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + for i, want := range []uint16{700, 900} { + if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want { + t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want) + } + } +} + +// TestCoalescerAtomicRandomIDsCoalesce guards the other direction: DF set +// makes the datagram atomic (RFC 6864), so arbitrary IDs must not block +// coalescing. +func TestCoalescerAtomicRandomIDsCoalesce(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + p1 := buildTCPv4(1000, tcpAck, pay) + setIPv4ID(p1, 0x1234, true) + p2 := buildTCPv4(1000+uint32(len(pay)), tcpAck, pay) + setIPv4ID(p2, 0x0007, true) + + if err := c.Commit(p1); err != nil { + t.Fatal(err) + } + if err := c.Commit(p2); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 { + t.Fatalf("DF=1 chain with arbitrary IDs must coalesce: gso=%d", len(w.gsoWrites)) + } +} + +// TestCoalescerSeqWrapAroundAppends pins the serial-number arithmetic on the +// append path: a chain crossing the 2^32 seq wrap must keep extending when +// contiguous. +func TestCoalescerSeqWrapAroundAppends(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + + payA := bytes.Repeat([]byte{'A'}, 32) + payB := bytes.Repeat([]byte{'B'}, 32) + seqA := uint32(0xffffffe0) // 32 before the wrap: nextSeq lands exactly on 0 + + if err := c.Commit(buildTCPv4(seqA, tcpAck, payA)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(0, tcpAck, payB)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write across the wrap, got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + g := w.gsoWrites[0] + const ipHdrLen = 20 + if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != seqA { + t.Errorf("seed seq=%#x want %#x", seedSeq, seqA) + } + if len(g.pays) != 2 { + t.Fatalf("segs=%d want 2", len(g.pays)) + } + if !bytes.Equal(g.pays[0], payA) || !bytes.Equal(g.pays[1], payB) { + t.Errorf("payload order wrong across the wrap: got %q then %q", g.pays[0][:1], g.pays[1][:1]) + } +} + +// TestCoalescerUnparseableSealsAllChains: an unparseable packet's flow is +// unknowable, so it must close every open chain. Later data — even data +// seq-contiguous with a pre-existing chain — seeds a fresh slot and emits +// after the unparseable packet, exactly as transmitted. +func TestCoalescerUnparseableSealsAllChains(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestTCPCoalescer(t, w) + pay := make([]byte, 1200) + + if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { + t.Fatal(err) + } + // IHL=6 fakes IP options: the parse bails, flow key unknown. + opts := buildTCPv4(5000, tcpAck, make([]byte, 500)) + opts[0] = 0x46 + if err := c.Commit(opts); err != nil { + t.Fatal(err) + } + // Contiguous with the first chain (nextSeq 3400), but that chain is + // sealed now: must not append, must not emit before the unparseable. + if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write (pre-fragment pair), got %d", len(w.gsoWrites)) + } + if len(w.gsoWrites[0].pays) != 2 { + t.Fatalf("pre-fragment chain segs=%d want 2", len(w.gsoWrites[0].pays)) + } + if len(w.writes) != 2 { + t.Fatalf("want 2 plain writes (unparseable + post-fragment seed), got %d", len(w.writes)) + } + if w.writes[0][0] != 0x46 { + t.Errorf("first plain write must be the unparseable packet") + } + if seq := binary.BigEndian.Uint32(w.writes[1][24:28]); seq != 3400 { + t.Errorf("post-fragment data seq=%d want 3400", seq) + } + if len(w.order) != 3 || w.order[0] != "gso" || w.order[1] != "write" || w.order[2] != "write" { + t.Fatalf("emission order = %v, want [gso write write]", w.order) + } +} diff --git a/overlay/batch/tx_batch.go b/overlay/batch/tx_batch.go new file mode 100644 index 00000000..4f6f7da2 --- /dev/null +++ b/overlay/batch/tx_batch.go @@ -0,0 +1,59 @@ +package batch + +import "net/netip" + +const SendBatchCap = 128 + +// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush. +type batchWriter interface { + WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) +} + +// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch. +// One SendBatch is owned by each listenIn goroutine; no locking is needed. +// Slots are backed by an Arena (see its docs) +type SendBatch struct { + out batchWriter + bufs [][]byte + dsts []netip.AddrPort + arena *Arena +} + +// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots +func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch { + return &SendBatch{ + out: out, + bufs: make([][]byte, 0, batchCap), + dsts: make([]netip.AddrPort, 0, batchCap), + arena: NewArena(arenaSize), + } +} + +func (b *SendBatch) Reserve(sz int) []byte { + return b.arena.Reserve(sz) +} + +// Len reports how many packets are queued for the next Flush. Callers use +// it to flush incrementally once a full sendmmsg batch has accumulated, +// bounding how long the first packet of a large read batch waits. +func (b *SendBatch) Len() int { return len(b.bufs) } + +func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) { + b.bufs = append(b.bufs, pkt) + b.dsts = append(b.dsts, dst) +} + +// Flush writes every queued packet and reports how many actually went out. A short count means some destinations +// were undeliverable; the batch is drained either way. +func (b *SendBatch) Flush() (int, error) { + var err error + written := 0 + if len(b.bufs) > 0 { + written, err = b.out.WriteBatch(b.bufs, b.dsts) + } + clear(b.bufs) + b.bufs = b.bufs[:0] + b.dsts = b.dsts[:0] + b.arena.Reset() + return written, err +} diff --git a/overlay/batch/tx_batch_test.go b/overlay/batch/tx_batch_test.go new file mode 100644 index 00000000..9a2a75b4 --- /dev/null +++ b/overlay/batch/tx_batch_test.go @@ -0,0 +1,122 @@ +package batch + +import ( + "net/netip" + "testing" +) + +type fakeBatchWriter struct { + bufs [][]byte + addrs []netip.AddrPort +} + +func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + // Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch + // returns, so tests must capture data before that happens. + w.bufs = make([][]byte, len(bufs)) + for i, b := range bufs { + cp := make([]byte, len(b)) + copy(cp, b) + w.bufs[i] = cp + } + w.addrs = append(w.addrs[:0], addrs...) + return len(bufs), nil +} + +func TestSendBatchReserveCommitFlush(t *testing.T) { + fw := &fakeBatchWriter{} + b := NewSendBatch(fw, 4, 32) + + ap := netip.MustParseAddrPort("10.0.0.1:4242") + for i := 0; i < 4; i++ { + slot := b.Reserve(32) + if cap(slot) != 32 { + t.Fatalf("slot %d: cap=%d want 32", i, cap(slot)) + } + pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2)) + b.Commit(pkt, ap) + } + if _, err := b.Flush(); err != nil { + t.Fatalf("Flush: %v", err) + } + if len(fw.bufs) != 4 { + t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs)) + } + for i, buf := range fw.bufs { + if len(buf) != 3 || buf[0] != byte(i) { + t.Errorf("buf %d: %x", i, buf) + } + if fw.addrs[i] != ap { + t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap) + } + } + + // Flush again with nothing committed — should be a no-op. + fw.bufs = nil + if _, err := b.Flush(); err != nil { + t.Fatalf("empty Flush: %v", err) + } + if fw.bufs != nil { + t.Fatalf("empty Flush triggered WriteBatch") + } + + // Reuse after Flush. + slot := b.Reserve(32) + if cap(slot) != 32 { + t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot)) + } +} + +func TestSendBatchSlotsDoNotOverlap(t *testing.T) { + fw := &fakeBatchWriter{} + b := NewSendBatch(fw, 3, 8) + ap := netip.MustParseAddrPort("10.0.0.1:80") + + for i := 0; i < 3; i++ { + s := b.Reserve(8) + pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i)) + b.Commit(pkt, ap) + } + if _, err := b.Flush(); err != nil { + t.Fatalf("Flush: %v", err) + } + + for i, buf := range fw.bufs { + if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) { + t.Errorf("slot %d corrupted: %x", i, buf) + } + } +} + +func TestSendBatchGrowPreservesCommitted(t *testing.T) { + fw := &fakeBatchWriter{} + // Tiny initial backing forces a grow on the second Reserve. + b := NewSendBatch(fw, 1, 4) + ap := netip.MustParseAddrPort("10.0.0.1:80") + + s1 := b.Reserve(4) + pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44) + b.Commit(pkt1, ap) + + s2 := b.Reserve(8) // exceeds remaining cap, triggers grow + pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE) + b.Commit(pkt2, ap) + + // pkt1 must still be intact even though backing reallocated. + if pkt1[0] != 0x11 || pkt1[3] != 0x44 { + t.Fatalf("first packet corrupted by grow: %x", pkt1) + } + + if _, err := b.Flush(); err != nil { + t.Fatalf("Flush: %v", err) + } + if len(fw.bufs) != 2 { + t.Fatalf("got %d bufs want 2", len(fw.bufs)) + } + if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 { + t.Errorf("first packet on the wire: %x", fw.bufs[0]) + } + if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE { + t.Errorf("second packet on the wire: %x", fw.bufs[1]) + } +} diff --git a/overlay/batch/udp_coalesce.go b/overlay/batch/udp_coalesce.go new file mode 100644 index 00000000..851bb59b --- /dev/null +++ b/overlay/batch/udp_coalesce.go @@ -0,0 +1,345 @@ +package batch + +import ( + "bytes" + "encoding/binary" + "io" + + "github.com/slackhq/nebula/overlay/tio" +) + +// ipProtoUDP is the IANA protocol number for UDP. +const ipProtoUDP = 17 + +// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the +// kernel's gso_max_size; payloads beyond this are emitted as-is. +const udpCoalesceBufSize = 65535 + +// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO +// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that. +const udpCoalesceMaxSegs = 64 + +// udpSlot is one entry in the UDPCoalescer's ordered event queue. +type udpSlot struct { + verbatim bool + // rawPkt is borrowed: the whole packet for verbatim slots, the seed + // packet for coalesce slots. A coalesce slot that never grows past one + // segment is emitted from rawPkt so its original (already valid) L4 + // checksum ships DATA_VALID instead of making the kernel recompute it. + // A multi-segment slot's superpacket header is rawPkt's, patched in place at flush. + rawPkt []byte + + fk flowKey + hdrLen int + ipHdrLen int + isV6 bool + gsoSize int // per-segment UDP payload length + numSeg int + totalPay int + payIovs [][]byte +} + +// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple +// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter. +// Preserves the in-flow order of packets as they are Commit-ed +// +// Owns no locks; one coalescer per TUN write queue. +type UDPCoalescer struct { + w tio.GSOWriter + slots []*udpSlot + openSlots map[flowKey]*udpSlot + // lastSlot caches the most recently touched open slot; see the + // TCPCoalescer field of the same name. Single-flow QUIC bulk is the + // dominant USO workload, and multi-flow arrival comes in GRO runs, so + // the fk compare beats the map's 38-byte key hash on most packets. + // Kept in lockstep with openSlots: nil whenever the slot it pointed at + // is removed. + lastSlot *udpSlot + pool []*udpSlot +} + +func NewUDPCoalescer(w io.Writer) *UDPCoalescer { + gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP) + if !ok { + return nil + } + return &UDPCoalescer{ + w: gw, + slots: make([]*udpSlot, 0, initialSlots), + openSlots: make(map[flowKey]*udpSlot, initialSlots), + pool: make([]*udpSlot, 0, initialSlots), + } +} + +// parsedUDP holds the fields extracted from a single parse so later steps +// (admission, slot lookup, canAppend) don't re-walk the header. +type parsedUDP struct { + fk flowKey + ipHdrLen int + hdrLen int // ipHdrLen + 8 + payLen int +} + +// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is +// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on +// entry and is filled in place. Returns false for malformed input or any shape that must not +// coalesce (IPv4 options/fragmentation, IPv6 extension headers). +func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool { + trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen) + if !ok { + return false + } + return p.parseTail(trimmed, ipHdrLen) +} + +// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet; +// fk's addresses are already filled. +func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool { + if len(pkt) < ipHdrLen+8 { + return false + } + // UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload. + udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6])) + if udpLen < 8 || udpLen > len(pkt)-ipHdrLen { + return false + } + p.ipHdrLen = ipHdrLen + p.hdrLen = ipHdrLen + 8 + p.payLen = udpLen - 8 + p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2]) + p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4]) + return true +} + +// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips +// hashing the 38-byte key when no chains are open. +func (c *UDPCoalescer) sealFlow(fk flowKey) { + if len(c.openSlots) == 0 { + return + } + if last := c.lastSlot; last != nil && last.fk == fk { + c.lastSlot = nil + } + delete(c.openSlots, fk) +} + +// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot +// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown — +// and rides the lane as an in-lane verbatim, still in transmission order. +func (c *UDPCoalescer) commitStaged(sp stagedPacket) error { + if sp.fragAny { + c.sealAllOpen() + c.addVerbatim(sp.pkt) + return nil + } + var info parsedUDP + if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) { + c.sealAllOpen() + c.addVerbatim(sp.pkt) + return nil + } + return c.commitParsed(sp.pkt, &info) +} + +// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a +// valid parse so the header is not re-walked here. +func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error { + // A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be + // coalesced. + if info.payLen == 0 { + c.sealFlow(info.fk) + c.addVerbatim(pkt) + return nil + } + // Cached-slot fast path; see the TCPCoalescer equivalent. + var open *udpSlot + if last := c.lastSlot; last != nil && last.fk == info.fk { + open = last + } else { + open = c.openSlots[info.fk] + } + if open != nil { + if c.canAppend(open, pkt, info) { + if c.appendPayload(open, pkt, info) { + // Chain closed (short segment): stop extending it. + c.sealFlow(info.fk) + } else { + c.lastSlot = open + } + return nil + } + // Can't extend: evict it from openSlots and fall through to seed a + // fresh slot. + c.sealFlow(info.fk) + } + c.seed(pkt, info) + return nil +} + +func (c *UDPCoalescer) Flush() error { + var first error + for _, s := range c.slots { + var err error + if s.verbatim || s.numSeg == 1 { + // A slot that never grew is byte-identical to the packet it was + // seeded from; ship the original so its valid checksum rides the + // DATA_VALID path instead of paying a kernel software csum. + _, err = c.w.Write(s.rawPkt) + } else { + err = c.flushSlot(s) + } + if err != nil && first == nil { + first = err + } + c.release(s) + } + clear(c.slots) + c.slots = c.slots[:0] + clear(c.openSlots) + c.lastSlot = nil + return first +} + +// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is +// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet. +func (c *UDPCoalescer) sealAllOpen() { + clear(c.openSlots) + c.lastSlot = nil +} + +func (c *UDPCoalescer) addVerbatim(pkt []byte) { + s := c.take() + s.verbatim = true + s.rawPkt = pkt + c.slots = append(c.slots, s) +} + +func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) { + if info.hdrLen+info.payLen > udpCoalesceBufSize { + // Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can + // be open here (commitParsed evicts before seeding), so sealFlow is defense in depth + // against a stale cache entry absorbing later data. + c.sealFlow(info.fk) + c.addVerbatim(pkt) + return + } + s := c.take() + s.verbatim = false + // rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is + // the superpacket header flushSlot patches in place. + s.rawPkt = pkt + s.hdrLen = info.hdrLen + s.ipHdrLen = info.ipHdrLen + s.isV6 = info.fk.isV6 + s.fk = info.fk + s.gsoSize = info.payLen + s.numSeg = 1 + s.totalPay = info.payLen + s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen]) + c.slots = append(c.slots, s) + c.openSlots[info.fk] = s + c.lastSlot = s +} + +// canAppend reports whether info's packet extends the slot's seed. +// Kernel UDP-GSO requires every segment except possibly the last to be +// exactly gsoSize, and the last may be shorter (≤ gsoSize). +func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool { + if info.hdrLen != s.hdrLen { + return false + } + if s.numSeg >= udpCoalesceMaxSegs { + return false + } + if info.payLen > s.gsoSize { + return false + } + if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize { + return false + } + // Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches + // here; closing removes the slot from openSlots, the only path in. + if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) { + return false + } + if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) { + return false + } + return true +} + +// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel +// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be +// the final one. The caller must deregister a closed slot from openSlots. +func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool { + s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen]) + s.numSeg++ + s.totalPay += info.payLen + return info.payLen < s.gsoSize +} + +func (c *UDPCoalescer) take() *udpSlot { + if n := len(c.pool); n > 0 { + s := c.pool[n-1] + c.pool[n-1] = nil + c.pool = c.pool[:n-1] + return s + } + return &udpSlot{} +} + +func (c *UDPCoalescer) release(s *udpSlot) { + // Reset every field, identity ones included; see TCPCoalescer.release. + clear(s.payIovs) + *s = udpSlot{payIovs: s.payIovs[:0]} + c.pool = append(c.pool, s) +} + +// flushSlot patches the IP header total length / IPv6 payload length and +// the UDP length to the *total* across all coalesced segments, then seeds +// the UDP checksum field with the pseudo-header partial (single-fold, not +// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the +// slot is released right after, so nothing re-reads the patched header. +func (c *UDPCoalescer) flushSlot(s *udpSlot) error { + hdr := s.rawPkt[:s.hdrLen] + total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes + l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads) + + if s.isV6 { + binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len)) + } else { + binary.BigEndian.PutUint16(hdr[2:4], uint16(total)) + hdr[10] = 0 + hdr[11] = 0 + binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen])) + } + + // UDP length field (offset 4 inside the UDP header) = total UDP size. + binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len)) + + var psum uint32 + if s.isV6 { + psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len) + } else { + psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len) + } + udpCsumOff := s.ipHdrLen + 6 + binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum)) + + return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP) +} + +// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on +// every field that must be identical across coalesced segments +func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool { + if len(a) != len(b) { + return false + } + if !ipHeadersMatch(a, b, isV6) { + return false + } + // UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]: + // length varies (we rewrite at flush) and the checksum will be redone. + udp := ipHdrLen + return bytes.Equal(a[udp:udp+4], b[udp:udp+4]) +} diff --git a/overlay/batch/udp_coalesce_bench_test.go b/overlay/batch/udp_coalesce_bench_test.go new file mode 100644 index 00000000..53426b64 --- /dev/null +++ b/overlay/batch/udp_coalesce_bench_test.go @@ -0,0 +1,72 @@ +package batch + +import ( + "testing" +) + +// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the +// steady state for single-flow QUIC bulk, the workload USO exists for. +func buildUDPv4BulkFlow(n, payloadLen int) [][]byte { + pay := make([]byte, payloadLen) + pkts := make([][]byte, n) + for i := range pkts { + pkts[i] = buildUDPv4(40000, 443, pay) + } + return pkts +} + +// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow +// datagrams arriving in GRO-burst runs of runLen per flow. +func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte { + pay := make([]byte, payloadLen) + pkts := make([][]byte, 0, nFlows*perFlow) + for done := 0; done < perFlow; done += runLen { + for f := range nFlows { + sport := uint16(40000 + f) + for range runLen { + pkts = append(pkts, buildUDPv4(sport, 443, pay)) + } + } + } + return pkts +} + +// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a +// time, flushing between batches, and reports per-packet cost. +func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) { + b.Helper() + c := newTestUDPCoalescer(b, nopTunWriter{}) + b.ReportAllocs() + b.SetBytes(int64(len(pkts[0]))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + pkt := pkts[i%len(pkts)] + if err := c.Commit(pkt); err != nil { + b.Fatal(err) + } + if (i+1)%batchSize == 0 { + if err := c.Flush(); err != nil { + b.Fatal(err) + } + } + } + _ = c.Flush() +} + +// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state. +func BenchmarkUDPCommitSingleFlow(b *testing.B) { + pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200) + runUDPCommitBench(b, pkts, udpCoalesceMaxSegs) +} + +// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin. +func BenchmarkUDPCommitInterleaved4(b *testing.B) { + pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200) + runUDPCommitBench(b, pkts, len(pkts)) +} + +// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16. +func BenchmarkUDPCommitRunInterleaved4(b *testing.B) { + pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200) + runUDPCommitBench(b, pkts, len(pkts)) +} diff --git a/overlay/batch/udp_coalesce_test.go b/overlay/batch/udp_coalesce_test.go new file mode 100644 index 00000000..c2be8e49 --- /dev/null +++ b/overlay/batch/udp_coalesce_test.go @@ -0,0 +1,536 @@ +package batch + +import ( + "bytes" + "encoding/binary" + "io" + "testing" +) + +// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports. +func buildUDPv4(sport, dport uint16, payload []byte) []byte { + const ipHdrLen = 20 + const udpHdrLen = 8 + total := ipHdrLen + udpHdrLen + len(payload) + pkt := make([]byte, total) + + pkt[0] = 0x45 + pkt[1] = 0x00 + binary.BigEndian.PutUint16(pkt[2:4], uint16(total)) + binary.BigEndian.PutUint16(pkt[4:6], 0) + binary.BigEndian.PutUint16(pkt[6:8], 0x4000) + pkt[8] = 64 + pkt[9] = ipProtoUDP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + + binary.BigEndian.PutUint16(pkt[20:22], sport) + binary.BigEndian.PutUint16(pkt[22:24], dport) + binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload))) + binary.BigEndian.PutUint16(pkt[26:28], 0) + + copy(pkt[28:], payload) + return pkt +} + +// buildUDPv6 builds a minimal IPv6+UDP packet. +func buildUDPv6(sport, dport uint16, payload []byte) []byte { + const ipHdrLen = 40 + const udpHdrLen = 8 + total := ipHdrLen + udpHdrLen + len(payload) + pkt := make([]byte, total) + + pkt[0] = 0x60 + binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload))) + pkt[6] = ipProtoUDP + pkt[7] = 64 + pkt[8] = 0xfe + pkt[9] = 0x80 + pkt[23] = 1 + pkt[24] = 0xfe + pkt[25] = 0x80 + pkt[39] = 2 + + binary.BigEndian.PutUint16(pkt[40:42], sport) + binary.BigEndian.PutUint16(pkt[42:44], dport) + binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload))) + binary.BigEndian.PutUint16(pkt[46:48], 0) + + copy(pkt[48:], payload) + return pkt +} + +// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't +// do USO. See newTestTCPCoalescer. +func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer { + tb.Helper() + c := NewUDPCoalescer(w) + if c == nil { + tb.Fatal("NewUDPCoalescer: writer does not support USO") + } + return c +} + +// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition: +// no USO, no coalescer. +func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) { + if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil { + t.Fatalf("want nil for a non-USO writer, got %v", c) + } + if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil { + t.Fatalf("want nil for a plain writer, got %v", c) + } +} + +func TestUDPCoalescerNonUDPPassthrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + // ICMP packet + pkt := make([]byte, 28) + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], 28) + pkt[9] = 1 + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pkt := buildUDPv4(1000, 53, make([]byte, 800)) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // A slot that never grew past one datagram flushes as a plain Write of + // the original packet bytes: the original (already valid) checksum + // ships via the DATA_VALID path, so the kernel does no csum work. + // WriteGSO is reserved for slots that actually coalesced (>=2 segs). + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if !bytes.Equal(w.writes[0], pkt) { + t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt)) + } +} + +func TestUDPCoalescerCoalescesEqualSized(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 1200) + for i := 0; i < 3; i++ { + if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { + t.Fatal(err) + } + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) + } + g := w.gsoWrites[0] + if g.gsoSize != 1200 { + t.Errorf("gsoSize=%d want 1200", g.gsoSize) + } + if len(g.pays) != 3 { + t.Errorf("pay count=%d want 3", len(g.pays)) + } + if g.csumStart != 20 { + t.Errorf("csumStart=%d want 20", g.csumStart) + } + // IP totalLen and UDP length must be the TOTAL across all segments — + // the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment + // value would silently drop everything but the first segment. Total = + // IP(20) + UDP(8) + 3*1200 = 3628. + gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4]) + if gotTotalLen != 3628 { + t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen) + } + gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6]) + if gotUDPLen != 8+3*1200 { + t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200) + } +} + +// Last segment may be shorter, sealing the chain. +func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + full := make([]byte, 1200) + tail := make([]byte, 600) + if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil { + t.Fatal(err) + } + // A 4th packet, even same-sized, must NOT join — chain is sealed. + if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // The sealed 3-datagram chain is a real superpacket; the re-seed stays + // single-segment and flushes as a plain write of the original packet. + if len(w.gsoWrites) != 1 || len(w.writes) != 1 { + t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } + if len(w.gsoWrites[0].pays) != 3 { + t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays)) + } + if got, want := len(w.writes[0]), 20+8+1200; got != want { + t.Errorf("re-seed plain write len=%d want %d", got, want) + } +} + +// A larger-than-gsoSize packet cannot extend the slot — it reseeds. +func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Both seeds stay single-segment → two plain writes in arrival order. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} { + if len(w.writes[i]) != want { + t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want) + } + } +} + +// Different 5-tuples must not coalesce. +func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 800) + if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Two flows × 2 datagrams each = 2 superpackets of 2 segments. + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites)) + } + for i, g := range w.gsoWrites { + if len(g.pays) != 2 { + t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays)) + } + } +} + +// Caps at udpCoalesceMaxSegs. +func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 100) + for i := 0; i < udpCoalesceMaxSegs+5; i++ { + if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { + t.Fatal(err) + } + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // First superpacket holds udpCoalesceMaxSegs segments; the spillover + // reseeds a new one. + if len(w.gsoWrites) != 2 { + t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites)) + } + if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs { + t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs) + } + if len(w.gsoWrites[1].pays) != 5 { + t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays)) + } +} + +// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares +// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run +// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram +// reseeds again. All three stay single-segment, so each ships as a plain +// write of its original bytes, keeping its own codepoint. +func TestUDPCoalescerDifferingECNReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 800) + pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT) + pkt1 := buildUDPv4(1000, 53, pay) + pkt1[1] = 0x03 // CE + pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again + for _, p := range [][]byte{pkt0, pkt1, pkt2} { + if err := c.Commit(p); err != nil { + t.Fatal(err) + } + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 3 || len(w.gsoWrites) != 0 { + t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + wantECN := []byte{0x00, 0x03, 0x00} + for i, p := range w.writes { + if got := p[1] & 0x03; got != wantECN[i] { + t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i]) + } + } +} + +// IPv6 path: same flow, equal-sized → coalesced. +func TestUDPCoalescerIPv6Coalesces(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 1200) + for i := 0; i < 3; i++ { + if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil { + t.Fatal(err) + } + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 { + t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites)) + } + g := w.gsoWrites[0] + if !g.isV6 { + t.Errorf("expected v6 write") + } + if g.csumStart != 40 { + t.Errorf("csumStart=%d want 40", g.csumStart) + } + // IPv6 payload_len and UDP length must be TOTAL — kernel's + // ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 + + // 3*1200 = 3608. + gotPlen := binary.BigEndian.Uint16(g.hdr[4:6]) + if gotPlen != 8+3*1200 { + t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200) + } + gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6]) + if gotUDPLen != 8+3*1200 { + t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200) + } +} + +// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte. +func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 800) + pkt0 := buildUDPv4(1000, 53, pay) + pkt1 := buildUDPv4(1000, 53, pay) + pkt1[1] = 0xb8 // EF DSCP, ECN=0 + if err := c.Commit(pkt0); err != nil { + t.Fatal(err) + } + if err := c.Commit(pkt1); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // Both seeds stay single-segment → two plain writes, no gso. + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +// Fragmented IPv4 must not be coalesced. +func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pkt := buildUDPv4(1000, 53, make([]byte, 200)) + binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1 + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +// A zero-length UDP datagram (UDP length == 8, no payload) is legal and +// must be delivered as a plain single datagram — never coalesced. Seeding +// it into a GSO slot stores an empty payload iovec that panics WriteGSO +// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it +// reach the GSO path. Regression: must not panic and must be written. +func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if len(w.writes[0]) != len(pkt) { + t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt)) + } +} + +// IPv6 zero-length UDP datagram: same verbatim contract as v4. +func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + if len(w.writes[0]) != len(pkt) { + t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt)) + } +} + +// A zero-length datagram arriving mid-flow must seal the open chain so the +// datagram after it seeds a fresh superpacket *after* the empty one on the +// wire — per-flow arrival order (full, empty, full) must be preserved. +func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + full := make([]byte, 800) + if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length + t.Fatal(err) + } + if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + // The empty datagram sealed the first slot, so the trailing full packet + // can't join it. All three emit as plain writes (the two full datagrams + // stayed single-segment; the empty one is verbatim) in per-flow + // arrival order: full, empty, full. + if len(w.writes) != 3 || len(w.gsoWrites) != 0 { + t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) + } + for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} { + if len(w.writes[i]) != want { + t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want) + } + } +} + +// IPv4 with options is not admissible (we require IHL=5). +func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pkt := buildUDPv4(1000, 53, make([]byte, 200)) + pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 1 || len(w.gsoWrites) != 0 { + t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } +} + +// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF +// clear is fine as long as the IDs already run seed+1 per datagram, so +// kernel USO's re-stamp reproduces them. +func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 1200) + + for i := range 2 { + pkt := buildUDPv4(40000, 443, pay) + setIPv4ID(pkt, uint16(40+i), false) + if err := c.Commit(pkt); err != nil { + t.Fatal(err) + } + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 { + t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites)) + } +} + +// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks +// the chain; each datagram stays a single-segment slot and flushes as a +// plain write that keeps its own (meaningful) ID. +func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) { + w := &fakeTunWriter{gsoEnabled: true} + c := newTestUDPCoalescer(t, w) + pay := make([]byte, 1200) + + p1 := buildUDPv4(40000, 443, pay) + setIPv4ID(p1, 40, false) + p2 := buildUDPv4(40000, 443, pay) + setIPv4ID(p2, 50, false) + + if err := c.Commit(p1); err != nil { + t.Fatal(err) + } + if err := c.Commit(p2); err != nil { + t.Fatal(err) + } + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(w.writes) != 2 || len(w.gsoWrites) != 0 { + t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) + } + for i, want := range []uint16{40, 50} { + if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want { + t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want) + } + } +} diff --git a/overlay/checksum/checksum_amd64.go b/overlay/checksum/checksum_amd64.go new file mode 100644 index 00000000..d504e73e --- /dev/null +++ b/overlay/checksum/checksum_amd64.go @@ -0,0 +1,23 @@ +package checksum + +import ( + "golang.org/x/sys/cpu" + gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum" +) + +//go:noescape +func checksumAVX2(buf []byte, initial uint16) uint16 + +var hasAVX2 = cpu.X86.HasAVX2 + +// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with +// initial. It is a drop-in replacement for gvisor's checksum.Checksum that +// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it, +// falling back to gvisor's pure-Go implementation otherwise. The result +// matches gvisor's bit-for-bit for any buffer length and initial seed. +func Checksum(buf []byte, initial uint16) uint16 { + if hasAVX2 { + return checksumAVX2(buf, initial) + } + return gvisorchecksum.Checksum(buf, initial) +} diff --git a/overlay/checksum/checksum_amd64.s b/overlay/checksum/checksum_amd64.s new file mode 100644 index 00000000..5ee864d6 --- /dev/null +++ b/overlay/checksum/checksum_amd64.s @@ -0,0 +1,157 @@ +#include "textflag.h" + +// func checksumAVX2(buf []byte, initial uint16) uint16 +// +// Computes the RFC 1071 ones-complement sum of buf, seeded with initial. +// +// Algorithm: sum the buffer treating it as a stream of uint32s in machine +// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits +// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of +// headroom which is far more than the 16 KB/64 KB max practical inputs). +// At the end we fold to 16 bits and byte-swap once to recover the on-wire +// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this +// equivalent to summing as 16-bit big-endian words. +// +// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel +// partial sums. The main loop loads 64 bytes per iter as four 16-byte +// chunks, zero-extending each chunk's four uint32s into a ymm via +// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per +// chunk to break the dep chain. After the vector loop the lane sums are +// horizontally reduced and merged with a scalar accumulator that handles +// the trailing 0..63 bytes plus the (byte-swapped) initial seed. +TEXT ·checksumAVX2(SB), NOSPLIT, $0-34 + MOVQ buf_base+0(FP), SI + MOVQ buf_len+8(FP), CX + MOVWQZX initial+24(FP), AX + + // Pre-byteswap initial into the LE-summing space so it merges directly + // with the rest of the accumulator. The final fold's bswap16 will undo + // this and convert the whole result back to BE. + XCHGB AH, AL + + CMPQ CX, $32 + JLT scalar_tail + + VPXOR Y4, Y4, Y4 + VPXOR Y5, Y5, Y5 + VPXOR Y6, Y6, Y6 + VPXOR Y7, Y7, Y7 + + CMPQ CX, $64 + JLT loop32 + +loop64: + VPMOVZXDQ (SI), Y0 + VPMOVZXDQ 16(SI), Y1 + VPMOVZXDQ 32(SI), Y2 + VPMOVZXDQ 48(SI), Y3 + VPADDQ Y0, Y4, Y4 + VPADDQ Y1, Y5, Y5 + VPADDQ Y2, Y6, Y6 + VPADDQ Y3, Y7, Y7 + ADDQ $64, SI + SUBQ $64, CX + CMPQ CX, $64 + JGE loop64 + +loop32: + CMPQ CX, $32 + JLT reduce_vec + VPMOVZXDQ (SI), Y0 + VPMOVZXDQ 16(SI), Y1 + VPADDQ Y0, Y4, Y4 + VPADDQ Y1, Y5, Y5 + ADDQ $32, SI + SUBQ $32, CX + JMP loop32 + +reduce_vec: + // Combine the four ymm accumulators into Y4. + VPADDQ Y5, Y4, Y4 + VPADDQ Y7, Y6, Y6 + VPADDQ Y6, Y4, Y4 + + // Horizontally reduce Y4's four uint64 lanes to a single scalar. + VEXTRACTI128 $1, Y4, X5 + VPADDQ X5, X4, X4 + VPSHUFD $0x4e, X4, X5 + VPADDQ X5, X4, X4 + VMOVQ X4, R8 + VZEROUPPER + + ADDQ R8, AX + ADCQ $0, AX + +scalar_tail: + // Handle remaining 0..63 bytes (or the entire buffer if it was < 32). + CMPQ CX, $8 + JLT tail4 + +loop8: + ADDQ (SI), AX + ADCQ $0, AX + ADDQ $8, SI + SUBQ $8, CX + CMPQ CX, $8 + JGE loop8 + +tail4: + CMPQ CX, $4 + JLT tail2 + MOVL (SI), R8 + ADDQ R8, AX + ADCQ $0, AX + ADDQ $4, SI + SUBQ $4, CX + +tail2: + CMPQ CX, $2 + JLT tail1 + MOVWQZX (SI), R8 + ADDQ R8, AX + ADCQ $0, AX + ADDQ $2, SI + SUBQ $2, CX + +tail1: + TESTQ CX, CX + JZ fold + MOVBQZX (SI), R8 + ADDQ R8, AX + ADCQ $0, AX + +fold: + // Fold the 64-bit accumulator to 16 bits via four rounds, mirroring + // gvisor's reduce(). Each pair (split, add) halves the live width; + // the truncation steps absorb the single bit that may be left over + // after each add so the next round's bound holds. + + // 64 → 33 bits. + MOVQ AX, R8 + SHRQ $32, R8 + MOVL AX, AX + ADDQ R8, AX + + // 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF. + MOVQ AX, R8 + SHRQ $32, R8 + ADDQ R8, AX + MOVL AX, AX + + // 32 → 17 bits. + MOVQ AX, R8 + SHRQ $16, R8 + MOVWQZX AX, AX + ADDQ R8, AX + + // 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16. + MOVQ AX, R8 + SHRQ $16, R8 + ADDQ R8, AX + + // AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip + // to big-endian to match the gvisor API contract. + XCHGB AH, AL + + MOVW AX, ret+32(FP) + RET diff --git a/overlay/checksum/checksum_arm64.go b/overlay/checksum/checksum_arm64.go new file mode 100644 index 00000000..561ba712 --- /dev/null +++ b/overlay/checksum/checksum_arm64.go @@ -0,0 +1,12 @@ +package checksum + +//go:noescape +func checksumNEON(buf []byte, initial uint16) uint16 + +// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with +// initial. It is a drop-in replacement for gvisor's checksum.Checksum +// that dispatches to a hand-written NEON routine. NEON is mandatory in +// armv8 so no feature check is needed. +func Checksum(buf []byte, initial uint16) uint16 { + return checksumNEON(buf, initial) +} diff --git a/overlay/checksum/checksum_arm64.s b/overlay/checksum/checksum_arm64.s new file mode 100644 index 00000000..11499820 --- /dev/null +++ b/overlay/checksum/checksum_arm64.s @@ -0,0 +1,143 @@ +#include "textflag.h" + +// func checksumNEON(buf []byte, initial uint16) uint16 +// +// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as +// a stream of uint32s in machine (little-endian) byte order, accumulating +// into 64-bit lanes that have ample carry headroom; fold and byte-swap once +// at the very end to recover the on-wire (big-endian) result. +// +// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs). +// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to +// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same +// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get +// updated twice per iter to break the dep chain. Tail bytes go through a +// scalar ADCS chain seeded with the byte-swapped initial. +TEXT ·checksumNEON(SB), NOSPLIT, $0-34 + MOVD buf_base+0(FP), R0 + MOVD buf_len+8(FP), R1 + MOVHU initial+24(FP), R2 + + // Pre-byteswap initial into the LE-summing space so it merges directly + // with the rest of the accumulator. + REV16W R2, R2 + + MOVD ZR, R3 // scalar accumulator + + CMP $32, R1 + BLT scalar_tail + + VEOR V8.B16, V8.B16, V8.B16 + VEOR V9.B16, V9.B16, V9.B16 + VEOR V10.B16, V10.B16, V10.B16 + VEOR V11.B16, V11.B16, V11.B16 + + CMP $64, R1 + BLT loop16_init + +loop64: + VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16] + VUADDW V0.S2, V8.D2, V8.D2 + VUADDW2 V0.S4, V9.D2, V9.D2 + VUADDW V1.S2, V10.D2, V10.D2 + VUADDW2 V1.S4, V11.D2, V11.D2 + VUADDW V2.S2, V8.D2, V8.D2 + VUADDW2 V2.S4, V9.D2, V9.D2 + VUADDW V3.S2, V10.D2, V10.D2 + VUADDW2 V3.S4, V11.D2, V11.D2 + SUB $64, R1, R1 + CMP $64, R1 + BGE loop64 + +loop16_init: + CMP $16, R1 + BLT reduce_vec + +loop16: + VLD1.P 16(R0), [V0.B16] + VUADDW V0.S2, V8.D2, V8.D2 + VUADDW2 V0.S4, V9.D2, V9.D2 + SUB $16, R1, R1 + CMP $16, R1 + BGE loop16 + +reduce_vec: + // Combine the four accumulators into V8. + VADD V9.D2, V8.D2, V8.D2 + VADD V11.D2, V10.D2, V10.D2 + VADD V10.D2, V8.D2, V8.D2 + + // Horizontal-add the two lanes of V8.D2 into a single uint64. + VADDP V8.D2, V8.D2, V8.D2 + VMOV V8.D[0], R8 + + ADDS R8, R3, R3 + ADC ZR, R3, R3 + +scalar_tail: + CMP $8, R1 + BLT tail4 + +loop8: + MOVD.P 8(R0), R8 + ADDS R8, R3, R3 + ADC ZR, R3, R3 + SUB $8, R1, R1 + CMP $8, R1 + BGE loop8 + +tail4: + CMP $4, R1 + BLT tail2 + MOVWU.P 4(R0), R8 + ADDS R8, R3, R3 + ADC ZR, R3, R3 + SUB $4, R1, R1 + +tail2: + CMP $2, R1 + BLT tail1 + MOVHU.P 2(R0), R8 + ADDS R8, R3, R3 + ADC ZR, R3, R3 + SUB $2, R1, R1 + +tail1: + CBZ R1, fold + MOVBU (R0), R8 + ADDS R8, R3, R3 + ADC ZR, R3, R3 + +fold: + // Merge the byte-swapped initial into our LE-form accumulator. + ADDS R2, R3, R3 + ADC ZR, R3, R3 + + // 64 → 33 bits. + LSR $32, R3, R8 + AND $0xffffffff, R3, R3 + ADD R8, R3, R3 + + // 33 → 32 (truncate after adding bit 32 back). + LSR $32, R3, R8 + ADD R8, R3, R3 + AND $0xffffffff, R3, R3 + + // 32 → 17. + LSR $16, R3, R8 + AND $0xffff, R3, R3 + ADD R8, R3, R3 + + // 17 → 16 (truncation absorbs bit 16 below). + LSR $16, R3, R8 + ADD R8, R3, R3 + + // AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip + // to big-endian to match the gvisor API contract. REV16W swaps bytes + // within each 16-bit halfword of the low 32 bits, so it acts as a + // 16-bit byte-swap on the live low 16. + REV16W R3, R3 + AND $0xffff, R3, R3 + + MOVH R3, ret+32(FP) + RET diff --git a/overlay/checksum/checksum_fallback.go b/overlay/checksum/checksum_fallback.go new file mode 100644 index 00000000..89ac90a5 --- /dev/null +++ b/overlay/checksum/checksum_fallback.go @@ -0,0 +1,10 @@ +//go:build !amd64 && !arm64 + +package checksum + +import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum" + +// Checksum delegates to gvisor on architectures without a hand-written body. +func Checksum(buf []byte, initial uint16) uint16 { + return gvisorchecksum.Checksum(buf, initial) +} diff --git a/overlay/checksum/checksum_test.go b/overlay/checksum/checksum_test.go new file mode 100644 index 00000000..c3c39b20 --- /dev/null +++ b/overlay/checksum/checksum_test.go @@ -0,0 +1,232 @@ +package checksum + +import ( + "fmt" + "math/rand/v2" + "testing" + + gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum" +) + +// archImpl names one checksum function under test. The per-arch +// export_*_test.go files enumerate the hand-written implementations so the +// suite compares each one against gvisor directly, regardless of which one +// the public Checksum dispatches to on the running CPU. Testing only the +// dispatcher was tautological wherever it resolved to the gvisor fallback +// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself, +// assembly untested, suite green. +type archImpl struct { + name string + fn func([]byte, uint16) uint16 + available bool +} + +// implsUnderTest is the public dispatcher plus every arch implementation. +func implsUnderTest() []archImpl { + return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...) +} + +// requireAvailable skips loudly when the running CPU can't execute an +// implementation — visible in test output, unlike the old silent tautology. +func requireAvailable(t *testing.T, impl archImpl) { + t.Helper() + if !impl.available { + t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name) + } +} + +// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial +// seeds and a handful of starting alignments, asserting that each local +// implementation matches gvisor's reference bit-for-bit. +func TestChecksumMatchesGvisor(t *testing.T) { + for _, impl := range implsUnderTest() { + t.Run(impl.name, func(t *testing.T) { + requireAvailable(t, impl) + rng := rand.New(rand.NewPCG(1, 2)) + const padFront = 16 + + // Random pool large enough for the longest case + alignment slop. + pool := make([]byte, 4096+padFront) + for i := range pool { + pool[i] = byte(rng.Uint32()) + } + + seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc} + offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16} + + for length := 0; length <= 4096; length++ { + for _, seed := range seeds { + for _, off := range offsets { + if off+length > len(pool) { + continue + } + buf := pool[off : off+length] + want := gvisorchecksum.Checksum(buf, seed) + got := impl.fn(buf, seed) + if got != want { + t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x", + length, off, seed, got, want) + } + } + } + } + }) + } +} + +// TestChecksumPatternedBuffers exercises specific byte patterns that have +// historically tripped up checksum implementations: all-zero, all-0xff, +// alternating, and ascending sequences. +func TestChecksumPatternedBuffers(t *testing.T) { + for _, impl := range implsUnderTest() { + t.Run(impl.name, func(t *testing.T) { + requireAvailable(t, impl) + for length := 0; length <= 256; length++ { + patterns := map[string][]byte{ + "zeros": make([]byte, length), + "ones": bytes(length, 0xff), + "alternating": pattern(length, []byte{0xa5, 0x5a}), + "ascending": ascending(length), + } + for name, buf := range patterns { + for _, seed := range []uint16{0, 0xffff, 0x8000} { + want := gvisorchecksum.Checksum(buf, seed) + got := impl.fn(buf, seed) + if got != want { + t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x", + name, length, seed, got, want) + } + } + } + } + }) + } +} + +func bytes(n int, v byte) []byte { + b := make([]byte, n) + for i := range b { + b[i] = v + } + return b +} + +func pattern(n int, p []byte) []byte { + b := make([]byte, n) + for i := range b { + b[i] = p[i%len(p)] + } + return b +} + +func ascending(n int) []byte { + b := make([]byte, n) + for i := range b { + b[i] = byte(i) + } + return b +} + +// TestChecksumTailPaths targets every combination of (SIMD body iterations, +// trailing tail bytes) the asm handlers walk through. The tail handlers +// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by +// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a +// representative spread of k values, including k=0 (no main loop, all tail) +// and k=1 (one main loop iter, then tail). It's explicit coverage for +// payload sizes that are odd, not divisible by 4, by 8, or by 32. +func TestChecksumTailPaths(t *testing.T) { + for _, impl := range implsUnderTest() { + t.Run(impl.name, func(t *testing.T) { + requireAvailable(t, impl) + rng := rand.New(rand.NewPCG(42, 17)) + const padFront = 16 + const maxK = 8 + + pool := make([]byte, 64*maxK+padFront+64) + for i := range pool { + pool[i] = byte(rng.Uint32()) + } + + seeds := []uint16{0, 0xffff, 0xabcd} + offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts + + for k := 0; k <= maxK; k++ { + for tail := 0; tail < 64; tail++ { + length := 64*k + tail + for _, seed := range seeds { + for _, off := range offsets { + if off+length > len(pool) { + continue + } + buf := pool[off : off+length] + want := gvisorchecksum.Checksum(buf, seed) + got := impl.fn(buf, seed) + if got != want { + t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x", + k, tail, length, off, seed, got, want) + } + } + } + } + } + }) + } +} + +// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples +// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail +// handler is meaningfully on the hot path. Sizes are picked to either exercise +// every tail branch (tiny lengths) or sit slightly off realistic packet +// boundaries (e.g. 1499 = MTU − 1). +func BenchmarkChecksumTailSizes(b *testing.B) { + sizes := []int{ + 1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail + 33, 35, 47, 63, // one loop32 + assorted tails + 65, 95, 127, // one loop64 + assorted tails + 1447, 1471, 1499, 1501, // around MTU + 8191, 8193, // around USO + 65531, 65533, // near the kernel max + } + for _, size := range sizes { + buf := make([]byte, size) + for i := range buf { + buf[i] = byte(i) + } + b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) { + b.SetBytes(int64(size)) + for i := 0; i < b.N; i++ { + _ = Checksum(buf, 0) + } + }) + b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) { + b.SetBytes(int64(size)) + for i := 0; i < b.N; i++ { + _ = gvisorchecksum.Checksum(buf, 0) + } + }) + } +} + +// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that +// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical +// USO size (8192), and the kernel's max GSO superpacket (65535). +func BenchmarkChecksum(b *testing.B) { + for _, size := range []int{60, 1448, 8192, 65535} { + buf := make([]byte, size) + for i := range buf { + buf[i] = byte(i) + } + b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) { + b.SetBytes(int64(size)) + for i := 0; i < b.N; i++ { + _ = Checksum(buf, 0) + } + }) + b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) { + b.SetBytes(int64(size)) + for i := 0; i < b.N; i++ { + _ = gvisorchecksum.Checksum(buf, 0) + } + }) + } +} diff --git a/overlay/checksum/export_amd64_test.go b/overlay/checksum/export_amd64_test.go new file mode 100644 index 00000000..9f158489 --- /dev/null +++ b/overlay/checksum/export_amd64_test.go @@ -0,0 +1,11 @@ +package checksum + +// archImpls exposes every hand-written implementation on this architecture +// so the tests exercise them directly, independent of what the public +// Checksum dispatches to on the running CPU. Without this, running the +// suite on a non-AVX2 machine compared gvisor against itself and left the +// assembly untested — silently. available=false makes the test skip loudly +// instead. +var archImpls = []archImpl{ + {name: "avx2", fn: checksumAVX2, available: hasAVX2}, +} diff --git a/overlay/checksum/export_arm64_test.go b/overlay/checksum/export_arm64_test.go new file mode 100644 index 00000000..db673350 --- /dev/null +++ b/overlay/checksum/export_arm64_test.go @@ -0,0 +1,8 @@ +package checksum + +// archImpls exposes every hand-written implementation on this architecture +// for direct testing; see export_amd64_test.go for the rationale. NEON is +// mandatory in armv8, so it is always available. +var archImpls = []archImpl{ + {name: "neon", fn: checksumNEON, available: true}, +} diff --git a/overlay/checksum/export_fallback_test.go b/overlay/checksum/export_fallback_test.go new file mode 100644 index 00000000..35cdb7b1 --- /dev/null +++ b/overlay/checksum/export_fallback_test.go @@ -0,0 +1,7 @@ +//go:build !amd64 && !arm64 + +package checksum + +// No hand-written implementations on this architecture; the dispatcher is +// pure gvisor and there is nothing separate to test. +var archImpls []archImpl diff --git a/overlay/device.go b/overlay/device.go index b6077aba..eca35c16 100644 --- a/overlay/device.go +++ b/overlay/device.go @@ -4,15 +4,25 @@ import ( "io" "net/netip" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" ) +// defaultBatchBufSize is the per-Queue scratch size for Read on backends +// that don't do TSO segmentation. 65535 covers any single IP packet. +const defaultBatchBufSize = 65535 + type Device interface { - io.ReadWriteCloser + io.Closer Activate() error Networks() []netip.Prefix Name() string RoutesFor(netip.Addr) routing.Gateways - SupportsMultiqueue() bool - NewMultiQueueReader() (io.ReadWriteCloser, error) + // Queues returns the device's packet queues, opening additional ones as + // needed until there are n. Platforms without multiqueue support return + // their single queue regardless of n, so callers must size reader loops + // to len(result), not n; implementations never return more than n. An + // error means a queue that should have opened could not; the caller owns + // cleanup via Close. Called once, during interface activation. + Queues(n int) ([]tio.Queue, error) } diff --git a/overlay/overlaytest/noop.go b/overlay/overlaytest/noop.go index 956da7dd..0268c9ec 100644 --- a/overlay/overlaytest/noop.go +++ b/overlay/overlaytest/noop.go @@ -3,10 +3,9 @@ package overlaytest import ( - "errors" - "io" "net/netip" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" ) @@ -31,20 +30,16 @@ func (NoopTun) Name() string { return "noop" } -func (NoopTun) Read([]byte) (int, error) { - return 0, nil +func (NoopTun) Read() ([]tio.Packet, error) { + return nil, nil } func (NoopTun) Write([]byte) (int, error) { return 0, nil } -func (NoopTun) SupportsMultiqueue() bool { - return false -} - -func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, errors.New("unsupported") +func (NoopTun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{NoopTun{}}, nil } func (NoopTun) Close() error { diff --git a/overlay/tio/blockon_linux.go b/overlay/tio/blockon_linux.go new file mode 100644 index 00000000..adea9d1b --- /dev/null +++ b/overlay/tio/blockon_linux.go @@ -0,0 +1,44 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "os" + + "golang.org/x/sys/unix" +) + +// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown. +// (events is POLLIN for reads, POLLOUT for writes) +// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage. +// +// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd) +// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR). +func blockOn(fd, shutdownFd int32, events int16) error { + const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR + pfds := [2]unix.PollFd{ + {Fd: fd, Events: events}, + {Fd: shutdownFd, Events: unix.POLLIN}, + } + var err error + for { + _, err = unix.Poll(pfds[:], -1) + if err != unix.EINTR { + break + } + } + tunEvents := pfds[0].Revents + shutdownEvents := pfds[1].Revents + // Check err before trusting the potentially bogus bits we just got. + if err != nil { + return err + } + if shutdownEvents&(unix.POLLIN|problemFlags) != 0 { + return os.ErrClosed + } + if tunEvents&problemFlags != 0 { + return os.ErrClosed + } + return nil +} diff --git a/overlay/tio/queueset_gso_linux.go b/overlay/tio/queueset_gso_linux.go new file mode 100644 index 00000000..4a26193d --- /dev/null +++ b/overlay/tio/queueset_gso_linux.go @@ -0,0 +1,101 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "encoding/binary" + "errors" + "fmt" + "log/slog" + "sync/atomic" + + "golang.org/x/sys/unix" +) + +type offloadQueueSet struct { + pq []*Offload + // pqi is exactly the same as pq, but stored as the interface type + pqi []Queue + shutdownFd int + // usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel. + // Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission. + usoEnabled bool + closed atomic.Bool + // l is handed to each queue for its bad-vnet-header drop logging. + l *slog.Logger +} + +// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation. +// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets. +func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) { + shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) + if err != nil { + return nil, fmt.Errorf("failed to create eventfd: %w", err) + } + + out := &offloadQueueSet{ + pq: []*Offload{}, + pqi: []Queue{}, + shutdownFd: shutdownFd, + usoEnabled: usoEnabled, + l: l, + } + + return out, nil +} + +func (c *offloadQueueSet) Queues() []Queue { + return c.pqi +} + +func (c *offloadQueueSet) Add(fd int) error { + if c.closed.Load() { + return errors.New("queue set already closed") + } + x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l) + if err != nil { + return err + } + c.pq = append(c.pq, x) + c.pqi = append(c.pqi, x) + + return nil +} + +func (c *offloadQueueSet) wakeForShutdown() error { + var buf [8]byte + binary.NativeEndian.PutUint64(buf[:], 1) + _, err := unix.Write(c.shutdownFd, buf[:]) + return err +} + +func (c *offloadQueueSet) Close() error { + if c.closed.Swap(true) { + return nil + } + + errs := []error{} + + // Signal all readers blocked in poll to wake up and exit. + // They observe POLLIN on the shutdown eventfd and return os.ErrClosed. + if err := c.wakeForShutdown(); err != nil { + errs = append(errs, err) + } + + // Close the per-queue tun fds; this also unblocks any in-flight reads. + for _, x := range c.pq { + if err := x.Close(); err != nil { + errs = append(errs, err) + } + } + + // Close the shutdown eventfd last: every reader's pollfd set references it, + // so it must outlive the wake + per-queue teardown above. + if err := unix.Close(c.shutdownFd); err != nil { + errs = append(errs, err) + } + c.shutdownFd = -1 + + return errors.Join(errs...) +} diff --git a/overlay/tio/queueset_poll_linux.go b/overlay/tio/queueset_poll_linux.go new file mode 100644 index 00000000..f73deef2 --- /dev/null +++ b/overlay/tio/queueset_poll_linux.go @@ -0,0 +1,91 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "encoding/binary" + "errors" + "fmt" + "sync/atomic" + + "golang.org/x/sys/unix" +) + +type pollQueueSet struct { + pq []*Poll + // pqi is exactly the same as pq, but stored as the interface type + pqi []Queue + shutdownFd int + closed atomic.Bool +} + +func NewPollQueueSet() (QueueSet, error) { + shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) + if err != nil { + return nil, fmt.Errorf("failed to create eventfd: %w", err) + } + + out := &pollQueueSet{ + pq: []*Poll{}, + pqi: []Queue{}, + shutdownFd: shutdownFd, + } + + return out, nil +} + +func (c *pollQueueSet) Queues() []Queue { + return c.pqi +} + +func (c *pollQueueSet) Add(fd int) error { + if c.closed.Load() { + return errors.New("queue set already closed") + } + x, err := newPoll(fd, c.shutdownFd) + if err != nil { + return err + } + c.pq = append(c.pq, x) + c.pqi = append(c.pqi, x) + + return nil +} + +func (c *pollQueueSet) wakeForShutdown() error { + var buf [8]byte + binary.NativeEndian.PutUint64(buf[:], 1) + _, err := unix.Write(int(c.shutdownFd), buf[:]) + return err +} + +func (c *pollQueueSet) Close() error { + if c.closed.Swap(true) { + return nil + } + + errs := []error{} + + // Signal all readers blocked in poll to wake up and exit. + // They observe POLLIN on the shutdown eventfd and return os.ErrClosed. + if err := c.wakeForShutdown(); err != nil { + errs = append(errs, err) + } + + // Close the per-queue tun fds; this also unblocks any in-flight reads. + for _, x := range c.pq { + if err := x.Close(); err != nil { + errs = append(errs, err) + } + } + + // Close the shutdown eventfd last: every reader's pollfd set references it, + // so it must outlive the wake + per-queue teardown above. + if err := unix.Close(c.shutdownFd); err != nil { + errs = append(errs, err) + } + c.shutdownFd = -1 + + return errors.Join(errs...) +} diff --git a/overlay/tio/segment_bench_test.go b/overlay/tio/segment_bench_test.go new file mode 100644 index 00000000..13713010 --- /dev/null +++ b/overlay/tio/segment_bench_test.go @@ -0,0 +1,65 @@ +//go:build linux && !android && !e2e_testing + +package tio + +import "testing" + +// fakeBatch stands in for batch.TxBatcher inside the bench — same shape +// of pointer-capturing closure that sendInsideMessage builds. +type fakeBatch struct{ buf [65536]byte } + +func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] } +func (b *fakeBatch) Commit([]byte) {} + +type fakeHostInfo struct { + remoteIndexId uint32 + counter uint64 +} +type fakeIface struct { + rebindCount uint8 + hi *fakeHostInfo +} + +// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per +// SegmentSuperpacket call when a closure captures pointer-bearing +// receivers — the realistic shape of sendInsideMessage's closure. +func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) { + const mss = 1400 + const numSeg = 32 + pkt := buildTSOv6(mss*numSeg, mss) + gso := GSOInfo{ + Size: mss, + HdrLen: 60, // 40 (IPv6) + 20 (TCP) + CsumStart: 40, + Proto: GSOProtoTCP, + } + p := Packet{Bytes: pkt, GSO: gso} + + hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef} + f := &fakeIface{rebindCount: 7, hi: hi} + fb := &fakeBatch{} + + // SegmentSuperpacket consumes pkt destructively; refresh from a master + // copy each iter (matches the production pattern where every TUN read + // hands the segmenter a fresh kernel-supplied buffer). + master := append([]byte(nil), pkt...) + work := make([]byte, len(pkt)) + p.Bytes = work + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + copy(work, master) + err := SegmentSuperpacket(p, func(seg []byte) error { + out := fb.Reserve(16 + len(seg) + 16) + out[0] = byte(f.rebindCount) + out[1] = byte(hi.counter) + hi.counter++ + fb.Commit(out) + return nil + }) + if err != nil { + b.Fatalf("SegmentSuperpacket: %v", err) + } + } +} diff --git a/overlay/tio/segment_other.go b/overlay/tio/segment_other.go new file mode 100644 index 00000000..0019067b --- /dev/null +++ b/overlay/tio/segment_other.go @@ -0,0 +1,16 @@ +//go:build !linux || android + +package tio + +import "fmt" + +func protoFromGSOType(_ uint8) (GSOProto, error) { + return 0, fmt.Errorf("GSO unsupported") +} + +func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error { + if pkt.GSO.IsSuperpacket() { + return fmt.Errorf("tio: GSO superpacket on platform without segmentation support") + } + return fn(pkt.Bytes) +} diff --git a/overlay/tio/single.go b/overlay/tio/single.go new file mode 100644 index 00000000..a24c2242 --- /dev/null +++ b/overlay/tio/single.go @@ -0,0 +1,49 @@ +package tio + +import "io" + +// singleQueue adapts a legacy one-datagram-per-Read source into a Queue. +// Read fills a private scratch buffer and returns exactly one Packet whose +// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract. +// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write. +type singleQueue struct { + rw io.ReadWriter + closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere) + buf []byte + ret [1]Packet +} + +// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue. +// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return. +// Close closes rwc. +func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue { + return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)} +} + +// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else, +// e.g. several queues sharing one device. Close on the returned Queue is a +// no-op so one queue can't tear the shared source out from under its +// siblings; the owner remains responsible for closing the source itself. +func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue { + return &singleQueue{rw: rw, buf: make([]byte, bufSize)} +} + +func (q *singleQueue) Read() ([]Packet, error) { + n, err := q.rw.Read(q.buf) + if err != nil { + return nil, err + } + q.ret[0] = Packet{Bytes: q.buf[:n]} + return q.ret[:], nil +} + +func (q *singleQueue) Write(p []byte) (int, error) { + return q.rw.Write(p) +} + +func (q *singleQueue) Close() error { + if q.closer == nil { + return nil + } + return q.closer.Close() +} diff --git a/overlay/tio/tio.go b/overlay/tio/tio.go new file mode 100644 index 00000000..c3c6b59d --- /dev/null +++ b/overlay/tio/tio.go @@ -0,0 +1,147 @@ +package tio + +import ( + "io" +) + +// QueueSet holds one or many Queue objects and helps close them in an orderly way. +type QueueSet interface { + io.Closer + Queues() []Queue + + // Add takes a tun fd, adds it to the set, and prepares it for use as a Queue. + Add(fd int) error +} + +// Capabilities advertises which kernel offload features a Queue successfully negotiated. +// Callers consult this to decide which coalescers to wire onto the write path. +type Capabilities struct { + // TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6, + // and WriteGSO with GSOProtoTCP is safe. + TSO bool + // USO means the kernel additionally agreed to TUN_F_USO4|USO6, + // so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2. + USO bool +} + +// Queue is a readable/writable Poll queue. +// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers; +// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue. +// +// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd +// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close, +// which first signals a shared shutdown eventfd every reader polls alongside its own fd. +// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns +// os.ErrClosed from Read, so it cannot be used to stop a single Queue. +type Queue interface { + io.Closer + + // Read returns one or more packets. + // The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid + // until the next Read or Close on this Queue. + // A Packet may carry a GSO/USO superpacket (see GSOInfo) + // Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call). + Read() ([]Packet, error) + + // Write emits a single packet on the plaintext (outside→inside) delivery path. + // Safe for concurrent use. + Write(p []byte) (int, error) +} + +// Packet is the unit Queue.Read returns. +// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it. +// GSO is the zero value for an already-segmented IP datagram; +// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming. +type Packet struct { + Bytes []byte + GSO GSOInfo +} + +// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes. +// The zero value means Bytes is one regular IP datagram and no segmentation is required. +type GSOInfo struct { + // Size is the GSO segment size: max payload bytes per segment + // (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket. + Size uint16 + // HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on). + HdrLen uint16 + // CsumStart is the L4 header offset inside Bytes (== L3 header length). + CsumStart uint16 + // Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply. + Proto GSOProto +} + +// IsSuperpacket reports whether g describes a multi-segment GSO/USO +// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire. +func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 } + +// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes, +// safe to retain past the next Read or Close on the originating Queue. +// GSO metadata is copied verbatim. +// Use this only when a caller needs the data to outlive the borrowed-slice contract. +func (p Packet) Clone() Packet { + if p.Bytes == nil { + return p + } + cp := make([]byte, len(p.Bytes)) + copy(cp, p.Bytes) + return Packet{Bytes: cp, GSO: p.GSO} +} + +// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time. +// Callers pick a write-path coalescer based on the result. +// Queues that don't implement it are treated as having no offload capability. +type CapsProvider interface { + Capabilities() Capabilities +} + +// GSOProto selects the L4 protocol for a GSO superpacket. +// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset +// inside the transport header virtio NEEDS_CSUM expects. +type GSOProto uint8 + +const ( + GSOProtoUnknown GSOProto = iota + GSOProtoTCP + GSOProtoUDP +) + +// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket +// assembled from a header prefix plus one or more borrowed payload fragments, +// in a single vectored write (writev with a leading virtio_net_hdr). +// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN. +// Backends without GSO support do not implement this interface and coalescing is skipped. +// +// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum). +// transportHdr is the TCP or UDP header +// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics). +// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload. +// They are read-only from the writer's perspective and must remain valid until the call returns. +// Every segment in pays except possibly the last must be exactly the same size. +// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set. +// +// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability: +// USO may not have been negotiated even when TSO was. +type GSOWriter interface { + io.Writer + CapsProvider + WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error +} + +// SupportsGSO reports whether w implements GSOWriter and the underlying +// queue advertises the negotiated capability for `want`. +func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) { + gw, ok := w.(GSOWriter) + if !ok { + return nil, false + } + caps := gw.Capabilities() + switch want { + case GSOProtoTCP: + return gw, caps.TSO + case GSOProtoUDP: + return gw, caps.USO + default: + return gw, false + } +} diff --git a/overlay/tio/tio_gso_linux.go b/overlay/tio/tio_gso_linux.go new file mode 100644 index 00000000..320ebf50 --- /dev/null +++ b/overlay/tio/tio_gso_linux.go @@ -0,0 +1,417 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "context" + "fmt" + "io" + "log/slog" + "os" + "sync/atomic" + "syscall" + "unsafe" + + "golang.org/x/sys/unix" + + "github.com/slackhq/nebula/overlay/tio/virtio" +) + +const maxSuperpacketLen = 65535 + +// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB. +// Segmentation happens at encrypt time on a per-routine MTU-sized scratch +// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes. +// We round up to give margin for the drain headroom check below. +const tunRxBufSize = 64 * 1024 + +// tunRxBufCap is the total size we allocate for the per-reader rx buffer. +// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes. +// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load, +// amortizing the wake and giving the sendmmsg planner longer same-destination runs. +// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain. +const tunRxBufCap = tunRxBufSize * 8 + +// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop. +// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off. +const tunDrainCap = 64 + +// gsoMaxIovs caps the iovec budget WriteGSO assembles per call: +// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments. +// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64) +// so realistic coalesced bursts never touch the limit. +// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB. +// WriteGSO returns an error rather than reallocating when a caller exceeds this budget. +const gsoMaxIovs = 256 + +// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write. +// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path +// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX +// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the +// L4 checksum anyway. What matters here is what the header does NOT say: +// no NEEDS_CSUM, so the kernel is never asked to finish a checksum. +// All packets that reach the plain Write paths already carry a valid L4 checksum. +var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID} + +// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking. +// A shared eventfd allows Close to wake all readers blocked in poll. +// +// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill +// the first cache line, and the state the reader mutates per packet (rxOff, pending, +// readIovs) all sits after it, so per-packet reader stores never invalidate the line +// concurrent Write callers load fd from. +type Offload struct { + fd int + shutdownFd int + // usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD, + // so writers can decide whether emitting GSO_UDP_L4 superpackets is safe. + usoEnabled bool + closed atomic.Bool + + // gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted + // by WriteGSO. Kept separate from the read-only package-level validVnetHdr + // so non-GSO Writes can ship that constant directly while WriteGSO + // rewrites this scratch on every call. + gsoHdrBuf [virtio.Size]byte + // gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to + // gsoMaxIovs at construction; never grown. WriteGSO returns an error + // (and drops the call) if a caller hands it more fragments than fit. + gsoIovs []unix.Iovec + + rxBuf []byte // backing store for kernel-handed packets read this drain + rxOff int // cursor into rxBuf for the current Read drain + pending []Packet // packets returned from the most recent Read + + // readVnetScratch holds the 10-byte virtio_net_hdr split off the front of + // every TUN read via readv(2). Decoupling the header from the packet body + // lets us read the body directly into rxBuf at the current rxOff with + // no userspace copy on the GSO_NONE fast path. + readVnetScratch [virtio.Size]byte + // readIovs is the readv(2) iovec scratch wired once at construction, + // iovec[0] points at readVnetScratch + // iovec[1].Base/Len is updated per read to address the current rxBuf slot. + readIovs [2]unix.Iovec + + // l is only consulted on the rare bad-vnet-header drop path; it lives + // after the hot state on purpose. May be nil (tests); drops go unlogged then. + l *slog.Logger +} + +func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) { + if err := unix.SetNonblock(fd, true); err != nil { + return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err) + } + + out := &Offload{ + fd: fd, + shutdownFd: shutdownFd, + usoEnabled: usoEnabled, + closed: atomic.Bool{}, + l: l, + + rxBuf: make([]byte, tunRxBufCap), + gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs), + } + + out.gsoIovs[0].Base = &out.gsoHdrBuf[0] + out.gsoIovs[0].SetLen(virtio.Size) + + // readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we + // only repoint readIovs[1] at the next rxBuf slot (see readPacket). + out.readIovs[0].Base = &out.readVnetScratch[0] + out.readIovs[0].SetLen(virtio.Size) + + return out, nil +} + +func (r *Offload) blockOnRead() error { + return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN) +} + +func (r *Offload) blockOnWrite() error { + return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT) +} + +// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch +// and reading the packet body directly into rxBuf at the current rxOff. +// Returns the body length (zero virtio header bytes, just the IP packet/superpacket). +// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not. +func (r *Offload) readPacket(block bool) (int, error) { + for { + r.readIovs[1].Base = &r.rxBuf[r.rxOff] + r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff) + n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs))) + if errno == 0 { + if int(n) < virtio.Size { + return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n) + } + return int(n) - virtio.Size, nil + } + if errno == unix.EAGAIN { + if !block { + return 0, errno + } + if err := r.blockOnRead(); err != nil { + return 0, err + } + continue + } + if errno == unix.EINTR { + continue + } + if errno == unix.EBADF { + return 0, os.ErrClosed + } + return 0, errno + } +} + +// Read returns one or more packets from the tun. +// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket). +// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until: +// - the kernel queue is empty (EAGAIN) +// - we've collected tunDrainCap packets, +// - or we're out of rxBuf headroom. +// +// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs). +// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue. +func (r *Offload) Read() ([]Packet, error) { + r.pending = r.pending[:0] + r.rxOff = 0 + + // Initial (blocking) read. + // Retry on decode errors so a single bad packet does not stall the reader. + for { + n, err := r.readPacket(true) + if err != nil { + return nil, err + } + if err := r.decodeRead(n); err != nil { + // Drop and read again. A bad packet should not kill the reader, + // but a systematic decode failure must not be invisible either. + r.logDroppedRead(err) + continue + } + break + } + + // Drain: non-blocking reads until the kernel queue is empty, the drain + // cap is reached, or rxBuf no longer has room for another worst-case + // kernel-supplied packet (tunRxBufSize). + for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize { + n, err := r.readPacket(false) + if err != nil { + // EAGAIN / EINTR / anything else: stop draining. We already + // have a valid batch from the first read. + break + } + if n <= 0 { + break + } + if err := r.decodeRead(n); err != nil { + // Drop this packet and stop the drain; we'd rather hand off + // what we have than keep spinning here. + r.logDroppedRead(err) + break + } + } + + return r.pending, nil +} + +// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio +// header. Debug-gated so the happy path never pays for attribute assembly. +func (r *Offload) logDroppedRead(err error) { + if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) { + r.l.Debug("dropping tun packet with bad virtio header", "error", err) + } +} + +// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen). +// The bytes stay in rxBuf: +// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set); +// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time. +// +// rxOff advances by pktLen on success +func (r *Offload) decodeRead(pktLen int) error { + if pktLen <= 0 { + return fmt.Errorf("short tun read: %d", pktLen) + } + var hdr virtio.Hdr + hdr.Decode(r.readVnetScratch[:]) + + body := r.rxBuf[r.rxOff : r.rxOff+pktLen] + + if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE { + if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 { + if err := virtio.FinishChecksum(body, hdr); err != nil { + return err + } + } + r.pending = append(r.pending, Packet{Bytes: body}) + r.rxOff += pktLen + return nil + } + + if err := virtio.CheckValid(body, hdr); err != nil { + return err + } + if err := virtio.CorrectHdrLen(body, &hdr); err != nil { + return err + } + proto, err := protoFromGSOType(hdr.GSOType()) + if err != nil { + return err + } + r.pending = append(r.pending, Packet{ + Bytes: body, + GSO: GSOInfo{ + Size: hdr.GSOSize, + HdrLen: hdr.HdrLen, + CsumStart: hdr.CsumStart, + Proto: proto, + }, + }) + r.rxOff += pktLen + return nil +} + +func (r *Offload) Write(buf []byte) (int, error) { + if len(buf) == 0 { + return 0, nil + } + iovs := [2]unix.Iovec{ + {Base: &validVnetHdr[0]}, + {Base: &buf[0]}, + } + iovs[0].SetLen(virtio.Size) + iovs[1].SetLen(len(buf)) + return r.rawWrite(unsafe.Slice(&iovs[0], 2)) +} + +func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) { + for { + n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs))) + if errno == 0 { + if int(n) < virtio.Size { + return 0, io.ErrShortWrite + } + return int(n) - virtio.Size, nil + } + if errno == unix.EAGAIN { + if err := r.blockOnWrite(); err != nil { + return 0, err + } + continue + } + if errno == unix.EINTR { + continue + } + if errno == unix.EBADF { + return 0, os.ErrClosed + } + return 0, errno + } +} + +// Capabilities reports the offload features negotiated for this Queue. TSO +// is always true for Offload (we only construct it on IFF_VNET_HDR FDs); +// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2). +func (r *Offload) Capabilities() Capabilities { + return Capabilities{TSO: true, USO: r.usoEnabled} +} + +func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error { + if len(pays) == 0 { + // There are no payload fragments. There is nothing to send. + return nil + } + var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr + switch proto { + case GSOProtoUDP: + csumOff = 6 + case GSOProtoTCP: + csumOff = 16 + default: + return fmt.Errorf("unknown GSO proto: %d", proto) + } + // Incorrect geometry must cause an error, not a silent drop. + // No sane packet should ever make it inside this branch. + if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 { + return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff) + } + // Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. + // The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change. + need := 3 + len(pays) + if need > cap(r.gsoIovs) { + return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs)) + } + r.gsoIovs = r.gsoIovs[:need] + r.gsoIovs[1].Base = &hdr[0] + r.gsoIovs[1].SetLen(len(hdr)) + r.gsoIovs[2].Base = &transportHdr[0] + r.gsoIovs[2].SetLen(len(transportHdr)) + + segSize := len(pays[0]) + total := len(hdr) + len(transportHdr) + for i, p := range pays { + if len(p) == 0 { + // The coalescers route zero-payload packets down the non-GSO path, + // so an empty fragment means the caller's accounting is broken. + return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays)) + } else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) { + // all segments must be the same size, except for the last one + return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize) + } + total += len(p) + r.gsoIovs[3+i].Base = &p[0] + r.gsoIovs[3+i].SetLen(len(p)) + } + // This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble. + if total > maxSuperpacketLen { + return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen) + } + + // A single segment ships as a plain checksummed packet (GSO_NONE, size 0). + // Multiple segments carry the real GSO type and segSize, which the loop + // above verified is the size of every fragment except possibly the last. + gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE) + if len(pays) > 1 { + gsoType = gsoTypeFromProto(proto, hdr[0]>>4) + if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE { + // gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble. + // A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet. + return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4) + } + } + var gsoSize uint16 + if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE { + gsoSize = uint16(segSize) + } + virtio.EncodeHeader( + r.gsoHdrBuf[:], + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + gsoType, /*gsoType*/ + uint16(len(hdr)+len(transportHdr)), /*hdrLen*/ + gsoSize, /*gsoSize*/ + uint16(len(hdr)), /*csumStart*/ + csumOff, /*csumOffset*/ + ) + + _, err := r.rawWrite(r.gsoIovs) + return err +} + +func (r *Offload) Close() error { + if r.closed.Swap(true) { + return nil + } + + // shutdownFd is owned by the container, so we should not close it + // Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load. + // That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in + // poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs). + // closed.Swap already guarantees we only close once. + return unix.Close(r.fd) +} diff --git a/overlay/tio/tio_poll_linux.go b/overlay/tio/tio_poll_linux.go new file mode 100644 index 00000000..50280f28 --- /dev/null +++ b/overlay/tio/tio_poll_linux.go @@ -0,0 +1,118 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "fmt" + "os" + "sync/atomic" + + "golang.org/x/sys/unix" +) + +type Poll struct { + fd int + shutdownFd int + closed atomic.Bool + + readBuf []byte + batchRet [1]Packet +} + +// newPoll wraps an existing tun fd. +// On failure it does NOT close fd: the caller owns fd and is the sole closer +// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error). +// This matches the newOffload convention and keeps closes at exactly one on every path. +func newPoll(fd int, shutdownFd int) (*Poll, error) { + if err := unix.SetNonblock(fd, true); err != nil { + return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err) + } + + out := &Poll{ + fd: fd, + shutdownFd: shutdownFd, + readBuf: make([]byte, 65535), // largest possible size Linux permits + } + return out, nil +} + +// blockOnRead waits until the Poll fd is readable or shutdown has been signaled. +// Returns os.ErrClosed if Close was called. +func (t *Poll) blockOnRead() error { + return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN) +} + +func (t *Poll) blockOnWrite() error { + return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT) +} + +// TODO: port Offload's post-wake drain loop here so one poll wake amortizes +// over a burst (up to tunDrainCap packets) instead of paying a syscall and a +// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd +// config currently lose that batching. blockOn and the EAGAIN plumbing are +// already shared; kept one-packet-per-Read for now to preserve behavior. +func (t *Poll) Read() ([]Packet, error) { + n, err := t.readOne(t.readBuf) + if err != nil { + return nil, err + } + t.batchRet[0] = Packet{Bytes: t.readBuf[:n]} + return t.batchRet[:], nil +} + +func (t *Poll) readOne(to []byte) (int, error) { + for { + n, errno := unix.Read(t.fd, to) + if errno == nil { + return n, nil + } + switch errno { + case unix.EAGAIN: + if err := t.blockOnRead(); err != nil { + return 0, err + } + case unix.EINTR: + // retry + case unix.EBADF: + return 0, os.ErrClosed + default: + return 0, errno + } + } +} + +// Write is safe for concurrent use +func (t *Poll) Write(from []byte) (int, error) { + for { + n, errno := unix.Write(t.fd, from) + if errno == nil { + return n, nil + } + switch errno { + case unix.EAGAIN: + if err := t.blockOnWrite(); err != nil { + return 0, err + } + case unix.EINTR: + // retry + case unix.EBADF: + return 0, os.ErrClosed + default: + return 0, errno + } + } +} + +func (t *Poll) Close() error { + if t.closed.Swap(true) { + return nil + } + + // shutdownFd is owned by the container, so we should not close it + // Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load. + // That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in + // poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs). + // closed.Swap already guarantees we only close once. + return unix.Close(t.fd) +} diff --git a/overlay/tio/tun_file_linux_test.go b/overlay/tio/tun_file_linux_test.go new file mode 100644 index 00000000..670da94f --- /dev/null +++ b/overlay/tio/tun_file_linux_test.go @@ -0,0 +1,228 @@ +//go:build linux && !android && !e2e_testing +// +build linux,!android,!e2e_testing + +package tio + +import ( + "errors" + "log/slog" + "os" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/sys/unix" +) + +// newReadPipe returns a read fd. The matching write fd is registered for cleanup. +// The caller takes ownership of the read fd (pass it into a QueueSet). +func newReadPipe(t *testing.T) int { + t.Helper() + var fds [2]int + if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil { + t.Fatalf("pipe2: %v", err) + } + t.Cleanup(func() { _ = unix.Close(fds[1]) }) + return fds[0] +} + +func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) { + pipe1 := newReadPipe(t) + pipe2 := newReadPipe(t) + parent, err := NewPollQueueSet() + require.NoError(t, err) + require.NoError(t, parent.Add(pipe1)) + require.NoError(t, parent.Add(pipe2)) + t.Cleanup(func() { + _ = unix.Close(pipe1) + _ = unix.Close(pipe2) + }) + + readers := parent.Queues() + errs := make([]error, len(readers)) + var wg sync.WaitGroup + for i, r := range readers { + wg.Add(1) + go func(i int, r Queue) { + defer wg.Done() + _, errs[i] = r.Read() + }(i, r) + } + + time.Sleep(50 * time.Millisecond) + + if err := parent.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + done := make(chan struct{}) + go func() { wg.Wait(); close(done) }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("readers did not wake") + } + + for i, err := range errs { + if !errors.Is(err, os.ErrClosed) { + t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err) + } + } +} + +// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer +// goroutines while a reader drains the other end of the pipe. The writers +// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same +// time — the exact scenario that raced on the old shared writePoll member +// array. Run under -race; a shared-array regression trips the detector here. +func TestPoll_ConcurrentWrite_NoRace(t *testing.T) { + var fds [2]int + require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC)) + readFd, writeFd := fds[0], fds[1] + + shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) + require.NoError(t, err) + t.Cleanup(func() { _ = unix.Close(shutdownFd) }) + + p, err := newPoll(writeFd, shutdownFd) + require.NoError(t, err) + + const writers = 2 + const perWriter = 4000 + payload := make([]byte, 100) + total := writers * perWriter * len(payload) + + // Reader: drain the read end (blocking) until every writer's bytes are + // consumed, so the writers keep making progress rather than wedging on a + // permanently full pipe. + readDone := make(chan struct{}) + go func() { + defer close(readDone) + buf := make([]byte, 4096) + got := 0 + for got < total { + n, rerr := unix.Read(readFd, buf) + got += n + if rerr != nil { + if rerr == unix.EINTR { + continue + } + return + } + if n == 0 { // EOF + return + } + } + }() + + var wg sync.WaitGroup + for w := 0; w < writers; w++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; i < perWriter; i++ { + if _, werr := p.Write(payload); werr != nil { + t.Errorf("write: %v", werr) + return + } + } + }() + } + wg.Wait() + + select { + case <-readDone: + case <-time.After(10 * time.Second): + t.Fatal("reader did not drain") + } + + require.NoError(t, p.Close()) + _ = unix.Close(readFd) +} + +// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when +// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's +// callers in tun_linux.go) is the sole closer. If newPoll also closed fd, +// the poll path would double-close on Add error. We force the failure with +// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not +// permitted on O_PATH fds and fails with EBADF, while the fd itself stays +// open so we can observe that newPoll left it alone. +func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) { + fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0) + require.NoError(t, err) + t.Cleanup(func() { _ = unix.Close(fd) }) + + p, err := newPoll(fd, 1) + require.Error(t, err, "SetNonblock on an O_PATH fd should fail") + require.Nil(t, p) + + // If newPoll had closed fd, F_GETFD would report it closed. It staying + // open proves newPoll left the fd for the caller to close exactly once. + require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer") +} + +func TestPoll_Close_Idempotent(t *testing.T) { + tf, err := newPoll(newReadPipe(t), 1) + require.NoError(t, err) + if err := tf.Close(); err != nil { + t.Fatalf("first Close: %v", err) + } + if err := tf.Close(); err != nil { + t.Fatalf("second Close should be a no-op, got %v", err) + } +} + +// fdOpen reports whether fd currently refers to an open file description. +// A closed (or never-allocated) fd makes F_GETFD fail with EBADF. +func fdOpen(t *testing.T, fd int) bool { + t.Helper() + _, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0) + if err == nil { + return true + } + if errors.Is(err, unix.EBADF) { + return false + } + t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err) + return false +} + +// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the +// leaked shutdown eventfd: the container that owns shutdownFd must close it in +// Close, and a second Close must be a safe no-op. +func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) { + qs, err := NewPollQueueSet() + require.NoError(t, err) + c, ok := qs.(*pollQueueSet) + require.True(t, ok) + require.NoError(t, qs.Add(newReadPipe(t))) + + shutdownFd := c.shutdownFd + require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close") + + require.NoError(t, qs.Close()) + require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close") + + // Second Close must not touch fds (shutdownFd is now -1) and must return nil. + require.NoError(t, qs.Close()) +} + +// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test +// for the GSO/offload queueset. +func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) { + qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler)) + require.NoError(t, err) + c, ok := qs.(*offloadQueueSet) + require.True(t, ok) + require.NoError(t, qs.Add(newReadPipe(t))) + + shutdownFd := c.shutdownFd + require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close") + + require.NoError(t, qs.Close()) + require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close") + + // Second Close must not touch fds (shutdownFd is now -1) and must return nil. + require.NoError(t, qs.Close()) +} diff --git a/overlay/tio/tun_linux_offload.go b/overlay/tio/tun_linux_offload.go new file mode 100644 index 00000000..2eb54b90 --- /dev/null +++ b/overlay/tio/tun_linux_offload.go @@ -0,0 +1,60 @@ +//go:build linux && !android +// +build linux,!android + +package tio + +import ( + "fmt" + + "golang.org/x/sys/unix" + + "github.com/slackhq/nebula/overlay/tio/virtio" +) + +// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the +// segment-time helpers use. Returns an error for GSO_NONE or any unknown +// value. The caller should only invoke this on a confirmed superpacket. +func protoFromGSOType(t uint8) (GSOProto, error) { + switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN { + case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6: + return GSOProtoTCP, nil + case unix.VIRTIO_NET_HDR_GSO_UDP_L4: + return GSOProtoUDP, nil + default: + return 0, fmt.Errorf("unsupported virtio gso type: %d", t) + } +} + +// gsoTypeFromProto is the reverse of protoFromGSOType +func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 { + switch { + case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6): + return unix.VIRTIO_NET_HDR_GSO_UDP_L4 + case ipVer == 6: + return unix.VIRTIO_NET_HDR_GSO_TCPV6 + case ipVer == 4: + return unix.VIRTIO_NET_HDR_GSO_TCPV4 + default: + return unix.VIRTIO_NET_HDR_GSO_NONE + } +} + +// SegmentSuperpacket invokes fn once per segment of pkt. +// For non-GSO pkts fn is called once with pkt.Bytes. +// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext +// (a freshly-patched L3+L4 header sliced in front of the original payload chunk). +// This slicing is destructive: pkt is consumed by this call. +// Aborts and returns the first error from fn or from per-segment construction. +func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error { + if !pkt.GSO.IsSuperpacket() { + return fn(pkt.Bytes) + } + switch pkt.GSO.Proto { + case GSOProtoTCP: + return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn) + case GSOProtoUDP: + return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn) + default: + return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto) + } +} diff --git a/overlay/tio/tun_linux_offload_test.go b/overlay/tio/tun_linux_offload_test.go new file mode 100644 index 00000000..84d44b49 --- /dev/null +++ b/overlay/tio/tun_linux_offload_test.go @@ -0,0 +1,1122 @@ +//go:build linux && !android && !e2e_testing +// +build linux,!android,!e2e_testing + +package tio + +import ( + "encoding/binary" + "os" + "testing" + + "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/tcpip/checksum" + + "github.com/slackhq/nebula/overlay/tio/virtio" +) + +// testSegScratchSize is a generous segmentation scratch sized to fit any +// of the synthetic TSO/USO superpackets these tests generate (one +// worst-case 64 KiB superpacket plus replicated per-segment headers). +const testSegScratchSize = 192 * 1024 + +// TestProtoFromGSOTypeMasksECN guards the CWR-superpacket drop bug: the +// kernel qualifies a TSO superpacket whose TCP header carries CWR with +// VIRTIO_NET_HDR_GSO_ECN (we negotiate TUN_F_TSO_ECN, so it WILL send +// them once ECN feedback flows), and the decoder must mask that bit +// rather than reject the packet as an unknown type. +func TestProtoFromGSOTypeMasksECN(t *testing.T) { + cases := []struct { + typ uint8 + want GSOProto + }{ + {unix.VIRTIO_NET_HDR_GSO_TCPV4, GSOProtoTCP}, + {unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP}, + {unix.VIRTIO_NET_HDR_GSO_TCPV6, GSOProtoTCP}, + {unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP}, + {unix.VIRTIO_NET_HDR_GSO_UDP_L4, GSOProtoUDP}, + } + for _, c := range cases { + got, err := protoFromGSOType(c.typ) + if err != nil || got != c.want { + t.Errorf("protoFromGSOType(%#x) = (%v, %v), want (%v, nil)", c.typ, got, err, c.want) + } + } + if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_NONE); err == nil { + t.Error("GSO_NONE must still be rejected") + } + if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_ECN); err == nil { + t.Error("a bare ECN bit with no base type must still be rejected") + } +} + +// verifyChecksum confirms that the one's-complement sum across `b`, seeded +// with a folded pseudo-header sum, equals all-ones (valid). +func verifyChecksum(b []byte, pseudo uint16) bool { + return checksum.Checksum(b, pseudo) == 0xffff +} + +// segmentForTest is the test-only counterpart to the production +// SegmentSuperpacket path. It handles GSO_NONE (with optional +// finishChecksum) inline and dispatches GSO superpackets through +// SegmentSuperpacket, draining each yielded segment into a +// freshly-copied [][]byte slot so callers can iterate after the call +// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not +// invoked here. +func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error { + if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE { + cp := append([]byte(nil), pkt...) + if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 { + if err := virtio.FinishChecksum(cp, hdr); err != nil { + return err + } + } + *out = append(*out, cp) + return nil + } + proto, err := protoFromGSOType(hdr.GSOType()) + if err != nil { + return err + } + gso := GSOInfo{ + Size: hdr.GSOSize, + HdrLen: hdr.HdrLen, + CsumStart: hdr.CsumStart, + Proto: proto, + } + return SegmentSuperpacket(Packet{Bytes: pkt, GSO: gso}, func(seg []byte) error { + *out = append(*out, append([]byte(nil), seg...)) + return nil + }) +} + +// pseudoHeaderIPv4 returns the folded pseudo-header sum used to verify a +// TCP/UDP segment's checksum in tests. src/dst are 4 bytes each. +func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 { + s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0)) + s += uint32(proto) + uint32(l4Len) + s = (s & 0xffff) + (s >> 16) + s = (s & 0xffff) + (s >> 16) + return uint16(s) +} + +// pseudoHeaderIPv6 returns the folded pseudo-header sum used to verify a +// TCP/UDP segment's checksum in tests. src/dst are 16 bytes each. +func pseudoHeaderIPv6(src, dst []byte, proto byte, l4Len int) uint16 { + s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0)) + s += uint32(l4Len>>16) + uint32(l4Len&0xffff) + uint32(proto) + s = (s & 0xffff) + (s >> 16) + s = (s & 0xffff) + (s >> 16) + return uint16(s) +} + +// buildTSOv4 builds a synthetic IPv4/TCP TSO superpacket with a payload of +// `payLen` bytes split at `mss`. +func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) { + t.Helper() + const ipLen = 20 + const tcpLen = 20 + pkt := make([]byte, ipLen+tcpLen+payLen) + + // IPv4 header + pkt[0] = 0x45 // version 4, IHL 5 + // total length is meaningless for TSO but set it anyway + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // original ID + pkt[8] = 64 // TTL + pkt[9] = unix.IPPROTO_TCP + copy(pkt[12:16], []byte{10, 0, 0, 1}) // src + copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst + + // TCP header + binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport + binary.BigEndian.PutUint16(pkt[22:24], 80) // dport + binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq + binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack + pkt[32] = 0x50 // data offset 5 words + pkt[33] = 0x18 // ACK | PSH + binary.BigEndian.PutUint16(pkt[34:36], 65535) // window + + // payload + for i := 0; i < payLen; i++ { + pkt[ipLen+tcpLen+i] = byte(i & 0xff) + } + return pkt, virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/ + uint16(ipLen+tcpLen), /*hdrLen*/ + uint16(mss), /*gsoSize*/ + uint16(ipLen), /*csumStart*/ + 16, /*csumOffset*/ + ) +} + +func TestSegmentTCPv4(t *testing.T) { + const mss = 100 + const numSeg = 3 + pkt, hdr := buildTSOv4(t, mss*numSeg, mss) + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != numSeg { + t.Fatalf("expected %d segments, got %d", numSeg, len(out)) + } + + for i, seg := range out { + if len(seg) != 40+mss { + t.Errorf("seg %d: unexpected len %d", i, len(seg)) + } + totalLen := binary.BigEndian.Uint16(seg[2:4]) + if totalLen != uint16(40+mss) { + t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 40+mss) + } + id := binary.BigEndian.Uint16(seg[4:6]) + if id != 0x4242+uint16(i) { + t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i)) + } + seq := binary.BigEndian.Uint32(seg[24:28]) + wantSeq := uint32(10000 + i*mss) + if seq != wantSeq { + t.Errorf("seg %d: seq=%d want %d", i, seq, wantSeq) + } + flags := seg[33] + wantFlags := byte(0x10) // ACK only, PSH cleared + if i == numSeg-1 { + wantFlags = 0x18 // ACK | PSH preserved on last + } + if flags != wantFlags { + t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags) + } + // IPv4 header checksum must verify against itself. + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + // TCP checksum must verify against the pseudo-header. + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+mss) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad TCP checksum", i) + } + } +} + +func TestSegmentTCPv4OddTail(t *testing.T) { + // Payload of 250 bytes with MSS 100 → segments of 100, 100, 50. + pkt, hdr := buildTSOv4(t, 250, 100) + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != 3 { + t.Fatalf("want 3 segments, got %d", len(out)) + } + wantPayLens := []int{100, 100, 50} + for i, seg := range out { + if len(seg)-40 != wantPayLens[i] { + t.Errorf("seg %d: pay len %d want %d", i, len(seg)-40, wantPayLens[i]) + } + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+wantPayLens[i]) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad TCP checksum", i) + } + } +} + +func TestSegmentTCPv6(t *testing.T) { + const ipLen = 40 + const tcpLen = 20 + const mss = 120 + const numSeg = 2 + payLen := mss * numSeg + pkt := make([]byte, ipLen+tcpLen+payLen) + + // IPv6 header + pkt[0] = 0x60 // version 6 + binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen)) + pkt[6] = unix.IPPROTO_TCP + pkt[7] = 64 + // src/dst fe80::1 / fe80::2 + pkt[8] = 0xfe + pkt[9] = 0x80 + pkt[23] = 1 + pkt[24] = 0xfe + pkt[25] = 0x80 + pkt[39] = 2 + + // TCP header + binary.BigEndian.PutUint16(pkt[40:42], 12345) + binary.BigEndian.PutUint16(pkt[42:44], 80) + binary.BigEndian.PutUint32(pkt[44:48], 7) + binary.BigEndian.PutUint32(pkt[48:52], 99) + pkt[52] = 0x50 + pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too + binary.BigEndian.PutUint16(pkt[54:56], 65535) + + for i := 0; i < payLen; i++ { + pkt[ipLen+tcpLen+i] = byte(i) + } + + hdr := virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/ + uint16(ipLen+tcpLen), /*hdrLen*/ + uint16(mss), /*gsoSize*/ + uint16(ipLen), /*csumStart*/ + 16, /*csumOffset*/ + ) + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != numSeg { + t.Fatalf("want %d segments, got %d", numSeg, len(out)) + } + + for i, seg := range out { + if len(seg) != ipLen+tcpLen+mss { + t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+tcpLen+mss) + } + pl := binary.BigEndian.Uint16(seg[4:6]) + if pl != uint16(tcpLen+mss) { + t.Errorf("seg %d: payload_length=%d want %d", i, pl, tcpLen+mss) + } + seq := binary.BigEndian.Uint32(seg[44:48]) + if seq != uint32(7+i*mss) { + t.Errorf("seg %d: seq=%d want %d", i, seq, 7+i*mss) + } + flags := seg[53] + // Original flags = 0x19 (FIN|ACK|PSH). FIN(0x01)+PSH(0x08) should be + // cleared on all but the last; ACK(0x10) always preserved. + wantFlags := byte(0x10) + if i == numSeg-1 { + wantFlags = 0x19 + } + if flags != wantFlags { + t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags) + } + psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpLen+mss) + if !verifyChecksum(seg[ipLen:], psum) { + t.Errorf("seg %d: bad TCP checksum", i) + } + } +} + +func TestSegmentGSONonePassesThrough(t *testing.T) { + pkt, hdr := buildTSOv4(t, 100, 100) + hdr.SetGSOType(unix.VIRTIO_NET_HDR_GSO_NONE) + hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != 1 { + t.Fatalf("want 1 segment, got %d", len(out)) + } + if len(out[0]) != len(pkt) { + t.Fatalf("unexpected length: %d vs %d", len(out[0]), len(pkt)) + } +} + +// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is +// still rejected; only modern GSO_UDP_L4 (USO) is supported. +func TestSegmentRejectsLegacyUDPGSO(t *testing.T) { + hdr := virtio.NewHeader(0, unix.VIRTIO_NET_HDR_GSO_UDP, 0, 0, 0, 0) + var out [][]byte + if err := segmentForTest(nil, hdr, &out, nil); err == nil { + t.Fatalf("expected rejection for legacy UDP GSO") + } +} + +// buildUSOv4 builds a synthetic IPv4/UDP USO superpacket with payload of +// payLen bytes, segmented at gsoSize. +func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) { + t.Helper() + const ipLen = 20 + const udpLen = 8 + pkt := make([]byte, ipLen+udpLen+payLen) + + // IPv4 header + pkt[0] = 0x45 // version 4, IHL 5 + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) + pkt[8] = 64 + pkt[9] = unix.IPPROTO_UDP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + + // UDP header. The kernel hands us a USO superpacket whose length field + // covers the WHOLE superpacket; the segmenter overwrites it per segment. + // Populating it here matters: leaving it zero makes the base-checksum path + // that must exclude it untestable, since excluding zero is a no-op. + binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport + binary.BigEndian.PutUint16(pkt[22:24], 53) // dport + binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length + + for i := 0; i < payLen; i++ { + pkt[ipLen+udpLen+i] = byte(i & 0xff) + } + + return pkt, virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/ + uint16(ipLen+udpLen), /*hdrLen*/ + uint16(gsoSize), /*gsoSize*/ + uint16(ipLen), /*csumStart*/ + 6, /*csumOffset*/ + ) +} + +func TestSegmentUDPv4(t *testing.T) { + const gso = 100 + const numSeg = 3 + pkt, hdr := buildUSOv4(t, gso*numSeg, gso) + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != numSeg { + t.Fatalf("expected %d segments, got %d", numSeg, len(out)) + } + + for i, seg := range out { + if len(seg) != 28+gso { + t.Errorf("seg %d: len %d want %d", i, len(seg), 28+gso) + } + totalLen := binary.BigEndian.Uint16(seg[2:4]) + if totalLen != uint16(28+gso) { + t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso) + } + // Software UDP GSO bumps the IPv4 ID per segment exactly like TSO + // (inet_gso_segment's fixed-ID case is TCP-only); wireguard-go's + // gsoSplit increments unconditionally too. + id := binary.BigEndian.Uint16(seg[4:6]) + if id != 0x4242+uint16(i) { + t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i)) + } + udpLen := binary.BigEndian.Uint16(seg[24:26]) + if udpLen != uint16(8+gso) { + t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+gso) + } + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, 8+gso) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad UDP checksum", i) + } + } +} + +func TestSegmentUDPv4OddTail(t *testing.T) { + // 250 bytes payload, gsoSize=100 → segments of 100, 100, 50. + pkt, hdr := buildUSOv4(t, 250, 100) + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != 3 { + t.Fatalf("want 3 segments, got %d", len(out)) + } + wantPay := []int{100, 100, 50} + for i, seg := range out { + if len(seg)-28 != wantPay[i] { + t.Errorf("seg %d: pay len %d want %d", i, len(seg)-28, wantPay[i]) + } + udpLen := binary.BigEndian.Uint16(seg[24:26]) + if udpLen != uint16(8+wantPay[i]) { + t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+wantPay[i]) + } + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, 8+wantPay[i]) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad UDP checksum", i) + } + } +} + +func TestSegmentUDPv6(t *testing.T) { + const ipLen = 40 + const udpLen = 8 + const gso = 120 + const numSeg = 2 + payLen := gso * numSeg + pkt := make([]byte, ipLen+udpLen+payLen) + + // IPv6 header + pkt[0] = 0x60 + binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+payLen)) + pkt[6] = unix.IPPROTO_UDP + pkt[7] = 64 + pkt[8] = 0xfe + pkt[9] = 0x80 + pkt[23] = 1 + pkt[24] = 0xfe + pkt[25] = 0x80 + pkt[39] = 2 + + binary.BigEndian.PutUint16(pkt[40:42], 12345) + binary.BigEndian.PutUint16(pkt[42:44], 53) + // Superpacket-wide length, as the kernel supplies it; see buildUSOv4. + binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen)) + + for i := 0; i < payLen; i++ { + pkt[ipLen+udpLen+i] = byte(i) + } + + hdr := virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/ + uint16(ipLen+udpLen), /*hdrLen*/ + uint16(gso), /*gsoSize*/ + uint16(ipLen), /*csumStart*/ + 6, /*csumOffset*/ + ) + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != numSeg { + t.Fatalf("want %d segments, got %d", numSeg, len(out)) + } + + for i, seg := range out { + if len(seg) != ipLen+udpLen+gso { + t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+udpLen+gso) + } + pl := binary.BigEndian.Uint16(seg[4:6]) + if pl != uint16(udpLen+gso) { + t.Errorf("seg %d: payload_length=%d want %d", i, pl, udpLen+gso) + } + ul := binary.BigEndian.Uint16(seg[ipLen+4 : ipLen+6]) + if ul != uint16(udpLen+gso) { + t.Errorf("seg %d: udp len=%d want %d", i, ul, udpLen+gso) + } + psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_UDP, udpLen+gso) + if !verifyChecksum(seg[ipLen:], psum) { + t.Errorf("seg %d: bad UDP checksum", i) + } + } +} + +// TestSegmentUDPCEPropagates confirms IP-level CE marks on the seed appear on +// every segment. UDP has no transport-level CWR/ECE: the IP TOS/TC byte is +// copied verbatim into every segment by the segment-prefix copy. +func TestSegmentUDPCEPropagates(t *testing.T) { + pkt, hdr := buildUSOv4(t, 200, 100) + pkt[1] = 0x03 // CE codepoint in IP-ECN + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != 2 { + t.Fatalf("want 2 segments, got %d", len(out)) + } + for i, seg := range out { + if seg[1]&0x03 != 0x03 { + t.Errorf("seg %d: CE missing (tos=%#x)", i, seg[1]) + } + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + } +} + +// TestSegmentTCPCwrFirstSegmentOnly confirms RFC 3168 §6.1.2: when a TSO +// burst's seed has CWR set, only the first emitted segment carries CWR. +// ECE is preserved on every segment (different signal, persistent state). +func TestSegmentTCPCwrFirstSegmentOnly(t *testing.T) { + const mss = 100 + const numSeg = 3 + pkt, hdr := buildTSOv4(t, mss*numSeg, mss) + // Seed flags: CWR | ECE | ACK | PSH. + pkt[33] = 0x80 | 0x40 | 0x10 | 0x08 + + scratch := make([]byte, testSegScratchSize) + var out [][]byte + if err := segmentForTest(pkt, hdr, &out, scratch); err != nil { + t.Fatalf("segmentForTest: %v", err) + } + if len(out) != numSeg { + t.Fatalf("expected %d segments, got %d", numSeg, len(out)) + } + for i, seg := range out { + flags := seg[33] + hasCwr := flags&0x80 != 0 + hasEce := flags&0x40 != 0 + hasPsh := flags&0x08 != 0 + wantCwr := i == 0 + wantPsh := i == numSeg-1 + if hasCwr != wantCwr { + t.Errorf("seg %d: CWR=%v want %v (flags=%#x)", i, hasCwr, wantCwr, flags) + } + if !hasEce { + t.Errorf("seg %d: ECE missing (flags=%#x)", i, flags) + } + if hasPsh != wantPsh { + t.Errorf("seg %d: PSH=%v want %v (flags=%#x)", i, hasPsh, wantPsh, flags) + } + // IP and TCP checksums must still verify after the flag rewrite. + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+mss) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad TCP checksum", i) + } + } +} + +func BenchmarkSegmentTCPv4(b *testing.B) { + sizes := []struct { + name string + payLen int + mss int + }{ + {"64KiB_MSS1460", 65000, 1460}, + {"16KiB_MSS1460", 16384, 1460}, + {"4KiB_MSS1460", 4096, 1460}, + } + for _, sz := range sizes { + b.Run(sz.name, func(b *testing.B) { + const ipLen = 20 + const tcpLen = 20 + pkt := make([]byte, ipLen+tcpLen+sz.payLen) + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+sz.payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) + pkt[8] = 64 + pkt[9] = unix.IPPROTO_TCP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + binary.BigEndian.PutUint16(pkt[20:22], 12345) + binary.BigEndian.PutUint16(pkt[22:24], 80) + binary.BigEndian.PutUint32(pkt[24:28], 10000) + binary.BigEndian.PutUint32(pkt[28:32], 20000) + pkt[32] = 0x50 + pkt[33] = 0x18 + binary.BigEndian.PutUint16(pkt[34:36], 65535) + for i := 0; i < sz.payLen; i++ { + pkt[ipLen+tcpLen+i] = byte(i) + } + hdr := virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/ + uint16(ipLen+tcpLen), /*hdrLen*/ + uint16(sz.mss), /*gsoSize*/ + uint16(ipLen), /*csumStart*/ + 16, /*csumOffset*/ + ) + + scratch := make([]byte, testSegScratchSize) + out := make([][]byte, 0, 64) + + // SegmentSuperpacket consumes its input destructively; restore + // pkt from a master copy each iteration. The restore mirrors the + // kernel→userspace copy that hands a fresh GSO blob to the + // segmenter in production, so it's representative cost rather + // than bench overhead. + master := append([]byte(nil), pkt...) + work := make([]byte, len(pkt)) + + b.SetBytes(int64(len(pkt))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + copy(work, master) + out = out[:0] + if err := segmentForTest(work, hdr, &out, scratch); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// TestTunFileWriteVnetHdrNoAlloc verifies the IFF_VNET_HDR fast-path write is +// allocation-free. We write to /dev/null so every call succeeds synchronously. +func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) { + fd, err := unix.Open("/dev/null", os.O_WRONLY, 0) + if err != nil { + t.Fatalf("open /dev/null: %v", err) + } + t.Cleanup(func() { _ = unix.Close(fd) }) + + tf := &Offload{fd: fd} + + payload := make([]byte, 1400) + // Warm up (first call may trigger one-time internal allocations elsewhere). + if _, err := tf.Write(payload); err != nil { + t.Fatalf("Write: %v", err) + } + + allocs := testing.AllocsPerRun(1000, func() { + if _, err := tf.Write(payload); err != nil { + t.Fatalf("Write: %v", err) + } + }) + if allocs != 0 { + t.Fatalf("Write allocated %.1f times per call, want 0", allocs) + } +} + +// TestSegmentSuperpacketNoAlloc pins the segmenters' zero-allocation +// contract. Both SegmentTCP and SegmentUDP derive their per-superpacket +// constants into fixed-size arrays (tmp/ipTmp/savedHdr) that must stay on +// the stack, and both take a yield closure that must not escape. Any of +// those escaping turns one allocation into one-per-superpacket on the +// hottest path in the reader, which BenchmarkSegmentSuperpacketAllocsTSO +// reports but nothing fails on. This does. +// +// The yield closure here only touches captured scalars: appending segments +// to a slice would allocate in the test itself and mask the measurement. +func TestSegmentSuperpacketNoAlloc(t *testing.T) { + const mss = 1400 + const numSeg = 8 + + cases := []struct { + name string + build func() ([]byte, virtio.Hdr) + }{ + {"tso-v4", func() ([]byte, virtio.Hdr) { return buildTSOv4(t, mss*numSeg, mss) }}, + {"uso-v4", func() ([]byte, virtio.Hdr) { return buildUSOv4(t, mss*numSeg, mss) }}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + master, hdr := tc.build() + proto, err := protoFromGSOType(hdr.GSOType()) + if err != nil { + t.Fatalf("protoFromGSOType: %v", err) + } + work := make([]byte, len(master)) + p := Packet{Bytes: work, GSO: GSOInfo{ + Size: hdr.GSOSize, + HdrLen: hdr.HdrLen, + CsumStart: hdr.CsumStart, + Proto: proto, + }} + + // Segmentation consumes its input destructively, so restore from + // the master copy each run; copy(2) into an existing slice does + // not allocate. seen/bytes keep the closure from being optimized + // away and double as a sanity check that work actually happened. + var seen, bytes int + run := func() { + copy(work, master) + seen, bytes = 0, 0 + if err := SegmentSuperpacket(p, func(seg []byte) error { + seen++ + bytes += len(seg) + return nil + }); err != nil { + t.Fatalf("SegmentSuperpacket: %v", err) + } + } + + run() // warm up: absorb any one-time allocation elsewhere + if seen != numSeg { + t.Fatalf("yielded %d segments, want %d", seen, numSeg) + } + + if allocs := testing.AllocsPerRun(200, run); allocs != 0 { + t.Fatalf("SegmentSuperpacket allocated %.1f times per call, want 0", allocs) + } + if seen != numSeg || bytes == 0 { + t.Fatalf("post-measure sanity: seen=%d bytes=%d", seen, bytes) + } + }) + } +} + +// buildTSOv6 builds a synthetic IPv6/TCP TSO superpacket with payLen bytes +// of payload, segmented at gso. Returns the packet bytes only; the +// virtio_net_hdr is the caller's responsibility. +func buildTSOv6(payLen, gso int) []byte { + const ipLen = 40 + const tcpLen = 20 + pkt := make([]byte, ipLen+tcpLen+payLen) + + pkt[0] = 0x60 // version 6 + binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen)) + pkt[6] = unix.IPPROTO_TCP + pkt[7] = 64 + pkt[8] = 0xfe + pkt[9] = 0x80 + pkt[23] = 1 + pkt[24] = 0xfe + pkt[25] = 0x80 + pkt[39] = 2 + + binary.BigEndian.PutUint16(pkt[40:42], 12345) + binary.BigEndian.PutUint16(pkt[42:44], 80) + binary.BigEndian.PutUint32(pkt[44:48], 7) + binary.BigEndian.PutUint32(pkt[48:52], 99) + pkt[52] = 0x50 + pkt[53] = 0x10 // ACK only + binary.BigEndian.PutUint16(pkt[54:56], 65535) + + for i := 0; i < payLen; i++ { + pkt[ipLen+tcpLen+i] = byte(i) + } + return pkt +} + +// TestDecodeReadFitsMaxTSOAtDrainThreshold proves the rxBuf sizing is +// correct: when rxOff is at the maximum value the drain headroom check +// allows, decodeRead must still be able to absorb a worst-case 64KiB +// TSO superpacket without dropping the burst. With segmentation deferred +// to encrypt time, decodeRead writes only the kernel-supplied bytes into +// rxBuf, so the size requirement is just "fit one worst-case input." +// +// Regression history: in a prior layout the rx buffer doubled as the +// segmentation output, a near-threshold drain read returned "scratch too +// small", the whole 45-segment TSO burst was dropped, and the remote's TCP +// fast-retransmit collapsed cwnd. Keeping this test in the new layout +// guards against re-introducing a drain headroom shortfall. +func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) { + const ipv6HdrLen = 40 + const tcpHdrLen = 20 + const headerLen = ipv6HdrLen + tcpHdrLen + // Maximum TUN read body at the drain threshold. readv bounds the body + // iovec by the space actually left in rxBuf, and the drain gate keeps that + // at >= tunRxBufSize, so that is the largest superpacket the kernel can + // hand back on the last permitted drain read. + pktLen := tunRxBufSize + payLen := pktLen - headerLen + const targetSegs = 64 + gsoSize := (payLen + targetSegs - 1) / targetSegs + + pkt := buildTSOv6(payLen, gsoSize) + if len(pkt) != pktLen { + t.Fatalf("buildTSOv6 produced %d bytes, want %d", len(pkt), pktLen) + } + + o := &Offload{ + rxBuf: make([]byte, tunRxBufCap), + } + // rxOff at the maximum value the drain headroom check permits before + // it would refuse another read. Any drain-time read up to this + // threshold MUST still process correctly. + o.rxOff = tunRxBufCap - tunRxBufSize + + // Stage the body in rxBuf as if readv(2) just placed it there. + copy(o.rxBuf[o.rxOff:], pkt) + + // Encode the matching virtio_net_hdr. + hdr := virtio.NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/ + uint16(headerLen), /*hdrLen*/ + uint16(gsoSize), /*gsoSize*/ + uint16(ipv6HdrLen), /*csumStart*/ + 16, /*csumOffset*/ + ) + hdr.Encode(o.readVnetScratch[:]) + + startRxOff := o.rxOff + if err := o.decodeRead(pktLen); err != nil { + t.Fatalf("decodeRead at drain threshold returned %v — rxBuf sizing regression: "+ + "tunRxBufSize=%d must hold one worst-case input (%d)", + err, tunRxBufSize, pktLen) + } + + if len(o.pending) != 1 { + t.Fatalf("got %d packets, want 1 superpacket entry", len(o.pending)) + } + got := o.pending[0] + if !got.GSO.IsSuperpacket() { + t.Fatalf("expected superpacket GSO metadata, got %+v", got.GSO) + } + if got.GSO.Proto != GSOProtoTCP { + t.Errorf("GSO.Proto=%d want TCP", got.GSO.Proto) + } + if got.GSO.Size != uint16(gsoSize) { + t.Errorf("GSO.Size=%d want %d", got.GSO.Size, gsoSize) + } + if got.GSO.HdrLen != uint16(headerLen) { + t.Errorf("GSO.HdrLen=%d want %d", got.GSO.HdrLen, headerLen) + } + if got.GSO.CsumStart != uint16(ipv6HdrLen) { + t.Errorf("GSO.CsumStart=%d want %d", got.GSO.CsumStart, ipv6HdrLen) + } + if len(got.Bytes) != pktLen { + t.Errorf("len(Bytes)=%d want %d", len(got.Bytes), pktLen) + } + + // rxOff advances exactly by the kernel-supplied body length — no + // segmentation output to account for any more. + if o.rxOff != startRxOff+pktLen { + t.Errorf("rxOff=%d want %d", o.rxOff, startRxOff+pktLen) + } + if o.rxOff > tunRxBufCap { + t.Fatalf("rxOff=%d overran rxBuf (cap=%d)", o.rxOff, tunRxBufCap) + } + + // Validate that segmenting the returned superpacket reproduces the + // expected per-segment IPv6 payload length and TCP checksum. + wantSegs := (payLen + gsoSize - 1) / gsoSize + gotSegs := 0 + if err := SegmentSuperpacket(got, func(seg []byte) error { + defer func() { gotSegs++ }() + if len(seg) < headerLen+1 { + t.Errorf("seg %d too short: %d", gotSegs, len(seg)) + return nil + } + if seg[0]>>4 != 6 { + t.Errorf("seg %d: bad IP version %#x", gotSegs, seg[0]) + } + segPay := len(seg) - headerLen + gotPL := binary.BigEndian.Uint16(seg[4:6]) + if gotPL != uint16(tcpHdrLen+segPay) { + t.Errorf("seg %d: payload_len=%d want %d", gotSegs, gotPL, tcpHdrLen+segPay) + } + psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpHdrLen+segPay) + if !verifyChecksum(seg[ipv6HdrLen:], psum) { + t.Errorf("seg %d: bad TCP checksum", gotSegs) + } + return nil + }); err != nil { + t.Fatalf("SegmentSuperpacket: %v", err) + } + if gotSegs != wantSegs { + t.Fatalf("got %d segments, want %d", gotSegs, wantSegs) + } +} + +// TestOffloadWriteZeroLength: a zero-length Write must be a no-op, not a +// panic. The guard used to live below the &buf[0] that tripped on it. +func TestOffloadWriteZeroLength(t *testing.T) { + tf := &Offload{fd: -1} // any write reaching the fd would fail loudly + for _, buf := range [][]byte{nil, {}} { + n, err := tf.Write(buf) + if n != 0 || err != nil { + t.Errorf("Write(len=0) = (%d, %v), want (0, nil)", n, err) + } + } +} + +// TestWriteGSOSuperpacketGeometry decodes the vnet header the kernel would see for a multi-segment write: +// the GSO type must match the proto and IP version +// gso_size must be the per-segment size (the kernel rejects a superpacket with gso_size == 0), +// and the csum fields must point at the transport header's checksum slot. +// Write through a pipe so the bytes can be read back and decoded. +func TestWriteGSOSuperpacketGeometry(t *testing.T) { + var pfds [2]int + if err := unix.Pipe(pfds[:]); err != nil { + t.Fatalf("pipe: %v", err) + } + t.Cleanup(func() { unix.Close(pfds[0]); unix.Close(pfds[1]) }) + + o := &Offload{fd: pfds[1], gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)} + o.gsoIovs[0].Base = &o.gsoHdrBuf[0] + o.gsoIovs[0].SetLen(virtio.Size) + + ipHdr := make([]byte, 20) + ipHdr[0] = 0x45 + udpHdr := make([]byte, 8) + seg := make([]byte, 1200) + + if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP); err != nil { + t.Fatalf("WriteGSO: %v", err) + } + + buf := make([]byte, virtio.Size+len(ipHdr)+len(udpHdr)+2*len(seg)+64) + n, err := unix.Read(pfds[0], buf) + if err != nil { + t.Fatalf("read pipe: %v", err) + } + var vhdr virtio.Hdr + vhdr.Decode(buf[:virtio.Size]) + if vhdr.GSOType() != unix.VIRTIO_NET_HDR_GSO_UDP_L4 { + t.Errorf("gsoType=%d want UDP_L4", vhdr.GSOType()) + } + if vhdr.GSOSize != 1200 { + t.Errorf("GSOSize=%d want 1200 (per-segment size from pays[0])", vhdr.GSOSize) + } + if vhdr.HdrLen != uint16(len(ipHdr)+len(udpHdr)) { + t.Errorf("HdrLen=%d want %d", vhdr.HdrLen, len(ipHdr)+len(udpHdr)) + } + if vhdr.CsumStart != uint16(len(ipHdr)) || vhdr.CsumOffset != 6 { + t.Errorf("csum start/offset = %d/%d want %d/6", vhdr.CsumStart, vhdr.CsumOffset, len(ipHdr)) + } + if want := virtio.Size + len(ipHdr) + len(udpHdr) + 2*len(seg); n != want { + t.Errorf("wrote %d bytes want %d", n, want) + } +} + +// TestWriteGSORejectsBadGeometry pins the length-check contracts +func TestWriteGSORejectsBadGeometry(t *testing.T) { + fd, err := unix.Open("/dev/null", os.O_WRONLY, 0) + if err != nil { + t.Fatalf("open /dev/null: %v", err) + } + t.Cleanup(func() { _ = unix.Close(fd) }) + + o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)} + o.gsoIovs[0].Base = &o.gsoHdrBuf[0] + o.gsoIovs[0].SetLen(virtio.Size) + + ipHdr := make([]byte, 20) + ipHdr[0] = 0x45 + udpHdr := make([]byte, 8) + tcpHdr := make([]byte, 20) + seg := make([]byte, 1200) + + cases := []struct { + name string + hdr, thdr []byte + pays [][]byte + proto GSOProto + wantErr bool + }{ + {"empty-ip-hdr-with-payload", nil, udpHdr, [][]byte{seg}, GSOProtoUDP, true}, + {"udp-transport-too-short-for-csum", ipHdr, udpHdr[:6], [][]byte{seg}, GSOProtoUDP, true}, + {"tcp-transport-too-short-for-csum", ipHdr, tcpHdr[:16], [][]byte{seg}, GSOProtoTCP, true}, + {"superpacket-over-65535", ipHdr, tcpHdr, [][]byte{make([]byte, 40000), make([]byte, 40000)}, GSOProtoTCP, true}, + {"sole-payload-empty", ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP, true}, + {"leading-empty-fragment", ipHdr, udpHdr, [][]byte{{}, seg, seg}, GSOProtoUDP, true}, + {"trailing-empty-fragment", ipHdr, tcpHdr, [][]byte{seg, {}}, GSOProtoTCP, true}, + {"oversize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 1201), seg}, GSOProtoUDP, true}, + {"undersize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 100), seg}, GSOProtoUDP, true}, + {"oversize-last-fragment", ipHdr, tcpHdr, [][]byte{seg, make([]byte, 1201)}, GSOProtoTCP, true}, + {"short-last-fragment-ok", ipHdr, udpHdr, [][]byte{seg, seg, make([]byte, 100)}, GSOProtoUDP, false}, + {"multi-segment-bad-ip-version", []byte{0x05}, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, true}, + {"single-segment-bad-ip-version-ok", []byte{0x05}, udpHdr, [][]byte{seg}, GSOProtoUDP, false}, + {"no-pays-noop", ipHdr, udpHdr, nil, GSOProtoUDP, false}, + {"valid-udp", ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, false}, + {"valid-tcp", ipHdr, tcpHdr, [][]byte{seg, seg}, GSOProtoTCP, false}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := o.WriteGSO(tc.hdr, tc.thdr, tc.pays, tc.proto) + if tc.wantErr && err == nil { + t.Errorf("WriteGSO = nil, want error") + } + if !tc.wantErr && err != nil { + t.Errorf("WriteGSO = %v, want nil", err) + } + }) + } +} + +// BenchmarkSegmentUDPv4 is the USO counterpart to BenchmarkSegmentTCPv4. The +// yield is a no-op so the measurement is segmentation plus checksum work only. +func BenchmarkSegmentUDPv4(b *testing.B) { + sizes := []struct { + name string + payLen int + gsoSize int + }{ + {"64KiB_GSO1400", 64000, 1400}, + {"16KiB_GSO1400", 16384, 1400}, + {"4KiB_GSO1400", 4096, 1400}, + } + for _, sz := range sizes { + b.Run(sz.name, func(b *testing.B) { + const ipLen = 20 + const udpLen = 8 + pkt := make([]byte, ipLen+udpLen+sz.payLen) + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+sz.payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) + pkt[8] = 64 + pkt[9] = unix.IPPROTO_UDP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + binary.BigEndian.PutUint16(pkt[20:22], 12345) + binary.BigEndian.PutUint16(pkt[22:24], 53) + binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+sz.payLen)) + for i := 0; i < sz.payLen; i++ { + pkt[ipLen+udpLen+i] = byte(i) + } + + master := append([]byte(nil), pkt...) + work := make([]byte, len(pkt)) + p := Packet{Bytes: work, GSO: GSOInfo{ + Size: uint16(sz.gsoSize), + HdrLen: ipLen + udpLen, + CsumStart: ipLen, + Proto: GSOProtoUDP, + }} + + b.SetBytes(int64(len(pkt))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + copy(work, master) + if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkSegmentUDPv6 mirrors BenchmarkSegmentUDPv4 for IPv6, where the +// pseudo-header address sum is 32 bytes rather than 8. +func BenchmarkSegmentUDPv6(b *testing.B) { + sizes := []struct { + name string + payLen int + gsoSize int + }{ + {"64KiB_GSO1400", 64000, 1400}, + {"16KiB_GSO1400", 16384, 1400}, + {"4KiB_GSO1400", 4096, 1400}, + } + for _, sz := range sizes { + b.Run(sz.name, func(b *testing.B) { + const ipLen = 40 + const udpLen = 8 + pkt := make([]byte, ipLen+udpLen+sz.payLen) + pkt[0] = 0x60 + binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+sz.payLen)) + pkt[6] = unix.IPPROTO_UDP + pkt[7] = 64 + pkt[8], pkt[9], pkt[23] = 0xfe, 0x80, 1 + pkt[24], pkt[25], pkt[39] = 0xfe, 0x80, 2 + binary.BigEndian.PutUint16(pkt[40:42], 12345) + binary.BigEndian.PutUint16(pkt[42:44], 53) + binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+sz.payLen)) + for i := 0; i < sz.payLen; i++ { + pkt[ipLen+udpLen+i] = byte(i) + } + + master := append([]byte(nil), pkt...) + work := make([]byte, len(pkt)) + p := Packet{Bytes: work, GSO: GSOInfo{ + Size: uint16(sz.gsoSize), + HdrLen: ipLen + udpLen, + CsumStart: ipLen, + Proto: GSOProtoUDP, + }} + + b.SetBytes(int64(len(pkt))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + copy(work, master) + if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/overlay/tio/virtio/header_linux.go b/overlay/tio/virtio/header_linux.go new file mode 100644 index 00000000..b3528f03 --- /dev/null +++ b/overlay/tio/virtio/header_linux.go @@ -0,0 +1,75 @@ +//go:build linux && !android +// +build linux,!android + +package virtio + +import ( + "encoding/binary" + + "golang.org/x/sys/unix" +) + +// Size is the on-wire length of struct virtio_net_hdr the kernel +// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ +// not set). +const Size = 10 + +// Hdr is the Go view of the legacy virtio_net_hdr. +type Hdr struct { + Flags uint8 + gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types" + HdrLen uint16 + GSOSize uint16 + CsumStart uint16 + CsumOffset uint16 +} + +func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr { + return Hdr{ + Flags: flags, + gsoType: gsoType, + HdrLen: hdrLen, + GSOSize: gsoSize, + CsumStart: csumStart, + CsumOffset: csumOffset, + } +} + +// Decode reads a virtio_net_hdr in host byte order (TUN default; we never +// call TUNSETVNETLE so the kernel matches our endianness). +func (h *Hdr) Decode(b []byte) { + h.Flags = b[0] + h.gsoType = b[1] + h.HdrLen = binary.NativeEndian.Uint16(b[2:4]) + h.GSOSize = binary.NativeEndian.Uint16(b[4:6]) + h.CsumStart = binary.NativeEndian.Uint16(b[6:8]) + h.CsumOffset = binary.NativeEndian.Uint16(b[8:10]) +} + +func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) { + b[0] = flags + b[1] = gsoType + binary.NativeEndian.PutUint16(b[2:4], hdrLen) + binary.NativeEndian.PutUint16(b[4:6], gsoSize) + binary.NativeEndian.PutUint16(b[6:8], csumStart) + binary.NativeEndian.PutUint16(b[8:10], csumOffset) +} + +// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b +// (must be at least Size bytes). Used to emit a TSO superpacket on egress. +func (h *Hdr) Encode(b []byte) { + EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset) +} + +// GSOType returns gsoType with the ECN-flag masked out +func (h *Hdr) GSOType() uint8 { + return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN +} + +func (h *Hdr) HasECNFlag() bool { + return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0 +} + +func (h *Hdr) SetGSOType(x uint8) { + h.gsoType = x +} diff --git a/overlay/tio/virtio/header_other.go b/overlay/tio/virtio/header_other.go new file mode 100644 index 00000000..fe7e38a5 --- /dev/null +++ b/overlay/tio/virtio/header_other.go @@ -0,0 +1,3 @@ +//go:build !linux || android + +package virtio diff --git a/overlay/tio/virtio/segment_linux.go b/overlay/tio/virtio/segment_linux.go new file mode 100644 index 00000000..0509b600 --- /dev/null +++ b/overlay/tio/virtio/segment_linux.go @@ -0,0 +1,441 @@ +//go:build linux && !android +// +build linux,!android + +// Package virtio implements the pure validation, header-correction, and +// per-segment slicing logic for kernel-supplied TSO/USO superpackets on +// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte +// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes, +// so it can be unit-tested in isolation from the tio Queue runtime. +package virtio + +import ( + "encoding/binary" + "errors" + "fmt" + + "golang.org/x/sys/unix" + + "github.com/slackhq/nebula/overlay/checksum" +) + +// Protocol header size bounds used to validate / cap kernel-supplied offsets. +const ( + ipv4HeaderMinLen = 20 // IHL=5, no options + ipv4HeaderMaxLen = 60 // IHL=15, max options + ipv6FixedLen = 40 // IPv6 base header; extensions would extend this + tcpHeaderMinLen = 20 // data-offset=5, no options + tcpHeaderMaxLen = 60 // data-offset=15, max options +) + +// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment. +// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes +const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120 + +// Byte offsets inside an IPv4 header. +const ( + ipv4TotalLenOff = 2 + ipv4IDOff = 4 + ipv4ChecksumOff = 10 + ipv4SrcOff = 12 + ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4) +) + +// Byte offsets inside an IPv6 header. +const ( + ipv6PayloadLenOff = 4 + ipv6SrcOff = 8 + ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16) +) + +// Byte offsets inside a TCP header (relative to its start, i.e. csumStart). +const ( + tcpSeqOff = 4 + tcpDataOffOff = 12 // upper nibble is header len in 32-bit words + tcpFlagsOff = 13 + tcpChecksumOff = 16 +) + +// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}. +const ( + udpHeaderLen = 8 + udpLengthOff = 4 + udpChecksumOff = 6 +) + +var errPacketTooShort = errors.New("packet too short") + +// tcpFinPshMask is cleared on every segment except the last of a TSO burst. +const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08) + +// tcpCwrFlag is cleared on every segment except the first. +// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window) +// and must appear on the first segment of a TSO burst only. +const tcpCwrFlag = 0x80 + +// CheckValid rejects packets whose virtio_net_hdr/IP combination would +// cause a downstream miscompute. The TUN should never emit RSC_INFO and +// the GSO type must agree with the IP version nibble. +func CheckValid(pkt []byte, hdr Hdr) error { + if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 { + return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads") + } + if len(pkt) < ipv4HeaderMinLen { + return errPacketTooShort + } + ipVersion := pkt[0] >> 4 + if ipVersion == 6 && len(pkt) < ipv6FixedLen { + return errPacketTooShort + } + + gsoType := hdr.GSOType() + if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 { + // A GSO type with no segment size would dodge IsSuperpacket() downstream and + // travel as a plain jumbo datagram with an unfinished checksum. + return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType) + } + if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) { + return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType) + } + switch gsoType { + case unix.VIRTIO_NET_HDR_GSO_TCPV4: + if ipVersion != 4 { + return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType) + } + case unix.VIRTIO_NET_HDR_GSO_TCPV6: + if ipVersion != 6 { + return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType) + } + case unix.VIRTIO_NET_HDR_GSO_UDP_L4: + // USO carries either v4 or v6; the leading nibble disambiguates. + if !(ipVersion == 4 || ipVersion == 6) { + return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType) + } + default: + if !(ipVersion == 6 || ipVersion == 4) { + return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType) + } + } + + return nil +} + +// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt. +// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it. +func CorrectHdrLen(pkt []byte, hdr *Hdr) error { + // Thank you wireguard-go for documenting these edge-cases + // Don't trust hdr.hdrLen from the kernel as it can be equal to the length + // of the entire first packet when the kernel is handling it as part of a FORWARD path. + // Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length. + + if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 { + hdr.HdrLen = hdr.CsumStart + 8 + } else { + if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) { + return errors.New("packet is too short") + } + + tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4) + if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen { + return fmt.Errorf("tcp header len is invalid: %d", tcpHLen) + } + hdr.HdrLen = hdr.CsumStart + tcpHLen + } + + if len(pkt) < int(hdr.HdrLen) { + return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen) + } + + if hdr.HdrLen < hdr.CsumStart { + return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart) + } + cSumAt := int(hdr.CsumStart + hdr.CsumOffset) + if cSumAt+1 >= len(pkt) { + return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt)) + } + return nil +} + +// segCount returns how many segments a payload of payLen bytes splits into at gsoSize, +// with a floor of one so a header-only superpacket still yields a single segment. +func segCount(payLen, gsoSize int) int { + n := (payLen + gsoSize - 1) / gsoSize + if n == 0 { + return 1 + } + return n +} + +// basePseudoSum folds the part of the L4 pseudo-header sum that is identical +// for every segment: the source and destination addresses plus the protocol +// number. The per-segment L4 length is added by the caller inside the loop. +func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 { + if isV4 { + return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto + } + return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto +} + +// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments. +// csumStart is the L3 header length, which bounds a valid IHL. +func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) { + ihl := int(pkt[0]&0x0f) * 4 + if ihl < ipv4HeaderMinLen || ihl > csumStart { + return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl) + } + // total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment. + sum := uint32(checksum.Checksum(pkt[:ihl], 0)) + sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2])) + sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2])) + sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])) + sum = (sum & 0xffff) + (sum >> 16) + sum = (sum & 0xffff) + (sum >> 16) + return sum, nil +} + +// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite +func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 { + seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4]) + flags := uint16(pkt[csumStart+tcpFlagsOff]) + + sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0)) + sum += uint32(^uint16(seq >> 16)) + sum += uint32(^uint16(seq)) + sum += uint32(^flags) + sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2])) + sum = (sum & 0xffff) + (sum >> 16) + sum = (sum & 0xffff) + (sum >> 16) + return sum +} + +// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt. +// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize, +// where it sits immediately before that segment's payload chunk in the original buffer. +// pkt is consumed by this call and must not be inspected by the caller after the final yield. +func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error { + if gsoSizeU == 0 { + return fmt.Errorf("gso_size is zero") + } + if csumStartU == 0 { + return fmt.Errorf("csum_start is zero") + } + + headerLen := int(hdrLenU) + csumStart := int(csumStartU) + if headerLen > maxSegHdrLen { + return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen) + } + isV4 := pkt[0]>>4 == 4 + + tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4 + payLen := len(pkt) - headerLen + gsoSize := int(gsoSizeU) + numSeg := segCount(payLen, gsoSize) + + origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4]) + origFlags := pkt[csumStart+tcpFlagsOff] + + baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP) + baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen) + + var origIPID uint16 + var baseIPHdrSum uint32 + if isV4 { + origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]) + var err error + // TSO bumps the ID per segment, so it stays out of the base sum. + baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart) + if err != nil { + return err + } + } + + // Snapshot the pristine L3+L4 header once. ' + // Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source. + var savedHdr [maxSegHdrLen]byte + copy(savedHdr[:headerLen], pkt[:headerLen]) + + for i := 0; i < numSeg; i++ { + segStart := i * gsoSize + segEnd := segStart + gsoSize + if segEnd > payLen { + segEnd = payLen + } + segPayLen := segEnd - segStart + segLen := headerLen + segPayLen + headerOff := i * gsoSize + + // Stamp the header into place immediately before this segment's payload, sourced from the snapshot. + // The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id) + if i > 0 { + // Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp + copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen]) + } + seg := pkt[headerOff : headerOff+segLen] + + segSeq := origSeq + uint32(segStart) + segFlags := origFlags + if i != 0 { + segFlags &^= tcpCwrFlag + } + if i != numSeg-1 { + segFlags &^= tcpFinPshMask + } + totalLen := segLen + + if isV4 { + segID := origIPID + uint16(i) + binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen)) + binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID) + ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID) + binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum)) + } else { + binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen)) + } + + binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq) + seg[csumStart+tcpFlagsOff] = segFlags + + tcpLen := tcpHdrLen + segPayLen + // Payload bytes still live at their original offset in pkt. + // The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed) + // and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize]. + paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0)) + wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum) + wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen) + wide = (wide & 0xffffffff) + (wide >> 32) + wide = (wide & 0xffffffff) + (wide >> 32) + binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide))) + + if err := yield(seg); err != nil { + return err + } + } + + return nil +} + +// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header +// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller. +// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. +// pkt is consumed destructively. +func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error { + if gsoSizeU == 0 { + return fmt.Errorf("gso_size is zero") + } + if csumStartU == 0 { + return fmt.Errorf("csum_start is zero") + } + + isV4 := pkt[0]>>4 == 4 + headerLen := int(hdrLenU) + csumStart := int(csumStartU) + if headerLen > maxSegHdrLen { + return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen) + } + if headerLen-csumStart != udpHeaderLen { + return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart) + } + + payLen := len(pkt) - headerLen + gsoSize := int(gsoSizeU) + numSeg := segCount(payLen, gsoSize) + + baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP) + + var origIPID uint16 + var baseIPHdrSum uint32 + if isV4 { + origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]) + var err error + // Software UDP GSO bumps the ID per segment just like TSO + // (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum. + baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart) + if err != nil { + return err + } + } + + // Snapshot the pristine L3+L4 header once and stamp every segment from it + var savedHdr [maxSegHdrLen]byte + copy(savedHdr[:headerLen], pkt[:headerLen]) + + for i := 0; i < numSeg; i++ { + segStart := i * gsoSize + segEnd := segStart + gsoSize + if segEnd > payLen { + segEnd = payLen + } + segPayLen := segEnd - segStart + segLen := headerLen + segPayLen + headerOff := i * gsoSize + + if i > 0 { + copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen]) + } + seg := pkt[headerOff : headerOff+segLen] + + totalLen := segLen + udpLen := udpHeaderLen + segPayLen + + if isV4 { + segID := origIPID + uint16(i) + binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen)) + binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID) + ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID) + binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum)) + } else { + binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen)) + } + + binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen)) + + // Sum the UDP header (length just written, checksum zeroed) together with + // this segment's payload in one pass, seeded with the pseudo-header sum. + seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0 + pseudo := baseProtoSum + uint32(udpLen) + pseudo = (pseudo & 0xffff) + (pseudo >> 16) + pseudo = (pseudo & 0xffff) + (pseudo >> 16) + csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo)) + if csum == 0 { + csum = 0xffff + } + binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum) + + if err := yield(seg); err != nil { + return err + } + } + + return nil +} + +// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set. +// CsumStart / CsumOffset point at the 16-bit checksum field. +// We zero it, fold a full sum from the partial one that the kernel provided, and store the result. +func FinishChecksum(seg []byte, hdr Hdr) error { + cs := int(hdr.CsumStart) + co := int(hdr.CsumOffset) + if cs+co+2 > len(seg) { + return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg)) + } + // The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the + // L4 region starting at cs, folding the prior partial in as the seed. + partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]) + seg[cs+co] = 0 + seg[cs+co+1] = 0 + csum := ^checksum.Checksum(seg[cs:], partial) + // RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value. + if co == udpChecksumOff && csum == 0 { + csum = 0xffff + } + binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum) + return nil +} + +// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and +// complements it, yielding the on-wire Internet checksum value. +func foldComplement(sum uint32) uint16 { + sum = (sum & 0xffff) + (sum >> 16) + sum = (sum & 0xffff) + (sum >> 16) + return ^uint16(sum) +} diff --git a/overlay/tio/virtio/segment_linux_test.go b/overlay/tio/virtio/segment_linux_test.go new file mode 100644 index 00000000..e7d29997 --- /dev/null +++ b/overlay/tio/virtio/segment_linux_test.go @@ -0,0 +1,602 @@ +//go:build linux && !android +// +build linux,!android + +package virtio + +import ( + "bytes" + "encoding/binary" + "testing" + + "golang.org/x/sys/unix" + + "github.com/slackhq/nebula/overlay/checksum" +) + +// verifyChecksum confirms that the one's-complement sum across b, seeded with +// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum). +// A corrupted header stamped into a segment makes this fail even when the +// checksum field itself was computed from the (pristine) base sums, because +// the bytes the receiver would sum no longer match what was checksummed. +func verifyChecksum(b []byte, pseudo uint16) bool { + return checksum.Checksum(b, pseudo) == 0xffff +} + +// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own +// address and length fields, used to independently verify its L4 checksum. +func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 { + s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0)) + s += uint32(proto) + uint32(l4Len) + s = (s & 0xffff) + (s >> 16) + s = (s & 0xffff) + (s >> 16) + return uint16(s) +} + +// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a +// payload of payLen bytes and returns it alongside the header fields the +// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP). +func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) { + const ipLen = 20 + const tcpLen = 20 + pkt = make([]byte, ipLen+tcpLen+payLen) + + // IPv4 header. + pkt[0] = 0x45 // version 4, IHL 5 + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID + pkt[8] = 64 // TTL + pkt[9] = unix.IPPROTO_TCP + copy(pkt[12:16], []byte{10, 0, 0, 1}) // src + copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst + + // TCP header. + binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport + binary.BigEndian.PutUint16(pkt[22:24], 80) // dport + binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq + binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack + pkt[32] = 0x50 // data offset 5 words + pkt[33] = 0x18 // ACK | PSH + binary.BigEndian.PutUint16(pkt[34:36], 65535) // window + + for i := 0; i < payLen; i++ { + pkt[ipLen+tcpLen+i] = byte(i & 0xff) + } + return pkt, ipLen + tcpLen, ipLen +} + +// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a +// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP). +func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) { + const ipLen = 20 + const udpLen = 8 + pkt = make([]byte, ipLen+udpLen+payLen) + + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen)) + binary.BigEndian.PutUint16(pkt[4:6], 0x4242) + pkt[8] = 64 + pkt[9] = unix.IPPROTO_UDP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + + binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport + binary.BigEndian.PutUint16(pkt[22:24], 53) // dport + + for i := 0; i < payLen; i++ { + pkt[ipLen+udpLen+i] = byte(i & 0xff) + } + return pkt, ipLen + udpLen, ipLen +} + +// collectTCP segments a fresh copy of pkt and returns each segment as an +// independent slice so assertions can run after segmentation completes. +func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte { + t.Helper() + work := append([]byte(nil), pkt...) + var out [][]byte + err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error { + out = append(out, append([]byte(nil), seg...)) + return nil + }) + if err != nil { + t.Fatalf("SegmentTCP: %v", err) + } + return out +} + +func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte { + t.Helper() + work := append([]byte(nil), pkt...) + var out [][]byte + err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error { + out = append(out, append([]byte(nil), seg...)) + return nil + }) + if err != nil { + t.Fatalf("SegmentUDP: %v", err) + } + return out +} + +// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place +// header-slide bug: when gsoSize < headerLen the old code stamped each +// segment's header from pkt[:headerLen], which had already been overwritten +// by the previous segment's overlapping stamp, so segments 2..n carried a +// corrupted header (garbage src/dst/ports/seq). Every segment must instead +// carry the ORIGINAL constant header fields with correct per-segment seq. +func TestSegmentTCPHeaderNotCorrupted(t *testing.T) { + const origSeq = 10000 + cases := []struct { + name string + payLen int + gsoSize uint16 + }{ + // gsoSize (8) < headerLen (40): the bug's trigger. Even split. + {"small-gso-even", 40, 8}, + // gsoSize (8) < headerLen (40) with a short final segment. + {"small-gso-odd-tail", 44, 8}, + // gsoSize (100) >= headerLen (40): the normal path, must still work. + {"normal-gso", 250, 100}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen) + gso := int(tc.gsoSize) + wantSeg := (tc.payLen + gso - 1) / gso + segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize) + if len(segs) != wantSeg { + t.Fatalf("got %d segments, want %d", len(segs), wantSeg) + } + + off := 0 + for i, seg := range segs { + // Constant header fields must be identical to the original in + // EVERY segment. These are exactly the bytes the old code + // corrupted in segments 2..n. + if got := seg[0]; got != 0x45 { + t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got) + } + if seg[9] != unix.IPPROTO_TCP { + t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP) + } + if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) { + t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16]) + } + if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) { + t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20]) + } + if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 { + t.Errorf("seg %d: sport=%d want 12345", i, sport) + } + if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 { + t.Errorf("seg %d: dport=%d want 80", i, dport) + } + if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 { + t.Errorf("seg %d: ack=%d want 20000", i, ack) + } + if seg[32] != 0x50 { + t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32]) + } + + // Per-segment seq must advance by the payload offset. + segStart := i * gso + if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) { + t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart) + } + + // Payload bytes must be the original contiguous slice. + segPayLen := len(seg) - int(hdrLen) + wantPay := make([]byte, segPayLen) + for k := 0; k < segPayLen; k++ { + wantPay[k] = byte((off + k) & 0xff) + } + if !bytes.Equal(seg[hdrLen:], wantPay) { + t.Errorf("seg %d: payload mismatch", i) + } + off += segPayLen + + // End-to-end: the stamped header must checksum-verify. A + // corrupted header fails here because the written checksum was + // derived from the pristine header. + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad TCP checksum", i) + } + } + }) + } +} + +// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in +// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check +// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a +// regression that doubled CsumStart and thus over-tightened the bound (since +// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus +// bound spuriously rejected valid small USO superpackets in decodeRead. +func TestCorrectHdrLenChecksumBound(t *testing.T) { + // A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments + // (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP + // checksum field lives at bytes 26..27, comfortably inside the 40-byte + // packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and + // rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset + // = 26) accepts. This case FAILS against the CsumStart+CsumStart regression. + t.Run("valid-small-uso-accepted", func(t *testing.T) { + pkt, _, csumStart := buildUDPv4Super(12) // total len 40 + hdr := NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/ + 0, /*hdrLen*/ + 6, /*gsoSize: two 6-byte segments*/ + csumStart, /*csumStart*/ + 6, /*csumOffset*/ + ) + if err := CorrectHdrLen(pkt, &hdr); err != nil { + t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err) + } + if hdr.HdrLen != csumStart+udpHeaderLen { + t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen) + } + }) + + // A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the + // checksum field would end at byte 27, but the packet is only 25 bytes + // (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it. + t.Run("too-short-rejected", func(t *testing.T) { + pkt := make([]byte, 25) + pkt[0] = 0x45 // IPv4, IHL 5 + hdr := NewHeader( + unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/ + unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/ + 0, /*hdrLen*/ + 6, /*gsoSize*/ + 20, /*csumStart*/ + 6, /*csumOffset*/ + ) + if err := CorrectHdrLen(pkt, &hdr); err == nil { + t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet") + } + }) +} + +// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs +// the same header stamp and must be correct when gsoSize < headerLen. +func TestSegmentUDPHeaderNotCorrupted(t *testing.T) { + cases := []struct { + name string + payLen int + gsoSize uint16 + }{ + {"small-gso-even", 40, 8}, + {"small-gso-odd-tail", 44, 8}, + {"normal-gso", 250, 100}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen) + gso := int(tc.gsoSize) + wantSeg := (tc.payLen + gso - 1) / gso + segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize) + if len(segs) != wantSeg { + t.Fatalf("got %d segments, want %d", len(segs), wantSeg) + } + + off := 0 + for i, seg := range segs { + if got := seg[0]; got != 0x45 { + t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got) + } + if seg[9] != unix.IPPROTO_UDP { + t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP) + } + if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) { + t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16]) + } + if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) { + t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20]) + } + if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 { + t.Errorf("seg %d: sport=%d want 12345", i, sport) + } + if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 { + t.Errorf("seg %d: dport=%d want 53", i, dport) + } + // Software UDP GSO bumps the IPv4 ID per segment just like TSO + // (inet_gso_segment's fixed-ID case is TCP-only). + if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) { + t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i)) + } + + segPayLen := len(seg) - int(hdrLen) + if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) { + t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen) + } + + wantPay := make([]byte, segPayLen) + for k := 0; k < segPayLen; k++ { + wantPay[k] = byte((off + k) & 0xff) + } + if !bytes.Equal(seg[hdrLen:], wantPay) { + t.Errorf("seg %d: payload mismatch", i) + } + off += segPayLen + + if !verifyChecksum(seg[:20], 0) { + t.Errorf("seg %d: bad IPv4 header checksum", i) + } + psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20) + if !verifyChecksum(seg[20:], psum) { + t.Errorf("seg %d: bad UDP checksum", i) + } + } + }) + } +} + +// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded +// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun. +func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) { + const ipLen, udpLen = 20, 8 + pkt = make([]byte, ipLen+udpLen+len(payload)) + + pkt[0] = 0x45 + binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt))) + pkt[8] = 64 + pkt[9] = unix.IPPROTO_UDP + copy(pkt[12:16], []byte{10, 0, 0, 1}) + copy(pkt[16:20], []byte{10, 0, 0, 2}) + + binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345) + binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53) + binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload))) + copy(pkt[ipLen+udpLen:], payload) + + pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload)) + binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo) + + return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff) +} + +// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as +// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright. +func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) { + var payload []byte + for i := 0; i < 0x10000; i++ { + p := []byte{byte(i >> 8), byte(i)} + pkt, hdr := buildUDPv4Single(p) + cs, co := int(hdr.CsumStart), int(hdr.CsumOffset) + partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2]) + pkt[cs+co], pkt[cs+co+1] = 0, 0 + if ^checksum.Checksum(pkt[cs:], partial) == 0 { + payload = p + break + } + } + if payload == nil { + t.Fatal("no 2-byte payload produced a zero checksum") + } + + pkt, hdr := buildUDPv4Single(payload) + if err := FinishChecksum(pkt, hdr); err != nil { + t.Fatalf("FinishChecksum: %v", err) + } + off := int(hdr.CsumStart) + int(hdr.CsumOffset) + if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff { + t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got) + } +} + +// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must +// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset. +func TestFinishChecksumTCPZeroPreserved(t *testing.T) { + const cs, co = 20, tcpChecksumOff + + // Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites. + seg := make([]byte, cs+co+2) + for i := range seg[cs:] { + seg[cs+i] = byte(i * 7) + } + var partial uint16 + for i := 0; i <= 0xffff; i++ { + binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i)) + probe := append([]byte(nil), seg...) + probe[cs+co], probe[cs+co+1] = 0, 0 + if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 { + partial = uint16(i) + break + } + } + binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial) + + hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co) + if err := FinishChecksum(seg, hdr); err != nil { + t.Fatalf("FinishChecksum: %v", err) + } + if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 { + t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got) + } +} + +// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts. +func TestFinishChecksumUDPValidates(t *testing.T) { + payload := []byte("the definitive tun offloads branch") + pkt, hdr := buildUDPv4Single(payload) + if err := FinishChecksum(pkt, hdr); err != nil { + t.Fatalf("FinishChecksum: %v", err) + } + pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload)) + if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) { + t.Fatal("completed UDP checksum does not validate") + } +} + +// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs +// into gso_type for TSO superpackets with CWR set. CheckValid must +// validate an ECN-qualified type as its base type — previously TCPV4|ECN +// fell into the default case and skipped the IP-version agreement check. +// The qualifier is TCP-only, so it must be rejected on UDP_L4. +func TestCheckValidMasksGSOECN(t *testing.T) { + v4pkt, _, _ := buildTCPv4Super(100) + v6pkt := make([]byte, len(v4pkt)) + copy(v6pkt, v4pkt) + v6pkt[0] = 0x60 // claim IPv6 + + cases := []struct { + name string + pkt []byte + gsoType uint8 + wantErr bool + }{ + {"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false}, + {"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true}, + {"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true}, + {"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true}, + {"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false}, + {"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0)) + if tc.wantErr && err == nil { + t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType) + } + if !tc.wantErr && err != nil { + t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err) + } + }) + } +} + +// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be +// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false, +// dodging both segmentation and FinishChecksum on its way downstream. +func TestCheckValidRejectsZeroGSOSize(t *testing.T) { + v4pkt, _, _ := buildTCPv4Super(100) + if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil { + t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0") + } +} + +// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert +// against an independent RFC 1071 reference fold, hitting the carry edge +// cases (values whose first fold produces another carry). +func TestFoldComplementMatchesReference(t *testing.T) { + refFold := func(s uint64) uint16 { + for s>>16 != 0 { + s = s&0xffff + s>>16 + } + return uint16(s) + } + cases := []uint32{ + 0, 1, 0xffff, + 0x10000, // single carry + 0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff + 0xffff0000, // high half only + 0xfffeffff, // first fold yields another carry + 0xffffffff, // worst case + } + for _, c := range cases { + if got, want := foldComplement(c), ^refFold(uint64(c)); got != want { + t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want) + } + } +} + +// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward +// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header +// into scratch, zero the fields the segment loop rewrites, sum. The production +// versions instead sum in place and subtract those fields via one's-complement +// arithmetic, which is faster but far less obvious — particularly for the TCP +// flags byte, which is only half of a 16-bit word. These references exist so +// that trade is checked rather than asserted. +func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 { + var ipTmp [ipv4HeaderMaxLen]byte + copy(ipTmp[:ihl], pkt[:ihl]) + ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0 + ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0 + ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0 + return uint32(checksum.Checksum(ipTmp[:ihl], 0)) +} + +func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 { + tcpLen := headerLen - csumStart + var tmp [tcpHeaderMaxLen]byte + copy(tmp[:tcpLen], pkt[csumStart:headerLen]) + tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0 + tmp[tcpFlagsOff] = 0 + tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0 + return uint32(checksum.Checksum(tmp[:tcpLen], 0)) +} + +// randSeed is a tiny deterministic PRNG so this test needs no imports beyond +// what the file already has and reproduces identically on every run. +func randByte(state *uint32) byte { + *state = *state*1664525 + 1013904223 + return byte(*state >> 24) +} + +func TestBaseSumsMatchZeroingReference(t *testing.T) { + state := uint32(12345) + + t.Run("ipv4", func(t *testing.T) { + for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 { + for iter := 0; iter < 5000; iter++ { + pkt := make([]byte, ihl) + for i := range pkt { + pkt[i] = randByte(&state) + } + pkt[0] = byte(0x40 | (ihl / 4)) + + want := referenceBaseIPv4HdrSum(pkt, ihl) + got, err := baseIPv4HdrSum(pkt, ihl) + if err != nil { + t.Fatalf("ihl=%d: %v", ihl, err) + } + // Compare the value that reaches the wire: the raw partial + // sums may legally differ by one's-complement -0 vs +0. + for _, tl := range []uint32{20, 1500, 65535} { + for _, id := range []uint32{0, 0x4242, 0xffff} { + if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b { + t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b) + } + } + } + } + } + }) + + t.Run("tcp", func(t *testing.T) { + const csumStart = 20 + for dataOff := 5; dataOff <= 15; dataOff++ { + tcpLen := dataOff * 4 + headerLen := csumStart + tcpLen + for iter := 0; iter < 5000; iter++ { + pkt := make([]byte, headerLen+64) + for i := range pkt { + pkt[i] = randByte(&state) + } + pkt[0] = 0x45 + pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4) + + want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen) + got := baseTCPHdrSum(pkt, csumStart, headerLen) + for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} { + for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} { + for _, l4 := range []uint32{20, 1460, 65535} { + a := foldComplement(want + seq + fl + l4) + b := foldComplement(got + seq + fl + l4) + if a != b { + t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x", + dataOff, seq, fl, l4, a, b) + } + } + } + } + } + } + }) +} diff --git a/overlay/tun_android.go b/overlay/tun_android.go index e4080b41..f7ab417a 100644 --- a/overlay/tun_android.go +++ b/overlay/tun_android.go @@ -13,6 +13,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" ) @@ -63,7 +64,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways { return r } -func (t tun) Activate() error { +func (t *tun) Activate() error { return nil } @@ -96,10 +97,6 @@ func (t *tun) Name() string { return "android" } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for android") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } diff --git a/overlay/tun_darwin.go b/overlay/tun_darwin.go index d30148b9..1076b687 100644 --- a/overlay/tun_darwin.go +++ b/overlay/tun_darwin.go @@ -6,7 +6,6 @@ package overlay import ( "errors" "fmt" - "io" "log/slog" "net/netip" "os" @@ -16,6 +15,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" netroute "golang.org/x/net/route" @@ -550,7 +550,9 @@ func (t *tun) Read(to []byte) (int, error) { return n - 4, nil } -// Write pushes one IP packet onto the utun device. +// Write pushes one IP packet onto the utun device. Safe for concurrent use: +// the AF prefix and iovecs are per-call stack state, and the fd write itself +// serializes on the runtime's fd mutex (see the Queue contract in tio.go). func (t *tun) Write(from []byte) (int, error) { if len(from) == 0 { return 0, syscall.EIO @@ -606,10 +608,6 @@ func (t *tun) Name() string { return t.Device } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } diff --git a/overlay/tun_disabled.go b/overlay/tun_disabled.go index f47880dd..82204ad6 100644 --- a/overlay/tun_disabled.go +++ b/overlay/tun_disabled.go @@ -10,6 +10,7 @@ import ( "github.com/rcrowley/go-metrics" "github.com/slackhq/nebula/iputil" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" ) @@ -23,6 +24,23 @@ type disabledTun struct { l *slog.Logger } +// Read hands the next queued packet to a reader, copying it into b. Reads +// from concurrent queues are safe: the channel receive serializes them and +// each queue copies into its own private scratch buffer. +func (t *disabledTun) Read(b []byte) (int, error) { + r, ok := <-t.read + if !ok { + return 0, io.EOF + } + + t.tx.Inc(1) + if t.l.Enabled(context.Background(), slog.LevelDebug) { + t.l.Debug("Write payload", "raw", prettyPacket(r)) + } + + return copy(b, r), nil +} + func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun { tun := &disabledTun{ vpnNetworks: vpnNetworks, @@ -57,24 +75,6 @@ func (*disabledTun) Name() string { return "disabled" } -func (t *disabledTun) Read(b []byte) (int, error) { - r, ok := <-t.read - if !ok { - return 0, io.EOF - } - - if len(r) > len(b) { - return 0, fmt.Errorf("packet larger than mtu: %d > %d bytes", len(r), len(b)) - } - - t.tx.Inc(1) - if t.l.Enabled(context.Background(), slog.LevelDebug) { - t.l.Debug("Write payload", "raw", prettyPacket(r)) - } - - return copy(b, r), nil -} - func (t *disabledTun) handleICMPEchoRequest(b []byte) bool { out := make([]byte, len(b)) out = iputil.CreateICMPEchoResponse(b, out) @@ -106,12 +106,14 @@ func (t *disabledTun) Write(b []byte) (int, error) { return len(b), nil } -func (t *disabledTun) SupportsMultiqueue() bool { - return true -} - -func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return t, nil +func (t *disabledTun) Queues(n int) ([]tio.Queue, error) { + out := make([]tio.Queue, n) + for i := range out { + // NoClose: the shared channel and metrics are owned by the + // disabledTun; Close on the device tears them down once for everybody. + out[i] = tio.NewSingleQueueNoClose(t, defaultBatchBufSize) + } + return out, nil } func (t *disabledTun) Close() error { diff --git a/overlay/tun_file_linux_test.go b/overlay/tun_file_linux_test.go deleted file mode 100644 index 5ab87e05..00000000 --- a/overlay/tun_file_linux_test.go +++ /dev/null @@ -1,120 +0,0 @@ -//go:build linux && !android && !e2e_testing -// +build linux,!android,!e2e_testing - -package overlay - -import ( - "errors" - "os" - "sync" - "testing" - "time" - - "golang.org/x/sys/unix" -) - -// newReadPipe returns a read fd. The matching write fd is registered for cleanup. -// The caller takes ownership of the read fd (pass it to newTunFd / newFriend). -func newReadPipe(t *testing.T) int { - t.Helper() - var fds [2]int - if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil { - t.Fatalf("pipe2: %v", err) - } - t.Cleanup(func() { _ = unix.Close(fds[1]) }) - return fds[0] -} - -func TestTunFile_WakeForShutdown_UnblocksRead(t *testing.T) { - tf, err := newTunFd(newReadPipe(t)) - if err != nil { - t.Fatalf("newTunFd: %v", err) - } - t.Cleanup(func() { _ = tf.Close() }) - - done := make(chan error, 1) - go func() { - _, err := tf.Read(make([]byte, 64)) - done <- err - }() - - // Verify Read is actually blocked in poll. - select { - case err := <-done: - t.Fatalf("Read returned before shutdown signal: %v", err) - case <-time.After(50 * time.Millisecond): - } - - if err := tf.wakeForShutdown(); err != nil { - t.Fatalf("wakeForShutdown: %v", err) - } - - select { - case err := <-done: - if !errors.Is(err, os.ErrClosed) { - t.Fatalf("expected os.ErrClosed, got %v", err) - } - case <-time.After(2 * time.Second): - t.Fatal("Read did not wake on shutdown") - } -} - -func TestTunFile_WakeForShutdown_WakesFriends(t *testing.T) { - parent, err := newTunFd(newReadPipe(t)) - if err != nil { - t.Fatalf("newTunFd: %v", err) - } - friend, err := parent.newFriend(newReadPipe(t)) - if err != nil { - _ = parent.Close() - t.Fatalf("newFriend: %v", err) - } - t.Cleanup(func() { - _ = friend.Close() - _ = parent.Close() - }) - - readers := []*tunFile{parent, friend} - errs := make([]error, len(readers)) - var wg sync.WaitGroup - for i, r := range readers { - wg.Add(1) - go func(i int, r *tunFile) { - defer wg.Done() - _, errs[i] = r.Read(make([]byte, 64)) - }(i, r) - } - - time.Sleep(50 * time.Millisecond) - - if err := parent.wakeForShutdown(); err != nil { - t.Fatalf("wakeForShutdown: %v", err) - } - - done := make(chan struct{}) - go func() { wg.Wait(); close(done) }() - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("readers did not wake") - } - - for i, err := range errs { - if !errors.Is(err, os.ErrClosed) { - t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err) - } - } -} - -func TestTunFile_Close_Idempotent(t *testing.T) { - tf, err := newTunFd(newReadPipe(t)) - if err != nil { - t.Fatalf("newTunFd: %v", err) - } - if err := tf.Close(); err != nil { - t.Fatalf("first Close: %v", err) - } - if err := tf.Close(); err != nil { - t.Fatalf("second Close should be a no-op, got %v", err) - } -} diff --git a/overlay/tun_freebsd.go b/overlay/tun_freebsd.go index 79f55697..e6479001 100644 --- a/overlay/tun_freebsd.go +++ b/overlay/tun_freebsd.go @@ -7,7 +7,6 @@ import ( "bytes" "errors" "fmt" - "io" "io/fs" "log/slog" "net/netip" @@ -20,7 +19,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" - + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" netroute "golang.org/x/net/route" @@ -561,12 +560,8 @@ func (t *tun) Name() string { return t.Device } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } func (t *tun) addRoutes(logErrors bool) error { diff --git a/overlay/tun_ios.go b/overlay/tun_ios.go index 27bf558b..56603b02 100644 --- a/overlay/tun_ios.go +++ b/overlay/tun_ios.go @@ -16,6 +16,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" "golang.org/x/sys/unix" @@ -159,10 +160,6 @@ func (t *tun) Name() string { return "iOS" } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for ios") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } diff --git a/overlay/tun_linux.go b/overlay/tun_linux.go index 52633842..9b39458b 100644 --- a/overlay/tun_linux.go +++ b/overlay/tun_linux.go @@ -4,10 +4,8 @@ package overlay import ( - "encoding/binary" "errors" "fmt" - "io" "log/slog" "net" "net/netip" @@ -20,188 +18,25 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" "github.com/vishvananda/netlink" "golang.org/x/sys/unix" ) -// tunFile wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking. -// A shared eventfd allows Close to wake all readers blocked in poll. -type tunFile struct { - fd int - shutdownFd int - lastOne bool - readPoll [2]unix.PollFd - writePoll [2]unix.PollFd - closed bool -} - -// newFriend makes a tunFile for a MultiQueueReader that copies the shutdown eventfd from the parent tun -func (r *tunFile) newFriend(fd int) (*tunFile, error) { - if err := unix.SetNonblock(fd, true); err != nil { - return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err) - } - return &tunFile{ - fd: fd, - shutdownFd: r.shutdownFd, - readPoll: [2]unix.PollFd{ - {Fd: int32(fd), Events: unix.POLLIN}, - {Fd: int32(r.shutdownFd), Events: unix.POLLIN}, - }, - writePoll: [2]unix.PollFd{ - {Fd: int32(fd), Events: unix.POLLOUT}, - {Fd: int32(r.shutdownFd), Events: unix.POLLIN}, - }, - }, nil -} - -func newTunFd(fd int) (*tunFile, error) { - if err := unix.SetNonblock(fd, true); err != nil { - return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err) - } - - shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) - if err != nil { - return nil, fmt.Errorf("failed to create eventfd: %w", err) - } - - out := &tunFile{ - fd: fd, - shutdownFd: shutdownFd, - lastOne: true, - readPoll: [2]unix.PollFd{ - {Fd: int32(fd), Events: unix.POLLIN}, - {Fd: int32(shutdownFd), Events: unix.POLLIN}, - }, - writePoll: [2]unix.PollFd{ - {Fd: int32(fd), Events: unix.POLLOUT}, - {Fd: int32(shutdownFd), Events: unix.POLLIN}, - }, - } - - return out, nil -} - -func (r *tunFile) blockOnRead() error { - const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR - var err error - for { - _, err = unix.Poll(r.readPoll[:], -1) - if err != unix.EINTR { - break - } - } - //always reset these! - tunEvents := r.readPoll[0].Revents - shutdownEvents := r.readPoll[1].Revents - r.readPoll[0].Revents = 0 - r.readPoll[1].Revents = 0 - //do the err check before trusting the potentially bogus bits we just got - if err != nil { - return err - } - if shutdownEvents&(unix.POLLIN|problemFlags) != 0 { - return os.ErrClosed - } else if tunEvents&problemFlags != 0 { - return os.ErrClosed - } - return nil -} - -func (r *tunFile) blockOnWrite() error { - const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR - var err error - for { - _, err = unix.Poll(r.writePoll[:], -1) - if err != unix.EINTR { - break - } - } - //always reset these! - tunEvents := r.writePoll[0].Revents - shutdownEvents := r.writePoll[1].Revents - r.writePoll[0].Revents = 0 - r.writePoll[1].Revents = 0 - //do the err check before trusting the potentially bogus bits we just got - if err != nil { - return err - } - if shutdownEvents&(unix.POLLIN|problemFlags) != 0 { - return os.ErrClosed - } else if tunEvents&problemFlags != 0 { - return os.ErrClosed - } - return nil -} - -func (r *tunFile) Read(buf []byte) (int, error) { - for { - if n, err := unix.Read(r.fd, buf); err == nil { - return n, nil - } else if err == unix.EAGAIN { - if err = r.blockOnRead(); err != nil { - return 0, err - } - continue - } else if err == unix.EINTR { - continue - } else if err == unix.EBADF { - return 0, os.ErrClosed - } else { - return 0, err - } - } -} - -func (r *tunFile) Write(buf []byte) (int, error) { - for { - if n, err := unix.Write(r.fd, buf); err == nil { - return n, nil - } else if err == unix.EAGAIN { - if err = r.blockOnWrite(); err != nil { - return 0, err - } - continue - } else if err == unix.EINTR { - continue - } else if err == unix.EBADF { - return 0, os.ErrClosed - } else { - return 0, err - } - } -} - -func (r *tunFile) wakeForShutdown() error { - var buf [8]byte - binary.NativeEndian.PutUint64(buf[:], 1) - _, err := unix.Write(int(r.readPoll[1].Fd), buf[:]) - return err -} - -func (r *tunFile) Close() error { - if r.closed { // avoid closing more than once. Technically a fd could get re-used, which would be a problem - return nil - } - r.closed = true - if r.lastOne { - _ = unix.Close(r.shutdownFd) - } - return unix.Close(r.fd) -} - type tun struct { - *tunFile - readers []*tunFile - closeLock sync.Mutex - Device string - vpnNetworks []netip.Prefix - MaxMTU int - DefaultMTU int - TXQueueLen int - deviceIndex int - ioctlFd uintptr + readers tio.QueueSet + closeLock sync.Mutex + Device string + vpnNetworks []netip.Prefix + MaxMTU int + DefaultMTU int + TXQueueLen int + deviceIndex int + ioctlFd uintptr + vnetHdr bool + offloadFlags uint Routes atomic.Pointer[[]Route] routeTree atomic.Pointer[bart.Table[routing.Gateways]] @@ -240,56 +75,112 @@ type ifreqQLEN struct { } func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) { - t, err := newTunGeneric(c, l, deviceFd, vpnNetworks) - if err != nil { - return nil, err + // We don't know what flags the caller opened this fd with and can't turn + // on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds. + return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0") +} + +// openTunDev opens /dev/net/tun, creating the device node first if it's +// missing (docker containers occasionally omit it). +func openTunDev() (int, error) { + fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0) + if err == nil { + return fd, nil } + if !os.IsNotExist(err) { + return -1, err + } + if err = os.MkdirAll("/dev/net", 0755); err != nil { + return -1, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err) + } + if err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))); err != nil { + return -1, fmt.Errorf("failed to create /dev/net/tun: %w", err) + } + fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0) + if err != nil { + return -1, fmt.Errorf("created /dev/net/tun, but still failed: %w", err) + } + return fd, nil +} - t.Device = "tun0" +// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success. +func tunSetIff(fd int, name string, flags uint16) (string, error) { + var req ifReq + req.Flags = flags + copy(req.Name[:], name) + if err := ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil { + return "", err + } + return strings.Trim(string(req.Name[:]), "\x00"), nil +} - return t, nil +// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available. +const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN + +// usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. +// Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO +const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6 + +func offloadUSOEnabled(offloadFlags uint) bool { + return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0 } func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) { - fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0) - if err != nil { - // If /dev/net/tun doesn't exist, try to create it (will happen in docker) - if os.IsNotExist(err) { - err = os.MkdirAll("/dev/net", 0755) - if err != nil { - return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err) - } - err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))) - if err != nil { - return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err) - } - - fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0) - if err != nil { - return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err) - } - } else { - return nil, err - } - } - - var req ifReq - req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI) + var err error + // IFF_TUN_EXCL prevents us from attaching to an already-running tun + baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL) if multiqueue { - req.Flags |= unix.IFF_MULTI_QUEUE + baseFlags |= unix.IFF_MULTI_QUEUE } nameStr := c.GetString("tun.dev", "") - copy(req.Name[:], nameStr) - if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil { - _ = unix.Close(fd) - return nil, &NameError{ - Name: nameStr, - Underlying: err, + useOffloads := c.GetBool("tun.use_offloads", true) + + var fd int + var name string + var offloadFlags uint + if useOffloads { + fd, err = openTunDev() + if err != nil { + return nil, err + } + // First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_* offloads + // We try TSO+USO first, fall back to TSO-only on kernels without USO (Linux < 6.2), + // and finally give up on virtio headers entirely and reopen as a plain TUN if neither offload mask is accepted. + + // offloadFlags is the exact TUN_F_* mask the kernel accepted. + // We save it so addQueue can replay the identical device-wide mask on added queues + name, err = tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR) + if err != nil { + _ = unix.Close(fd) + useOffloads = false + } else { + if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil { + offloadFlags = usoAndTSOOffloadFlags + } else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil { + offloadFlags = tsoOffloadFlags + } else { + l.Warn("Failed to enable TUN offload (TSO); proceeding without virtio headers", "error", err) + _ = unix.Close(fd) + useOffloads = false + } } } - name := strings.Trim(string(req.Name[:]), "\x00") - t, err := newTunGeneric(c, l, fd, vpnNetworks) + if !useOffloads { + fd, err = openTunDev() + if err != nil { + return nil, err + } + name, err = tunSetIff(fd, nameStr, baseFlags) + if err != nil { + _ = unix.Close(fd) + return nil, &NameError{Name: nameStr, Underlying: err} + } + } + + l.Info("TUN offload status", "tso", useOffloads, "uso", offloadUSOEnabled(offloadFlags)) + + t, err := newTunGeneric(c, l, fd, useOffloads, offloadFlags, vpnNetworks, name) if err != nil { return nil, err } @@ -299,17 +190,37 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue return t, nil } -// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error. -func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) { - tfd, err := newTunFd(fd) +// newTunGeneric does all the stuff common to different tun initialization paths. +// It will close your files on error. +// offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false) +func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) { + var qs tio.QueueSet + var err error + if vnetHdr { + qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l) + } else { + qs, err = tio.NewPollQueueSet() + } + if err != nil { _ = unix.Close(fd) return nil, err } + err = qs.Add(fd) + if err != nil { + // Add only appends on success, so closing the set here can't + // double-close fd; it releases the set's shutdown eventfd. + _ = unix.Close(fd) + _ = qs.Close() + return nil, err + } + t := &tun{ - tunFile: tfd, - readers: []*tunFile{tfd}, + Device: name, + readers: qs, closeLock: sync.Mutex{}, + vnetHdr: vnetHdr, + offloadFlags: offloadFlags, vpnNetworks: vpnNetworks, TXQueueLen: c.GetInt("tun.tx_queue", 500), useSystemRoutes: c.GetBool("tun.use_system_route_table", false), @@ -407,36 +318,49 @@ func (t *tun) reload(c *config.C, initial bool) error { return nil } -func (t *tun) SupportsMultiqueue() bool { - return true +// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all. +func (t *tun) Queues(n int) ([]tio.Queue, error) { + for len(t.readers.Queues()) < n { + if err := t.addQueue(); err != nil { + return nil, err + } + } + return t.readers.Queues(), nil } -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { +// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set. +func (t *tun) addQueue() error { t.closeLock.Lock() defer t.closeLock.Unlock() fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0) if err != nil { - return nil, err + return err } - var req ifReq - req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE) - copy(req.Name[:], t.Device) - if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil { + flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE) + if t.vnetHdr { + flags |= unix.IFF_VNET_HDR + } + if _, err = tunSetIff(fd, t.Device, flags); err != nil { _ = unix.Close(fd) - return nil, err + return err } - out, err := t.tunFile.newFriend(fd) + if t.vnetHdr { + if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil { + _ = unix.Close(fd) + return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err) + } + } + + err = t.readers.Add(fd) if err != nil { _ = unix.Close(fd) - return nil, err + return err } - t.readers = append(t.readers, out) - - return out, nil + return nil } func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways { @@ -613,6 +537,13 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error { Table: unix.RT_TABLE_MAIN, Type: unix.RTN_UNICAST, } + // Match the metric the kernel uses for its auto-installed connected route, + // so RouteReplace overwrites it in place instead of adding a second route at a worse metric. + // IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. + // Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6. + if cidr.Addr().Is6() { + nr.Priority = 256 + } err := netlink.RouteReplace(&nr) if err != nil { t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr) @@ -888,32 +819,10 @@ func (t *tun) Close() error { t.routeChan = nil } - // Signal all readers blocked in poll to wake up and exit - _ = t.tunFile.wakeForShutdown() - if t.ioctlFd > 0 { _ = unix.Close(int(t.ioctlFd)) t.ioctlFd = 0 } - for i := range t.readers { - if i == 0 { - continue //we want to close the zeroth reader last - } - err := t.readers[i].Close() - if err != nil { - t.l.Error("error closing tun reader", "reader", i, "error", err) - } else { - t.l.Info("closed tun reader", "reader", i) - } - } - - //this is t.readers[0] too - err := t.tunFile.Close() - if err != nil { - t.l.Error("error closing tun reader", "reader", 0, "error", err) - } else { - t.l.Info("closed tun reader", "reader", 0) - } - return err + return t.readers.Close() } diff --git a/overlay/tun_linux_test.go b/overlay/tun_linux_test.go index 1c1842da..e074ee0a 100644 --- a/overlay/tun_linux_test.go +++ b/overlay/tun_linux_test.go @@ -3,7 +3,9 @@ package overlay -import "testing" +import ( + "testing" +) var runAdvMSSTests = []struct { name string @@ -32,3 +34,65 @@ func TestTunAdvMSS(t *testing.T) { }) } } + +// TestOffloadUSOEnabled pins the single source of truth for the per-queue USO +// capability: it is derived from the negotiated offload mask, so the mask +// stored on the tun and the capability reported to coalescers cannot drift. +func TestOffloadUSOEnabled(t *testing.T) { + // usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise + // the TSO-only fallback (and the historic hardcoded-mask bug in + // addQueue) would not actually be a downgrade. + if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags { + t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags) + } + if usoAndTSOOffloadFlags == tsoOffloadFlags { + t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags") + } + + cases := []struct { + name string + offloadFlags uint + wantUSO bool + }{ + {"uso-negotiated", usoAndTSOOffloadFlags, true}, + {"tso-fallback", tsoOffloadFlags, false}, + {"no-vnet-hdr", 0, false}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := offloadUSOEnabled(tc.offloadFlags); got != tc.wantUSO { + t.Fatalf("offloadUSOEnabled(%#x) = %v, want %v", tc.offloadFlags, got, tc.wantUSO) + } + }) + } +} + +// TestAddQueueReplaysNegotiatedMask guards the device-wide TUNSETOFFLOAD +// downgrade bug: addQueue must issue the exact mask newTun negotiated +// (t.offloadFlags), not a hardcoded TSO-only mask. Because TUNSETOFFLOAD is +// per-netdev, a narrower mask on an added queue silently disables USO for +// every queue on a USO-capable kernel while the queues keep advertising it. +// +// A full multi-queue exercise needs /dev/net/tun and CAP_NET_ADMIN, which are +// not available in CI/sandbox, so this asserts on the struct field that the +// TUNSETOFFLOAD argument is read from. +func TestAddQueueReplaysNegotiatedMask(t *testing.T) { + t.Run("uso-negotiated", func(t *testing.T) { + tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags} + // The ioctl argument in addQueue is uintptr(t.offloadFlags); + // it must equal the negotiated USO mask, and must NOT be the TSO-only + // mask (the original bug). + if tn.offloadFlags != usoAndTSOOffloadFlags { + t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags) + } + if tn.offloadFlags == tsoOffloadFlags { + t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated") + } + }) + t.Run("tso-fallback", func(t *testing.T) { + tn := &tun{vnetHdr: true, offloadFlags: tsoOffloadFlags} + if tn.offloadFlags != tsoOffloadFlags { + t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, tsoOffloadFlags) + } + }) +} diff --git a/overlay/tun_netbsd.go b/overlay/tun_netbsd.go index c971bb6e..97691543 100644 --- a/overlay/tun_netbsd.go +++ b/overlay/tun_netbsd.go @@ -6,7 +6,6 @@ package overlay import ( "errors" "fmt" - "io" "log/slog" "net/netip" "os" @@ -17,6 +16,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" netroute "golang.org/x/net/route" @@ -390,12 +390,8 @@ func (t *tun) Name() string { return t.Device } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } func (t *tun) addRoutes(logErrors bool) error { diff --git a/overlay/tun_openbsd.go b/overlay/tun_openbsd.go index 41224777..23816b0d 100644 --- a/overlay/tun_openbsd.go +++ b/overlay/tun_openbsd.go @@ -6,7 +6,6 @@ package overlay import ( "errors" "fmt" - "io" "log/slog" "net/netip" "os" @@ -17,6 +16,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" netroute "golang.org/x/net/route" @@ -138,8 +138,8 @@ func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error) //go:noescape func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error) -// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from the -// packet so the payload lands directly in to. +// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from +// the packet so the payload lands directly in to. func (t *tun) Read(to []byte) (int, error) { var head [4]byte @@ -369,12 +369,8 @@ func (t *tun) Name() string { return t.Device } -func (t *tun) SupportsMultiqueue() bool { - return false -} - -func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd") +func (t *tun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } func (t *tun) addRoutes(logErrors bool) error { diff --git a/overlay/tun_tester.go b/overlay/tun_tester.go index 8acd83f0..4b2685e0 100644 --- a/overlay/tun_tester.go +++ b/overlay/tun_tester.go @@ -14,6 +14,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/udp" ) @@ -177,10 +178,6 @@ func (t *TestTun) Read(b []byte) (int, error) { return n, nil } -func (t *TestTun) SupportsMultiqueue() bool { - return false -} - -func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented") +func (t *TestTun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, udp.MTU)}, nil } diff --git a/overlay/tun_windows.go b/overlay/tun_windows.go index cf01615f..6be85ffc 100644 --- a/overlay/tun_windows.go +++ b/overlay/tun_windows.go @@ -6,7 +6,6 @@ package overlay import ( "crypto" "fmt" - "io" "log/slog" "net/netip" "os" @@ -18,6 +17,7 @@ import ( "github.com/gaissmai/bart" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" "github.com/slackhq/nebula/wintun" @@ -47,6 +47,10 @@ type winTun struct { tun *wintun.NativeTun } +func (t *winTun) Read(b []byte) (int, error) { + return t.tun.Read(b, 0) +} + func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) { return nil, fmt.Errorf("newTunFromFd not supported in Windows") } @@ -255,20 +259,12 @@ func (t *winTun) Name() string { return t.Device } -func (t *winTun) Read(b []byte) (int, error) { - return t.tun.Read(b, 0) -} - func (t *winTun) Write(b []byte) (int, error) { return t.tun.Write(b, 0) } -func (t *winTun) SupportsMultiqueue() bool { - return false -} - -func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return nil, fmt.Errorf("TODO: multiqueue not implemented for windows") +func (t *winTun) Queues(int) ([]tio.Queue, error) { + return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil } func (t *winTun) Close() error { diff --git a/overlay/user.go b/overlay/user.go index e5f27f37..2d775bde 100644 --- a/overlay/user.go +++ b/overlay/user.go @@ -6,6 +6,7 @@ import ( "net/netip" "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/routing" ) @@ -46,12 +47,16 @@ func (d *UserDevice) RoutesFor(ip netip.Addr) routing.Gateways { return routing.Gateways{routing.NewGateway(ip, 1)} } -func (d *UserDevice) SupportsMultiqueue() bool { - return true -} - -func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) { - return d, nil +func (d *UserDevice) Queues(n int) ([]tio.Queue, error) { + out := make([]tio.Queue, n) + for i := range out { + // All queues share the underlying pipes (the io.Pipe serializes + // concurrent callers) but each owns a private scratch buffer so + // concurrent Reads across queues never alias. NoClose: the pipes are + // owned by the UserDevice and torn down once by UserDevice.Close. + out[i] = tio.NewSingleQueueNoClose(d, defaultBatchBufSize) + } + return out, nil } func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) { @@ -61,9 +66,11 @@ func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) { func (d *UserDevice) Read(p []byte) (n int, err error) { return d.outboundReader.Read(p) } + func (d *UserDevice) Write(p []byte) (n int, err error) { return d.inboundWriter.Write(p) } + func (d *UserDevice) Close() error { d.inboundWriter.Close() d.outboundWriter.Close() diff --git a/overlay/user_test.go b/overlay/user_test.go new file mode 100644 index 00000000..9e0e9c9c --- /dev/null +++ b/overlay/user_test.go @@ -0,0 +1,163 @@ +package overlay + +import ( + "fmt" + "net/netip" + "sync" + "testing" + + "github.com/slackhq/nebula/overlay/tio" +) + +// newTestUserDevice returns the concrete *UserDevice so tests can reach Pipe() +// and the internal queue plumbing. +func newTestUserDevice(t *testing.T) *UserDevice { + t.Helper() + dev, err := NewUserDevice([]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}) + if err != nil { + t.Fatalf("NewUserDevice: %v", err) + } + ud, ok := dev.(*UserDevice) + if !ok { + t.Fatalf("NewUserDevice returned %T, want *UserDevice", dev) + } + return ud +} + +// TestUserDeviceReadersDistinctBuffers ensures each Queue is actually different +func TestUserDeviceReadersDistinctBuffers(t *testing.T) { + d := newTestUserDevice(t) + + readers, err := d.Queues(2) + if err != nil { + t.Fatalf("Queues: %v", err) + } + if len(readers) != 2 { + t.Fatalf("Queues(2) returned %d queues, want 2", len(readers)) + } + + // Distinct queue objects. + if readers[0] == readers[1] { + t.Fatal("Queues(2) returned the same queue object twice") + } + + // Drive one packet through each queue and confirm the borrowed bytes from + // the first read are NOT clobbered by the second read. With a shared + // buffer, reading pkt1 into q1 would corrupt q0's still-borrowed slice. + _, ow := d.Pipe() + + pkt0 := []byte("packet-zero-aaaaaaaa") + pkt1 := []byte("packet-one-bbbbbbbbb") + + // The pipe is unbuffered, so writes block until a reader consumes them. + // Serialize: write pkt0 (read on q0), then write pkt1 (read on q1). + go func() { + if _, err := ow.Write(pkt0); err != nil { + t.Errorf("write pkt0: %v", err) + } + if _, err := ow.Write(pkt1); err != nil { + t.Errorf("write pkt1: %v", err) + } + }() + + got0, err := readers[0].Read() + if err != nil { + t.Fatalf("q0.Read: %v", err) + } + if len(got0) != 1 || string(got0[0].Bytes) != string(pkt0) { + t.Fatalf("q0 first read = %q, want %q", firstBytes(got0), pkt0) + } + // Hold onto q0's borrowed slice across q1's read. + borrowed := got0[0].Bytes + + got1, err := readers[1].Read() + if err != nil { + t.Fatalf("q1.Read: %v", err) + } + if len(got1) != 1 || string(got1[0].Bytes) != string(pkt1) { + t.Fatalf("q1 read = %q, want %q", firstBytes(got1), pkt1) + } + + // q0's borrowed bytes must still hold pkt0 - a shared buffer would now + // show pkt1's contents. + if string(borrowed) != string(pkt0) { + t.Fatalf("q0 borrowed bytes were clobbered by q1's read: got %q, want %q", borrowed, pkt0) + } +} + +// TestUserDeviceReadersConcurrentRace exercises two queues reading distinct +// packets concurrently. Run it under `go test -race`: with the old +// shared-buffer implementation the concurrent Reads raced on readBuf/batchRet +// and corrupted each other's returned slices. +func TestUserDeviceReadersConcurrentRace(t *testing.T) { + d := newTestUserDevice(t) + readers, err := d.Queues(2) + if err != nil { + t.Fatalf("Queues: %v", err) + } + _, ow := d.Pipe() + + const iterations = 200 + + errs := make(chan error, 3) + + // Each reader parks in Read on the shared outboundReader; io.Pipe hands + // each write to whichever reader is currently waiting. We only care that + // concurrent Reads into distinct buffers are race-free, so any parked + // reader may serve any write. + var wg sync.WaitGroup + run := func(idx int) { + defer wg.Done() + for i := 0; i < iterations; i++ { + pkts, err := readers[idx].Read() + if err != nil { + errs <- err + return + } + if len(pkts) != 1 { + errs <- fmt.Errorf("reader %d: got %d packets, want 1", idx, len(pkts)) + return + } + // Touch every byte of the borrowed slice while the other reader + // may be mid-Read; a shared buffer would race here. + total := 0 + for _, c := range pkts[0].Bytes { + total += int(c) + } + _ = total + } + } + + wg.Add(2) + go run(0) + go run(1) + + // Feed 2*iterations packets. io.Pipe copies each write straight into the + // waiting reader's private buffer, so reusing buf between writes is safe. + go func() { + buf := make([]byte, 32) + for i := 0; i < 2*iterations; i++ { + for j := range buf { + buf[j] = byte(i + j) + } + if _, err := ow.Write(buf); err != nil { + errs <- err + return + } + } + }() + + wg.Wait() + select { + case err := <-errs: + t.Fatalf("concurrent reader failed: %v", err) + default: + } +} + +func firstBytes(p []tio.Packet) []byte { + if len(p) == 0 { + return nil + } + return p[0].Bytes +} diff --git a/pprof_debug.go b/pprof_debug.go new file mode 100644 index 00000000..49ae5151 --- /dev/null +++ b/pprof_debug.go @@ -0,0 +1,35 @@ +//go:build debug + +package nebula + +import ( + "context" + "errors" + "log/slog" + "net/http" + _ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux +) + +// startPprofServer serves net/http/pprof on localhost:6060 for the life of +// ctx. It is only compiled into debug builds (`-tags debug`, `make debug`), +// so a debug build announces itself with the Info line below. Loopback only: +// a wildcard bind would expose profiles (peer addresses, config-derived +// state) to anything that can reach the host, the overlay included. +func startPprofServer(ctx context.Context, l *slog.Logger) { + server := &http.Server{Addr: "localhost:6060", Handler: nil} + l.Info("Starting pprof debug server (debug build)", "addr", server.Addr) + + go func() { + if err := server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { + l.Error("pprof debug server stopped", "error", err) + } + }() + + // Shut down the server when the context is cancelled. + go func() { + <-ctx.Done() + if err := server.Shutdown(context.Background()); err != nil { + l.Debug("Error shutting down pprof debug server", "error", err) + } + }() +} diff --git a/pprof_nodebug.go b/pprof_nodebug.go new file mode 100644 index 00000000..f74be0d1 --- /dev/null +++ b/pprof_nodebug.go @@ -0,0 +1,11 @@ +//go:build !debug + +package nebula + +import ( + "context" + "log/slog" +) + +// startPprofServer is a no-op unless built with `-tags debug` (see make debug). +func startPprofServer(_ context.Context, _ *slog.Logger) {} diff --git a/relay_manager.go b/relay_manager.go index 1ae382a3..46d7a2bb 100644 --- a/relay_manager.go +++ b/relay_manager.go @@ -161,7 +161,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak switch existingRelay.State { case Established: hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String()) - f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false) + f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false, 0) case Disestablished: // Mark this relay as 'requested' relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested) diff --git a/service/service_test.go b/service/service_test.go index 4bcc8437..add09057 100644 --- a/service/service_test.go +++ b/service/service_test.go @@ -4,6 +4,8 @@ import ( "bytes" "context" "errors" + "fmt" + "net" "net/netip" "os" "testing" @@ -89,7 +91,23 @@ func newSimpleService(caCrt cert.Certificate, caKey []byte, name string, udpIp n return s } +// ephemeralUDPPort reserves a free UDP port by binding port 0, then releases +// it for the caller to use. A fixed port would collide with concurrent test +// runs or an unrelated process (say, a real nebula) already listening on it. +func ephemeralUDPPort(t *testing.T) int { + t.Helper() + pc, err := net.ListenPacket("udp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + port := pc.LocalAddr().(*net.UDPAddr).Port + _ = pc.Close() + return port +} + func TestService(t *testing.T) { + lighthousePort := ephemeralUDPPort(t) + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) a := newSimpleService(ca, caKey, "a", netip.MustParseAddr("10.0.0.1"), m{ "static_host_map": m{}, @@ -98,12 +116,12 @@ func TestService(t *testing.T) { }, "listen": m{ "host": "0.0.0.0", - "port": 4243, + "port": lighthousePort, }, }) b := newSimpleService(ca, caKey, "b", netip.MustParseAddr("10.0.0.2"), m{ "static_host_map": m{ - "10.0.0.1": []string{"localhost:4243"}, + "10.0.0.1": []string{fmt.Sprintf("localhost:%d", lighthousePort)}, }, "lighthouse": m{ "hosts": []string{"10.0.0.1"}, diff --git a/udp/conn.go b/udp/conn.go index 30d89dec..6d20cbbf 100644 --- a/udp/conn.go +++ b/udp/conn.go @@ -8,16 +8,41 @@ import ( const MTU = 9001 +// MaxWriteBatch is the largest batch any Conn.WriteBatch implementation is +// required to accept. Callers SHOULD NOT pass more than this per call; Linux +// backends preallocate sendmmsg scratch sized to this value, so exceeding it +// only costs additional sendmmsg chunks within a single WriteBatch call. +const MaxWriteBatch = 128 + type EncReader func( addr netip.AddrPort, payload []byte, ) +type Settings struct { + Listen netip.AddrPort + Multi bool + Batch int + Offloads bool +} + type Conn interface { Rebind() error LocalAddr() (netip.AddrPort, error) - ListenOut(r EncReader) error + // ListenOut invokes r for each received packet. + // On batch-capable backends (recvmmsg), flush is called after each batch is fully delivered. + // Callers use it to flush per-batch accumulators such as TUN write coalescers. + // Single-packet backends call flush after each packet. flush must not be nil. + ListenOut(r EncReader, flush func()) error WriteTo(b []byte, addr netip.AddrPort) error + // WriteBatch sends a contiguous batch of packets, each with its own + // destination. bufs and addrs must have the same length. Linux uses + // sendmmsg(2) for a single syscall. + // + // Returns the number of packets successfully written. A destination the kernel rejects costs only + // its own packet, so a short count means some peers were undeliverable, not that the batch failed. + // Not safe for concurrent use on the same Conn. + WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) ReloadConfig(c *config.C) SupportsMultipleReaders() bool Close() error @@ -31,7 +56,7 @@ func (NoopConn) Rebind() error { func (NoopConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil } -func (NoopConn) ListenOut(_ EncReader) error { +func (NoopConn) ListenOut(_ EncReader, _ func()) error { return nil } func (NoopConn) SupportsMultipleReaders() bool { @@ -40,6 +65,9 @@ func (NoopConn) SupportsMultipleReaders() bool { func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil } +func (NoopConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) { + return len(bufs), nil +} func (NoopConn) ReloadConfig(_ *config.C) { return } diff --git a/udp/udp_android.go b/udp/udp_android.go index 213ab422..9de6de2c 100644 --- a/udp/udp_android.go +++ b/udp/udp_android.go @@ -7,14 +7,13 @@ import ( "fmt" "log/slog" "net" - "net/netip" "syscall" "golang.org/x/sys/unix" ) -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - return NewGenericListener(l, ip, port, multi, batch) +func NewListener(l *slog.Logger, s Settings) (Conn, error) { + return NewGenericListener(l, s) } func NewListenConfig(multi bool) net.ListenConfig { diff --git a/udp/udp_bsd.go b/udp/udp_bsd.go index 31ae9c5a..833fc236 100644 --- a/udp/udp_bsd.go +++ b/udp/udp_bsd.go @@ -10,14 +10,13 @@ import ( "fmt" "log/slog" "net" - "net/netip" "syscall" "golang.org/x/sys/unix" ) -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - return NewGenericListener(l, ip, port, multi, batch) +func NewListener(l *slog.Logger, s Settings) (Conn, error) { + return NewGenericListener(l, s) } func NewListenConfig(multi bool) net.ListenConfig { diff --git a/udp/udp_darwin.go b/udp/udp_darwin.go index 574e4494..6d89d61b 100644 --- a/udp/udp_darwin.go +++ b/udp/udp_darwin.go @@ -27,9 +27,9 @@ type StdConn struct { var _ Conn = &StdConn{} -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - lc := NewListenConfig(multi) - pc, err := lc.ListenPacket(context.TODO(), "udp", net.JoinHostPort(ip.String(), fmt.Sprintf("%v", port))) +func NewListener(l *slog.Logger, s Settings) (Conn, error) { + lc := NewListenConfig(s.Multi) + pc, err := lc.ListenPacket(context.TODO(), "udp", s.Listen.String()) if err != nil { return nil, err } @@ -140,6 +140,22 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error { } } +func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + // An un-sendable destination costs its own packet, never the ones behind it in the batch. + // TODO: WriteTo maps EWOULDBLOCK to an error, so a full send buffer + // silently drops the rest of a burst (linux blocks instead). Poll for + // writability on EAGAIN before giving up on the remainder. + written := 0 + for i, b := range bufs { + if err := u.WriteTo(b, addrs[i]); err == nil { + written++ + } else { + u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err) + } + } + return written, nil +} + func (u *StdConn) LocalAddr() (netip.AddrPort, error) { a := u.UDPConn.LocalAddr() @@ -165,7 +181,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() { return func() {} } -func (u *StdConn) ListenOut(r EncReader) error { +func (u *StdConn) ListenOut(r EncReader, flush func()) error { buffer := make([]byte, MTU) for { @@ -179,7 +195,8 @@ func (u *StdConn) ListenOut(r EncReader) error { continue } - r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n]) + r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n]) + flush() } } diff --git a/udp/udp_generic.go b/udp/udp_generic.go index 131eb73b..0ba1d412 100644 --- a/udp/udp_generic.go +++ b/udp/udp_generic.go @@ -27,9 +27,9 @@ type GenericConn struct { var _ Conn = &GenericConn{} -func NewGenericListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - lc := NewListenConfig(multi) - pc, err := lc.ListenPacket(context.TODO(), "udp", net.JoinHostPort(ip.String(), fmt.Sprintf("%v", port))) +func NewGenericListener(l *slog.Logger, s Settings) (Conn, error) { + lc := NewListenConfig(s.Multi) + pc, err := lc.ListenPacket(context.TODO(), "udp", s.Listen.String()) if err != nil { return nil, err } @@ -44,6 +44,19 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error { return err } +func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + // An un-sendable destination costs its own packet, never the ones behind it in the batch. + written := 0 + for i, b := range bufs { + if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err == nil { + written++ + } else { + u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err) + } + } + return written, nil +} + func (u *GenericConn) LocalAddr() (netip.AddrPort, error) { a := u.UDPConn.LocalAddr() @@ -73,7 +86,7 @@ type rawMessage struct { Len uint32 } -func (u *GenericConn) ListenOut(r EncReader) error { +func (u *GenericConn) ListenOut(r EncReader, flush func()) error { buffer := make([]byte, MTU) var lastRecvErr time.Time @@ -93,7 +106,8 @@ func (u *GenericConn) ListenOut(r EncReader) error { continue } - r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n]) + r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n]) + flush() } } diff --git a/udp/udp_linux.go b/udp/udp_linux.go index 3920342c..fca713e5 100644 --- a/udp/udp_linux.go +++ b/udp/udp_linux.go @@ -1,5 +1,4 @@ //go:build !android && !e2e_testing -// +build !android,!e2e_testing package udp @@ -25,11 +24,18 @@ type StdConn struct { isV4 bool l *slog.Logger batch int + + // bw owns the sendmmsg/UDP-GSO transmit path: the per-queue write + // scratch and the GSO capability state probed at socket creation. See + // udp_linux_writebatch.go. + bw *batchWriter + + groSupported bool } -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { +func NewListener(l *slog.Logger, s Settings) (Conn, error) { af := unix.AF_INET6 - if ip.Is4() { + if s.Listen.Addr().Is4() { af = unix.AF_INET } syscall.ForkLock.RLock() @@ -42,7 +48,7 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) return nil, fmt.Errorf("unable to open socket: %w", err) } - if multi { + if s.Multi { if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil { _ = unix.Close(fd) return nil, fmt.Errorf("unable to set SO_REUSEPORT: %w", err) @@ -50,13 +56,14 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) } var sa unix.Sockaddr - if ip.Is4() { + port := int(s.Listen.Port()) + if s.Listen.Addr().Is4() { sa4 := &unix.SockaddrInet4{Port: port} - sa4.Addr = ip.As4() + sa4.Addr = s.Listen.Addr().As4() sa = sa4 } else { sa6 := &unix.SockaddrInet6{Port: port} - sa6.Addr = ip.As16() + sa6.Addr = s.Listen.Addr().As16() sa = sa6 } if err = unix.Bind(fd, sa); err != nil { @@ -64,7 +71,60 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) return nil, fmt.Errorf("unable to bind to socket: %w", err) } - return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, nil + out := &StdConn{sysFd: fd, isV4: s.Listen.Addr().Is4(), l: l, batch: s.Batch} + + out.bw = newBatchWriter(fd, out.isV4, l, s.Offloads) + + // GRO coalesces same-flow datagrams into superpackets that must be split back apart via the delivered gso_size cmsg + // batch == 1 means the caller wants plain single-datagram reads with MTU-sized buffers, so leave it off. + if s.Batch > 1 && s.Offloads { + out.prepareGRO() + } + + return out, nil +} + +// udpGROBufferSize sizes the per-entry recvmmsg buffer when UDP_GRO is on. +// The kernel stitches a run of same-flow datagrams into a single skb whose +// length is bounded by sk_gso_max_size (65535) +const udpGROBufferSize = 65535 + +// udpGROCmsgPayload is the size of the UDP_GRO cmsg data delivered by the +// kernel: a single int (gso_size in bytes). See udp_cmsg_recv() in net/ipv4/udp.c. +const udpGROCmsgPayload = 4 + +// prepareGRO turns on UDP_GRO so the kernel coalesces consecutive same-flow +// datagrams into one recvmmsg entry, with a cmsg carrying the gso_size used +// to split them back apart on the application side. +func (u *StdConn) prepareGRO() { + err := unix.SetsockoptInt(u.sysFd, unix.IPPROTO_UDP, unix.UDP_GRO, 1) + if err != nil { + u.l.Info("udp: GRO disabled", "reason", "kernel rejected probe", "error", err) + recordCapability("udp.gro.enabled", false) + return + } + u.groSupported = true + u.l.Info("udp: GRO enabled") + recordCapability("udp.gro.enabled", true) +} + +// recordCapability registers (or updates) a boolean gauge for one of the +// kernel-feature probes. Gauges go to 1 when the feature is enabled, 0 when +// it is not — dashboards can show degraded state on partially-supported +// kernels at a glance. Calling repeatedly with the same name updates the +// existing gauge rather than registering a duplicate. +// +// Caveat: the gauge is process-global while the capability state it reports +// is per-socket. With multiple listen routines the last probe wins, and a +// runtime downgrade on one socket (e.g. the GSO EIO disable) flips the gauge +// for all of them. Treat it as "at least one socket looks like this." +func recordCapability(name string, enabled bool) { + g := metrics.GetOrRegisterGauge(name, nil) + if enabled { + g.Update(1) + } else { + g.Update(0) + } } func (u *StdConn) SupportsMultipleReaders() bool { @@ -114,7 +174,7 @@ func (u *StdConn) LocalAddr() (netip.AddrPort, error) { } } -// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams +// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams. func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) { r, _, errno := unix.Syscall6( unix.SYS_RECVMMSG, @@ -138,40 +198,70 @@ func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) { return n, nil } -// recvmsg does one blocking recvmsg into msgs[0] -func (u *StdConn) recvmsg(msgs []rawMessage) (int, error) { - r, _, errno := unix.Syscall6( - unix.SYS_RECVMSG, - uintptr(u.sysFd), - uintptr(unsafe.Pointer(&msgs[0].Hdr)), - 0, - 0, - 0, - 0, - ) - if errno != 0 { - if u.closed.Load() { - return 0, net.ErrClosed +// prepareRawMessages allocates the recvmmsg scratch: +// n rawMessages, each wired to its own bufSize receive buffer, sockaddr name slot +// and, when cmsgSpace > 0, a slice of one contiguous ancillary-data slab. +// All iovecs share a single slab kept alive by the msghdrs that point into it. +func prepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) { + msgs := make([]rawMessage, n) + buffers := make([][]byte, n) + names := make([][]byte, n) + iovs := make([]iovec, n) + + var cmsgs []byte + if cmsgSpace > 0 { + cmsgs = make([]byte, n*cmsgSpace) + } + + for i := range msgs { + buffers[i] = make([]byte, bufSize) + names[i] = make([]byte, unix.SizeofSockaddrInet6) + + iovs[i].Base = &buffers[i][0] + setIovLen(&iovs[i], bufSize) + msgs[i].Hdr.Iov = &iovs[i] + setMsgIovlen(&msgs[i].Hdr, 1) + + msgs[i].Hdr.Name = &names[i][0] + msgs[i].Hdr.Namelen = uint32(len(names[i])) + + if cmsgSpace > 0 { + msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace] + setMsgControllen(&msgs[i].Hdr, cmsgSpace) } - return 0, &net.OpError{Op: "recvmsg", Err: errno} } - if r == 0 && u.closed.Load() { - return 0, net.ErrClosed - } - msgs[0].Len = uint32(r) - return 1, nil + + return msgs, buffers, names, cmsgs } -func (u *StdConn) ListenOut(r EncReader) error { +func getFrom(names [][]byte, i int, isV4 bool) netip.AddrPort { var ip netip.Addr - msgs, buffers, names := u.PrepareRawMessages(u.batch) - read := u.recvmmsg - if u.batch == 1 { - read = u.recvmsg + // Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic + if isV4 { + ip, _ = netip.AddrFromSlice(names[i][4:8]) + } else { + ip, _ = netip.AddrFromSlice(names[i][8:24]) } + return netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])) +} + +func (u *StdConn) ListenOut(r EncReader, flush func()) error { + bufSize := MTU + cmsgSpace := 0 + if u.groSupported { + bufSize = udpGROBufferSize + cmsgSpace = unix.CmsgSpace(udpGROCmsgPayload) + } + msgs, buffers, names, _ := prepareRawMessages(u.batch, bufSize, cmsgSpace) for { - n, err := read(msgs) + if cmsgSpace > 0 { + for i := range msgs { + setMsgControllen(&msgs[i].Hdr, cmsgSpace) + } + } + + n, err := u.recvmmsg(msgs) if err != nil { if errors.Is(err, unix.EINTR) { continue // interrupted by a signal, retry the read @@ -181,73 +271,128 @@ func (u *StdConn) ListenOut(r EncReader) error { return err } - for i := 0; i < n; i++ { - // Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic - if u.isV4 { - ip, _ = netip.AddrFromSlice(names[i][4:8]) - } else { - ip, _ = netip.AddrFromSlice(names[i][8:24]) + for i := range n { + from := getFrom(names, i, u.isV4) + payload := buffers[i][:msgs[i].Len] + + segSize := 0 + if cmsgSpace > 0 { + segSize = parseRecvCmsg(&msgs[i].Hdr) } - r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len]) + + deliverSegments(r, from, payload, segSize) } + + flush() } } +// deliverSegments hands a received superdatagram to r, splitting it back into pre-coalesce packets +func deliverSegments(r EncReader, from netip.AddrPort, payload []byte, segSize int) { + if segSize <= 0 || segSize >= len(payload) { //avoid bogus values + r(from, payload[:len(payload):len(payload)]) + return + } + for off := 0; off < len(payload); off += segSize { + end := off + segSize + if end > len(payload) { + end = len(payload) + } + r(from, payload[off:end:end]) + } +} + +// parseRecvCmsg walks the per-slot ancillary buffer and extracts the UDP_GRO +// gso_size, or 0 when no UDP_GRO cmsg is present. +func parseRecvCmsg(hdr *msghdr) (gso int) { + controllen := int(hdr.Controllen) + if controllen < unix.SizeofCmsghdr || hdr.Control == nil { + return 0 + } + ctrl := unsafe.Slice(hdr.Control, controllen) + off := 0 + for off+unix.SizeofCmsghdr <= len(ctrl) { + ch := (*unix.Cmsghdr)(unsafe.Pointer(&ctrl[off])) + clen := int(ch.Len) + // Compare against the remaining bytes rather than off+clen + if clen < unix.SizeofCmsghdr || clen > len(ctrl)-off { + return gso + } + dataOff := off + unix.CmsgLen(0) + if ch.Level == unix.SOL_UDP && ch.Type == unix.UDP_GRO { + if dataOff+udpGROCmsgPayload <= len(ctrl) { + gso = int(int32(binary.NativeEndian.Uint32(ctrl[dataOff : dataOff+udpGROCmsgPayload]))) + } + } + // Advance by the aligned cmsg space. + off += unix.CmsgSpace(clen - unix.CmsgLen(0)) + } + return gso +} + func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error { - if u.isV4 { - return u.writeTo4(b, ip) - } - return u.writeTo6(b, ip) + return sendto(u.sysFd, b, ip, u.isV4) } -func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error { - var rsa unix.RawSockaddrInet6 - rsa.Family = unix.AF_INET6 - rsa.Addr = ip.Addr().As16() - binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port()) - - for { - _, _, err := unix.Syscall6( - unix.SYS_SENDTO, - uintptr(u.sysFd), - uintptr(unsafe.Pointer(&b[0])), - uintptr(len(b)), - uintptr(0), - uintptr(unsafe.Pointer(&rsa)), - uintptr(unix.SizeofSockaddrInet6), - ) - if err != 0 { - return &net.OpError{Op: "sendto", Err: err} - } - return nil +func sendto(fd int, b []byte, addr netip.AddrPort, isV4 bool) error { + var rsa [unix.SizeofSockaddrInet6]byte + nlen, err := writeSockaddr(rsa[:], addr, isV4) + if err != nil { + return err } + var base *byte + if len(b) > 0 { + base = &b[0] + } + _, _, errno := unix.Syscall6( + unix.SYS_SENDTO, + uintptr(fd), + uintptr(unsafe.Pointer(base)), + uintptr(len(b)), + 0, + uintptr(unsafe.Pointer(&rsa[0])), + uintptr(nlen), + ) + if errno != 0 { + return &net.OpError{Op: "sendto", Err: errno} + } + return nil } -func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error { - if !ip.Addr().Is4() { - return ErrInvalidIPv6RemoteForSocket - } +// WriteBatch sends bufs via sendmmsg(2), coalescing same-destination runs into UDP-GSO superpackets when supported. +// See batchWriter in udp_linux_writebatch.go for the mechanics. +func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + return u.bw.WriteBatch(bufs, addrs) +} - var rsa unix.RawSockaddrInet4 - rsa.Family = unix.AF_INET - rsa.Addr = ip.Addr().As4() - binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port()) - - for { - _, _, err := unix.Syscall6( - unix.SYS_SENDTO, - uintptr(u.sysFd), - uintptr(unsafe.Pointer(&b[0])), - uintptr(len(b)), - uintptr(0), - uintptr(unsafe.Pointer(&rsa)), - uintptr(unix.SizeofSockaddrInet4), - ) - if err != 0 { - return &net.OpError{Op: "sendto", Err: err} +// writeSockaddr encodes addr into buf (which must be at least SizeofSockaddrInet6 bytes). +// Returns the number of bytes used. +// If isV4 is true and addr is not a v4 (or v4-in-v6) address, returns an error. +func writeSockaddr(buf []byte, addr netip.AddrPort, isV4 bool) (int, error) { + ap := addr.Addr().Unmap() + if isV4 { + if !ap.Is4() { + return 0, ErrInvalidIPv6RemoteForSocket } - return nil + // struct sockaddr_in: { sa_family_t(2), in_port_t(2, BE), in_addr(4), zero(8) } + // sa_family is host endian. + binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET) + binary.BigEndian.PutUint16(buf[2:4], addr.Port()) + ip4 := ap.As4() + copy(buf[4:8], ip4[:]) + for j := 8; j < 16; j++ { + buf[j] = 0 + } + return unix.SizeofSockaddrInet4, nil } + // struct sockaddr_in6: { sa_family_t(2), in_port_t(2, BE), flowinfo(4), in6_addr(16), scope_id(4) } + binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET6) + binary.BigEndian.PutUint16(buf[2:4], addr.Port()) + binary.NativeEndian.PutUint32(buf[4:8], 0) + ip6 := addr.Addr().As16() + copy(buf[8:24], ip6[:]) + binary.NativeEndian.PutUint32(buf[24:28], 0) + return unix.SizeofSockaddrInet6, nil } func (u *StdConn) ReloadConfig(c *config.C) { @@ -303,7 +448,7 @@ func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error { func (u *StdConn) Close() error { u.closed.Store(true) - // Wake the reader parked in recvmmsg/recvmsg. shutdown(2) on an unconnected socket + // Wake the reader parked in recvmmsg. shutdown(2) on an unconnected socket // returns ENOTCONN but still wakes it, so ignore the error. // The reader then sees closed and stops touching the fd, making the Close below safe. _ = unix.Shutdown(u.sysFd, unix.SHUT_RDWR) diff --git a/udp/udp_linux_32.go b/udp/udp_linux_32.go index de8f1cdf..efa2dba8 100644 --- a/udp/udp_linux_32.go +++ b/udp/udp_linux_32.go @@ -30,25 +30,18 @@ type rawMessage struct { Len uint32 } -func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) { - msgs := make([]rawMessage, n) - buffers := make([][]byte, n) - names := make([][]byte, n) +func setIovLen(v *iovec, n int) { + v.Len = uint32(n) +} - for i := range msgs { - buffers[i] = make([]byte, MTU) - names[i] = make([]byte, unix.SizeofSockaddrInet6) +func setMsgIovlen(m *msghdr, n int) { + m.Iovlen = uint32(n) +} - vs := []iovec{ - {Base: &buffers[i][0], Len: uint32(len(buffers[i]))}, - } +func setMsgControllen(m *msghdr, n int) { + m.Controllen = uint32(n) +} - msgs[i].Hdr.Iov = &vs[0] - msgs[i].Hdr.Iovlen = uint32(len(vs)) - - msgs[i].Hdr.Name = &names[i][0] - msgs[i].Hdr.Namelen = uint32(len(names[i])) - } - - return msgs, buffers, names +func setCmsgLen(h *unix.Cmsghdr, n int) { + h.Len = uint32(n) } diff --git a/udp/udp_linux_64.go b/udp/udp_linux_64.go index 48c5a978..65b72fd4 100644 --- a/udp/udp_linux_64.go +++ b/udp/udp_linux_64.go @@ -33,25 +33,18 @@ type rawMessage struct { Pad0 [4]byte } -func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) { - msgs := make([]rawMessage, n) - buffers := make([][]byte, n) - names := make([][]byte, n) +func setIovLen(v *iovec, n int) { + v.Len = uint64(n) +} - for i := range msgs { - buffers[i] = make([]byte, MTU) - names[i] = make([]byte, unix.SizeofSockaddrInet6) +func setMsgIovlen(m *msghdr, n int) { + m.Iovlen = uint64(n) +} - vs := []iovec{ - {Base: &buffers[i][0], Len: uint64(len(buffers[i]))}, - } +func setMsgControllen(m *msghdr, n int) { + m.Controllen = uint64(n) +} - msgs[i].Hdr.Iov = &vs[0] - msgs[i].Hdr.Iovlen = uint64(len(vs)) - - msgs[i].Hdr.Name = &names[i][0] - msgs[i].Hdr.Namelen = uint32(len(names[i])) - } - - return msgs, buffers, names +func setCmsgLen(h *unix.Cmsghdr, n int) { + h.Len = uint64(n) } diff --git a/udp/udp_linux_fixes_test.go b/udp/udp_linux_fixes_test.go new file mode 100644 index 00000000..8824b31f --- /dev/null +++ b/udp/udp_linux_fixes_test.go @@ -0,0 +1,708 @@ +//go:build linux && !android && !e2e_testing + +package udp + +import ( + "fmt" + "log/slog" + "net" + "net/netip" + "slices" + "testing" + "time" + "unsafe" + + "golang.org/x/sys/unix" +) + +// TestGSOMaxSegmentsKernelGate pins the corrected kernel-version gate: the +// 128-segment cap (127 usable) only lands in Linux v6.9 (commit 1382e3b6a350), +// not 5.5. Everything older stays at the conservative 63. +func TestGSOMaxSegmentsKernelGate(t *testing.T) { + cases := []struct { + release string + want int + }{ + {"5.4.0", 63}, + {"5.5.0-generic", 63}, // the old bug bumped here — it must not now + {"5.15.0", 63}, + {"6.1.0", 63}, + {"6.8.0-generic", 63}, + {"6.9.0", 127}, + {"6.10.1-arch1-1", 127}, + {"7.0.5-arch1-1", 127}, + {"garbage", 63}, + {"", 63}, + } + for _, c := range cases { + if got := gsoMaxSegments(c.release); got != c.want { + t.Errorf("gsoMaxSegments(%q) = %d, want %d", c.release, got, c.want) + } + } +} + +// buildCmsg lays out a single ancillary cmsg (header + data) in a fresh buffer +// the way the kernel would deliver it, so parseRecvCmsg can be exercised +// without a live socket. +func buildCmsg(level, typ int32, data []byte) []byte { + buf := make([]byte, unix.CmsgSpace(len(data))) + h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[0])) + h.Level = level + h.Type = typ + setCmsgLen(h, unix.CmsgLen(len(data))) + copy(buf[unix.CmsgLen(0):], data) + return buf +} + +func testLogger() *slog.Logger { + return slog.New(slog.DiscardHandler) +} + +// TestWriteBatchBadFamilyDeliversOthers is the H3 regression: a batch that +// contains one destination the socket can't reach (an IPv6 remote on a +// v4-bound socket) must still deliver every other packet. Before the fix the +// writeSockaddr error returned early and dropped the whole chunk. +func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) { + rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Skipf("cannot open v4 receiver (sandbox?): %v", err) + } + defer rx.Close() + rxPort := rx.LocalAddr().(*net.UDPAddr).Port + + // Bind a *non-wildcard* v4 address so Go gives us a genuine AF_INET + // socket. A wildcard v4 bind (0.0.0.0) via network "udp" comes up as a + // dual-stack AF_INET6 socket on Linux, for which a v6 dest is not a bad + // family — which would defeat the point of this test. + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: false, + Batch: 1, + Offloads: false, + } + c, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Skipf("cannot open v4 sender (sandbox?): %v", err) + } + defer c.Close() + sender := c.(*StdConn) + if !sender.isV4 { + t.Fatalf("expected a v4-bound sender socket, got isV4=false") + } + + good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort)) + bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // genuine v6, unreachable on v4 socket + + bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")} + addrs := []netip.AddrPort{good, bad, good} + + n, err := sender.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err) + } + if n != 2 { + t.Errorf("WriteBatch wrote %d packets, want 2 of 3 (the bad-family dest is the only casualty)", n) + } + + got := map[string]bool{} + rx.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + for i := 0; i < 2; i++ { + n, _, rerr := rx.ReadFromUDPAddrPort(buf) + if rerr != nil { + t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr) + } + got[string(buf[:n])] = true + } + if !got["AAA"] || !got["CCC"] { + t.Errorf("delivered set = %v, want AAA and CCC both present", got) + } + if got["BBB"] { + t.Errorf("the bad-family packet BBB was somehow delivered") + } +} + +// TestWriteBatchUnreachableDestDeliversOthers is the kernel-rejection twin of +// TestWriteBatchBadFamilyDeliversOthers. A destination the kernel refuses outright (240.0.0.0/4 is reserved, so +// the send returns EINVAL) fails its sendmmsg entry; WriteBatch must drop only that entry and still deliver +// every other packet rather than abandoning the batch at the first failure. +func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) { + rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Skipf("cannot open v4 receiver (sandbox?): %v", err) + } + defer rx.Close() + rxPort := rx.LocalAddr().(*net.UDPAddr).Port + + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: false, + Batch: 1, + Offloads: false, + } + c, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Skipf("cannot open v4 sender (sandbox?): %v", err) + } + defer c.Close() + sender := c.(*StdConn) + + good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort)) + bad := netip.MustParseAddrPort("240.0.0.1:9999") // reserved space, the kernel refuses it + + bufs := [][]byte{[]byte("P0"), []byte("P1"), []byte("BAD"), []byte("P3"), []byte("P4")} + addrs := []netip.AddrPort{good, good, bad, good, good} + + // The bad destination is reported, but only after every other packet has been attempted. + if _, err := sender.WriteBatch(bufs, addrs); err == nil { + t.Log("WriteBatch returned nil; kernel accepted the reserved address, delivery assertions still apply") + } + + got := map[string]bool{} + rx.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + for i := 0; i < 4; i++ { + n, _, rerr := rx.ReadFromUDPAddrPort(buf) + if rerr != nil { + t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got) + } + got[string(buf[:n])] = true + } + for _, want := range []string{"P0", "P1", "P3", "P4"} { + if !got[want] { + t.Errorf("packet %s was not delivered; delivered set = %v", want, got) + } + } +} + +// TestParseRecvCmsgCorruptLenNoPanic: a cmsg Len near max-int used to wrap +// off+clen negative, slip past the bounds check, and drive the walk offset +// negative -- a panic on the next ctrl[off]. The guard must compare Len +// against the remaining bytes instead. Also pins the plain truncated-Len +// cases (too small, larger than the buffer) to a clean early return. +func TestParseRecvCmsgCorruptLenNoPanic(t *testing.T) { + // First cmsg: a valid empty one so the walk advances past off=0 + // (off+clen can't overflow while off is still zero). + valid := buildCmsg(int32(unix.SOL_UDP), int32(unix.UDP_GRO), make([]byte, 4)) + + corrupt := func(lenVal int) []byte { + buf := make([]byte, len(valid)+unix.CmsgSpace(4)) + copy(buf, valid) + h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[len(valid)])) + h.Level = int32(unix.IPPROTO_IP) + h.Type = int32(unix.IP_TOS) + setCmsgLen(h, lenVal) + return buf + } + + cases := []struct { + name string + ctrl []byte + }{ + {"len_near_max_int", corrupt(int(^uint(0)>>1) - 8)}, + {"len_too_small", corrupt(unix.SizeofCmsghdr - 1)}, + {"len_past_buffer", corrupt(1 << 20)}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + hdr := &msghdr{Control: &c.ctrl[0]} + setMsgControllen(hdr, len(c.ctrl)) + gso := parseRecvCmsg(hdr) + // The valid leading UDP_GRO cmsg (payload 0) must still parse; + // the corrupt trailer just ends the walk. + if gso != 0 { + t.Errorf("parseRecvCmsg = %d, want 0", gso) + } + }) + } +} + +// TestDeliverSegments pins the GRO RX splitting: a kernel-coalesced buffer +// must come back out as the exact pre-coalesce packets -- every boundary +// error here shreds encrypted packets and every decrypt downstream fails. +func TestDeliverSegments(t *testing.T) { + from := netip.MustParseAddrPort("192.0.2.1:4242") + // Spare backing capacity mimics the recvmmsg row a real payload sits in; + // the cap checks below prove none of it leaks to a delivered segment. + pay := func(n int) []byte { + b := make([]byte, n, n+512) + for i := range b { + b[i] = byte(i) + } + return b + } + + cases := []struct { + name string + payload []byte + segSize int + wantLens []int + }{ + {"no-gro", pay(1400), 0, []int{1400}}, + {"negative-segsize", pay(1400), -5, []int{1400}}, + {"segsize-equals-payload", pay(1400), 1400, []int{1400}}, + {"segsize-past-payload", pay(1400), 2000, []int{1400}}, + {"even-split", pay(4200), 1400, []int{1400, 1400, 1400}}, + {"short-tail", pay(3000), 1400, []int{1400, 1400, 200}}, + {"single-byte-tail", pay(2801), 1400, []int{1400, 1400, 1}}, + {"segsize-one", pay(3), 1, []int{1, 1, 1}}, + {"empty-payload", pay(0), 1400, []int{0}}, + {"max-coalesce", pay(65500), 1372, nil}, // lens derived below + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + wantLens := c.wantLens + if wantLens == nil { + for rem := len(c.payload); rem > 0; rem -= c.segSize { + wantLens = append(wantLens, min(c.segSize, rem)) + } + } + + var got [][]byte + deliverSegments(func(a netip.AddrPort, seg []byte) { + if a != from { + t.Errorf("from = %v, want %v", a, from) + } + got = append(got, seg) + }, from, c.payload, c.segSize) + + if len(got) != len(wantLens) { + t.Fatalf("delivered %d segments, want %d", len(got), len(wantLens)) + } + // Segments must tile the payload in order with no gap, overlap, + // or copy: each must alias the payload at the right offset. + off := 0 + for i, seg := range got { + if len(seg) != wantLens[i] { + t.Fatalf("segment %d len=%d want %d", i, len(seg), wantLens[i]) + } + if cap(seg) != len(seg) { + // EncReader contract: an append into spare capacity would + // scribble into the next segment of the shared row. + t.Errorf("segment %d cap=%d, want %d (capacity must not reach into the row)", i, cap(seg), len(seg)) + } + if len(seg) > 0 && &seg[0] != &c.payload[off] { + t.Errorf("segment %d does not alias payload at offset %d", i, off) + } + off += len(seg) + } + if off != len(c.payload) { + t.Errorf("segments cover %d bytes, payload has %d", off, len(c.payload)) + } + }) + } +} + +// newRewindTestWriter builds a batchWriter with no socket: GSO planning on, +// sendFn left for the test to script. fd is invalid on purpose -- any path +// that actually hits the kernel fails loudly. +func newRewindTestWriter() *batchWriter { + w := &batchWriter{fd: -1, isV4: true, l: testLogger()} + // gsoSupported must be set before prepareWriteMessages: the cmsg slab is + // only allocated when GSO is already known to be supported. + w.gsoSupported = true + w.maxGSOSegments = 63 + w.prepareWriteMessages(MaxWriteBatch, true) + return w +} + +// capturePrepared decodes n prepared mmsghdr entries beginning at start +// straight from their iovecs -- ground truth, deliberately not the entryEnd +// bookkeeping the resume logic itself relies on. Returns one []byte per +// packed packet, in entry order. +func capturePrepared(w *batchWriter, start, n int) [][]byte { + var out [][]byte + for e := start; e < start+n; e++ { + hdr := &w.msgs[e].Hdr + iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen)) + for _, iov := range iovs { + b := make([]byte, int(iov.Len)) + if iov.Len > 0 { + copy(b, unsafe.Slice(iov.Base, int(iov.Len))) + } + out = append(out, b) + } + } + return out +} + +// TestWriteBatchPartialSendRewind drives WriteBatch through scripted +// partial sendmmsg results and asserts the rewind resumes exactly where +// the kernel stopped: every packet on the wire exactly once, in order, +// no duplicate, no loss. This is the hairiest logic in the write path +// and a rewind bug means silent packet duplication or loss under EAGAIN- +// style backpressure. +func TestWriteBatchPartialSendRewind(t *testing.T) { + dstA := netip.MustParseAddrPort("127.0.0.1:4242") + dstB := netip.MustParseAddrPort("127.0.0.2:4242") + + mkBuf := func(tag byte, n int) []byte { + b := make([]byte, n) + for i := range b { + b[i] = tag + } + b[0] = tag // tag identifies the packet uniquely below + return b + } + + // Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run + // tail), then two to B. The planner packs this as multiple entries with + // multi-iovec runs, which is what makes the rewind arithmetic hairy. + bufs := [][]byte{ + mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A + mkBuf(4, 600), // short tail to A + mkBuf(5, 900), mkBuf(6, 900), // run to B + } + addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB} + + scripts := [][]int{ + {99}, // accept everything first call + {1, 99}, // one entry per call, then the rest + {1, 1, 1, 99}, // strictly one entry per call + {2, 99}, // two entries, then the rest + } + for si, script := range scripts { + t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) { + w := newRewindTestWriter() + var wire [][]byte + call := 0 + w.sendFn = func(start, n int) (int, error) { + accept := n + if call < len(script) && script[call] < n { + accept = script[call] + } + call++ + wire = append(wire, capturePrepared(w, start, accept)...) + return accept, nil + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != len(bufs) { + t.Errorf("written = %d, want %d", written, len(bufs)) + } + if len(wire) != len(bufs) { + t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs)) + } + for i, b := range wire { + if len(b) != len(bufs[i]) || b[0] != bufs[i][0] { + t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)", + i, b[0], len(b), bufs[i][0], len(bufs[i])) + } + } + }) + } +} + +// TestWriteBatchSkipUnroutableRunAccounting: an unroutable destination mid- +// batch is skipped without committing an entry, leaving a hole in the bufs +// index space. The written count must tally packets per sent entry -- the +// index span would count the hole -- across both full and partial sendmmsg +// success. +func TestWriteBatchSkipUnroutableRunAccounting(t *testing.T) { + dstA := netip.MustParseAddrPort("127.0.0.1:4242") + dstB := netip.MustParseAddrPort("127.0.0.2:4242") + bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // v6 dest, v4 writer + + mk := func(tag byte, n int) []byte { + b := make([]byte, n) + b[0] = tag + return b + } + bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 500), mk(4, 900), mk(5, 900)} + addrs := []netip.AddrPort{dstA, dstA, bad, dstB, dstB} + + for si, script := range [][]int{{99}, {1, 99}} { + t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) { + w := newRewindTestWriter() + var wire [][]byte + call := 0 + w.sendFn = func(start, n int) (int, error) { + accept := n + if call < len(script) && script[call] < n { + accept = script[call] + } + call++ + wire = append(wire, capturePrepared(w, start, accept)...) + return accept, nil + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != 4 { + t.Errorf("written = %d, want 4 (the unroutable run is the only casualty)", written) + } + wantTags := []byte{1, 2, 4, 5} + if len(wire) != len(wantTags) { + t.Fatalf("wire got %d packets, want %d (dup or loss around the skip)", len(wire), len(wantTags)) + } + for i, b := range wire { + if b[0] != wantTags[i] { + t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i]) + } + } + }) + } +} + +// TestWriteBatchMidChunkRejectResumes: after a partial success, a zero-sent +// error on the FIRST REMAINING entry (done > 0) must drop only that entry's +// run and resume the rest of the chunk in place -- no repacking, no packets +// lost from entries before or after the rejected one. +func TestWriteBatchMidChunkRejectResumes(t *testing.T) { + dstA := netip.MustParseAddrPort("127.0.0.1:4242") + dstB := netip.MustParseAddrPort("127.0.0.2:4242") + dstC := netip.MustParseAddrPort("127.0.0.3:4242") + + mk := func(tag byte, n int) []byte { + b := make([]byte, n) + b[0] = tag + return b + } + // Three entries: a 2-packet GSO run to A, a 2-packet run to B, one to C. + bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 900), mk(4, 900), mk(5, 600)} + addrs := []netip.AddrPort{dstA, dstA, dstB, dstB, dstC} + + w := newRewindTestWriter() + var wire [][]byte + var starts []int + call := 0 + w.sendFn = func(start, n int) (int, error) { + starts = append(starts, start) + call++ + switch call { + case 1: // accept only entry 0 (the run to A) + wire = append(wire, capturePrepared(w, start, 1)...) + return 1, nil + case 2: // reject entry 1 (the run to B) outright + return -1, &net.OpError{Op: "sendmmsg", Err: unix.EPERM} + default: // accept the rest + wire = append(wire, capturePrepared(w, start, n)...) + return n, nil + } + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != 3 { + t.Errorf("written = %d, want 3 (B's rejected run is the only casualty)", written) + } + wantTags := []byte{1, 2, 5} + if len(wire) != len(wantTags) { + t.Fatalf("wire got %d packets, want %d (dup or loss around the mid-chunk reject)", len(wire), len(wantTags)) + } + for i, b := range wire { + if b[0] != wantTags[i] { + t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i]) + } + } + // The resume must reuse the prepared entries: same chunk, advancing + // start offsets, no repack (which would restart at 0 with fresh entries). + if want := []int{0, 1, 2}; !slices.Equal(starts, want) { + t.Errorf("sendFn start offsets = %v, want %v", starts, want) + } +} + +// TestWriteBatchMidChunkEIODisablesGSOWithoutDup: an EIO on a GSO entry +// after earlier entries in the chunk already went out must replay ONLY from +// the failed run (replanned as single-packet entries) -- the already-sent +// entries must not be duplicated. +func TestWriteBatchMidChunkEIODisablesGSOWithoutDup(t *testing.T) { + dstA := netip.MustParseAddrPort("127.0.0.1:4242") + dstB := netip.MustParseAddrPort("127.0.0.2:4242") + + mk := func(tag byte, n int) []byte { + b := make([]byte, n) + b[0] = tag + return b + } + // Entry 0: single packet to A. Entry 1: 2-packet GSO run to B. + bufs := [][]byte{mk(1, 600), mk(2, 1200), mk(3, 1200)} + addrs := []netip.AddrPort{dstA, dstB, dstB} + + w := newRewindTestWriter() + var wire [][]byte + call := 0 + w.sendFn = func(start, n int) (int, error) { + call++ + switch call { + case 1: // accept entry 0 only + wire = append(wire, capturePrepared(w, start, 1)...) + return 1, nil + case 2: // EIO on the GSO run to B + return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO} + default: // replanned single-packet replay + wire = append(wire, capturePrepared(w, start, n)...) + return n, nil + } + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if w.gsoSupported { + t.Error("gsoSupported still true after EIO on a GSO entry") + } + if written != len(bufs) { + t.Errorf("written = %d, want %d", written, len(bufs)) + } + wantTags := []byte{1, 2, 3} + if len(wire) != len(wantTags) { + t.Fatalf("wire got %d packets, want %d (packet 1 duplicated, or B's run lost)", len(wire), len(wantTags)) + } + for i, b := range wire { + if b[0] != wantTags[i] { + t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i]) + } + } +} + +// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an +// error rather than spin forever replaying the same chunk. +func TestWriteBatchZeroProgress(t *testing.T) { + w := newRewindTestWriter() + w.sendFn = func(start, n int) (int, error) { return 0, nil } + bufs := [][]byte{make([]byte, 100)} + addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")} + if _, err := w.WriteBatch(bufs, addrs); err == nil { + t.Fatal("WriteBatch = nil error on zero progress, want error") + } +} + +// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a +// sendmmsg rejected with EIO on a GSO superpacket entry must clear +// gsoSupported and replay the same packets as per-packet entries through +// sendmmsg (keeping batching), not fall back to per-packet sendto. +func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) { + dst := netip.MustParseAddrPort("127.0.0.1:4242") + bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)} + addrs := []netip.AddrPort{dst, dst, dst} + + w := newRewindTestWriter() + var entryCounts []int + call := 0 + w.sendFn = func(start, n int) (int, error) { + entryCounts = append(entryCounts, n) + call++ + if call == 1 { + return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO} + } + return n, nil + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if w.gsoSupported { + t.Error("gsoSupported still true after EIO on a GSO entry") + } + if written != len(bufs) { + t.Errorf("written = %d, want %d", written, len(bufs)) + } + // First call: one GSO entry carrying the whole run. Replay: one entry + // per packet, still via sendmmsg. + want := []int{1, 3} + if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] { + t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want) + } +} + +// TestGSOEngagesOnLoopback is the offload smoke test: real sockets, real +// UDP_SEGMENT cmsg, real kernel segmentation over loopback. It asserts +// both that GSO *engaged* (the whole batch left in a single sendmmsg +// entry -- a silent fallback to per-packet entries fails the test) and +// that the kernel carved the superpacket back into the exact original +// datagrams on the receive side. Runs in CI (make test on ubuntu-latest), +// which is what guards against the offload path silently degrading. +func TestGSOEngagesOnLoopback(t *testing.T) { + rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("listen rx: %v", err) + } + defer rx.Close() + dst := rx.LocalAddr().(*net.UDPAddr).AddrPort() + + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: false, + Batch: 8, + Offloads: true, + } + uc, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + sc := uc.(*StdConn) + defer sc.Close() + + if !sc.bw.gsoSupported { + var un unix.Utsname + _ = unix.Uname(&un) + release := string(un.Release[:]) + if major, minor := parseRelease(release); major > 4 || (major == 4 && minor >= 18) { + t.Fatalf("kernel %q supports UDP_SEGMENT but the GSO probe failed", release) + } + t.Skipf("kernel %q predates UDP_SEGMENT (4.18)", release) + } + + // Spy on the real syscall to count entries per sendmmsg without + // changing what hits the kernel. + var entryCounts []int + real := sc.bw.sendFn + sc.bw.sendFn = func(start, n int) (int, error) { + entryCounts = append(entryCounts, n) + return real(start, n) + } + + const numPkts = 8 + const pktLen = 1200 + bufs := make([][]byte, numPkts) + addrs := make([]netip.AddrPort, numPkts) + for i := range bufs { + bufs[i] = make([]byte, pktLen) + for j := range bufs[i] { + bufs[i][j] = byte(i) + } + addrs[i] = dst + } + + written, err := sc.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != numPkts { + t.Fatalf("written = %d, want %d", written, numPkts) + } + // GSO engaged means the run went out as ONE sendmmsg entry carrying a + // UDP_SEGMENT superpacket. Per-packet entries mean it silently fell + // back -- exactly the regression this test exists to catch. + if len(entryCounts) != 1 || entryCounts[0] != 1 { + t.Fatalf("sendmmsg entry counts = %v, want [1]: GSO did not engage", entryCounts) + } + + // The kernel must deliver the original datagram boundaries and bytes. + _ = rx.SetReadDeadline(time.Now().Add(5 * time.Second)) + got := make([]byte, pktLen+1) + for i := 0; i < numPkts; i++ { + n, _, err := rx.ReadFromUDP(got) + if err != nil { + t.Fatalf("rx read %d: %v", i, err) + } + if n != pktLen { + t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen) + } + for j := 0; j < n; j++ { + if got[j] != byte(i) { + t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i)) + } + } + } +} diff --git a/udp/udp_linux_offloads_test.go b/udp/udp_linux_offloads_test.go new file mode 100644 index 00000000..80648ca0 --- /dev/null +++ b/udp/udp_linux_offloads_test.go @@ -0,0 +1,277 @@ +//go:build linux && !android && !e2e_testing + +package udp + +import ( + "net" + "net/netip" + "slices" + "testing" + "time" + + "golang.org/x/sys/unix" +) + +// These tests pin the listen.udp_offloads=false behavior: no GSO/GRO probes, +// no cmsg scratch, and — critically — a still-functional send/receive path. +// The sockaddr name buffers are needed for every sendmmsg entry whether or +// not offloads are on, so prepareWriteMessages must allocate them even when +// it skips the cmsg slab (a nil name buffer panics in writeSockaddr on the +// first WriteBatch). + +// TestPrepareWriteMessagesAlwaysAllocatesNames covers all four +// (offloadsEnabled, gsoSupported) combinations: the sockaddr name buffers +// must exist in every one, and the cmsg slab only when both are true. +// gsoSupported=false with offloads enabled is the old-kernel path where the +// UDP_SEGMENT probe fails — not just a config choice. +func TestPrepareWriteMessagesAlwaysAllocatesNames(t *testing.T) { + cases := []struct { + name string + offloads bool + gso bool + }{ + {"offloads-off", false, false}, + {"offloads-on-probe-failed", true, false}, + {"offloads-off-gso-flag-set", false, true}, + {"offloads-on", true, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + w := &batchWriter{fd: -1, isV4: true, l: testLogger()} + w.gsoSupported = tc.gso + w.prepareWriteMessages(MaxWriteBatch, tc.offloads) + + for i := range w.msgs { + if len(w.names[i]) != unix.SizeofSockaddrInet6 { + t.Fatalf("names[%d] len=%d, want %d", i, len(w.names[i]), unix.SizeofSockaddrInet6) + } + if w.msgs[i].Hdr.Name == nil { + t.Fatalf("msgs[%d].Hdr.Name is nil", i) + } + } + + wantCmsg := tc.offloads && tc.gso + if (w.cmsg != nil) != wantCmsg { + t.Errorf("cmsg allocated = %v, want %v", w.cmsg != nil, wantCmsg) + } + }) + } +} + +// TestWriteBatchOffloadsDisabledScripted drives WriteBatch through a +// batchWriter built with offloads disabled and a scripted sendFn: every +// packet must become its own sendmmsg entry (no GSO coalescing to plan), +// packed correctly despite the missing cmsg slab. +func TestWriteBatchOffloadsDisabledScripted(t *testing.T) { + w := &batchWriter{fd: -1, isV4: true, l: testLogger()} + w.prepareWriteMessages(MaxWriteBatch, false) + + var entryCounts []int + w.sendFn = func(start, n int) (int, error) { + entryCounts = append(entryCounts, n) + return n, nil + } + + // Same destination, equal sizes: prime coalescing bait that must not + // coalesce with offloads off. + dst := netip.MustParseAddrPort("127.0.0.1:4242") + const numPkts = 4 + bufs := make([][]byte, numPkts) + addrs := make([]netip.AddrPort, numPkts) + for i := range bufs { + bufs[i] = make([]byte, 1200) + for j := range bufs[i] { + bufs[i][j] = byte(i) + } + addrs[i] = dst + } + + written, err := w.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != numPkts { + t.Errorf("written = %d, want %d", written, numPkts) + } + if len(entryCounts) != 1 || entryCounts[0] != numPkts { + t.Errorf("sendmmsg entry counts = %v, want [%d]: packets must be one entry each", entryCounts, numPkts) + } + + // Each prepared entry must carry exactly its own packet's bytes. + got := capturePrepared(w, 0, numPkts) + if len(got) != numPkts { + t.Fatalf("prepared %d packets, want %d", len(got), numPkts) + } + for i, pkt := range got { + if !slices.Equal(pkt, bufs[i]) { + t.Errorf("entry %d bytes differ from bufs[%d]", i, i) + } + } +} + +// TestOffloadsDisabledOnLoopback is the offloads-off smoke test, the mirror +// of TestGSOEngagesOnLoopback: a real socket built with Offloads=false must +// skip the GSO/GRO probes entirely (even on kernels that support them) and +// still deliver a same-destination batch as plain per-packet datagrams. +func TestOffloadsDisabledOnLoopback(t *testing.T) { + rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("listen rx: %v", err) + } + defer rx.Close() + dst := rx.LocalAddr().(*net.UDPAddr).AddrPort() + + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: false, + Batch: 8, // batch > 1 would enable GRO if Offloads did not gate it + Offloads: false, + } + uc, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + sc := uc.(*StdConn) + defer sc.Close() + + if sc.bw.gsoSupported { + t.Error("gsoSupported true with offloads disabled: probe was not skipped") + } + if sc.bw.cmsg != nil { + t.Error("cmsg slab allocated with offloads disabled") + } + if sc.groSupported { + t.Error("groSupported true with offloads disabled: probe was not skipped") + } + + var entryCounts []int + real := sc.bw.sendFn + sc.bw.sendFn = func(start, n int) (int, error) { + entryCounts = append(entryCounts, n) + return real(start, n) + } + + const numPkts = 8 + const pktLen = 1200 + bufs := make([][]byte, numPkts) + addrs := make([]netip.AddrPort, numPkts) + for i := range bufs { + bufs[i] = make([]byte, pktLen) + for j := range bufs[i] { + bufs[i][j] = byte(i) + } + addrs[i] = dst + } + + written, err := sc.WriteBatch(bufs, addrs) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != numPkts { + t.Fatalf("written = %d, want %d", written, numPkts) + } + // One sendmmsg call with one entry per packet: a single-entry call here + // means GSO engaged despite being disabled. + if len(entryCounts) != 1 || entryCounts[0] != numPkts { + t.Fatalf("sendmmsg entry counts = %v, want [%d]", entryCounts, numPkts) + } + + _ = rx.SetReadDeadline(time.Now().Add(5 * time.Second)) + got := make([]byte, pktLen+1) + for i := range numPkts { + n, _, err := rx.ReadFromUDP(got) + if err != nil { + t.Fatalf("rx read %d: %v", i, err) + } + if n != pktLen { + t.Fatalf("rx read %d: len=%d want %d", i, n, pktLen) + } + for j := range n { + if got[j] != byte(i) { + t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i)) + } + } + } +} + +// TestOffloadsDisabledRxDelivers exercises the receive path with GRO gated +// off but batch reads still on: ListenOut must deliver plain datagrams via +// the MTU-sized buffer layout (no cmsg slots). +func TestOffloadsDisabledRxDelivers(t *testing.T) { + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: false, + Batch: 8, + Offloads: false, + } + uc, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + sc := uc.(*StdConn) + + addr, err := sc.LocalAddr() + if err != nil { + t.Fatalf("LocalAddr: %v", err) + } + + type rxPkt struct { + from netip.AddrPort + payload []byte + } + rxCh := make(chan rxPkt, 16) + listenDone := make(chan struct{}) + go func() { + defer close(listenDone) + _ = sc.ListenOut(func(from netip.AddrPort, payload []byte) { + // payload aliases the shared recv buffer row; copy before handing off. + rxCh <- rxPkt{from, slices.Clone(payload)} + }, func() {}) + }() + + tx, err := net.DialUDP("udp4", nil, net.UDPAddrFromAddrPort(addr)) + if err != nil { + t.Fatalf("dial tx: %v", err) + } + defer tx.Close() + + want := [][]byte{ + []byte("one"), + make([]byte, 1200), + make([]byte, 9000), // near-MTU datagram must fit the non-GRO buffer size + } + for i := range want[1] { + want[1][i] = 0xAB + } + for i := range want[2] { + want[2][i] = 0xCD + } + for i, p := range want { + if _, err := tx.Write(p); err != nil { + t.Fatalf("tx write %d: %v", i, err) + } + } + + for i, p := range want { + select { + case got := <-rxCh: + if !slices.Equal(got.payload, p) { + t.Errorf("packet %d: payload differs (len=%d want %d)", i, len(got.payload), len(p)) + } + if got.from.Port() != tx.LocalAddr().(*net.UDPAddr).AddrPort().Port() { + t.Errorf("packet %d: from=%v, want sender port %d", i, got.from, tx.LocalAddr().(*net.UDPAddr).Port) + } + case <-time.After(5 * time.Second): + t.Fatalf("timed out waiting for packet %d", i) + } + } + + if err := sc.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + select { + case <-listenDone: + case <-time.After(5 * time.Second): + t.Fatal("ListenOut did not return after Close") + } +} diff --git a/udp/udp_linux_test.go b/udp/udp_linux_test.go index f9e7b3d8..b99f89eb 100644 --- a/udp/udp_linux_test.go +++ b/udp/udp_linux_test.go @@ -5,10 +5,8 @@ package udp import ( "errors" "fmt" - "log/slog" "net" "net/netip" - "os" "runtime" "sync/atomic" "testing" @@ -17,16 +15,18 @@ import ( "golang.org/x/sys/unix" ) -func testLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError})) -} - // TestShutdownWakesAfterRx_Mechanism exercises the kernel quirk our teardown // relies on: once a socket has received a packet, shutdown(2) wakes a blocked // recvmmsg with n>=1/Len==0 (not n==0). recvmmsg must turn that into net.ErrClosed // once Close set closed, so a parked reader exits instead of spinning. func TestShutdownWakesAfterRx_Mechanism(t *testing.T) { - c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, 64) + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: true, + Batch: 64, + Offloads: true, + } + c, err := NewListener(testLogger(), udpSettings) if err != nil { t.Fatalf("NewListener: %v", err) } @@ -35,7 +35,7 @@ func TestShutdownWakesAfterRx_Mechanism(t *testing.T) { if err != nil { t.Fatalf("LocalAddr: %v", err) } - msgs, _, _ := sc.PrepareRawMessages(sc.batch) + msgs, _, _, _ := prepareRawMessages(sc.batch, 0xffff, 16) // Receive a real packet so the socket has carried data. send, err := net.Dial("udp", addr.String()) @@ -109,8 +109,8 @@ func TestListenOutTeardown_TrafficPatterns(t *testing.T) { }}, } - // batch 1 exercises the recvmsg path, batch 64 the recvmmsg path; both must - // tear down cleanly. + // batch 1 exercises single-message reads, batch 64 a full recvmmsg batch; + // both must tear down cleanly. for _, batch := range []int{1, 64} { for _, tc := range cases { t.Run(fmt.Sprintf("batch%d/%s", batch, tc.name), func(t *testing.T) { @@ -121,7 +121,13 @@ func TestListenOutTeardown_TrafficPatterns(t *testing.T) { } func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net.Conn, stop <-chan struct{})) { - c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, batch) + udpSettings := Settings{ + Listen: netip.MustParseAddrPort("127.0.0.1:0"), + Multi: true, + Batch: batch, + Offloads: true, + } + c, err := NewListener(testLogger(), udpSettings) if err != nil { t.Fatalf("NewListener: %v", err) } @@ -136,7 +142,7 @@ func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net go func() { loopDone <- sc.ListenOut(func(netip.AddrPort, []byte) { received.Add(1) - }) + }, func() {}) }() send, err := net.Dial("udp", addr.String()) diff --git a/udp/udp_linux_writebatch.go b/udp/udp_linux_writebatch.go new file mode 100644 index 00000000..d53f3fb9 --- /dev/null +++ b/udp/udp_linux_writebatch.go @@ -0,0 +1,387 @@ +//go:build linux && !android && !e2e_testing + +package udp + +import ( + "context" + "encoding/binary" + "errors" + "fmt" + "log/slog" + "net" + "net/netip" + "strconv" + "strings" + "unsafe" + + "golang.org/x/sys/unix" +) + +// batchWriter owns the sendmmsg(2)/UDP-GSO transmit path for a StdConn: the +// scratch WriteBatch packs mmsghdr entries into, plus the GSO capability +// state probed at socket creation. Each queue has its own StdConn and +// batchWriter, so no locking is needed. +// +// Terminology, smallest to largest: +// +// packet one element of bufs: a single UDP datagram. The unit of the +// returned written count. +// run consecutive packets planRun groups into one entry: same +// destination, equal sizes (a shorter packet only last), within +// maxGSOBytes and maxGSOSegments. Without GSO a run is always one +// packet. Runs are atomic: packed whole into one entry, or +// skipped whole if the socket cannot address their destination, +// leaving a hole (bufs indices covered by no entry). +// entry one mmsghdr slot of the sendmmsg array; the kernel's unit of +// success and failure. A multi-packet entry carries a UDP_SEGMENT +// cmsg and is sent as one superpacket the kernel segments into +// gso_size-byte datagrams. Entries never split. +// chunk the entries packed for one sendmmsg call, at most MaxWriteBatch. +// batch the caller's whole bufs/addrs pair, processed as one or more chunks. +type batchWriter struct { + fd int + isV4 bool + + // UDP GSO (sendmsg with UDP_SEGMENT cmsg) support, probed once at + // socket creation and cleared by WriteBatch if the kernel later rejects + // a GSO send (the setsockopt probe cannot see per-route limitations). + // When true, WriteBatch coalesces runs into UDP_SEGMENT entries; + // otherwise each packet is its own entry. + gsoSupported bool + maxGSOSegments int + + // sendmmsg scratch, sized to MaxWriteBatch at construction; WriteBatch + // chunks larger inputs. + msgs []rawMessage + iovs []iovec + names [][]byte + l *slog.Logger + + // sendFn sends n prepared entries beginning at w.msgs[start] + // This is a function pointer to facilitate testing. + sendFn func(start, n int) (int, error) + + // Per-entry cmsg scratch: one contiguous slab of + // MaxWriteBatch * cmsgSpace bytes holding one UDP_SEGMENT cmsg per entry. + cmsg []byte + cmsgSpace int + + // entryEnd[e] is the bufs index after the last packet packed into entry e. + // entryEnd[e]-entryPkts[e] recovers the bufs index the entry's run started at, + // used to rewind i for the GSO-disable replay. + entryEnd []int + + // entryPkts[e] is the number of packets packed into entry e. + entryPkts []int +} + +func newBatchWriter(fd int, isV4 bool, l *slog.Logger, offloadsEnabled bool) *batchWriter { + w := &batchWriter{fd: fd, isV4: isV4, l: l} + w.sendFn = w.sendmmsg + if offloadsEnabled { + w.prepareGSO() + } + w.prepareWriteMessages(MaxWriteBatch, offloadsEnabled) + return w +} + +// prepareWriteMessages allocates the per-entry mmsghdr/iovec/sockaddr/cmsg +// scratch. Hdr.Iov/Iovlen/Control/Controllen are wired per call, since an +// entry spans a variable number of iovecs and may or may not carry a cmsg. +// +// Each entry's cmsg slot holds one UDP_SEGMENT (gso_size, uint16) header, +// pre-filled here; only its payload is rewritten per call. +// Hdr.Control/Controllen select whether it applies (none / segment). +func (w *batchWriter) prepareWriteMessages(n int, offloadsEnabled bool) { + w.msgs = make([]rawMessage, n) + w.iovs = make([]iovec, n) + w.names = make([][]byte, n) + w.entryEnd = make([]int, n) + w.entryPkts = make([]int, n) + + w.cmsgSpace = unix.CmsgSpace(2) + + for i := range w.msgs { + w.names[i] = make([]byte, unix.SizeofSockaddrInet6) + w.msgs[i].Hdr.Name = &w.names[i][0] + } + + if !offloadsEnabled || !w.gsoSupported { + return //avoid allocating cmsg space if we will never use it + } + + w.cmsg = make([]byte, n*w.cmsgSpace) + + for k := 0; k < n; k++ { + base := k * w.cmsgSpace + seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base])) + seg.Level = unix.SOL_UDP + seg.Type = unix.UDP_SEGMENT + setCmsgLen(seg, unix.CmsgLen(2)) + } +} + +// maxGSOBytes bounds the total payload of one UDP_SEGMENT send. The kernel +// builds a single skb, which must fit the 16-bit UDP length field and +// sk_gso_max_size (65536 on most devices); 65000 leaves headroom for headers. +const maxGSOBytes = 65000 + +// prepareGSO probes UDP_SEGMENT support and sets w.gsoSupported on success. +// Best-effort; failure leaves it false. +func (w *batchWriter) prepareGSO() { + w.maxGSOSegments = 63 // pre-6.9 cap; see gsoMaxSegments + + if err := unix.SetsockoptInt(w.fd, unix.IPPROTO_UDP, unix.UDP_SEGMENT, 0); err != nil { + w.l.Info("udp: GSO disabled", "reason", "rawconn control failed", "error", err) + recordCapability("udp.gso.enabled", false) + return + } + + var un unix.Utsname + if err := unix.Uname(&un); err != nil { + w.l.Warn("udp: kernel version probe failed, capping GSO at 63 segments", "error", err) + } else { + w.maxGSOSegments = gsoMaxSegments(string(un.Release[:])) + } + + w.gsoSupported = true + w.l.Info("udp: GSO enabled", "maxGSOSegments", w.maxGSOSegments) + recordCapability("udp.gso.enabled", true) +} + +// gsoMaxSegments returns the most segments one UDP_SEGMENT send may carry: +// the kernel cap (UDP_MAX_SEGMENTS: 64 before 6.9, 128 after) minus one, +// because the kernel counts the 8-byte UDP header against the gso_size * UDP_MAX_SEGMENTS budget. +func gsoMaxSegments(release string) int { + major, minor := parseRelease(release) + if major > 6 || (major == 6 && minor >= 9) { + return 127 + } + return 63 +} + +func parseRelease(r string) (major, minor int) { + // strip anything after the second dot or any non-digit + parts := strings.SplitN(r, ".", 3) + if len(parts) < 2 { + return 0, 0 + } + major, _ = strconv.Atoi(parts[0]) + // minor may have trailing junk like "15-generic" + mp := parts[1] + for i, c := range mp { + if c < '0' || c > '9' { + mp = mp[:i] + break + } + } + minor, _ = strconv.Atoi(mp) + return +} + +// WriteBatch sends bufs via sendmmsg(2), coalescing runs into UDP_SEGMENT +// entries, so one syscall can mix GSO superpackets and plain datagrams. +// Without GSO support every packet is its own entry. +// Callers shall deliver same-destination packets contiguously and in counter order +// +// Batches larger than the scratch take one sendmmsg per chunk. +// A partial success resumes the same prepared entries at the first unsent entry. +// A zero-sent error means the kernel rejected the first remaining entry: +// its packets are dropped and the rest of the chunk resumes in place. +// +// Returns the number of packets sent. An error means the call itself failed. +// A short count means some destinations were undeliverable. +func (w *batchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + if len(bufs) != len(addrs) { + return 0, fmt.Errorf("WriteBatch: len(bufs)=%d != len(addrs)=%d", len(bufs), len(addrs)) + } + + // A destination the kernel rejects results in us dropping that entry (one packet, or one same-destination GSO run). + // We count what actually made it out rather than returning an error. + written := 0 + + i := 0 + for i < len(bufs) { + entry := 0 + iovIdx := 0 + for entry < len(w.msgs) && i < len(bufs) { + iovBudget := len(w.iovs) - iovIdx + if iovBudget < 1 { + break + } + runLen, segSize := w.planRun(bufs, addrs, i, iovBudget) + if runLen == 0 { + break + } + + for k := 0; k < runLen; k++ { + b := bufs[i+k] + if len(b) == 0 { + w.iovs[iovIdx+k].Base = nil + setIovLen(&w.iovs[iovIdx+k], 0) + } else { + w.iovs[iovIdx+k].Base = &b[0] + setIovLen(&w.iovs[iovIdx+k], len(b)) + } + } + + nlen, err := writeSockaddr(w.names[entry], addrs[i], w.isV4) + if err != nil { + // The destination's address family does not match the socket + // (e.g. an IPv6 remote on a v4-bound socket). The packets are + // undeliverable and no entry is committed yet: skip the run. + if w.l.Enabled(context.Background(), slog.LevelDebug) { + w.l.Debug("skipping unroutable batch destination", "udpAddr", addrs[i], "packets", runLen, "error", err) + } + i += runLen + continue + } + + hdr := &w.msgs[entry].Hdr + hdr.Iov = &w.iovs[iovIdx] + setMsgIovlen(hdr, runLen) + hdr.Namelen = uint32(nlen) + + w.writeEntryCmsg(entry, runLen, segSize) + + i += runLen + iovIdx += runLen + w.entryEnd[entry] = i + w.entryPkts[entry] = runLen + entry++ + } + + if entry == 0 { + // Every remaining packet was skipped; i reached len(bufs). + break + } + + // Drain the packed entries without repacking: everything the packing + // loop wired (iovecs, names, cmsgs) stays intact until the next chunk + // overwrites it, so a partial success resumes the same sendmmsg array + // at the first unsent entry, and a rejected entry is skipped in place. + // Only the GSO-disable path replans, since its entries change shape. + done := 0 + for done < entry { + sent, serr := w.sendFn(done, entry-done) + if sent > 0 { + // Count packets per entry; the bufs index span would + // overcount across holes left by skipped runs. + for e := done; e < done+sent; e++ { + written += w.entryPkts[e] + } + done += sent + continue + } + if serr == nil { + return written, fmt.Errorf("sendmmsg made no progress") + } + // sent<=0 means the first remaining entry itself failed. + // EIO on a superpacket means the route cannot carry a GSO send even though the setsockopt probe passed: + // udp_send_skb() returns EIO when: + // * the egress device lacks TX checksum offload (kernels through 6.10) + // * or when an xfrm policy covers the route. + // Persistent, so disable GSO and replay from the failed run as one-packet entries. + if w.gsoSupported && w.entryPkts[done] >= 2 && errors.Is(serr, unix.EIO) { + w.gsoSupported = false + w.l.Warn("udp: kernel rejected GSO send, disabling GSO", "error", serr) + recordCapability("udp.gso.enabled", false) + i = w.entryEnd[done] - w.entryPkts[done] + break + } + // Any other zero-sent error is a per-entry failure. + // Transient errnos (EINTR, ENOBUFS) were already retried inside sendFn. + // These packets are doomed. Log them and move on. + if w.l.Enabled(context.Background(), slog.LevelDebug) { + w.l.Debug("sendmmsg rejected entry", + "error", serr, + "udpAddr", addrs[w.entryEnd[done]-w.entryPkts[done]], + "packets", w.entryPkts[done], + "gso", w.gsoSupported, + ) + } + done++ + } + } + return written, nil +} + +// planRun returns the length of the run starting at start and its segment +// size (len(bufs[start])). A run of length 1 carries no UDP_SEGMENT cmsg +// and is sent as a plain datagram; without GSO support planRun always returns 1. +func (w *batchWriter) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovBudget int) (int, int) { + if start >= len(bufs) || iovBudget < 1 { + return 0, 0 + } + segSize := len(bufs[start]) + if !w.gsoSupported || segSize == 0 || segSize > maxGSOBytes { + return 1, segSize + } + dst := addrs[start] + maxLen := w.maxGSOSegments + if iovBudget < maxLen { + maxLen = iovBudget + } + runLen := 1 + total := segSize + for runLen < maxLen && start+runLen < len(bufs) { + nextLen := len(bufs[start+runLen]) + if nextLen == 0 || nextLen > segSize { + break + } + if addrs[start+runLen] != dst { + break + } + if total+nextLen > maxGSOBytes { + break + } + total += nextLen + runLen++ + if nextLen < segSize { + // A short packet must be the last in the run. + break + } + } + return runLen, segSize +} + +// writeEntryCmsg writes one entry's UDP_SEGMENT payload when runLen >= 2 and +// points Hdr.Control at it; a single-packet entry carries no cmsg. +func (w *batchWriter) writeEntryCmsg(entry, runLen, segSize int) { + hdr := &w.msgs[entry].Hdr + base := entry * w.cmsgSpace + + if runLen >= 2 { + dataOff := base + unix.CmsgLen(0) + binary.NativeEndian.PutUint16(w.cmsg[dataOff:dataOff+2], uint16(segSize)) + hdr.Control = &w.cmsg[base] + setMsgControllen(hdr, w.cmsgSpace) + } else { + hdr.Control = nil + setMsgControllen(hdr, 0) + } +} + +// sendmmsg issues sendmmsg(2) against n entries of w.msgs starting at start. +// +// EINTR is automatically retried and will never be returned. +// ENOBUFS is retried enobufsRetries times, and should be treated like any other error +func (w *batchWriter) sendmmsg(start, n int) (int, error) { + const enobufsRetries = 3 + for enobufs := 0; ; { + r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, uintptr(w.fd), + uintptr(unsafe.Pointer(&w.msgs[start])), uintptr(n), + 0, 0, 0, + ) + switch { + case errno == unix.EINTR: //similar to stdlib's ignoringEINTRIO + continue + case errno == unix.ENOBUFS && enobufs < enobufsRetries: + enobufs++ //worth a retry or three + continue + case errno != 0: + return int(r1), &net.OpError{Op: "sendmmsg", Err: errno} + } + return int(r1), nil + } +} diff --git a/udp/udp_linux_writebatch_alloc_test.go b/udp/udp_linux_writebatch_alloc_test.go new file mode 100644 index 00000000..e1d96530 --- /dev/null +++ b/udp/udp_linux_writebatch_alloc_test.go @@ -0,0 +1,99 @@ +//go:build linux && !android && !e2e_testing + +package udp + +import ( + "net/netip" + "testing" +) + +// TestWriteBatchNoAllocs verifies the sendmmsg/UDP-GSO transmit path performs +// no per-packet heap allocations on the happy path: all mmsghdr/iovec/cmsg +// scratch is preallocated in newBatchWriter and WriteBatch may only rewrite +// it. The batch deliberately mixes a GSO-eligible run, a short tail segment, +// destination changes, so the planner, sockaddr, +// and cmsg paths are all exercised. +func TestWriteBatchNoAllocs(t *testing.T) { + for _, tc := range []struct { + name string + addr string + }{ + {"v4", "127.0.0.1"}, + {"v6", "::1"}, + } { + t.Run(tc.name, func(t *testing.T) { + ip := netip.MustParseAddr(tc.addr) + newConn := func() Conn { + udpSettings := Settings{ + Listen: netip.AddrPortFrom(ip, 0), + Multi: false, + Batch: 8, + Offloads: true, + } + c, err := NewListener(testLogger(), udpSettings) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + t.Cleanup(func() { _ = c.Close() }) + return c + } + tx := newConn() + rxA := newConn() + rxB := newConn() + if sc, ok := tx.(*StdConn); ok { + // Records which planner path the measurement covered; GSO + // support depends on the running kernel. + t.Logf("gsoSupported=%v maxGSOSegments=%d", sc.bw.gsoSupported, sc.bw.maxGSOSegments) + } + dstA, err := rxA.LocalAddr() + if err != nil { + t.Fatalf("LocalAddr: %v", err) + } + dstB, err := rxB.LocalAddr() + if err != nil { + t.Fatalf("LocalAddr: %v", err) + } + + payload := make([]byte, 1200) + short := make([]byte, 900) + + var bufs [][]byte + var addrs []netip.AddrPort + add := func(b []byte, dst netip.AddrPort) { + bufs = append(bufs, b) + addrs = append(addrs, dst) + } + // GSO-eligible run with a short tail. + for k := 0; k < 8; k++ { + add(payload, dstA) + } + add(short, dstA) + add(payload, dstA) + // Alternating destinations defeat coalescing entirely. + for k := 0; k < 4; k++ { + dst := dstA + if k%2 == 0 { + dst = dstB + } + add(payload, dst) + } + + var werr error + // Warm-up outside the measured runs. + if _, err := tx.WriteBatch(bufs, addrs); err != nil { + t.Fatalf("WriteBatch warm-up: %v", err) + } + allocs := testing.AllocsPerRun(100, func() { + if _, err := tx.WriteBatch(bufs, addrs); err != nil { + werr = err + } + }) + if werr != nil { + t.Fatalf("WriteBatch: %v", werr) + } + if allocs != 0 { + t.Fatalf("WriteBatch allocated %.1f times per call, want 0", allocs) + } + }) + } +} diff --git a/udp/udp_netbsd.go b/udp/udp_netbsd.go index b0c81393..b45222eb 100644 --- a/udp/udp_netbsd.go +++ b/udp/udp_netbsd.go @@ -9,14 +9,13 @@ import ( "fmt" "log/slog" "net" - "net/netip" "syscall" "golang.org/x/sys/unix" ) -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - return NewGenericListener(l, ip, port, multi, batch) +func NewListener(l *slog.Logger, s Settings) (Conn, error) { + return NewGenericListener(l, s) } func NewListenConfig(multi bool) net.ListenConfig { diff --git a/udp/udp_rio_windows.go b/udp/udp_rio_windows.go index d110af19..b04770af 100644 --- a/udp/udp_rio_windows.go +++ b/udp/udp_rio_windows.go @@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *slog.Logger, sa windows.Sockaddr) error { return nil } -func (u *RIOConn) ListenOut(r EncReader) error { +func (u *RIOConn) ListenOut(r EncReader, flush func()) error { buffer := make([]byte, MTU) var lastRecvErr time.Time @@ -161,7 +161,8 @@ func (u *RIOConn) ListenOut(r EncReader) error { continue } - r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n]) + r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n:n]) + flush() } } @@ -316,6 +317,19 @@ func (u *RIOConn) WriteTo(buf []byte, ip netip.AddrPort) error { return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0) } +func (u *RIOConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + // An un-sendable destination costs its own packet, never the ones behind it in the batch. + written := 0 + for i, b := range bufs { + if err := u.WriteTo(b, addrs[i]); err == nil { + written++ + } else { + u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err) + } + } + return written, nil +} + func (u *RIOConn) LocalAddr() (netip.AddrPort, error) { sa, err := windows.Getsockname(u.sock) if err != nil { diff --git a/udp/udp_tester.go b/udp/udp_tester.go index 9c0d989f..af6b09be 100644 --- a/udp/udp_tester.go +++ b/udp/udp_tester.go @@ -84,14 +84,14 @@ type TesterConn struct { l *slog.Logger } -func NewListener(l *slog.Logger, ip netip.Addr, port int, _ bool, _ int) (Conn, error) { +func NewListener(l *slog.Logger, s Settings) (Conn, error) { c := &TesterConn{ RxPackets: make(chan *Packet, 10), TxPackets: make(chan *Packet, 10), done: make(chan struct{}), l: l, } - c.SetAddr(netip.AddrPortFrom(ip, uint16(port))) + c.SetAddr(s.Listen) return c, nil } @@ -171,14 +171,28 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error { return nil } } +func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) { + written := 0 + for i, b := range bufs { + if err := u.WriteTo(b, addrs[i]); err == nil { + written++ + } else { + u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err) + } + } + return written, nil +} -func (u *TesterConn) ListenOut(r EncReader) error { +func (u *TesterConn) ListenOut(r EncReader, flush func()) error { for { select { case <-u.done: return os.ErrClosed case p := <-u.RxPackets: - r(p.From, p.Data) + r(p.From, p.Data[:len(p.Data):len(p.Data)]) + // The batcher borrows plaintext decrypted in place inside p.Data + // until Flush, so the packet must stay alive across flush() + flush() p.Release() } } diff --git a/udp/udp_windows.go b/udp/udp_windows.go index 1f34f0bc..57c14b59 100644 --- a/udp/udp_windows.go +++ b/udp/udp_windows.go @@ -7,12 +7,11 @@ import ( "fmt" "log/slog" "net" - "net/netip" "syscall" ) -func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - if multi { +func NewListener(l *slog.Logger, s Settings) (Conn, error) { + if s.Multi { //NOTE: Technically we can support it with RIO but it wouldn't be at the socket level // The udp stack would need to be reworked to hide away the implementation differences between // Windows and Linux @@ -20,12 +19,12 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) } var conn Conn - rc, err := NewRIOListener(l, ip, port) + rc, err := NewRIOListener(l, s.Listen.Addr(), int(s.Listen.Port())) if err == nil { conn = rc } else { l.Error("Falling back to standard udp sockets", "error", err) - conn, err = NewGenericListener(l, ip, port, multi, batch) + conn, err = NewGenericListener(l, s) if err != nil { return nil, err } diff --git a/util/cpupin_linux.go b/util/cpupin_linux.go new file mode 100644 index 00000000..d975b4af --- /dev/null +++ b/util/cpupin_linux.go @@ -0,0 +1,49 @@ +//go:build linux && !android && !e2e_testing + +package util + +import ( + "runtime" + + "golang.org/x/sys/unix" +) + +// PinThreadToCPU restricts the calling OS thread to the given CPU via +// sched_setaffinity(2). Combined with runtime.LockOSThread on the +// goroutine, this prevents the kernel from migrating us across CPUs and +// in turn keeps every sendmmsg from this goroutine going through the +// same XPS-selected TX ring, eliminating the wire-side reorder that +// otherwise fragments one nebula flow across multiple rings. +func PinThreadToCPU(cpu int) error { + runtime.LockOSThread() + var set unix.CPUSet + set.Zero() + set.Set(cpu) + if err := unix.SchedSetaffinity(0, &set); err != nil { + // Without the affinity the thread lock buys no TX-ring stability; + // don't leave the goroutine wedded to one OS thread for nothing. + runtime.UnlockOSThread() + return err + } + return nil +} + +// AllowedCPUs returns the CPU IDs the calling process is currently allowed to +// run on, as reported by sched_getaffinity(2). Under a cgroup cpuset or a +// `taskset` mask the allowed IDs are frequently not the contiguous range +// 0..NumCPU-1 (e.g. pinned to CPUs 4-7: NumCPU reports 4 while the valid IDs +// are 4,5,6,7). Callers that need a real CPU to pin to must choose from this +// set rather than assuming i % NumCPU is runnable, or every pin fails. +func AllowedCPUs() ([]int, error) { + var set unix.CPUSet + if err := unix.SchedGetaffinity(0, &set); err != nil { + return nil, err + } + cpus := make([]int, 0, set.Count()) + for cpu := 0; cpu < len(set)*64; cpu++ { + if set.IsSet(cpu) { + cpus = append(cpus, cpu) + } + } + return cpus, nil +} diff --git a/util/cpupin_other.go b/util/cpupin_other.go new file mode 100644 index 00000000..50522b9a --- /dev/null +++ b/util/cpupin_other.go @@ -0,0 +1,18 @@ +//go:build !linux || android || e2e_testing + +package util + +// PinThreadToCPU is a no-op outside Linux: only Linux exposes a stable +// per-thread CPU affinity API and only Linux has XPS-driven TX ring +// selection in the first place. On every other platform there's nothing +// to fix here. +func PinThreadToCPU(_ int) error { + return nil +} + +// AllowedCPUs has no meaningful answer off Linux (no sched_getaffinity), so it +// reports "unknown" by returning a nil slice and nil error. Callers treat an +// empty result as "fall back to the default CPU choice". +func AllowedCPUs() ([]int, error) { + return nil, nil +}