Compare commits

..

37 Commits

Author SHA1 Message Date
JackDoan 7fe4fab167 tun: default pin CPUs avoid NIC IRQ cores
When tun.pin_threads is on and tun.cpu_affinity is unset, pick pin CPUs
from the allowed set that do not service any up physical NIC's MSI vectors
(read-only walk of /sys/class/net/*/device/msi_irqs and
/proc/irq/*/effective_affinity_list). The old allowed[i] default pinned
encrypt threads onto exactly the cores drivers affine their first RX queue
IRQs to; a flow whose RSS queue fired there had NAPI fighting encrypt for
the core (measured 8.4 vs 10.2 Gbps REV bimodality). Falls back to the old
spread, with a log, when there aren't enough IRQ-free CPUs - e.g. drivers
that allocate one queue per core until the admin narrows them (ethtool -L).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-14 15:20:04 -05:00
JackDoan 37b924945d unslop some comments 2026-07-14 13:38:39 -05:00
JackDoan 5631346b07 batch: back SendBatch with Arena instead of a hand-rolled slab
SendBatch.Reserve duplicated Arena's grow-on-demand logic byte for byte.
Use an Arena for the slot backing so the borrow/grow/recycle semantics
live in one place.
2026-07-14 12:19:50 -05:00
JackDoan 97eb3c635a batch: move shared-arena Reset ownership from lanes to their owner 2026-07-14 12:19:50 -05:00
JackDoan 05f7923860 re-align to master 2026-07-14 11:38:43 -05:00
JackDoan 49028cb755 some tests 2026-07-14 11:00:55 -05:00
JackDoan adf71d1458 simplify making new Queues 2026-07-14 11:00:55 -05:00
JackDoan cefda6524c make service test less annoying 2026-07-14 11:00:55 -05:00
JackDoan d0f14de739 checkpt 2026-07-14 11:00:55 -05:00
JackDoan 9e7646ee62 checkpt 2026-07-14 11:00:55 -05:00
JackDoan 6e6cfc89db more ram -> more speed 2026-07-14 11:00:55 -05:00
JackDoan 3719f135e3 more fixes! 2026-07-14 11:00:55 -05:00
JackDoan 2724b4a96c more fixes! 2026-07-14 11:00:55 -05:00
JackDoan e386e290ab lint 2026-07-14 11:00:55 -05:00
JackDoan 0a44376403 datapath: fix 12 correctness findings from tun/UDP offload review
Multi-disciplinary correctness review of the batched tun / GSO-GRO / sendmmsg
rework. Each fix has a regression test; the merged tree builds on
linux/darwin/openbsd/windows/freebsd/netbsd, vets clean, passes the unit and
e2e suites, and is -race clean.

Critical:
- C1 zero-length inner UDP datagram no longer panics the process (remote DoS):
  the UDP coalescer routes payLen==0 to passthrough instead of seeding a GSO
  slot, and WriteGSO skips empty payload iovecs as defense in depth.
- C2 segmenter no longer corrupts inner headers when gsoSize < headerLen: the
  L3+L4 header is snapshotted once and each segment stamped from the copy,
  replacing the destructive overlapping in-place slide (SegmentTCP + SegmentUDP).

High:
- H1 applyOuterECN updates the IPv4 header checksum (RFC 1624 incremental) when
  folding outer CE into the inner ToS, so passthrough packets are no longer
  dropped by the peer stack.
- H2 the GRO reject path caps the borrowed RX segment ([:n:n]) so a reject can
  no longer overrun into the next coalesced segment's Nebula header. Note:
  oversized ICMPv6 rejects that need >16B beyond the segment are now refused
  rather than sent under GRO (safe; see TOFIX.md for the scratch-buffer follow-up).
- H3 WriteBatch falls back to per-packet WriteTo for a chunk when writeSockaddr
  fails, so one bad-family destination costs only its own packet, not the batch.
- H4 UserDevice.Readers returns N distinct queue wrappers with private buffers
  (sharing the pipes) so concurrent readers no longer race/overwrite borrowed
  packet bytes.
- H5 Poll.Close / Offload.Close no longer null t.fd (matching master's
  tunFile.Close), removing the data race with a concurrent readOne load.

Medium/Low:
- M1 the UDP GSO 127-segment gate moved from kernel >=5.5 to >=6.9 (the real
  UDP_MAX_SEGMENTS 64->128 threshold), avoiding EINVAL + per-packet fallback on
  5.5-6.8 kernels.
- M2 NewMultiQueueReader replays the offload mask newTun actually negotiated
  instead of the TSO-only mask, so adding a queue no longer disables USO
  device-wide; the advertised USO capability derives from the same mask.
- M3 the shutdown eventfd is closed in pollQueueSet.Close / offloadQueueSet.Close
  (double-close guarded), fixing the per-lifecycle fd leak.
- M4 dual-stack ECN selects the cmsg by address family, not socket family: RX
  parseRecvCmsg reads both IP_TOS and IPV6_TCLASS; TX writeEntryCmsg stamps
  IP_TOS for v4/v4-mapped dests and IPV6_TCLASS for v6 (on-host verified).
- L1 newPoll no longer closes the fd on failure (matching newOffload), removing
  the double-close on QueueSet.Add error.
2026-07-14 11:00:55 -05:00
JackDoan 9e61269935 make mobile happy 2026-07-14 11:00:55 -05:00
JackDoan a0836aa819 correctly shutdown the pprofserver 2026-07-14 11:00:55 -05:00
JackDoan cf2800f7bd SendVia: don't emit a zero-length packet when prepareSendVia fails 2026-07-14 11:00:55 -05:00
JackDoan fc89b9d14c adapt Control lifecycle tests to the batched tio.Queue Device interface 2026-07-14 11:00:55 -05:00
JackDoan 3e5326e60a udp setsockopt correctness fixes 2026-07-14 11:00:55 -05:00
JackDoan 840841a53c use less ram pls 2026-07-14 11:00:55 -05:00
JackDoan cfb4ab24c3 clean up a comment a bit 2026-07-14 11:00:55 -05:00
JackDoan 67f9cfad91 drop in a logger 2026-07-14 11:00:55 -05:00
JackDoan 1c78f5e500 go mod tidy 2026-07-14 11:00:55 -05:00
JackDoan ad8ff2e45e lint 2026-07-14 11:00:55 -05:00
JackDoan c13a1ff4ec fix 2026-07-14 11:00:55 -05:00
JackDoan 8b66ebaf72 faster
grr heap usage!
2026-07-14 11:00:55 -05:00
JackDoan bbdaf5c3f3 no 2026-07-14 11:00:55 -05:00
JackDoan a515aeaf39 use clear() 2026-07-14 11:00:55 -05:00
JackDoan 0f6f14eaf6 remove udp-level RX reorder buf 2026-07-14 11:00:55 -05:00
JackDoan dafdd34af9 make relays take the fast path maybe 2026-07-14 11:00:55 -05:00
JackDoan f6791130df scoot pinning around 2026-07-14 11:00:55 -05:00
JackDoan 145a6267fa scoot stuff around for e2e 2026-07-14 11:00:55 -05:00
JackDoan 7716f1da23 disable sort-on-RX, CPU pinning seems to work for now 2026-07-14 11:00:55 -05:00
JackDoan beb1d7d89f switch to ASM vector checksum 2026-07-14 11:00:55 -05:00
JackDoan fa0d593f28 GSO/GRO offloads, with TCP+ECN and UDP support 2026-07-14 11:00:55 -05:00
JackDoan ff48040c78 better and batched tun interface 2026-07-14 11:00:55 -05:00
109 changed files with 4197 additions and 8016 deletions
+6 -6
View File
@@ -12,9 +12,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -38,9 +38,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -78,9 +78,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Import certificates
+6 -6
View File
@@ -32,9 +32,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: add hashicorp source
@@ -64,9 +64,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: add hashicorp source
@@ -90,9 +90,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
+2 -2
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: build
+7 -7
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Install goimports
@@ -42,7 +42,7 @@ jobs:
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
version: v2.12
version: v2.5
test:
name: Test ${{ matrix.name }}
@@ -80,9 +80,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -125,9 +125,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build ${{ matrix.name }}
-82
View File
@@ -7,88 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
## [1.11.0] - 2026-07-23
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
### Breaking
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
one today and likely want to swap them before upgrading. (#1798)
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
`tun.network_category` to `unset` to keep the old behavior. (#1710)
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
directory set. The directory is not created for you. (#1622)
### Added
- Sign the Windows release binaries. (#1718)
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
- Add version labels to the Docker/OCI images. (#1772)
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
### Changed
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
- Update a static host's addresses when they change on reload. (#1713)
- Don't require a port on ICMP firewall rules. (#1609)
- Connection track ICMP traffic. (#1602)
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
- Record the local host's details in the DNS server. (#1716)
- Install Windows unsafe routes as link routes. (#1709)
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
changes. (#1733, #1765, #1810)
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
instead of leaking them. (#1794)
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
- Update to build against go v1.26. (#1818)
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
### Fixed
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
- Don't panic in the DNS server on a short or empty query name. (#1635)
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
- Fix a race in relay state handling. (#1753)
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
- Properly handle `closetunnel` packets. (#1638)
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
- Don't build new relay state on a tunnel that was just discarded. (#1796)
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
- Don't call the packet reader after a UDP error on Darwin. (#1755)
- Open the FreeBSD tun device non blocking. (#1666)
## [1.10.3] - 2026-02-06
### Security
-96
View File
@@ -1,96 +0,0 @@
//go:build linux && !android && !e2e_testing
package main
import (
"fmt"
"net/netip"
"os"
"path/filepath"
"runtime"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
cert_test "github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/require"
)
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
// a library, and on a config update dnclient calls Stop() in-process to tear the
// old instance down before starting a new one. This boots a real nebula (real
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
// dump instead of relying on a process signal to unstick them.
func TestControlStopClosesOnTimer(t *testing.T) {
l := test.NewLogger()
dir := t.TempDir()
before := time.Now().Add(-time.Hour)
after := time.Now().Add(time.Hour)
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
caPath := filepath.Join(dir, "ca.pem")
certPath := filepath.Join(dir, "cert.pem")
keyPath := filepath.Join(dir, "key.pem")
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
// tun disabled so no device/root is needed; routines: 2 so we exercise the
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
configBody := fmt.Sprintf(`
pki:
ca: %s
cert: %s
key: %s
listen:
host: 127.0.0.1
port: 0
tun:
disabled: true
firewall:
outbound:
- port: any
proto: any
host: any
inbound:
- port: any
proto: any
host: any
routines: 2
`, caPath, certPath, keyPath)
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
c := config.NewC(l)
require.NoError(t, c.Load(dir))
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
require.NoError(t, err)
require.NoError(t, ctrl.Start())
// Run like a live nebula, then close on a timer, exactly as dnclient does.
<-time.NewTimer(5 * time.Second).C
stopped := make(chan struct{})
go func() {
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
ctrl.Wait() // blocks until every reader goroutine has returned
close(stopped)
}()
select {
case <-stopped:
t.Log("nebula closed cleanly on timer")
case <-time.After(10 * time.Second):
buf := make([]byte, 1<<20)
n := runtime.Stack(buf, true)
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
}
}
-1
View File
@@ -25,7 +25,6 @@ func newTestLighthouse() *LightHouse {
lighthouses := []netip.Addr{}
staticList := map[netip.Addr]struct{}{}
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
lh.lighthouses.Store(&lighthouses)
lh.staticList.Store(&staticList)
-65
View File
@@ -2,26 +2,16 @@ package nebula
import (
"encoding/json"
"log/slog"
"sync"
"sync/atomic"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/handshake"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/noiseutil"
)
const ReplayWindow = 8192
// 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
@@ -30,10 +20,7 @@ type ConnectionState struct {
initiator bool
messageCounter atomic.Uint64
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
@@ -48,7 +35,6 @@ func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
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++ {
@@ -68,54 +54,3 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
func (cs *ConnectionState) Curve() cert.Curve {
return cs.myCert.Curve()
}
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()
if !result {
return nil, ErrAlreadySeen
}
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
if err != nil {
return nil, err
}
cs.decryptLock.Lock()
result = cs.window.Update(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return nil, ErrAlreadySeen
}
return out, nil
}
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
cs.decryptLock.Lock()
result := cs.window.Check(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
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)
if err != nil {
return err
}
cs.decryptLock.Lock()
result = cs.window.Update(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return ErrAlreadySeen
}
return nil
}
+1 -9
View File
@@ -53,7 +53,6 @@ type Control struct {
statsStart func()
dnsStart func()
lighthouseStart func()
networkChangeStart func(rebind func())
connectionManagerStart func(context.Context)
}
@@ -105,9 +104,6 @@ func (c *Control) Start() error {
if c.dnsStart != nil {
go c.dnsStart()
}
if c.networkChangeStart != nil {
go c.networkChangeStart(c.RebindUDPServer)
}
if c.connectionManagerStart != nil {
go c.connectionManagerStart(c.ctx)
}
@@ -202,11 +198,7 @@ func (c *Control) RebindUDPServer() {
return
}
// A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
// unlikely to help. Say so instead of silently carrying on as if we rebound.
if err := c.f.outside.Rebind(); err != nil {
c.l.Error("Failed to rebind udp socket", "error", err)
}
_ = c.f.outside.Rebind()
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
c.f.lightHouse.SendUpdate()
+4 -4
View File
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]*batch.MultiCoalescer, 1),
batchers: make([]batch.RxBatcher, 1),
routines: 1,
hostMap: newHostMap(l),
lightHouse: lh,
@@ -148,8 +148,8 @@ func (c *fakeConn) Rebind() error { c.rebinds++; ret
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) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
return nil
}
func (c *fakeConn) ReloadConfig(_ *config.C) {}
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
@@ -177,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]*batch.MultiCoalescer, 2),
batchers: make([]batch.RxBatcher, 2),
routines: 2,
l: test.NewLogger(),
}
+1 -13
View File
@@ -108,19 +108,7 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
}
func (c *Control) GetUDPAddr() netip.AddrPort {
return c.f.outside.(*udp.TesterConn).GetAddr()
}
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
// network. Register the new address with the router as well or nothing will route back.
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
c.f.outside.(*udp.TesterConn).SetAddr(addr)
}
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
c.f.lightHouse.localAddrsFn = fn
return c.f.outside.(*udp.TesterConn).Addr
}
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
-187
View File
@@ -1,187 +0,0 @@
// 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)
}
-171
View File
@@ -1,171 +0,0 @@
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)
}
}
-154
View File
@@ -1,154 +0,0 @@
//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)))
}
-163
View File
@@ -1,163 +0,0 @@
//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)
}
}
}
-10
View File
@@ -1,10 +0,0 @@
//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, ""
}
-118
View File
@@ -1,118 +0,0 @@
//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
}
-111
View File
@@ -1,111 +0,0 @@
//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)
}
}
-10
View File
@@ -1,10 +0,0 @@
//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)
}
+7 -16
View File
@@ -97,7 +97,8 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
newAddr := getDnsServerAddr(c)
d.serverMu.Lock()
running := d.server != nil
running := d.server
runningStarted := d.started
sameAddr := d.addr == newAddr
d.addr = newAddr
d.enabled.Store(enabled)
@@ -111,7 +112,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
}
if !enabled {
if running {
if running != nil {
d.Stop()
}
// Drop any records that accumulated while enabled; a later re-enable
@@ -120,12 +121,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
return nil
}
if !running {
if running == nil {
// Was disabled (or never started); bring it up now.
go d.Start()
} else if !sameAddr {
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
d.Stop()
d.shutdownServer(running, runningStarted, "reload")
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
go d.Start()
}
@@ -161,9 +162,7 @@ func (d *dnsServer) Start() {
started := make(chan struct{})
d.serverMu.Lock()
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
if d.ctx.Err() != nil {
d.serverMu.Unlock()
return
}
@@ -201,14 +200,6 @@ func (d *dnsServer) Start() {
close(started)
}
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
d.serverMu.Lock()
if d.server == server {
d.server = nil
d.started = nil
}
d.serverMu.Unlock()
if err != nil {
d.l.Warn("Failed to run the DNS responder", "error", err)
}
+4 -206
View File
@@ -194,51 +194,14 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
}
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
setDnsConfig(c, "127.0.0.1", "0", true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
before := ds.server
ds.serverMu.Unlock()
require.NotNil(t, before)
// Same address, so the running listener must be left alone rather than rebuilt under live queries
// No server running yet, no addr change. Reload should not spawn anything.
require.NoError(t, ds.reload(c, false))
assert.True(t, ds.enabled.Load())
ds.serverMu.Lock()
after := ds.server
ds.serverMu.Unlock()
assert.Same(t, before, after, "a same-address reload must not restart the listener")
ds.Stop()
}
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
// initial only records config, it never starts anything
require.NoError(t, ds.reload(c, true))
ds.serverMu.Lock()
assert.Nil(t, ds.server, "the initial reload must not start a listener")
ds.serverMu.Unlock()
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
ds.serverMu.Unlock()
ds.Stop()
assert.Nil(t, ds.server)
}
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
@@ -464,168 +427,3 @@ func waitFor(t *testing.T, cond func() bool) {
}
t.Fatal("timed out waiting for condition")
}
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
func TestDnsServer_Start_isIdempotent(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
first := ds.server
ds.serverMu.Unlock()
require.NotNil(t, first)
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("second Start never returned")
}
ds.serverMu.Lock()
second := ds.server
ds.serverMu.Unlock()
assert.Same(t, first, second, "a second Start must not replace the running server")
// The real proof, after Stop the port must actually be free
ds.Stop()
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
if err != nil {
return false
}
_ = pc.Close()
return true
})
}
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
// installed, so reload has to clear the slot before shutting the old one down.
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
first := freeUDPPort(t)
second := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", first, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
for i := range 8 {
want := second
if i%2 == 1 {
want = first
}
setDnsConfig(c, "127.0.0.1", want, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
srv := ds.server
ds.serverMu.Unlock()
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
}
// Land back on second so the port assertions below are meaningful
setDnsConfig(c, "127.0.0.1", second, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
// The old port must be released and the new one actually held
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
if err != nil {
return false
}
_ = pc.Close()
return true
})
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
require.Error(t, err, "the new address should be bound by the DNS responder")
ds.Stop()
}
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
port := freeUDPPort(t)
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
ds.Start() // returns once the bind fails
ds.serverMu.Lock()
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
ds.serverMu.Unlock()
// With the slot released, a reload can retry once the port frees up
require.NoError(t, blocker.Close())
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
ds.serverMu.Unlock()
ds.Stop()
}
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
require.True(t, ds.enabled.Load())
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
ds.serverMu.Lock()
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
ds.serverMu.Unlock()
t.Fatal("Start returned early, the test never exercised the window")
case <-time.After(time.Millisecond * 100):
}
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
ds.enabled.Store(false)
ds.serverMu.Unlock()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start never returned")
}
ds.serverMu.Lock()
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
ds.serverMu.Unlock()
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err, "an orphaned listener is still holding the port")
_ = pc.Close()
}
-64
View File
@@ -725,70 +725,6 @@ func TestReestablishRelays(t *testing.T) {
}
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
t.Parallel()
// If them tears down the tunnel while me keeps Established relay state, me's next
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
// them's Disestablished terminal relay entry. them must re-establish that entry, or
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
// them can receive but every send is silently dropped.
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
// Teach my how to get to the relay and that their can be reached via the relay
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
// Build a router so we don't have to reason who gets which packet
r := router.NewR(t, myControl, relayControl, theirControl)
defer r.RenderFlow()
// Start the servers
myControl.Start()
relayControl.Start()
theirControl.Start()
t.Log("Trigger a handshake from me to them via the relay")
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
p := r.RouteForAllUntilTxTun(theirControl)
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
t.Log("Re-handshake from me, riding the still-Established relay state")
myControl.ReHandshake(theirVpnIpNet[0].Addr())
for {
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
break
}
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
return router.RouteAndExit
})
}
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
t.Log("Send from them to me; their only relay entry must survive the transmit")
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
require.Never(t, func() bool {
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
return h == nil || len(h.CurrentRelaysToMe) == 0
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
p = r.RouteForAllUntilTxTun(myControl)
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
}
func TestStage1RaceRelays(t *testing.T) {
t.Parallel()
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
-225
View File
@@ -1,225 +0,0 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"net/netip"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
t.Helper()
cm := lh.QueryLighthouse(vpnAddr)
if cm == nil {
return nil
}
var out []netip.AddrPort
for _, c := range *cm {
out = append(out, c.Reported...)
out = append(out, c.Learned...)
}
return out
}
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
t.Helper()
h := &header.H{}
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c != lh {
return router.KeepRouting
}
// Punches are a single byte and never parse, they are just not what we are after
if err := h.Parse(p.Data); err != nil {
return router.KeepRouting
}
if h.Type == header.LightHouse {
return router.RouteAndExit
}
return router.KeepRouting
})
}
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
// so we call RebindUDPServer directly, which is the same thing the monitor does.
func TestRebindSendsLighthouseUpdate(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
// Let the startup registration finish, then clear everything it left behind
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
// Nothing should be talking to the lighthouse on its own now
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
"nothing should reach the lighthouse before the rebind")
myControl.RebindUDPServer()
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
"a rebind should push an update to the lighthouse rather than waiting out the interval")
lhControl.Stop()
myControl.Stop()
}
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
// whose remote NAT state died while we were on a different network.
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
lhCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
// Without this the peers advertise this machine's real addresses and then try to punch at them,
// which the router has no route for.
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
r.RouteFor(time.Second)
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
r.RouteFor(time.Millisecond * 300)
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
// so this cannot be satisfied by the update the rebind itself pushes.
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
"an ordinary send should not requery the lighthouse")
myControl.RebindUDPServer()
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
// is picked up.
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
return []netip.Addr{myControl.GetUDPAddr().Addr()}
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
"the lighthouse should know the address we started on")
// Wake up somewhere else
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
myControl.SetUDPAddr(newAddr)
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
r.RouteFor(time.Millisecond * 400)
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"the lighthouse should still be handing out the old address before the rebind")
myControl.RebindUDPServer()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
r.RouteFor(time.Millisecond * 400)
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"after the rebind the lighthouse should hand peers our new address")
lhControl.Stop()
myControl.Stop()
}
-136
View File
@@ -1,136 +0,0 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/udp"
)
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
// we are now, so everything we send disappears.
//
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
// since that is when its NAT opens and the tunnel is usable again.
//
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
func TestRecoveryTiming(t *testing.T) {
for _, tc := range []struct {
name string
rebind bool
}{
{"no trigger", false},
{"rebind counter", true},
} {
t.Run(tc.name, func(t *testing.T) {
d, lost := measureRecovery(t, tc.rebind)
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
})
}
}
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
t.Helper()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
peerCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
defer func() {
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
r.RouteFor(time.Second)
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
t.Fatal("failed to establish the tunnel we are measuring")
}
r.RouteFor(time.Millisecond * 500)
// From here the peer's NAT has no state for us, everything we send it disappears
start := time.Now()
blackholed := 0
var recovered time.Duration
if rebind {
myControl.RebindUDPServer()
}
// Keep the tun busy the way someone retrying a stalled connection would
stop := make(chan struct{})
defer close(stop)
go func() {
tick := time.NewTicker(time.Millisecond * 200)
defer tick.Stop()
for {
select {
case <-stop:
return
case <-tick.C:
myControl.InjectTunPacket(BuildTunUDPPacket(
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
}
}
}()
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c == theirControl && p.From == myControl.GetUDPAddr() {
blackholed++
return router.Drop
}
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
if c == myControl && p.From == theirUdpAddr {
recovered = time.Since(start)
return router.RouteAndExit
}
return router.KeepRouting
})
if recovered == 0 {
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
}
return recovered, blackholed
}
+30 -149
View File
@@ -114,28 +114,6 @@ type packet struct {
packet *udp.Packet
tun bool // a packet pulled off a tun device
rx bool // the packet was received by a udp device
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
h header.H
parseErr error
}
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
// addresses, so they fall back to the control.
func (p *packet) fromAddr() netip.AddrPort {
if p.tun || !p.packet.From.IsValid() {
return p.from.GetUDPAddr()
}
return p.packet.From
}
func (p *packet) toAddr() netip.AddrPort {
if p.tun || !p.packet.To.IsValid() {
return p.to.GetUDPAddr()
}
return p.packet.To
}
func (p *packet) WasReceived() {
@@ -153,9 +131,6 @@ const (
ExitNow ExitType = 1
// RouteAndExit routes this packet and exits immediately afterwards
RouteAndExit ExitType = 2
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
// a restrictive NAT refusing traffic from an address it has not seen.
Drop ExitType = 3
)
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
@@ -166,9 +141,7 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
func NewR(t testing.TB, controls ...*nebula.Control) *R {
ctx, cancel := context.WithCancel(context.Background())
// t.Name() contains a slash for subtests, so the flow log can land in a nested directory
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
if err := os.MkdirAll("mermaid", 0755); err != nil {
panic(err)
}
@@ -179,7 +152,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
outNat: make(map[outNatKey]netip.AddrPort),
flow: []flowEntry{},
ignoreFlows: []ignoreFlow{},
fn: fn,
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
t: t,
cancelRender: cancel,
}
@@ -276,7 +249,7 @@ func (r *R) renderFlow() {
continue
}
addr := e.packet.fromAddr()
addr := e.packet.from.GetUDPAddr()
if _, ok := participants[addr]; ok {
continue
}
@@ -295,6 +268,7 @@ func (r *R) renderFlow() {
}
// Print packets
h := &header.H{}
for _, e := range r.flow {
if e.packet == nil {
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
@@ -306,22 +280,21 @@ func (r *R) renderFlow() {
fmt.Fprintln(f, r.formatUdpPacket(p))
} else {
if err := h.Parse(p.packet.Data); err != nil {
panic(err)
}
line := "--x"
if p.rx {
line = "->>"
}
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
if p.parseErr != nil {
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
}
fmt.Fprintf(f, " %s%s%s: %s\n",
normalizeName(p.fromAddr().String()),
fmt.Fprintf(f,
" %s%s%s: %s(%s), index %v, counter: %v\n",
normalizeName(p.from.GetUDPAddr().String()),
line,
normalizeName(p.toAddr().String()),
detail,
normalizeName(p.to.GetUDPAddr().String()),
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
)
}
}
@@ -435,34 +408,29 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
var h header.H
var parseErr error
if !tun {
parseErr = h.Parse(p.Data)
}
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
for _, i := range r.ignoreFlows {
if tun {
if i.tun.HasValue && i.tun.IsTrue {
return nil
}
continue
if len(r.ignoreFlows) > 0 {
var h header.H
err := h.Parse(p.Data)
if err != nil {
panic(err)
}
// A packet we could not parse has no type to match against, so no rule can ignore it
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
return nil
for _, i := range r.ignoreFlows {
if !tun {
if i.messageType == h.Type && i.subType == h.Subtype {
return nil
}
} else if i.tun.HasValue && i.tun.IsTrue {
return nil
}
}
}
fp := &packet{
from: from,
to: to,
packet: p.Copy(),
tun: tun,
h: h,
parseErr: parseErr,
from: from,
to: to,
packet: p.Copy(),
tun: tun,
}
r.flow = append(r.flow, flowEntry{packet: fp})
@@ -692,10 +660,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
p.Release()
return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(sender, receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(sender, receiver, p, false)
receiver.InjectUDPPacket(p)
@@ -726,85 +690,6 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
})
}
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
// more packets right behind it.
func (r *R) RouteFor(d time.Duration) {
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
return KeepRouting
})
}
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
// assert that something does NOT happen, or to route for a fixed settling period.
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
cm := make([]*nebula.Control, 0, len(r.controls))
for _, c := range r.controls {
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(c.GetUDPTxChan()),
Send: reflect.Value{},
})
cm = append(cm, c)
}
timer := time.NewTimer(timeout)
defer timer.Stop()
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(timer.C),
Send: reflect.Value{},
})
for {
x, rx, _ := reflect.Select(sc)
if x == len(cm) {
return false
}
r.Lock()
p := rx.Interface().(*udp.Packet)
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
if receiver == nil {
r.Unlock()
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
}
e := whatDo(p, receiver)
switch e {
case ExitNow:
r.Unlock()
p.Release()
return true
case RouteAndExit:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
r.Unlock()
p.Release()
return true
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
default:
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
}
r.Unlock()
p.Release()
}
}
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
h := &header.H{}
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
@@ -897,10 +782,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
p.Release()
return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
+188
View File
@@ -0,0 +1,188 @@
package nebula
import (
"encoding/binary"
"log/slog"
"testing"
"golang.org/x/net/ipv4"
)
func TestInnerECN(t *testing.T) {
cases := []struct {
name string
pkt []byte
want byte
}{
{"empty", nil, 0},
{"v4_NotECT", v4WithToS(0x00), 0x00},
{"v4_ECT0", v4WithToS(0x02), 0x02},
{"v4_ECT1", v4WithToS(0x01), 0x01},
{"v4_CE", v4WithToS(0x03), 0x03},
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
{"v6_NotECT", v6WithTC(0x00), 0x00},
{"v6_ECT0", v6WithTC(0x02), 0x02},
{"v6_CE", v6WithTC(0x03), 0x03},
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
{"unknown_version", []byte{0xa5, 0xff}, 0},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
got := innerECN(c.pkt)
if got != c.want {
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
}
})
}
}
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
// the DSCP and ECN portions through the byte 1 mask.
func v4WithToS(tos byte) []byte {
return []byte{0x45, tos}
}
// v6WithTC builds a 2-byte slice that places a known traffic class value
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
func v6WithTC(tc byte) []byte {
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
}
func TestApplyOuterECN(t *testing.T) {
silent := slog.New(slog.DiscardHandler)
hi := &HostInfo{}
// Build a v4 packet helper with a given inner ECN field.
v4 := func(innerECN byte) []byte {
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
return []byte{
0x45, innerECN, 0, 28,
0, 0, 0x40, 0,
64, 6, 0, 0,
10, 0, 0, 1,
10, 0, 0, 2,
}
}
// Build a v6 packet helper with a given inner ECN field. ECN occupies
// TC[1:0] which sit at byte 1 mask 0x30.
v6 := func(innerECN byte) []byte {
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
pkt := make([]byte, 40)
pkt[0] = 0x60 // version=6, TC[7:4]=0
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
return pkt
}
type cell struct {
outer byte
inner byte
wantECN byte
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
}
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
table := []cell{
{ecnNotECT, ecnNotECT, ecnNotECT, true},
{ecnNotECT, ecnECT0, ecnECT0, true},
{ecnNotECT, ecnECT1, ecnECT1, true},
{ecnNotECT, ecnCE, ecnCE, true},
{ecnECT0, ecnNotECT, ecnNotECT, true},
{ecnECT0, ecnECT0, ecnECT0, true},
{ecnECT0, ecnECT1, ecnECT1, true},
{ecnECT0, ecnCE, ecnCE, true},
{ecnECT1, ecnNotECT, ecnNotECT, true},
{ecnECT1, ecnECT0, ecnECT0, true},
{ecnECT1, ecnECT1, ecnECT1, true},
{ecnECT1, ecnCE, ecnCE, true},
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
{ecnCE, ecnECT1, ecnCE, false},
{ecnCE, ecnCE, ecnCE, true},
}
for _, c := range table {
t.Run("v4", func(t *testing.T) {
pkt := v4(c.inner)
applyOuterECN(pkt, c.outer, hi, silent)
got := pkt[1] & 0x03
if got != c.wantECN {
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
}
})
t.Run("v6", func(t *testing.T) {
pkt := v6(c.inner)
applyOuterECN(pkt, c.outer, hi, silent)
got := (pkt[1] >> 4) & 0x03
if got != c.wantECN {
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
}
})
}
}
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
// The passthrough emit paths write the packet verbatim, so a stale checksum
// turns an underlay congestion mark into packet loss at the receiver.
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
silent := slog.New(slog.DiscardHandler)
hi := &HostInfo{}
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
// flips only the low two bits of the ToS byte while leaving DSCP intact.
pkt := []byte{
0x45, 0x88 | ecnECT0, 0, 40,
0x1c, 0x46, 0x40, 0x00,
64, 6, 0, 0,
10, 0, 0, 1,
10, 0, 0, 2,
}
// Stamp a correct header checksum before the fold.
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatal("test setup: initial header checksum invalid")
}
applyOuterECN(pkt, ecnCE, hi, silent)
// CE folded in, DSCP preserved.
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
}
// The incremental RFC 1624 update must leave the checksum valid and equal
// to a full recompute over the mutated header.
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
}
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
}
}
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
// treating the checksum field (bytes 10:12) as zero.
func ipv4HeaderChecksum(hdr []byte) uint16 {
var sum uint32
for i := 0; i+1 < len(hdr); i += 2 {
if i == 10 {
continue // checksum field
}
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
}
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
return ^uint16(sum)
}
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
// computation over the header.
func ipv4HeaderChecksumValid(hdr []byte) bool {
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
}
+20 -18
View File
@@ -131,9 +131,6 @@ 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)
@@ -149,14 +146,6 @@ listen:
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
#windows_bypass_wdf: true
# On macOS only
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
# the routing socket and rebinds the listener once the change settles.
# iOS does not use this, the host app drives the same rebind itself.
# Default true. Not reloadable.
#rebind_on_network_change: true
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
@@ -267,19 +256,22 @@ tun:
# 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.
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
#
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
# equal N`) or set cpu_affinity explicitly to benefit.
#pin_threads: true
# 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 and falls back to spreading queues across the allowed
# CPUs. Only meaningful while pin_threads is true. Not reloadable.
# When unset, 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.
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
#cpu_affinity:
# - 2
# - 4
@@ -420,6 +412,16 @@ logging:
# This setting is reloadable
#inactivity_timeout: 10m
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
# the route half of this setting to take effect.
#ecn: true
# Nebula security group configuration
firewall:
# Action to take when a packet is not allowed by the firewall rules.
-9
View File
@@ -8,15 +8,6 @@ Before=sshd.service
Type=notify
NotifyAccess=main
SyslogIdentifier=nebula
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
#User=nebula
#Group=nebula
#CapabilityBoundingSet=CAP_NET_ADMIN
#AmbientCapabilities=CAP_NET_ADMIN
ExecReload=/bin/kill -HUP $MAINPID
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
Restart=always
-9
View File
@@ -65,12 +65,3 @@ 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
}
+6 -6
View File
@@ -1,6 +1,6 @@
module github.com/slackhq/nebula
go 1.26.0
go 1.25.0
require (
dario.cat/mergo v1.0.2
@@ -24,12 +24,12 @@ require (
github.com/vishvananda/netlink v1.3.1
go.uber.org/goleak v1.3.0
go.yaml.in/yaml/v3 v3.0.4
golang.org/x/crypto v0.54.0
golang.org/x/crypto v0.53.0
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
golang.org/x/net v0.57.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/term v0.45.0
golang.org/x/net v0.56.0
golang.org/x/sync v0.21.0
golang.org/x/sys v0.46.0
golang.org/x/term v0.44.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
golang.zx2c4.com/wireguard/windows v1.0.1
+10 -10
View File
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
+5 -15
View File
@@ -295,13 +295,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
err := hm.outside.WriteTo(stage0, addr)
if err != nil {
// These repeat every attempt, so match the success log below and only shout when the remotes changed
level := slog.LevelDebug
if remotesHaveChanged {
level = slog.LevelError
}
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
hostinfo.logger(hm.l).Error("Failed to send handshake message",
"udpAddr", addr,
"initiatorIndex", hostinfo.localIndexId,
"handshake", hsFields,
@@ -535,9 +529,7 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
for _, addr := range hostinfo.vpnAddrs {
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
delete(hm.vpnIps, addr)
}
delete(hm.vpnIps, addr)
}
if len(hm.vpnIps) == 0 {
@@ -975,9 +967,7 @@ 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.
//todo use a sendbatcher
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
}
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
@@ -1088,8 +1078,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
// 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, 0)
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
}
}
+1 -1
View File
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
return
}
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
return
}
+6 -11
View File
@@ -190,18 +190,13 @@ func SubTypeName(t MessageType, s MessageSubType) string {
}
func IsValidSubType(t MessageType, s MessageSubType) bool {
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
if n, ok := subTypeMap[t]; ok {
if _, ok := (*n)[s]; ok {
return true
}
}
return false
}
// NewHeader turns bytes into a header
-51
View File
@@ -102,57 +102,6 @@ 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,
+1 -11
View File
@@ -287,6 +287,7 @@ type HostInfo struct {
type ViaSender struct {
UdpAddr netip.AddrPort
relayHI *HostInfo // relayHI is the host info object of the relay
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
}
@@ -543,17 +544,6 @@ 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 {
+59 -31
View File
@@ -15,7 +15,7 @@ import (
"github.com/slackhq/nebula/routing"
)
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, 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
@@ -74,7 +74,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
return
}
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, 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.
@@ -105,9 +105,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
return
}
dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache)
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
if dropReason == nil {
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
} else {
f.rejectInside(packet, rejectBuf, q)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
@@ -126,6 +126,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
c := ci.messageCounter.Add(1)
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
f.connectionManager.Out(hostinfo)
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
if noiseutil.EncryptLockNeeded {
@@ -137,7 +138,8 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
"udpAddr", hostinfo.GetRemote(),
"counter", c,
)
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
// Skip this segment; the rest of the superpacket can still
// go out — TCP will retransmit anything we drop here.
return nil
}
@@ -149,19 +151,16 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
// 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) {
// 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.TxBatcher, rejectBuf []byte, q int) {
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()
ecnEnabled := f.ecnEnabled.Load()
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.
@@ -212,7 +211,11 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
return nil
}
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
var ecn byte
if ecnEnabled {
ecn = innerECN(seg)
}
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
return nil
})
if err != nil {
@@ -230,14 +233,36 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
return nil
}
sendBatch.Commit(out, remote)
var ecn byte
if ecnEnabled {
ecn = innerECN(seg)
}
sendBatch.Commit(out, remote, ecn)
return nil
})
if err != nil {
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
"error", err,
)
}
}
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
func innerECN(pkt []byte) byte {
if len(pkt) < 2 {
return 0
}
switch pkt[0] >> 4 {
case 4:
return pkt[1] & 0x03
case 6:
return (pkt[1] >> 4) & 0x03
}
return 0
}
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
if !f.firewall.OutboundSendReject {
return
@@ -254,30 +279,27 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
}
}
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
if !f.firewall.InboundSendReject {
return
}
// 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)
out = iputil.CreateRejectPacket(packet, out)
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, encryptBuf, q)
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, 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
@@ -371,7 +393,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.ParsedPacket{}
fp := &firewall.Packet{}
err := newPacket(p, false, fp)
if err != nil {
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
@@ -379,7 +401,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
}
// check if packet is in outbound fw rules
dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil)
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
if dropReason != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("dropping cached packet",
@@ -492,14 +514,20 @@ func (f *Interface) prepareSendVia(via *HostInfo,
// 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) {
func (f *Interface) SendVia(via *HostInfo,
relay *Relay,
ad,
nb,
out []byte,
nocopy bool,
) {
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
if err != nil {
// already logged by prepareSendVia
return
}
err = f.writers[q].WriteTo(toSend, via.GetRemote())
err = f.writers[0].WriteTo(toSend, via.GetRemote())
if err != nil {
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
}
@@ -573,7 +601,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
if err != nil {
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
"error", err,
"udpAddr", hr,
"udpAddr", remote,
)
}
} else {
@@ -588,7 +616,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
)
continue
}
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
break
}
}
+79 -89
View File
@@ -96,7 +96,12 @@ type Interface struct {
// 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
pinThreads bool
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
// inside.go copies the inner ECN onto the outer carrier on encap and
// decryptToTun folds outer CE into the inner header on decap. Toggle
// via tunnels.ecn (default true).
ecnEnabled atomic.Bool
relayManager *relayManager
tryPromoteEvery atomic.Uint32
@@ -115,12 +120,10 @@ type Interface struct {
ctx context.Context
writers []udp.Conn
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
// batchers is one per tun queue, wrapping queues[i].
// decryptToTun sends plaintext into the batch.RxBatcher;
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
batchers []batch.RxBatcher
wg sync.WaitGroup
// fatalErr holds the first unexpected reader error that caused shutdown.
@@ -132,13 +135,18 @@ 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, q int)
SendVia(via *HostInfo,
relay *Relay,
ad,
nb,
out []byte,
nocopy bool,
)
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)
@@ -197,10 +205,6 @@ 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,
@@ -218,7 +222,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
routines: c.routines,
version: c.version,
writers: make([]udp.Conn, c.routines),
batchers: make([]*batch.MultiCoalescer, c.routines),
batchers: make([]batch.RxBatcher, c.routines),
myVpnNetworks: cs.myVpnNetworks,
myVpnNetworksTable: cs.myVpnNetworksTable,
myVpnAddrs: cs.myVpnAddrs,
@@ -231,7 +235,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
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),
@@ -285,13 +288,6 @@ func (f *Interface) activate() error {
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)
@@ -301,7 +297,18 @@ func (f *Interface) activate() error {
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
for i := range f.queues {
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
caps := tio.QueueCapabilities(f.queues[i])
if caps.TSO || caps.USO {
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
// is on, everything else (and either lane disabled) falls
// through to passthrough so non-IP / non-TCP-UDP traffic still
// reaches the TUN.
arena := batch.NewArena(batch.DefaultMultiArenaCap)
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
} else {
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
}
}
// On error the caller owns the cleanup, Control.Start cancels the service context
@@ -349,31 +356,6 @@ 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 {
@@ -382,17 +364,21 @@ func (f *Interface) listenOut(i int) {
li = f.outside
}
rxc := newRxContext(f, i)
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
lhh := f.lightHouse.NewRequestHandler()
h := &header.H{}
fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12)
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
plaintext := f.batchers[i].Reserve(len(payload))
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
}
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)
@@ -408,36 +394,32 @@ func (f *Interface) listenOut(i int) {
f.l.Debug("underlay reader is done", "reader", i)
}
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)
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)
}
}
rejectBuf := make([]byte, mtu)
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
fwPacket := &firewall.ParsedPacket{}
fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12)
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
@@ -459,35 +441,26 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
// 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)
if err := sb.Flush(); err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
}
}
}
f.flushSendBatch(sb, i)
if err := sb.Flush(); err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", 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)
c.RegisterReloadCallback(f.reloadAcceptRecvError)
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
c.RegisterReloadCallback(f.reloadMisc)
c.RegisterReloadCallback(f.reloadEcn)
for _, udpConn := range f.writers {
c.RegisterReloadCallback(udpConn.ReloadConfig)
@@ -620,6 +593,23 @@ func (f *Interface) reloadMisc(c *config.C) {
}
}
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
// Default is enabled (RFC 6040 normal mode); set false on the rare path
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
func (f *Interface) reloadEcn(c *config.C) {
initial := c.InitialLoad()
if initial || c.HasChanged("tunnels.ecn") {
v := c.GetBool("tunnels.ecn", true)
changed := f.ecnEnabled.Swap(v) != v
if !initial {
f.l.Info("tunnels.ecn changed", "enabled", v)
if changed {
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
}
}
}
}
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
ticker := time.NewTicker(i)
defer ticker.Stop()
+3 -27
View File
@@ -199,7 +199,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
}
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
proto, offset, isFragment := IPv6FindUpperProtocol(packet)
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
if isFragment {
return nil
}
@@ -333,34 +333,11 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
return out
}
// maxIPv6ExtHeaders caps the extension-header walk in IPv6FindUpperProtocol.
// RFC 8200 legal chains are shorter (each header at most once, Destination
// Options at most twice), so the cap only bites crafted packets, which would
// otherwise make us walk their whole payload 8 bytes at a time.
const maxIPv6ExtHeaders = 8
// IPv6FindUpperProtocol walks packet's IPv6 extension-header chain and
// returns the terminal (upper-layer) protocol number, the byte offset where
// that protocol's header begins, and whether the packet is a non-first
// fragment. It steps over Hop-by-Hop (0), Routing (43), Fragment (44),
// AH (51), and Destination Options (60); anything else — including ESP,
// whose payload is encrypted — terminates the walk.
//
// For a non-first fragment, nextHeader still names the flow's upper
// protocol (copied from the fragment header) but offset points at fragment
// payload, not a real transport header: consult isFragment before
// dereferencing offset. If the chain is truncated, over-long, or the packet
// is shorter than an IPv6 header, the walk stops early and nextHeader is
// whatever it stopped on (59, IPPROTO_NONE, for the too-short case) —
// callers treat any non-transport result as unclassifiable.
func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
if len(packet) < ipv6.HeaderLen {
return 59, 0, false // IPPROTO_NONE: nothing to classify
}
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
nextHeader = packet[6]
offset = ipv6.HeaderLen
for range maxIPv6ExtHeaders {
for {
switch nextHeader {
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
if len(packet) < offset+2 {
@@ -390,7 +367,6 @@ func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm
return nextHeader, offset, isFragment
}
}
return nextHeader, offset, isFragment
}
func CreateICMPEchoResponse(packet, out []byte) []byte {
-118
View File
@@ -515,121 +515,3 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
result := CreateICMPEchoResponse(packet, out)
assert.Nil(t, result)
}
func TestIPv6FindUpperProtocol(t *testing.T) {
src := net.ParseIP("fd00::1")
dst := net.ParseIP("fd00::2")
// extHdr builds one 8-byte-unit extension header: next, hdrExtLen
// ((extra+1)*8 bytes total), padded to size.
extHdr := func(next uint8, extra int) []byte {
b := make([]byte, (extra+1)*8)
b[0] = next
b[1] = uint8(extra)
return b
}
t.Run("no extension headers", func(t *testing.T) {
for _, proto := range []uint8{6, 17, 58} {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, proto, make([]byte, 20)))
assert.Equal(t, proto, nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
}
})
t.Run("hop-by-hop then TCP", func(t *testing.T) {
payload := append(extHdr(6, 0), make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, frag)
})
t.Run("chained headers honor length units", func(t *testing.T) {
// Hop-by-Hop (8B) -> Dest Options (16B) -> Routing (8B) -> UDP.
payload := extHdr(60, 0)
payload = append(payload, extHdr(43, 1)...)
payload = append(payload, extHdr(17, 0)...)
payload = append(payload, make([]byte, 8)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8+16+8, offset)
assert.False(t, frag)
})
t.Run("AH length is in 4-byte units plus 2", func(t *testing.T) {
// AH payload-len byte 4 -> (4+2)*4 = 24 bytes on the wire.
ah := make([]byte, 24)
ah[0] = 6
ah[1] = 4
payload := append(ah, make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 51, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+24, offset)
assert.False(t, frag)
})
t.Run("first fragment walks to the transport header", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 0x0001) // offset 0, M=1
payload := append(frag, make([]byte, 8)...)
nh, offset, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, isFrag, "first fragment carries the real transport header")
})
t.Run("non-first fragment is flagged", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 1<<3) // offset 1, M=0
payload := append(frag, make([]byte, 8)...)
nh, _, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh, "fragment header still names the flow's L4")
assert.True(t, isFrag, "offset points at fragment payload, not a header")
})
t.Run("ESP terminates the walk", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 50, make([]byte, 16)))
assert.Equal(t, uint8(50), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("unknown protocol terminates the walk", func(t *testing.T) {
nh, offset, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 132, make([]byte, 16))) // SCTP
assert.Equal(t, uint8(132), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
})
t.Run("truncated extension header stops the walk", func(t *testing.T) {
// Next header says Hop-by-Hop but the packet ends at the IPv6 header.
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, nil))
assert.Equal(t, uint8(0), nh, "unresolvable chain returns the extension header it stopped on")
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("crafted over-long chain hits the cap", func(t *testing.T) {
// Ten chained Hop-by-Hop headers, then TCP. Illegal per RFC 8200
// (Hop-by-Hop may only appear first); the cap must stop the walk
// before it resolves rather than crawling arbitrary crafted chains.
var payload []byte
for i := 0; i < 9; i++ {
payload = append(payload, extHdr(0, 0)...)
}
payload = append(payload, extHdr(6, 0)...)
payload = append(payload, make([]byte, 20)...)
nh, _, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(0), nh, "walk must stop at the cap, not resolve to TCP")
})
t.Run("packet shorter than an IPv6 header", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(make([]byte, 39))
assert.Equal(t, uint8(59), nh) // IPPROTO_NONE
assert.Equal(t, 0, offset)
assert.False(t, frag)
})
}
+1 -9
View File
@@ -36,10 +36,6 @@ type LightHouse struct {
myVpnNetworksTable *bart.Lite
punchy *Punchy
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
localAddrsFn func(*LocalAllowList) []netip.Addr
// Local cache of answers from light houses
// map of vpn addr to answers
addrMap map[netip.Addr]*RemoteList
@@ -111,10 +107,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
l: l,
}
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
return localAddrs(h.l, al)
}
lighthouses := make([]netip.Addr, 0)
h.lighthouses.Store(&lighthouses)
staticList := make(map[netip.Addr]struct{})
@@ -926,7 +918,7 @@ func (lh *LightHouse) SendUpdate() {
}
lal := lh.GetLocalAllowList()
for _, e := range lh.localAddrsFn(lal) {
for _, e := range localAddrs(lh.l, lal) {
if lh.myVpnNetworksTable.Contains(e) {
continue
}
+1 -1
View File
@@ -498,7 +498,7 @@ type testEncWriter struct {
protocolVersion cert.Version
}
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
}
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
}
+55 -23
View File
@@ -6,14 +6,12 @@ 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/overlay"
"github.com/slackhq/nebula/sshd"
"github.com/slackhq/nebula/udp"
@@ -167,13 +165,7 @@ 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)))
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)
}
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, batchSize)
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
if err != nil {
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
}
@@ -224,17 +216,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
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())
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
key = uint64(ap.Port())
}
cpuAffinity = cpupick.Default(routines, key, l)
if pinThreads && len(cpuAffinity) == 0 && !configTest {
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
}
ifConfig := &InterfaceConfig{
@@ -277,6 +260,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
ifce.reloadDisconnectInvalid(c)
ifce.reloadSendRecvError(c)
ifce.reloadAcceptRecvError(c)
ifce.reloadEcn(c)
handshakeManager.f = ifce
go handshakeManager.Run(ctx)
@@ -297,8 +281,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
attachCommands(l, c, ssh, ifce)
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
return &Control{
state: StateReady,
f: ifce,
@@ -309,7 +291,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
statsStart: stats.Start,
dnsStart: ds.Start,
lighthouseStart: lightHouse.StartUpdateWorker,
networkChangeStart: networkChanges.Start,
connectionManagerStart: connManager.Start,
}, nil
}
@@ -378,6 +359,57 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
return cpus
}
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
// any physical NIC's interrupts. The stock allowed[i] spread pins the
// encrypt threads onto exactly the cores most drivers affine their first RX
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
//
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
// info is unavailable or when there aren't enough IRQ-free CPUs to give
// every routine its own core: silently doubling readers up on fewer cores
// is worse than the occasional IRQ collision. NICs whose vectors blanket
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
// effective.
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
irq, err := util.NICIRQCPUs()
if err != nil || len(irq) == 0 {
return nil
}
allowed, err := util.AllowedCPUs()
if err != nil {
return nil
}
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
if cpus == nil {
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
return nil
}
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
return cpus
}
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
// irq, or nil if fewer than `routines` qualify.
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
free := make([]int, 0, routines)
for _, cpu := range allowed {
if irq[cpu] {
continue
}
free = append(free, cpu)
if len(free) == routines {
return free
}
}
return nil
}
func moduleVersion() string {
info, ok := debug.ReadBuildInfo()
if !ok {
+20
View File
@@ -9,6 +9,26 @@ import (
"github.com/stretchr/testify/assert"
)
func TestChooseIRQFreeCPUs(t *testing.T) {
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
// Exactly enough.
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
// than doubling readers up on shared cores.
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
// No IRQ info at all behaves like a plain prefix of allowed.
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
// Non-contiguous allowed set (cgroup cpuset) with holes.
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
}
func TestParseCpuAffinity(t *testing.T) {
l := test.NewLogger()
-45
View File
@@ -164,48 +164,3 @@ 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")
}
+155 -76
View File
@@ -13,7 +13,7 @@ import (
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/udp"
"golang.org/x/net/ipv4"
)
@@ -23,11 +23,7 @@ const (
var ErrOutOfWindow = errors.New("out of window packet")
// 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
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, meta udp.RxMeta) {
err := h.Parse(packet)
if err != nil {
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
@@ -95,7 +91,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
if isMessageRelay {
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
} else {
hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
}
// At this point we should have a valid existing tunnel, verify and send
@@ -107,32 +103,26 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
return
}
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
f.messageMetrics.RxInvalid(1)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("packet too small", "from", via, "length", len(packet))
}
return
}
// All remaining packets are encrypted
if isMessageRelay {
// Relay packets are special, this branch should always early-return
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, packet, rxc)
ci := hostinfo.ConnectionState
if !ci.window.Check(f.l, h.MessageCounter) {
return
}
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
// Relay packets are special
if isMessageRelay {
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
return
}
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, 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)
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
"error", err,
"from", via,
"header", h,
)
}
return
}
@@ -145,7 +135,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
case header.Message:
switch h.Subtype {
case header.MessageNone:
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta)
default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
return
@@ -153,23 +143,15 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
case header.LightHouse:
//TODO: assert via is not relayed
rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
lhf.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:
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])
//recycle the input packet ciphertext as our output buffer
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
return
@@ -187,10 +169,28 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
}
}
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()]
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, meta udp.RxMeta) {
// 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)-hostinfo.ConnectionState.dKey.Overhead()]
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
var err error
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
if err != nil {
return
}
// Advance the replay window now that the frame is authenticated
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
}
return
}
// Successfully validated the thing. Get rid of the Relay header.
signedPayload = signedPayload[header.Len:]
// Pull the Roaming parts up here, and return in all call paths.
f.handleHostRoaming(hostinfo, via)
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
@@ -201,7 +201,9 @@ 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
}
@@ -212,10 +214,11 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
via = ViaSender{
UdpAddr: via.UdpAddr,
relayHI: hostinfo,
remoteIdx: relay.RemoteIndex,
relay: relay,
IsRelayed: true,
}
f.readOutsidePackets(via, signedPayload, rxc)
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta)
case ForwardingType:
// Find the target HostInfo relay object
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
@@ -232,11 +235,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
if targetRelay.State == Established {
switch targetRelay.Type {
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]
//todo it would potentially be nice to batch these
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
// Forward this packet through the relay tunnel
// Find the target HostInfo //todo it would potentially be nice to batch these
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
case TerminalType:
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
return
@@ -317,11 +318,7 @@ 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.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
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
if len(data) < 1 {
return ErrPacketTooShort
}
@@ -336,7 +333,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return ErrUnknownIPVersion
}
func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
dataLen := len(data)
if dataLen < ipv6.HeaderLen {
return ErrIPv6PacketTooShort
@@ -362,7 +359,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
switch proto {
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
fp.Protocol = uint8(proto)
fp.IPHdrLen = offset
fp.RemotePort = 0
fp.LocalPort = 0
fp.Fragment = false
@@ -373,7 +369,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return ErrIPv6PacketTooShort
}
fp.Protocol = uint8(proto)
fp.IPHdrLen = offset
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
icmptype := data[offset+1]
switch icmptype {
@@ -391,9 +386,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
}
fp.Protocol = uint8(proto)
// offset is the L4 header start: 40 for a plain packet, past the extension chain
// otherwise. The coalescer only accepts 40.
fp.IPHdrLen = offset
if incoming {
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
@@ -411,9 +403,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return ErrIPv6PacketTooShort
}
// A fragment shape the coalescer must not touch either way, first fragment included.
fp.FragAny = true
// Check if this is the first fragment
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
if fragmentOffset != 0 {
@@ -455,7 +444,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return ErrIPv6CouldNotFindPayload
}
func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
// Do we at least have an ipv4 header worth of data?
if len(data) < ipv4.HeaderLen {
return ErrIPv4PacketTooShort
@@ -472,10 +461,6 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) 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]
@@ -519,23 +504,117 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return nil
}
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
err := newPacket(out, true, rxc.fwPacket)
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
var err error
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
if err != nil {
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
return nil, err
}
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
return nil, ErrOutOfWindow
}
return out, nil
}
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
const (
ecnNotECT = 0x00
ecnECT1 = 0x01
ecnECT0 = 0x02
ecnCE = 0x03
)
// applyOuterECN folds an outer CE mark from the underlay into the inner
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
// codepoints are advisory only and leave the inner unchanged.
//
// Merge cases (outer × inner → action):
//
// outer != CE : no-op (inner is authoritative)
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
return
}
switch pkt[0] >> 4 {
case 4:
switch pkt[1] & 0x03 {
case ecnNotECT:
if l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
}
case ecnCE:
// Already CE.
default:
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
// checksum lives at pkt[10:12]. A header too short to carry a
// checksum can't be fixed up here, so leave it for newPacket to
// reject rather than emit a mangled packet.
if len(pkt) < ipv4.HeaderLen {
return
}
m := binary.BigEndian.Uint16(pkt[0:2])
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
mNew := binary.BigEndian.Uint16(pkt[0:2])
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
}
case 6:
switch (pkt[1] >> 4) & 0x03 {
case ecnNotECT:
if l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
}
case ecnCE:
// Already CE.
default:
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
}
}
}
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
// underlay into the inner header before firewall + TUN write. Other
// outer codepoints are advisory only — we keep the inner unchanged.
if f.ecnEnabled.Load() {
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
}
err := newPacket(out, true, fwPacket)
if err != nil {
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
"error", err,
"packet", out,
)
return
}
dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
if dropReason != nil {
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
// 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. With UDP GRO this is a single segment of a shared
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
hostinfo.logger(f.l).Debug("dropping inbound packet",
"fwPacket", fwPacket,
"reason", dropReason,
)
}
return
}
err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
err = f.batchers[q].Commit(out)
if err != nil {
f.l.Error("Failed to write to tun", "error", err)
}
+5 -90
View File
@@ -17,7 +17,7 @@ import (
)
func Test_newPacket(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
// length fails
err := newPacket([]byte{}, true, p)
@@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) {
}
func Test_newPacket_v6(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
// invalid ipv6
ip := layers.IPv6{
@@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
}
func Test_newPacket_ipv6Fragment(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
ip := &layers.IPv6{
Version: 6,
@@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
secondFrag = append(secondFrag, fragHeader...)
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
fp := &firewall.ParsedPacket{}
fp := &firewall.Packet{}
b.Run("Normal", func(b *testing.B) {
for i := 0; i < b.N; i++ {
@@ -649,7 +649,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.ParsedPacket{}
p := &firewall.Packet{}
const (
hdrLen = 40 // IPv6 header
@@ -675,88 +675,3 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
// the host delivers to, not the forged 443 at the overflowed offset.
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
}
// 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)
}
+25 -8
View File
@@ -1,11 +1,28 @@
package batch
// SortKey identifies a packet's position in its sender's transmission order.
// 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.
// Counter is the packet's AEAD message counter within that tunnel.
type SortKey struct {
Epoch uint64
Counter uint64
import "net/netip"
type RxBatcher interface {
// Reserve creates a pkt to borrow
Reserve(sz int) []byte
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
Commit(pkt []byte) error
// Flush emits every queued packet in arrival order.
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
// After Flush returns, borrowed payload slices may be recycled.
Flush() error
}
type TxBatcher interface {
// Reserve creates a pkt to borrow
Reserve(sz int) []byte
// Commit borrows pkt and records its destination plus the 2-bit
// IP-level ECN codepoint to set on the outer (carrier) header. The
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
// to leave the outer ECN field unset.
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
// Flush emits every queued packet via the underlying batch writer in arrival order.
// Returns an errors.Join of one or more errors.
// After Flush returns, borrowed payload slices may be recycled.
Flush() error
}
-187
View File
@@ -1,187 +0,0 @@
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)
}
}
}
+122 -102
View File
@@ -8,125 +8,135 @@ import (
// 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.
// 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.
// 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 encrypted 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.
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
// L4-specific parsing (TCP / UDP) on top.
type parsedIP struct {
fk flowKey
ipHdrLen int
// pkt is the original buffer trimmed to the IP-declared total length.
// Anything below the IP layer (transport parsers) should slice into
// pkt rather than the unbounded original.
pkt []byte
}
// parseIPPrologue extracts the IP-level fields the coalescers care about:
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
// or IPv6 with extension headers (all rejected by both coalescers in
// identical ways before this refactor).
//
// 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) {
// On success, p.pkt is len-trimmed to the IP-declared length so callers
// don't have to repeat the trim. wantProto is the IANA protocol number to
// require (6 for TCP, 17 for UDP); ok=false for any other value.
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
var p parsedIP
if len(pkt) < 20 {
return nil, false
return p, false
}
switch pkt[0] >> 4 {
v := pkt[0] >> 4
switch v {
case 4:
if ipHdrLen != 20 {
return nil, false
ihl := int(pkt[0]&0x0f) * 4
if ihl != 20 {
return p, false
}
return fk.parseIPv4Prologue(pkt)
if pkt[9] != wantProto {
return p, false
}
// Reject actual fragmentation (MF or non-zero frag offset).
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
return p, false
}
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
if totalLen > len(pkt) || totalLen < ihl {
return p, false
}
p.ipHdrLen = 20
p.fk.isV6 = false
copy(p.fk.src[:4], pkt[12:16])
copy(p.fk.dst[:4], pkt[16:20])
p.pkt = pkt[:totalLen]
case 6:
if ipHdrLen != 40 || len(pkt) < 40 {
return nil, false
if len(pkt) < 40 {
return p, false
}
return fk.parseIPv6Prologue(pkt)
if pkt[6] != wantProto {
return p, false
}
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
if 40+payloadLen > len(pkt) {
return p, false
}
p.ipHdrLen = 40
p.fk.isV6 = true
copy(p.fk.src[:], pkt[8:24])
copy(p.fk.dst[:], pkt[24:40])
p.pkt = pkt[:40+payloadLen]
default:
return p, false
}
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
return p, 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.
// 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.
// 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 {
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
if a[0] != b[0] {
return false
}
if a[1] != b[1] {
return false
}
if !bytes.Equal(a[2:4], b[2:4]) {
return false
}
if !bytes.Equal(a[6:40], b[6:40]) {
return false
}
return true
}
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
// Compare byte 1 fully so ECN must match.
// Skip [2:4] total len, [4:6] id, [10:12] csum.
if a[0] != b[0] {
return false
}
if a[1] != b[1] {
return false
}
if !bytes.Equal(a[6:10], b[6:10]) {
return false
}
if !bytes.Equal(a[12:20], b[12:20]) {
return false
}
return true
}
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
@@ -135,14 +145,17 @@ type Arena struct {
buf []byte
}
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
// NewArena returns an Arena with a pre-allocated backing of the given
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
// only feeds the coalescer pre-made []byte packets via Commit).
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.
// 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)
@@ -153,9 +166,16 @@ func (a *Arena) Reserve(sz int) []byte {
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.
// 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]
}
// Reserver hands out an sz-byte slice valid until its Resetter runs.
type Reserver func(sz int) []byte
// Resetter clears all reservations held by a Reserver. Only the arena's
// owner holds one; lanes inside a MultiCoalescer get nil.
type Resetter func()
-112
View File
@@ -1,112 +0,0 @@
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))
}
-76
View File
@@ -1,76 +0,0 @@
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)
}
+81 -82
View File
@@ -1,121 +1,119 @@
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.
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
// across lanes so the caller's allocation pattern is unchanged.
//
// 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.
// Lanes are processed independently: the TCP coalescer only sees TCP, the
// UDP coalescer only sees UDP, and the passthrough lane handles everything
// else. Per-flow arrival order is preserved because a single 5-tuple only
// ever lands in one lane and each lane preserves its own slot order.
//
// 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.
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
// This is acceptable because the carrier-side recvmmsg path already
// stable-sorts by (peer, message counter) before delivering plaintext
// here, so replay-window invariants are unaffected, and apps observe
// correct per-flow ordering — which is all the IP layer guarantees anyway.
// Do not "fix" this by interleaving lane outputs at flush time; that
// negates the entire point of coalescing (each lane needs to see runs of
// adjacent same-flow packets to coalesce them).
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
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
// and Flush resets it exactly once after every lane has drained.
arena *Arena
}
// 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
}
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
// burst worth of MTU-sized packets without the arena growing.
const DefaultMultiArenaCap = initialSlots * 65535
// 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 {
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
// likewise gates UDP coalescing (only enable when USO was negotiated).
// Either lane disabled redirects its traffic into the passthrough lane.
// arena is the single backing slab shared across every lane; the caller
// pre-sizes it via NewArena so the hot path never allocates.
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
m := &MultiCoalescer{
pt: NewPassthrough(w),
staged: make([]stagedPacket, 0, initialSlots),
pt: NewPassthrough(w, arena.Reserve, nil),
arena: arena,
}
if tcpEnabled {
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
}
if udpEnabled {
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
}
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
func (m *MultiCoalescer) Reserve(sz int) []byte {
return m.arena.Reserve(sz)
}
// 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
// Commit dispatches pkt to the appropriate lane based on IP version + L4
// proto. Borrowed slice contract is identical to the single-lane batchers,
// pkt must remain valid until the next Flush.
//
// On the success path the IP/TCP-or-UDP parse happens here once and the
// parsed struct is handed to the lane via commitParsed so the lane doesn't
// re-walk the header.
func (m *MultiCoalescer) Commit(pkt []byte) error {
if len(pkt) < 20 {
return m.pt.Commit(pkt)
}
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 {
v := pkt[0] >> 4
var proto byte
switch v {
case 4:
proto = pkt[9]
case 6:
if len(pkt) < 40 {
return m.pt.Commit(pkt)
}
proto = pkt[6]
default:
return m.pt.Commit(pkt)
}
switch proto {
case ipProtoTCP:
if m.tcp != nil {
return m.tcp.commitStaged(sp)
info, ok := parseTCPBase(pkt)
if !ok {
// Malformed/unsupported TCP shape (IP options, fragments, ...).
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
m.tcp.addPassthrough(pkt)
return nil
}
return m.tcp.commitParsed(pkt, info)
}
case ipProtoUDP:
if m.udp != nil {
return m.udp.commitStaged(sp)
info, ok := parseUDP(pkt)
if !ok {
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
return nil
}
return m.udp.commitParsed(pkt, info)
}
}
return m.pt.enqueue(sp.pkt)
return m.pt.Commit(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.
// Flush drains every lane in a fixed order, then resets the shared arena once.
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
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)
@@ -129,5 +127,6 @@ func (m *MultiCoalescer) Flush() error {
if err := m.pt.Flush(); err != nil {
errs = append(errs, err)
}
m.arena.Reset()
return errors.Join(errs...)
}
+18 -359
View File
@@ -1,39 +1,17 @@
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}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
tcpPay := make([]byte, 1200)
udpPay := make([]byte, 1200)
@@ -43,19 +21,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
icmp[3] = 28
icmp[9] = 1
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
if err := m.Commit(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 {
if err := m.Commit(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 {
if err := m.Commit(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 {
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
t.Fatal(err)
}
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
if err := m.Commit(icmp); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
@@ -70,162 +48,17 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
}
}
// 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) {
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
// the kernel via the passthrough lane rather than being lost.
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
pay := make([]byte, 1200)
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
// 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 {
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); 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 {
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
@@ -239,164 +72,16 @@ func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
}
}
// 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) {
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
func TestMultiCoalescerDisabledTCPFallsThrough(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")
}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
pay := make([]byte, 1200)
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
if err := m.Commit(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 {
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
@@ -409,29 +94,3 @@ func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
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
}
+34 -9
View File
@@ -2,29 +2,54 @@ package batch
import (
"io"
"github.com/slackhq/nebula/udp"
)
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
// order enqueued.
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
type Passthrough struct {
out io.Writer
slots [][]byte
out io.Writer
slots [][]byte
reserver Reserver
resetter Resetter
cursor int
}
func NewPassthrough(w io.Writer) *Passthrough {
const passthroughBaseNumSlots = 128
// DefaultPassthroughArenaCap is the recommended arena capacity for a
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
return &Passthrough{
out: w,
slots: make([][]byte, 0, 128),
out: w,
slots: make([][]byte, 0, passthroughBaseNumSlots),
reserver: reserver,
resetter: resetter,
}
}
// enqueue accepts one packet, already sorted into transmission order by dispatch.
func (p *Passthrough) enqueue(pkt []byte) error {
func (p *Passthrough) Reserve(sz int) []byte {
return p.reserver(sz)
}
func (p *Passthrough) Commit(pkt []byte) error {
p.slots = append(p.slots, pkt)
return nil
}
// Flush drains every queued packet and calls the configured Resetter
func (p *Passthrough) Flush() error {
firstErr := p.drain()
if p.resetter != nil {
p.resetter()
}
return firstErr
}
// drain writes out every queued packet and clears the slot list.
func (p *Passthrough) drain() error {
var firstErr error
for _, s := range p.slots {
_, err := p.out.Write(s)
+432 -176
View File
@@ -2,9 +2,12 @@ package batch
import (
"bytes"
"context"
"encoding/binary"
"io"
"log/slog"
"net/netip"
"slices"
"github.com/slackhq/nebula/overlay/tio"
)
@@ -20,19 +23,24 @@ const tcpCoalesceBufSize = 65535
// 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.
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
const tcpCoalesceHdrCap = 100
// coalesceSlot is one entry in the coalescer's ordered event queue. When
// passthrough is true the slot holds a single borrowed packet that must be
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
// passthrough is false the slot is an in-progress coalesced superpacket:
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
// length and pseudo-header partial at flush), and payIovs are *borrowed*
// slices from the caller's plaintext buffers — no payload is ever copied.
// The caller (listenOut) must keep those buffers 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
passthrough bool
rawPkt []byte // borrowed when passthrough
fk flowKey
hdrBuf [tcpCoalesceHdrCap]byte
hdrLen int
ipHdrLen int
isV6 bool
@@ -40,203 +48,216 @@ type coalesceSlot struct {
numSeg int
totalPay int
nextSeq uint32
payIovs [][]byte
// psh closes the chain: set when the last-accepted segment had PSH or
// was sub-gsoSize. No further appends after that.
psh bool
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.
// 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. All output — coalesced or not — is
// deferred until Flush so arrival order is preserved on the wire. Owns
// no locks; one coalescer per TUN write queue.
type TCPCoalescer struct {
w tio.GSOWriter
plainW io.Writer
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
// slots is the ordered event queue. Flush walks it once and emits each
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
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 maps a flow key to its most recent non-sealed slot, so new
// segments can extend an in-progress superpacket in O(1). Slots are
// removed from this map when they close (PSH or short-last-segment),
// when a non-admissible packet for that flow arrives, 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.
// lastSlot caches the most recently touched open slot. Steady-state
// bulk traffic is dominated by a single 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) entirely.
// Kept in lockstep with openSlots: nil whenever the slot it pointed
// at is removed.
// at is removed/sealed.
lastSlot *coalesceSlot
pool []*coalesceSlot // free list for reuse
reserver Reserver
resetter Resetter
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,
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
c := &TCPCoalescer{
plainW: w,
slots: make([]*coalesceSlot, 0, initialSlots),
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
pool: make([]*coalesceSlot, 0, initialSlots),
reserver: reserver,
resetter: resetter,
l: l,
}
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
c.gsoW = gw
}
return c
}
// 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
fk flowKey
ipHdrLen int
tcpHdrLen 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)
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
// regardless of whether it's admissible for coalescing. Returns ok=false
// for non-TCP or malformed input.
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
var p parsedTCP
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
if !ok {
return false
return p, false
}
return p.parseTail(trimmed, ipHdrLen)
}
pkt = ip.pkt
p.fk = ip.fk
p.ipHdrLen = ip.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
if len(pkt) < p.ipHdrLen+20 {
return p, false
}
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
if tcpOff < 20 || tcpOff > 60 {
return false
return p, false
}
if len(pkt) < ipHdrLen+tcpOff {
return false
if len(pkt) < p.ipHdrLen+tcpOff {
return p, false
}
p.ipHdrLen = ipHdrLen
p.hdrLen = ipHdrLen + tcpOff
p.tcpHdrLen = tcpOff
p.hdrLen = p.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
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
p.flags = pkt[p.ipHdrLen+13]
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
return p, 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.
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
// negative mask in coalesceable, not by name.
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
// coalesceable reports whether a parsed TCP segment is eligible for
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
// non-empty payload. CWR is excluded because it marks a one-shot
// congestion-window-reduced transition the receiver must observe at a
// segment boundary.
func (p parsedTCP) coalesceable() bool {
if p.flags&tcpFlagAck == 0 {
return false
}
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
return false
}
return p.payLen > 0
}
// 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)
func (c *TCPCoalescer) Reserve(sz int) []byte {
return c.reserver(sz)
}
// 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)
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
func (c *TCPCoalescer) Commit(pkt []byte) error {
if c.gsoW == nil {
c.addPassthrough(pkt)
return nil
}
var info parsedTCP
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
info, ok := parseTCPBase(pkt)
if !ok {
c.addPassthrough(pkt)
return nil
}
return c.commitParsed(sp.pkt, &info)
return c.commitParsed(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)
// commitParsed is the post-parse half of Commit. The caller must have
// already verified parseTCPBase succeeded (info is a valid TCP parse).
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
// after the dispatcher has already done so.
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
if c.gsoW == nil {
c.addPassthrough(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)
if !info.coalesceable() {
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
// Seal this flow's open slot so later in-flow packets don't extend
// it and accidentally reorder past this passthrough.
if last := c.lastSlot; last != nil && last.fk == info.fk {
c.lastSlot = nil
}
delete(c.openSlots, info.fk)
c.addPassthrough(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.
// Single-flow fast path: with only one open flow the cache hits every
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
// when there are multiple flows in flight (where the hit rate would
// be ~0 and the compare is pure overhead).
var open *coalesceSlot
if last := c.lastSlot; last != nil && last.fk == info.fk {
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && 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)
c.appendPayload(open, pkt, info)
if open.psh {
delete(c.openSlots, info.fk)
c.lastSlot = nil
} 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)
// Can't extend — seal it and fall through to seed a fresh slot.
delete(c.openSlots, info.fk)
if c.lastSlot == open {
c.lastSlot = nil
}
}
c.seed(pkt, info)
return nil
}
// Flush emits every queued event in (per-flow) seq order.
func (c *TCPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot (reordering/merging coalesced runs first)
// and clears the slot state.
func (c *TCPCoalescer) drain() error {
c.reorderForFlush()
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)
if s.passthrough {
_, err = c.plainW.Write(s.rawPkt)
} else {
err = c.flushSlot(s)
}
@@ -253,27 +274,23 @@ func (c *TCPCoalescer) Flush() error {
return first
}
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
s := c.take()
s.verbatim = true
s.passthrough = 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)
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
// Pathological shape can't fit our scratch, emit as-is.
c.addPassthrough(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.passthrough = false
s.rawPkt = nil
copy(s.hdrBuf[:], pkt[:info.hdrLen])
s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6
@@ -282,23 +299,26 @@ func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
s.numSeg = 1
s.totalPay = info.payLen
s.nextSeq = info.seq + uint32(info.payLen)
s.psh = info.flags&tcpFlagPsh != 0
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
c.slots = append(c.slots, s)
if info.flags&tcpFlagPsh == 0 {
if !s.psh {
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)
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
// PSH-on-seed seals the slot immediately. Any prior cached open
// slot for this flow has just been sealed-and-replaced by this
// passthrough-shaped seed, so drop the cache too.
c.lastSlot = nil
}
}
// 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 {
// canAppend reports whether info's packet extends the slot's seed: same
// header shape and stable contents, adjacent seq, not oversized, chain not closed.
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
if s.psh {
return false
}
if info.hdrLen != s.hdrLen {
return false
}
@@ -314,35 +334,31 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) b
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]
// 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.hdrBuf[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) {
if !headersMatch(s.hdrBuf[: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 {
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
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
// Propagate PSH into the seed header so kernel TSO sets it on the
// last segment. Without this the sender's push signal is dropped.
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
}
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
s.psh = true
}
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
}
func (c *TCPCoalescer) take() *coalesceSlot {
@@ -356,18 +372,21 @@ func (c *TCPCoalescer) take() *coalesceSlot {
}
func (c *TCPCoalescer) release(s *coalesceSlot) {
s.passthrough = false
s.rawPkt = nil
clear(s.payIovs)
*s = coalesceSlot{payIovs: s.payIovs[:0]}
s.payIovs = s.payIovs[:0]
s.numSeg = 0
s.totalPay = 0
s.psh = false
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.
// flushSlot patches the header and calls WriteGSO. 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]
hdr := s.hdrBuf[:s.hdrLen]
if s.isV6 {
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
@@ -387,7 +406,7 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
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)
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
}
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
@@ -418,6 +437,242 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
return true
}
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
// this pass a small wire reorder — counter 250 arriving in batch K when
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
// receiver as a much larger reorder than the wire actually had.
//
// Two phases:
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
// Cross-flow ordering inside a segment isn't preserved (it never was
// and doesn't matter for any single flow's TCP correctness).
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
// matters because the kernel TSO splitter chops at gsoSize from the
// start of the merged payload — a short segment in the middle would
// desynchronize every later segment.
//
// Passthrough slots act as barriers: the merge check skips them on either
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
// data.
func (c *TCPCoalescer) reorderForFlush() {
if len(c.slots) <= 1 {
return
}
runStart := 0
for i := 0; i <= len(c.slots); i++ {
if i < len(c.slots) && !c.slots[i].passthrough {
continue
}
c.sortRun(c.slots[runStart:i])
runStart = i + 1
}
out := c.slots[:0]
logged := false
for _, s := range c.slots {
if n := len(out); n > 0 {
prev := out[n-1]
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
// Same-flow neighbors after sort. If they aren't seq-
// contiguous it's a real gap — packets the wire reordered
// across batches, or actual loss before nebula. Log it so
// the operator can quantify how often it happens; the data
// itself still emits in seq order, kernel TCP handles the
// gap via its OOO queue.
if c.l.Enabled(context.Background(), slog.LevelDebug) {
if prev.nextSeq != slotSeedSeq(s) {
logged = true
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
c.l.Debug("tcp coalesce: cross-slot seq gap",
"src", flowKeyAddr(s.fk, false),
"dst", flowKeyAddr(s.fk, true),
"sport", s.fk.sport,
"dport", s.fk.dport,
"prev_seed_seq", slotSeedSeq(prev),
"prev_next_seq", prev.nextSeq,
"this_seed_seq", slotSeedSeq(s),
"gap_bytes", gap,
"prev_seg_count", prev.numSeg,
"prev_total_pay", prev.totalPay,
)
}
}
if canMergeSlots(prev, s) {
mergeSlots(prev, s)
c.release(s)
continue
}
}
}
out = append(out, s)
}
if logged {
c.l.Warn("==== end of batch ====")
}
c.slots = out
}
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
// logging. Only used on the cold gap-log path so the netip allocation
// doesn't matter.
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
src := fk.src
if dst {
src = fk.dst
}
if fk.isV6 {
return netip.AddrFrom16(src)
}
var v4 [4]byte
copy(v4[:], src[:4])
return netip.AddrFrom4(v4)
}
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
// cluster together in seq order, ready for the merge sweep. Stable so
// equal-key slots keep their original relative position (defensive — a
// duplicate seedSeq would already mean something's wrong upstream).
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
if len(run) <= 1 {
return
}
// slices.SortStableFunc with a free, non-capturing comparator avoids the
// reflection + closure-escape allocations that sort.SliceStable forces.
slices.SortStableFunc(run, compareCoalesceSlots)
}
func compareCoalesceSlots(a, b *coalesceSlot) int {
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
return cmp
}
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
if aSeq == bSeq {
return 0
}
if tcpSeqLess(aSeq, bSeq) {
return -1
}
return 1
}
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
// nextSeq tracks the seq just past the last appended byte; subtracting
// totalPay walks back to the seed. uint32 wraparound is the right TCP
// arithmetic so no special-casing is needed.
func slotSeedSeq(s *coalesceSlot) uint32 {
return s.nextSeq - uint32(s.totalPay)
}
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
// into the right comparison even across the 2^32 wrap.
func tcpSeqLess(a, b uint32) bool {
return int32(a-b) < 0
}
// flowKeyCompare orders flowKeys deterministically. The exact ordering
// is irrelevant — only that same-flow slots cluster together so the
// post-sort sweep can merge contiguous pairs.
func flowKeyCompare(a, b flowKey) int {
// Cheap scalar fields first so most non-matching keys short-circuit
// without ever calling bytes.Compare. sport is the ephemeral port on
// egress flows and discriminates fastest. For matching keys (same
// flow), array equality on src/dst inlines to word-sized compares,
// so we only pay bytes.Compare when the arrays actually differ.
if a.sport != b.sport {
if a.sport < b.sport {
return -1
}
return 1
}
if a.dport != b.dport {
if a.dport < b.dport {
return -1
}
return 1
}
if a.dst != b.dst {
return bytes.Compare(a.dst[:], b.dst[:])
}
if a.src != b.src {
return bytes.Compare(a.src[:], b.src[:])
}
if a.isV6 != b.isV6 {
if !a.isV6 {
return -1
}
return 1
}
return 0
}
// canMergeSlots reports whether s can fold into prev as one merged TSO
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
// fits within the kernel TSO limits. The tail-of-prev check rejects any
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
// would split the merged skb at gsoSize boundaries from the start, so a
// short segment in the middle would corrupt every later segment. PSH and
// ECE state must agree across both slots: PSH is a semantic delimiter
// (preserving the sender's push boundary) and ECE state must be uniform
// across a window (the same rule canAppend enforces for in-flow appends).
// The IP-level ECN codepoint must also match: this check calls headersMatch
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
// differing ECN marks stay separate superpackets, each keeping its own mark.
//
// Note: a slot sealed by reorder (canAppend returned false on seq
// mismatch) keeps psh=false, so this restriction does not block the
// reorder-fix merge — only legitimate PSH-set seals.
func canMergeSlots(prev, s *coalesceSlot) bool {
if prev.psh {
return false
}
if prev.fk != s.fk {
return false
}
if prev.gsoSize != s.gsoSize {
return false
}
if prev.nextSeq != slotSeedSeq(s) {
return false
}
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
return false
}
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
return false
}
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
return false
}
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
sFlags := s.hdrBuf[s.ipHdrLen+13]
if (prevFlags^sFlags)&tcpFlagEce != 0 {
return false
}
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
return false
}
return true
}
// mergeSlots folds src into dst in place: payIovs concatenated, counters
// and totals updated, PSH OR'd into the seed header so the push signal is
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
// is responsible for releasing src (it's no longer in c.slots after this call).
func mergeSlots(dst, src *coalesceSlot) {
dst.payIovs = append(dst.payIovs, src.payIovs...)
dst.numSeg += src.numSeg
dst.totalPay += src.totalPay
dst.nextSeq = src.nextSeq
if src.psh {
dst.psh = true
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
}
}
// 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.
@@ -462,8 +717,9 @@ func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
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
// 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 — the kernel will add the payload sum and invert.
func foldOnceNoInvert(sum uint32) uint16 {
for sum>>16 != 0 {
sum = (sum & 0xffff) + (sum >> 16)
+78 -51
View File
@@ -2,9 +2,9 @@ package batch
import (
"encoding/binary"
"runtime"
"testing"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
)
@@ -55,31 +55,7 @@ func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
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
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
// branch in Commit.
func buildICMPv4() []byte {
pkt := make([]byte, 28)
@@ -95,7 +71,8 @@ func buildICMPv4() []byte {
// between batches, and reports per-packet cost.
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
c := newTestTCPCoalescer(b, nopTunWriter{})
arena := NewArena(0)
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
@@ -136,17 +113,8 @@ func BenchmarkCommitInterleaved16(b *testing.B) {
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.
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
// bails early and addPassthrough is the only work.
func BenchmarkCommitPassthrough(b *testing.B) {
pkt := buildICMPv4()
pkts := make([][]byte, 64)
@@ -158,7 +126,7 @@ func BenchmarkCommitPassthrough(b *testing.B) {
// 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.
// map delete + passthrough. Measures the seal-without-slot cost.
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
pay := make([]byte, 0)
pkts := make([][]byte, 64)
@@ -168,24 +136,18 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
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.
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
// the IP/L4 parse once and passes the parsed struct to the lane, so this
// is the bench that shows the savings of skipping the lane's re-parse.
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)
}
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true)
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 {
pkt := pkts[i%len(pkts)]
if err := m.Commit(pkt); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
@@ -212,3 +174,68 @@ func BenchmarkMultiCommitInterleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runMultiCommitBench(b, pkts, len(pkts))
}
// flowKeyPair is one comparison input for the flowKeyCompare bench.
type flowKeyPair struct{ a, b flowKey }
// makeFlowKey builds an IPv4 flowKey from compact inputs.
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
var fk flowKey
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
fk.sport = sport
fk.dport = dport
return fk
}
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
// repeatedly when many segments share a flow).
// - sportDiffers: same src/dst/dport, different sport — the typical
// "sibling flows from one host to one server" pattern.
// - dstDiffers: same src/sport/dport, different dst — outbound to many
// servers from a fixed local port.
// - allDiffer: every field differs; worst case for short-circuiting.
func flowKeyCases() map[string][]flowKeyPair {
const n = 64
cases := map[string][]flowKeyPair{
"sameFlow": make([]flowKeyPair, n),
"sportDiffers": make([]flowKeyPair, n),
"dstDiffers": make([]flowKeyPair, n),
"allDiffer": make([]flowKeyPair, n),
}
for i := range n {
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
cases["sportDiffers"][i] = flowKeyPair{
a: base,
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
}
cases["dstDiffers"][i] = flowKeyPair{
a: base,
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
}
cases["allDiffer"][i] = flowKeyPair{
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
}
}
return cases
}
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
// the sort step actually sees. Use this to compare reorderings.
func BenchmarkFlowKeyCompare(b *testing.B) {
for name, pairs := range flowKeyCases() {
b.Run(name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var sink int
for i := 0; i < b.N; i++ {
p := pairs[i&(len(pairs)-1)]
sink += flowKeyCompare(p.a, p.b)
}
runtime.KeepAlive(sink)
})
}
}
File diff suppressed because it is too large Load Diff
+9 -8
View File
@@ -6,7 +6,7 @@ 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)
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
}
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
@@ -16,6 +16,7 @@ type SendBatch struct {
out batchWriter
bufs [][]byte
dsts []netip.AddrPort
ecns []byte
arena *Arena
}
@@ -25,6 +26,7 @@ func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
out: out,
bufs: make([][]byte, 0, batchCap),
dsts: make([]netip.AddrPort, 0, batchCap),
ecns: make([]byte, 0, batchCap),
arena: NewArena(arenaSize),
}
}
@@ -38,22 +40,21 @@ func (b *SendBatch) Reserve(sz int) []byte {
// 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) {
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
b.bufs = append(b.bufs, pkt)
b.dsts = append(b.dsts, dst)
b.ecns = append(b.ecns, outerECN)
}
// 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) {
func (b *SendBatch) Flush() error {
var err error
written := 0
if len(b.bufs) > 0 {
written, err = b.out.WriteBatch(b.bufs, b.dsts)
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
}
clear(b.bufs)
b.bufs = b.bufs[:0]
b.dsts = b.dsts[:0]
b.ecns = b.ecns[:0]
b.arena.Reset()
return written, err
return err
}
+12 -10
View File
@@ -8,9 +8,10 @@ import (
type fakeBatchWriter struct {
bufs [][]byte
addrs []netip.AddrPort
ecns []byte
}
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) 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))
@@ -20,7 +21,8 @@ func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int
w.bufs[i] = cp
}
w.addrs = append(w.addrs[:0], addrs...)
return len(bufs), nil
w.ecns = append(w.ecns[:0], ecns...)
return nil
}
func TestSendBatchReserveCommitFlush(t *testing.T) {
@@ -34,9 +36,9 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
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)
b.Commit(pkt, ap, 0)
}
if _, err := b.Flush(); err != nil {
if err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
if len(fw.bufs) != 4 {
@@ -53,7 +55,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
// Flush again with nothing committed — should be a no-op.
fw.bufs = nil
if _, err := b.Flush(); err != nil {
if err := b.Flush(); err != nil {
t.Fatalf("empty Flush: %v", err)
}
if fw.bufs != nil {
@@ -75,9 +77,9 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
for i := 0; i < 3; i++ {
s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap)
b.Commit(pkt, ap, 0)
}
if _, err := b.Flush(); err != nil {
if err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
@@ -96,18 +98,18 @@ func TestSendBatchGrowPreservesCommitted(t *testing.T) {
s1 := b.Reserve(4)
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
b.Commit(pkt1, ap)
b.Commit(pkt1, ap, 0)
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
b.Commit(pkt2, ap)
b.Commit(pkt2, ap, 0)
// 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 {
if err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
if len(fw.bufs) != 2 {
+152 -142
View File
@@ -1,7 +1,6 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
@@ -19,56 +18,72 @@ const udpCoalesceBufSize = 65535
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
const udpCoalesceMaxSegs = 64
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
const udpCoalesceHdrCap = 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
passthrough bool
rawPkt []byte // borrowed when passthrough
fk flowKey
hdrBuf [udpCoalesceHdrCap]byte
hdrLen int
ipHdrLen int
isV6 bool
gsoSize int // per-segment UDP payload length
numSeg int
totalPay int
payIovs [][]byte
// sealed closes the chain: set when a sub-gsoSize segment is appended
// (kernel UDP-GSO requires every segment but the last to be exactly
// gsoSize) or when limits are hit. No further appends after.
sealed bool
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
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
// underlying writer doesn't support USO.
//
// All output — coalesced or not — is deferred until Flush so per-flow
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
// across the TCP/UDP/passthrough split when this coalescer runs alongside
// others — see multi_coalesce.go. Per-flow order is preserved because a
// single 5-tuple only ever lands in one lane and each lane preserves its
// own slot order.
//
// Owns no locks; one coalescer per TUN write queue.
type UDPCoalescer struct {
w tio.GSOWriter
plainW io.Writer
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
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
pool []*udpSlot
reserver Reserver
resetter Resetter
}
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
if !ok {
return nil
}
return &UDPCoalescer{
w: gw,
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
// this when the underlying Queue's Capabilities advertise USO; otherwise
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
// plain Writes — same defensive shape as the TCP coalescer.
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer {
c := &UDPCoalescer{
plainW: w,
slots: make([]*udpSlot, 0, initialSlots),
openSlots: make(map[flowKey]*udpSlot, initialSlots),
pool: make([]*udpSlot, 0, initialSlots),
reserver: reserver,
resetter: resetter,
}
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
c.gsoW = gw
}
return c
}
// parsedUDP holds the fields extracted from a single parse so later steps
@@ -80,111 +95,104 @@ type parsedUDP struct {
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)
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
// (IPv4 with options/fragmentation, IPv6 with extension headers).
func parseUDP(pkt []byte) (parsedUDP, bool) {
var p parsedUDP
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
if !ok {
return false
return p, false
}
return p.parseTail(trimmed, ipHdrLen)
}
pkt = ip.pkt
p.fk = ip.fk
p.ipHdrLen = ip.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
if len(pkt) < p.ipHdrLen+8 {
return p, false
}
p.hdrLen = p.ipHdrLen + 8
// 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
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
return p, 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
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
return p, 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)
func (c *UDPCoalescer) Reserve(sz int) []byte {
return c.reserver(sz)
}
// 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)
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
func (c *UDPCoalescer) Commit(pkt []byte) error {
if c.gsoW == nil {
c.addPassthrough(pkt)
return nil
}
var info parsedUDP
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
info, ok := parseUDP(pkt)
if !ok {
c.addPassthrough(pkt)
return nil
}
return c.commitParsed(sp.pkt, &info)
return c.commitParsed(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.
// commitParsed is the post-parse half of Commit. The caller must have
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
// avoid re-walking the IP/UDP header.
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
if c.gsoW == nil {
c.addPassthrough(pkt)
return nil
}
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
// reach the TUN, but it can't be coalesced: a GSO slot would store an
// empty payload iovec and the kernel has nothing to segment. Seal any
// open chain for this flow (so a later, non-empty datagram seeds fresh
// *after* this one and per-flow arrival order is preserved) and deliver
// it as a plain single datagram.
if info.payLen == 0 {
c.sealFlow(info.fk)
c.addVerbatim(pkt)
delete(c.openSlots, info.fk)
c.addPassthrough(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 open := c.openSlots[info.fk]; 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
c.appendPayload(open, pkt, info)
if open.sealed {
delete(c.openSlots, info.fk)
}
return nil
}
// Can't extend: evict it from openSlots and fall through to seed a
// fresh slot.
c.sealFlow(info.fk)
// Can't extend — seal it and fall through to seed a fresh slot.
delete(c.openSlots, info.fk)
}
c.seed(pkt, info)
return nil
}
// Flush drains every queued slot and calls the configured Resetter.
func (c *UDPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot in arrival order and clears the slot state.
// It does NOT reset the arena: borrowed payload slices stay valid until the
// arena's owner recycles it.
func (c *UDPCoalescer) drain() 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)
if s.passthrough {
_, err = c.plainW.Write(s.rawPkt)
} else {
err = c.flushSlot(s)
}
@@ -196,38 +204,25 @@ func (c *UDPCoalescer) Flush() error {
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) {
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
s := c.take()
s.verbatim = true
s.passthrough = 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)
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
c.addPassthrough(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.passthrough = false
s.rawPkt = nil
copy(s.hdrBuf[:], pkt[:info.hdrLen])
s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6
@@ -235,16 +230,19 @@ func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
s.gsoSize = info.payLen
s.numSeg = 1
s.totalPay = info.payLen
s.sealed = false
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 {
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
if s.sealed {
return false
}
if info.hdrLen != s.hdrLen {
return false
}
@@ -257,25 +255,20 @@ func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
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) {
if !udpHeadersMatch(s.hdrBuf[: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 {
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++
s.totalPay += info.payLen
return info.payLen < s.gsoSize
if info.payLen < s.gsoSize {
// Last-segment-can-be-shorter: this seals the chain.
s.sealed = true
}
}
func (c *UDPCoalescer) take() *udpSlot {
@@ -289,19 +282,30 @@ func (c *UDPCoalescer) take() *udpSlot {
}
func (c *UDPCoalescer) release(s *udpSlot) {
// Reset every field, identity ones included; see TCPCoalescer.release.
s.passthrough = false
s.rawPkt = nil
clear(s.payIovs)
*s = udpSlot{payIovs: s.payIovs[:0]}
s.payIovs = s.payIovs[:0]
s.numSeg = 0
s.totalPay = 0
s.sealed = false
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.
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
// values would silently drop everything but the first segment. The kernel
// then walks each segment in __udp_gso_segment, recomputing per-segment
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
// our seed's uh->check must be consistent with the seed's uh->len, which
// is what passing the total to both pseudoSum and the UDP length field
// guarantees.
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
hdr := s.rawPkt[:s.hdrLen]
hdr := s.hdrBuf[:s.hdrLen]
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
@@ -326,11 +330,14 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
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)
return c.gsoW.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
// every field that must be identical across coalesced segments. Length
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
// don't coalesce, matching kernel GRO.
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) {
return false
@@ -338,8 +345,11 @@ func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if !ipHeadersMatch(a, b, isV6) {
return false
}
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
// 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])
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
return false
}
return true
}
-72
View File
@@ -1,72 +0,0 @@
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))
}
+70 -134
View File
@@ -1,9 +1,7 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
"testing"
)
@@ -60,31 +58,29 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
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")
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: false}
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 100))
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
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 len(w.writes) != 0 || len(w.gsoWrites) != 0 {
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
t.Fatalf("want nil for a plain writer, got %v", c)
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
}
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
// ICMP packet
pkt := make([]byte, 28)
pkt[0] = 0x45
@@ -105,7 +101,8 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 800))
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
@@ -113,21 +110,17 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
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 {
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
if len(w.gsoWrites) != 1 || len(w.writes) != 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)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -167,7 +160,8 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
// Last segment may be shorter, sealing the chain.
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 1200)
tail := make([]byte, 600)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
@@ -186,23 +180,22 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
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) != 2 {
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
}
if len(w.gsoWrites[0].pays) != 3 {
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
t.Errorf("first 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)
if len(w.gsoWrites[1].pays) != 1 {
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
}
}
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err)
}
@@ -212,21 +205,16 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
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)
}
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
}
}
// Different 5-tuples must not coalesce.
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
@@ -257,7 +245,8 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
// Caps at udpCoalesceMaxSegs.
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 100)
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -282,12 +271,12 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
// 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.
// seals the Not-ECT chain and seeds a fresh superpacket that keeps CE; the
// trailing Not-ECT datagram seeds another.
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
pkt1 := buildUDPv4(1000, 53, pay)
@@ -301,13 +290,16 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
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))
if len(w.gsoWrites) != 3 {
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
}
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])
for i, g := range w.gsoWrites {
if len(g.pays) != 1 {
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
}
if got := g.hdr[1] & 0x03; got != wantECN[i] {
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
}
}
}
@@ -315,7 +307,8 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
// IPv6 path: same flow, equal-sized → coalesced.
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
@@ -351,7 +344,8 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay)
pkt1 := buildUDPv4(1000, 53, pay)
@@ -365,16 +359,16 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
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))
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
}
}
// Fragmented IPv4 must not be coalesced.
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 200))
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
if err := c.Commit(pkt); err != nil {
@@ -395,7 +389,8 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
// 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)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
@@ -411,10 +406,11 @@ func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
}
}
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
// IPv6 zero-length UDP datagram: same passthrough contract as v4.
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
@@ -435,7 +431,8 @@ func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
// wire — per-flow arrival order (full, empty, full) must be preserved.
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
@@ -450,23 +447,17 @@ func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
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)
}
// can't join it: two single-segment superpackets bracket one plain write.
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
}
}
// IPv4 with options is not admissible (we require IHL=5).
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
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 {
@@ -479,58 +470,3 @@ func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
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)
}
}
}
+63 -105
View File
@@ -8,69 +8,37 @@ import (
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.
// seeds and a handful of starting alignments, asserting that our local
// Checksum 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
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())
}
// 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}
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)
}
}
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 := Checksum(buf, seed)
if got != want {
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
length, off, seed, got, want)
}
}
})
}
}
}
@@ -78,28 +46,23 @@ func TestChecksumMatchesGvisor(t *testing.T) {
// 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)
}
}
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 := Checksum(buf, seed)
if got != want {
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
name, length, seed, got, want)
}
}
})
}
}
}
@@ -135,41 +98,36 @@ func ascending(n int) []byte {
// 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
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())
}
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
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)
}
}
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 := Checksum(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)
}
}
}
})
}
}
}
-11
View File
@@ -1,11 +0,0 @@
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},
}
-8
View File
@@ -1,8 +0,0 @@
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},
}
-7
View File
@@ -1,7 +0,0 @@
//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
+4 -3
View File
@@ -9,9 +9,10 @@ import (
"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.
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for
// reads, POLLOUT for writes) or shutdownFd signals teardown. 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).
+15 -17
View File
@@ -7,7 +7,6 @@ import (
"encoding/binary"
"errors"
"fmt"
"log/slog"
"sync/atomic"
"golang.org/x/sys/unix"
@@ -18,17 +17,18 @@ type offloadQueueSet struct {
// 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 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) {
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
// should fall back to per-packet writes when this is false.
func NewOffloadQueueSet(usoEnabled bool) (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)
@@ -39,7 +39,6 @@ func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
pqi: []Queue{},
shutdownFd: shutdownFd,
usoEnabled: usoEnabled,
l: l,
}
return out, nil
@@ -50,10 +49,7 @@ func (c *offloadQueueSet) Queues() []Queue {
}
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)
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
if err != nil {
return err
}
@@ -77,21 +73,23 @@ func (c *offloadQueueSet) Close() error {
errs := []error{}
// Signal all readers blocked in poll to wake up and exit.
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
// 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.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
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.
// 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)
}
+6 -7
View File
@@ -40,9 +40,6 @@ func (c *pollQueueSet) Queues() []Queue {
}
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
@@ -67,21 +64,23 @@ func (c *pollQueueSet) Close() error {
errs := []error{}
// Signal all readers blocked in poll to wake up and exit.
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
// Wake any reader blocked in poll so it observes POLLIN on the shutdown
// eventfd and returns 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.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
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.
// 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)
}
+7 -1
View File
@@ -1,4 +1,4 @@
//go:build !linux || android
//go:build !linux || android || e2e_testing
package tio
@@ -8,6 +8,12 @@ func protoFromGSOType(_ uint8) (GSOProto, error) {
return 0, fmt.Errorf("GSO unsupported")
}
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
// builds (and Android/e2e_testing) this package does not provide a Queue
// implementation, so any caller that does construct a Packet here can only
// be operating on non-superpacket bytes and the stub forwards them
// directly. A non-zero GSO field is a programming error from the caller
// and returns an explicit error rather than silently misbehaving.
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")
+6 -5
View File
@@ -4,8 +4,9 @@ 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.
// 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)
@@ -13,9 +14,9 @@ type singleQueue struct {
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.
// 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)}
}
+85 -60
View File
@@ -13,72 +13,79 @@ type QueueSet interface {
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.
// Capabilities advertises which kernel offload features a Queue
// successfully negotiated. Callers consult this to decide which coalescers
// to wire onto the write path — a Queue without TSO can't usefully accept a
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
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 means the FD was opened with IFF_VNET_HDR and the kernel agreed
// to TUN_F_TSO4|TSO6 — i.e. 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 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;
// 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 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 - callers must encrypt
// or copy each slice before the next call. A Packet may carry a
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
// true the caller must segment Bytes before treating it as a single
// IP datagram. 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 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.
// 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.
// The zero value means "not a superpacket" — 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.
// (== 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 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 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 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.
// 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.
// GSO metadata is copied verbatim. Use this only when a caller genuinely
// needs to outlive the borrowed-slice contract — the hot path reads should
// continue to consume the borrow synchronously to avoid the allocation.
func (p Packet) Clone() Packet {
if p.Bytes == nil {
return p
@@ -88,60 +95,78 @@ func (p Packet) Clone() Packet {
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.
// CapsProvider is an optional interface implemented by Queues that
// successfully negotiated 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 — callers must
// fall back to plain per-packet writes.
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
// QueueCapabilities returns q's negotiated offload capabilities, or the
// zero value when q does not advertise any.
func QueueCapabilities(q Queue) Capabilities {
if cp, ok := q.(CapsProvider); ok {
return cp.Capabilities()
}
return 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
GSOProtoTCP GSOProto = iota
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.
// 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.
// 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 is 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.
// Callers should also consult CapsProvider (via SupportsGSO or
// QueueCapabilities) for the per-protocol negotiated capability; an
// implementation of GSOWriter is necessary but not sufficient since 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) {
// queue advertises the negotiated capability for `want`. A writer that
// implements GSOWriter but not CapsProvider is treated as permissive
// (used by tests and fakes that don't negotiate).
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
gw, ok := w.(GSOWriter)
if !ok {
return nil, false
}
caps := gw.Capabilities()
cp, ok := w.(CapsProvider)
if !ok {
return gw, true
}
caps := cp.Capabilities()
switch want {
case GSOProtoTCP:
return gw, caps.TSO
case GSOProtoUDP:
return gw, caps.USO
default:
return gw, false
}
return gw, false
}
+150 -153
View File
@@ -4,7 +4,6 @@
package tio
import (
"context"
"fmt"
"io"
"log/slog"
@@ -18,56 +17,71 @@ import (
"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.
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
// 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.
// We round up to give comfortable 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.
// tunRxBufCap is the total size we allocate for the per-reader rx
// buffer. With reads landing directly in rxBuf, 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.
// 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.
// 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 — fine to keep resident on every queue. 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.
// 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, which marks the skb
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
// verification. All packets that reach the plain Write paths already carry
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during
// superpacket segmentation, or built locally by CreateRejectPacket), so
// trusting them is safe.
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
closed atomic.Bool
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
// 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
@@ -78,27 +92,9 @@ type Offload struct {
// 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) {
func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
if err := unix.SetNonblock(fd, true); err != nil {
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
}
@@ -108,7 +104,6 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offlo
shutdownFd: shutdownFd,
usoEnabled: usoEnabled,
closed: atomic.Bool{},
l: l,
rxBuf: make([]byte, tunRxBufCap),
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
@@ -133,18 +128,26 @@ 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.
// 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.
//
// The body iovec capacity is always tunReadBufSize; callers (the Read
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
// hold one worst-case kernel-supplied packet body. Without that gate the
// body iovec could be smaller than the next inbound packet and the
// kernel would truncate.
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)
r.readIovs[1].SetLen(tunReadBufSize)
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 0, io.ErrShortWrite
}
return int(n) - virtio.Size, nil
}
@@ -167,30 +170,29 @@ func (r *Offload) readPacket(block bool) (int, error) {
}
}
// 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.
// 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.
// 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)
// Drop and read again — a bad packet should not kill the reader.
continue
}
break
@@ -212,7 +214,6 @@ func (r *Offload) Read() ([]Packet, error) {
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
}
}
@@ -220,20 +221,13 @@ func (r *Offload) Read() ([]Packet, error) {
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
// 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 past the
// kernel-supplied body and nothing else, since segmentation no longer
// writes back into rxBuf.
func (r *Offload) decodeRead(pktLen int) error {
if pktLen <= 0 {
return fmt.Errorf("short tun read: %d", pktLen)
@@ -243,7 +237,7 @@ func (r *Offload) decodeRead(pktLen int) error {
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
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
@@ -254,13 +248,17 @@ func (r *Offload) decodeRead(pktLen int) error {
return nil
}
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
// the metadata. The bytes stay in rxBuf untouched, segmentation
// happens in SegmentSuperpacket at encrypt time.
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())
proto, err := protoFromGSOType(hdr.GSOType)
if err != nil {
return err
}
@@ -278,16 +276,22 @@ func (r *Offload) decodeRead(pktLen int) error {
}
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))
return r.writeWithScratch(buf, &iovs)
}
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
if len(buf) == 0 {
return 0, nil
}
iovs[1].Base = &buf[0]
iovs[1].SetLen(len(buf))
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
}
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
@@ -317,34 +321,57 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
// 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).
// 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.
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
return nil
}
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
var csumOff uint16
switch proto {
case GSOProtoUDP:
csumOff = 6
case GSOProtoTCP:
csumOff = 16
default:
return fmt.Errorf("unknown GSO proto: %d", proto)
csumOff = 16
}
// 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)
vhdr := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
HdrLen: uint16(len(hdr) + len(transportHdr)),
GSOSize: uint16(len(pays[0])),
CsumStart: uint16(len(hdr)),
CsumOffset: csumOff,
}
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
if len(pays) > 1 {
ipVer := hdr[0] >> 4
switch {
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
case ipVer == 6:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
case ipVer == 4:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
default:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
vhdr.GSOSize = 0
}
} else {
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
vhdr.GSOSize = 0
}
vhdr.Encode(r.gsoHdrBuf[:])
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
// wired to gsoHdrBuf at construction and never changes.
need := 3 + len(pays)
if need > cap(r.gsoIovs) {
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
}
r.gsoIovs = r.gsoIovs[:need]
@@ -352,52 +379,22 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
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 {
// Defense in depth: an empty payload fragment can't be a valid GSO
// segment and &p[0] would panic on it. Callers route zero-length
// datagrams through the plain path (see UDPCoalescer.commitParsed), so
// this should never fire, but skip empties rather than index into one.
// `n` tracks where the next payload iovec lands, since skips make it
// drift from 3+i.
n := 3
for _, 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)
continue
}
total += len(p)
r.gsoIovs[3+i].Base = &p[0]
r.gsoIovs[3+i].SetLen(len(p))
r.gsoIovs[n].Base = &p[0]
r.gsoIovs[n].SetLen(len(p))
n++
}
// 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*/
)
r.gsoIovs = r.gsoIovs[:n]
_, err := r.rawWrite(r.gsoIovs)
return err
@@ -408,10 +405,10 @@ func (r *Offload) Close() error {
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.
//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 readOne, and mutating the field would race that load.
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
// ppoll first). closed.Swap already guarantees we only close once.
return unix.Close(r.fd)
}
+15 -16
View File
@@ -11,6 +11,11 @@ import (
"golang.org/x/sys/unix"
)
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
// prefix plus the virtio header.
const tunReadBufSize = 65535
type Poll struct {
fd int
shutdownFd int
@@ -20,10 +25,10 @@ type Poll struct {
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.
// 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)
@@ -32,7 +37,7 @@ func newPoll(fd int, shutdownFd int) (*Poll, error) {
out := &Poll{
fd: fd,
shutdownFd: shutdownFd,
readBuf: make([]byte, 65535), // largest possible size Linux permits
readBuf: make([]byte, tunReadBufSize),
}
return out, nil
}
@@ -47,11 +52,6 @@ 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 {
@@ -108,11 +108,10 @@ 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.
//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.
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
// ppoll first). closed.Swap already guarantees we only close once.
return unix.Close(t.fd)
}
+1 -2
View File
@@ -5,7 +5,6 @@ package tio
import (
"errors"
"log/slog"
"os"
"sync"
"testing"
@@ -211,7 +210,7 @@ func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
// 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))
qs, err := NewOffloadQueueSet(false)
require.NoError(t, err)
c, ok := qs.(*offloadQueueSet)
require.True(t, ok)
+16 -25
View File
@@ -1,5 +1,5 @@
//go:build linux && !android
// +build linux,!android
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
@@ -11,11 +11,11 @@ import (
"github.com/slackhq/nebula/overlay/tio/virtio"
)
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
// 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.
// 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 {
switch t {
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
return GSOProtoTCP, nil
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
@@ -25,26 +25,17 @@ func protoFromGSOType(t uint8) (GSOProto, error) {
}
}
// 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.
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts
// fn is called once with pkt.Bytes (no segmentation, no copy). 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). The slide is destructive: pkt is
// consumed by this call and its bytes are in an undefined state when
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier
// seg slice past fn's return for that segment. The scratch parameter is
// unused on the destructive path and kept only for cross-platform
// signature compatibility. 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)
+88 -384
View File
@@ -19,36 +19,6 @@ import (
// 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 {
@@ -63,7 +33,7 @@ func verifyChecksum(b []byte, pseudo uint16) bool {
// 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 {
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 {
@@ -73,7 +43,7 @@ func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) e
*out = append(*out, cp)
return nil
}
proto, err := protoFromGSOType(hdr.GSOType())
proto, err := protoFromGSOType(hdr.GSOType)
if err != nil {
return err
}
@@ -140,14 +110,15 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
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*/
)
return pkt, virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
HdrLen: uint16(ipLen + tcpLen),
GSOSize: uint16(mss),
CsumStart: uint16(ipLen),
CsumOffset: 16,
}
}
func TestSegmentTCPv4(t *testing.T) {
@@ -261,14 +232,14 @@ func TestSegmentTCPv6(t *testing.T) {
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*/
)
hdr := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
HdrLen: uint16(ipLen + tcpLen),
GSOSize: uint16(mss),
CsumStart: uint16(ipLen),
CsumOffset: 16,
}
scratch := make([]byte, testSegScratchSize)
var out [][]byte
@@ -310,7 +281,7 @@ func TestSegmentTCPv6(t *testing.T) {
func TestSegmentGSONonePassesThrough(t *testing.T) {
pkt, hdr := buildTSOv4(t, 100, 100)
hdr.SetGSOType(unix.VIRTIO_NET_HDR_GSO_NONE)
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
scratch := make([]byte, testSegScratchSize)
@@ -329,7 +300,7 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
// 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)
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
var out [][]byte
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
t.Fatalf("expected rejection for legacy UDP GSO")
@@ -353,26 +324,22 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
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
// UDP header (length + checksum filled in per segment by segmentUDPYield)
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, 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*/
)
return pkt, virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
HdrLen: uint16(ipLen + udpLen),
GSOSize: uint16(gsoSize),
CsumStart: uint16(ipLen),
CsumOffset: 6,
}
}
func TestSegmentUDPv4(t *testing.T) {
@@ -397,12 +364,11 @@ func TestSegmentUDPv4(t *testing.T) {
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.
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every
// segment carries the same ID as the seed.
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))
if id != 0x4242 {
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242)
}
udpLen := binary.BigEndian.Uint16(seg[24:26])
if udpLen != uint16(8+gso) {
@@ -470,21 +436,19 @@ func TestSegmentUDPv6(t *testing.T) {
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*/
)
hdr := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
HdrLen: uint16(ipLen + udpLen),
GSOSize: uint16(gso),
CsumStart: uint16(ipLen),
CsumOffset: 6,
}
scratch := make([]byte, testSegScratchSize)
var out [][]byte
@@ -616,14 +580,14 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
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*/
)
hdr := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
HdrLen: uint16(ipLen + tcpLen),
GSOSize: uint16(sz.mss),
CsumStart: uint16(ipLen),
CsumOffset: 16,
}
scratch := make([]byte, testSegScratchSize)
out := make([][]byte, 0, 64)
@@ -676,72 +640,35 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
}
}
// 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) }},
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the
// zero-length UDP DoS: a payload fragment of length zero would make &p[0]
// panic (index-out-of-range) when building the iovec array. WriteGSO must
// skip empties instead. We write to /dev/null so the writev always succeeds
// synchronously; the point is simply that neither call panics.
func TestWriteGSOSkipsEmptyPayloads(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) })
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,
}}
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
o.gsoIovs[0].SetLen(virtio.Size)
// 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)
}
}
ipHdr := make([]byte, 20)
ipHdr[0] = 0x45 // IPv4, IHL 5
udpHdr := make([]byte, 8)
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)
}
})
// Sole payload empty: exercises the all-empty skip (n stays at 3).
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil {
t.Fatalf("WriteGSO with a single empty payload: %v", err)
}
// Empty mixed with a real fragment: exercises the index-drift skip so a
// later non-empty payload still lands in the right iovec slot.
real := make([]byte, 1200)
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil {
t.Fatalf("WriteGSO with a trailing empty payload: %v", err)
}
}
@@ -794,11 +721,9 @@ 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
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec
// is what bounds the kernel's superpacket length.
pktLen := tunReadBufSize
payLen := pktLen - headerLen
const targetSegs = 64
gsoSize := (payLen + targetSegs - 1) / targetSegs
@@ -820,14 +745,14 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
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 := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
HdrLen: uint16(headerLen),
GSOSize: uint16(gsoSize),
CsumStart: uint16(ipv6HdrLen),
CsumOffset: 16,
}
hdr.Encode(o.readVnetScratch[:])
startRxOff := o.rxOff
@@ -899,224 +824,3 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
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)
}
}
})
}
}
+9 -41
View File
@@ -3,11 +3,7 @@
package virtio
import (
"encoding/binary"
"golang.org/x/sys/unix"
)
import "encoding/binary"
// Size is the on-wire length of struct virtio_net_hdr the kernel
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
@@ -17,59 +13,31 @@ 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"
GSOType uint8
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.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
b[0] = h.Flags
b[1] = h.GSOType
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize)
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart)
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset)
}
+136 -143
View File
@@ -27,8 +27,11 @@ const (
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
// 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; the array is sized to that
// worst case so the snapshot lives on the stack with no per-call heap
// allocation.
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
// Byte offsets inside an IPv4 header.
@@ -62,72 +65,63 @@ const (
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.
// 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 {
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
// carry coalescing info rather than checksum offsets. A TUN writing via
// IFF_VNET_HDR should never emit this, but if it did we would silently
// miscompute the segment checksums — refuse the packet instead.
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
return fmt.Errorf("packet too short")
}
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 {
switch hdr.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)
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)
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)
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 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.
// 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.
// 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 {
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
hdr.HdrLen = hdr.CsumStart + 8
} else {
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
@@ -135,7 +129,8 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
}
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
if tcpHLen < 20 || tcpHLen > 60 {
// A TCP header must be between 20 and 60 bytes in length.
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
}
hdr.HdrLen = hdr.CsumStart + tcpHLen
@@ -155,62 +150,19 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
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.
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
// slice into pkt itself. 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. The stamp is destructive but harmless: iter i's header write lands
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers
// live payload — this holds even when gsoSize < hdrLen. The header bytes are
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
// overwrite the leading header in place and every stamp after the first would
// copy corrupted bytes. 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")
@@ -229,28 +181,49 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
payLen := len(pkt) - headerLen
gsoSize := int(gsoSizeU)
numSeg := segCount(payLen, gsoSize)
numSeg := (payLen + gsoSize - 1) / gsoSize
if numSeg == 0 {
numSeg = 1
}
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 tmp [tcpHeaderMaxLen]byte
copy(tmp[:tcpHdrLen], 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
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
var baseProtoSum uint32
if isV4 {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
} else {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
}
baseProtoSum += uint32(unix.IPPROTO_TCP)
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
ihl := int(pkt[0]&0x0f) * 4
if ihl < ipv4HeaderMinLen || ihl > csumStart {
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
}
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
}
// 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.
// 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. The variable fields (seq/flags/cksum/
// totalLen/id) captured here are stale but are overwritten per segment.
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
@@ -264,10 +237,12 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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)
// Stamp the header into place immediately before this segment's
// payload, sourced from the pristine snapshot. Iter 0's header is
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1
// needs the stamp. The per-segment patches below overwrite the
// variable fields.
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]
@@ -296,9 +271,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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].
// Payload bytes still live at their original offset in pkt. The
// header slide above only writes into pkt[i*G : i*G+H], which is
// the tail of seg_{i-1}'s payload (already consumed) and never
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
@@ -314,10 +290,17 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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.
// 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*G:i*G+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; see SegmentTCP for the layout reasoning, including
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
// (correctness when gsoSize < hdrLen).
//
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
// bump it), which is why the IP-level per-segment work is limited to
// total_len + IPv4 header checksum (v4) or payload_len (v6).
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero")
@@ -338,24 +321,41 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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
}
numSeg := (payLen + gsoSize - 1) / gsoSize
if numSeg == 0 {
numSeg = 1
}
// Snapshot the pristine L3+L4 header once and stamp every segment from it
var udpTmp [udpHeaderLen]byte
copy(udpTmp[:], pkt[csumStart:headerLen])
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
var baseProtoSum uint32
if isV4 {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
} else {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
}
baseProtoSum += uint32(unix.IPPROTO_UDP)
var baseIPHdrSum uint32
if isV4 {
ihl := int(pkt[0]&0x0f) * 4
if ihl < ipv4HeaderMinLen || ihl > csumStart {
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
}
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
}
// Snapshot the pristine L3+L4 header once and stamp every segment from
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
// segments when gsoSize < headerLen.
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
@@ -378,10 +378,8 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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)
ipSum := baseIPHdrSum + uint32(totalLen)
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
} else {
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
@@ -389,13 +387,12 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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))
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum)
wide += uint64(udpLen) + uint64(udpLen)
wide = (wide & 0xffffffff) + (wide >> 32)
wide = (wide & 0xffffffff) + (wide >> 32)
csum := foldComplement(uint32(wide))
if csum == 0 {
csum = 0xffff
}
@@ -409,9 +406,10 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
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.
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
// the pseudo-header partial sum by the kernel), and store the result.
func FinishChecksum(seg []byte, hdr Hdr) error {
cs := int(hdr.CsumStart)
co := int(hdr.CsumOffset)
@@ -423,12 +421,7 @@ func FinishChecksum(seg []byte, hdr Hdr) error {
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)
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
return nil
}
+17 -284
View File
@@ -226,14 +226,13 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
// = 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*/
)
hdr := Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
GSOSize: 6, // two 6-byte segments
CsumStart: csumStart,
CsumOffset: 6,
}
if err := CorrectHdrLen(pkt, &hdr); err != nil {
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
}
@@ -248,14 +247,13 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
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*/
)
hdr := Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
GSOSize: 6,
CsumStart: 20,
CsumOffset: 6,
}
if err := CorrectHdrLen(pkt, &hdr); err == nil {
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
}
@@ -305,10 +303,9 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
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))
// UDP-GSO keeps the same IPv4 ID across every segment.
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 {
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id)
}
segPayLen := len(seg) - int(hdrLen)
@@ -336,267 +333,3 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
})
}
}
// 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)
}
}
}
}
}
}
})
}
+1 -3
View File
@@ -550,9 +550,7 @@ func (t *tun) Read(to []byte) (int, error) {
return n - 4, nil
}
// 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).
// Write pushes one IP packet onto the utun device. Only valid for single threaded use.
func (t *tun) Write(from []byte) (int, error) {
if len(from) == 0 {
return 0, syscall.EIO
+92 -43
View File
@@ -25,17 +25,32 @@ import (
)
type tun struct {
readers tio.QueueSet
closeLock sync.Mutex
Device string
vpnNetworks []netip.Prefix
MaxMTU int
DefaultMTU int
TXQueueLen int
deviceIndex int
ioctlFd uintptr
vnetHdr bool
readers tio.QueueSet
closeLock sync.Mutex
Device string
vpnNetworks []netip.Prefix
MaxMTU int
DefaultMTU int
TXQueueLen int
deviceIndex int
ioctlFd uintptr
vnetHdr bool
// offloadFlags is the exact TUN_F_* offload mask newTun negotiated with
// the kernel: usoOffloadFlags when USO was accepted, tsoOffloadFlags on
// the TSO-only fallback, or 0 when vnetHdr is off. TUNSETOFFLOAD is
// device-wide (drivers/net/tun.c set_offload updates tun->set_features
// for the whole netdev), so addQueue must replay this exact
// mask on every added queue — issuing a narrower mask there would
// silently downgrade offloads (e.g. disable USO) for all queues while
// they still advertise the stale capability.
offloadFlags uint
// routeFeatureECN, when true, sets RTAX_FEATURE_ECN on every route we
// install for the tun. The kernel then actively negotiates ECN for
// connections destined to those prefixes (equivalent to `ip route
// change ... features ecn`) regardless of net.ipv4.tcp_ecn, so flows
// across the nebula mesh use ECN even when the host default is the
// passive setting (=2). Disable via tunnels.ecn=false.
routeFeatureECN bool
Routes atomic.Pointer[[]Route]
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
@@ -76,7 +91,14 @@ type ifreqQLEN struct {
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
// 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")
t, err := newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks)
if err != nil {
return nil, err
}
t.Device = "tun0"
return t, nil
}
// openTunDev opens /dev/net/tun, creating the device node first if it's
@@ -102,7 +124,8 @@ func openTunDev() (int, error) {
return fd, nil
}
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
// 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
@@ -113,45 +136,57 @@ func tunSetIff(fd int, name string, flags uint16) (string, error) {
return strings.Trim(string(req.Name[:]), "\x00"), nil
}
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
// TSO_ECN tells the kernel we propagate ECN correctly through coalesce and
// segmentation, so it can deliver superpackets whose seed has CWR/ECE set
// or whose IP-level codepoint is CE.
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
// usoOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. Requires
// Linux 6.2; older kernels reject it and we fall back to TCP-only TSO via
// tsoOffloadFlags.
const usoOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
// offloadUSOEnabled reports whether the negotiated offload mask includes UDP
// Segmentation Offload. It is the single source of truth for the usoEnabled
// capability surfaced by each queue, so the mask stored on the tun and the USO
// bit reported to coalescers can never drift apart.
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) {
// 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)
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
if multiqueue {
baseFlags |= unix.IFF_MULTI_QUEUE
}
nameStr := c.GetString("tun.dev", "")
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_*
// offloads via TUNSETOFFLOAD so we can receive TSO/USO superpackets.
// 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.
fd, err := openTunDev()
if err != nil {
return nil, err
}
vnetHdr := true
// 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
// offloadFlags is the exact TUN_F_* mask the kernel accepted. We remember
// it (rather than a plain bool) so addQueue can replay the
// identical device-wide mask on added queues instead of downgrading them.
var offloadFlags uint
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
if err != nil {
_ = unix.Close(fd)
vnetHdr = false
} else {
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
offloadFlags = usoAndTSOOffloadFlags
// Try TSO+USO first. On kernels without USO support (Linux < 6.2)
// the ioctl returns EINVAL; fall back to the TCP-only mask before
// giving up on VNET_HDR entirely.
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoOffloadFlags)); err == nil {
offloadFlags = usoOffloadFlags
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
offloadFlags = tsoOffloadFlags
} else {
@@ -177,7 +212,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
}
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks, name)
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks)
if err != nil {
return nil, err
}
@@ -187,14 +222,16 @@ 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.
// 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) {
// 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 (0 when vnetHdr is off); the queues' USO capability is
// derived from it so it can never disagree with the mask we replay on added
// multiqueue readers.
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix) (*tun, error) {
var qs tio.QueueSet
var err error
if vnetHdr {
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags))
} else {
qs, err = tio.NewPollQueueSet()
}
@@ -205,15 +242,11 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
}
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{
Device: name,
readers: qs,
closeLock: sync.Mutex{},
vnetHdr: vnetHdr,
@@ -222,6 +255,7 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
TXQueueLen: c.GetInt("tun.tx_queue", 500),
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
routeFeatureECN: c.GetBool("tunnels.ecn", true),
routesFromSystem: map[netip.Prefix]routing.Gateways{},
l: l,
}
@@ -315,7 +349,9 @@ func (t *tun) reload(c *config.C, initial bool) error {
return nil
}
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
// Queues opens additional kernel multiqueue fds until the device has n
// queues, then returns them all. The first queue was opened by newTun; each
// extra fd replays the negotiated offload state (see addQueue).
func (t *tun) Queues(n int) ([]tio.Queue, error) {
for len(t.readers.Queues()) < n {
if err := t.addQueue(); err != nil {
@@ -325,7 +361,8 @@ func (t *tun) Queues(n int) ([]tio.Queue, error) {
return t.readers.Queues(), nil
}
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
// 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()
@@ -345,6 +382,10 @@ func (t *tun) addQueue() error {
}
if t.vnetHdr {
// Replay the exact mask newTun negotiated. TUNSETOFFLOAD is
// device-wide, so issuing the TSO-only mask here would disable USO
// for every queue (including queue 0) on kernels where newTun
// successfully enabled it, while the queues keep advertising USO.
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)
@@ -525,13 +566,18 @@ 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.
// 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
}
if t.routeFeatureECN {
nr.Features |= unix.RTAX_FEATURE_ECN
}
err := netlink.RouteReplace(&nr)
if err != nil {
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
@@ -581,6 +627,9 @@ func (t *tun) addRoutes(logErrors bool) error {
if r.Metric > 0 {
nr.Priority = r.Metric
}
if t.routeFeatureECN {
nr.Features |= unix.RTAX_FEATURE_ECN
}
err := netlink.RouteReplace(&nr)
if err != nil {
+9 -9
View File
@@ -39,14 +39,14 @@ func TestTunAdvMSS(t *testing.T) {
// 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
// usoOffloadFlags 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 usoOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
t.Fatalf("usoOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoOffloadFlags, tsoOffloadFlags)
}
if usoAndTSOOffloadFlags == tsoOffloadFlags {
t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
if usoOffloadFlags == tsoOffloadFlags {
t.Fatal("usoOffloadFlags must add bits beyond tsoOffloadFlags")
}
cases := []struct {
@@ -54,7 +54,7 @@ func TestOffloadUSOEnabled(t *testing.T) {
offloadFlags uint
wantUSO bool
}{
{"uso-negotiated", usoAndTSOOffloadFlags, true},
{"uso-negotiated", usoOffloadFlags, true},
{"tso-fallback", tsoOffloadFlags, false},
{"no-vnet-hdr", 0, false},
}
@@ -78,12 +78,12 @@ func TestOffloadUSOEnabled(t *testing.T) {
// TUNSETOFFLOAD argument is read from.
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
t.Run("uso-negotiated", func(t *testing.T) {
tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
tn := &tun{vnetHdr: true, offloadFlags: usoOffloadFlags}
// 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 != usoOffloadFlags {
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoOffloadFlags)
}
if tn.offloadFlags == tsoOffloadFlags {
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
+4 -6
View File
@@ -10,13 +10,11 @@ import (
_ "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.
// startPprofServer serves net/http/pprof on :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.
func startPprofServer(ctx context.Context, l *slog.Logger) {
server := &http.Server{Addr: "localhost:6060", Handler: nil}
server := &http.Server{Addr: ":6060", Handler: nil}
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
go func() {
+1 -1
View File
@@ -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, 0)
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false)
case Disestablished:
// Mark this relay as 'requested'
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
+28 -13
View File
@@ -14,28 +14,43 @@ const MTU = 9001
// only costs additional sendmmsg chunks within a single WriteBatch call.
const MaxWriteBatch = 128
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
// data, kernel offload state, etc.) and passed to EncReader callbacks.
// Backends that do not produce a particular signal leave its zero value.
//
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
// means Not-ECT, which is also the value backends without ECN RX support
// supply on every packet.
type RxMeta struct {
OuterECN byte
}
type EncReader func(
addr netip.AddrPort,
payload []byte,
meta RxMeta,
)
type Conn interface {
Rebind() error
LocalAddr() (netip.AddrPort, 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 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)
// destination. bufs and addrs must have the same length. outerECNs may
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
// cmsg; other backends ignore it. Returns on the first error; callers
// may observe a partial send if some packets went out before the error.
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
ReloadConfig(c *config.C)
SupportsMultipleReaders() bool
Close() error
@@ -58,8 +73,8 @@ 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) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
return nil
}
func (NoopConn) ReloadConfig(_ *config.C) {
return
-61
View File
@@ -1,61 +0,0 @@
package udp
import (
"context"
"log/slog"
"github.com/slackhq/nebula/config"
)
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
//
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
// keeps this package from referencing whatever owns the rebind.
//
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
type NetworkChangeMonitor struct {
l *slog.Logger
ctx context.Context
enabled bool
}
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
return &NetworkChangeMonitor{
l: l,
ctx: ctx,
enabled: c.GetBool("listen.rebind_on_network_change", true),
}
}
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
func (m *NetworkChangeMonitor) Start(rebind func()) {
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
return
}
changes, err := watchNetworkChanges(m.ctx, m.l)
if err != nil {
// Not fatal. Everything else still works, we just won't notice a network change on our own.
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
"error", err,
)
return
}
if changes == nil {
// This platform's sockets don't go stale, so there is nothing to watch for.
return
}
m.l.Info("Watching for network changes to rebind the udp listener")
for range changes {
m.l.Info("Local network changed, rebinding the udp listener")
rebind()
}
}
-164
View File
@@ -1,164 +0,0 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"errors"
"log/slog"
"os"
"time"
"golang.org/x/sys/unix"
)
const (
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
netChangeSettleWindow = time.Second
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
// the kernel, so being generous here is how we avoid missing a message.
netChangeReadBuffer = 4096
)
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
//
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
//
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
sock, err := openRouteSocket()
if err != nil {
return nil, err
}
changes := make(chan struct{}, 1)
go func() {
defer close(changes)
defer func() { _ = sock.Close() }()
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
_ = sock.Close()
case <-done:
}
}()
watchRouteSocket(l, sock, changes)
}()
return changes, nil
}
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
// the socket is closed, which is how cancellation gets us out of here.
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
buf := make([]byte, netChangeReadBuffer)
for {
n, err := sock.Read(buf)
if err != nil {
logRouteSocketError(l, err)
return
}
if !isNetworkChange(buf[:n]) {
continue
}
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
// simply produce another report, which is the correct outcome anyway.
deadline := time.Now().Add(netChangeSettleWindow)
for {
if err = sock.SetReadDeadline(deadline); err != nil {
logRouteSocketError(l, err)
return
}
if _, err = sock.Read(buf); err != nil {
if os.IsTimeout(err) {
break
}
logRouteSocketError(l, err)
return
}
}
if err = sock.SetReadDeadline(time.Time{}); err != nil {
logRouteSocketError(l, err)
return
}
select {
case changes <- struct{}{}:
default:
// One already pending, and a second "the network moved" tells the reader nothing new.
}
}
}
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
func logRouteSocketError(l *slog.Logger, err error) {
if errors.Is(err, os.ErrClosed) {
return
}
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
}
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
func openRouteSocket() (*os.File, error) {
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
if err != nil {
return nil, err
}
if err = unix.SetNonblock(fd, true); err != nil {
_ = unix.Close(fd)
return nil, err
}
return os.NewFile(uintptr(fd), "route"), nil
}
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
//
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
func isNetworkChange(msg []byte) bool {
if len(msg) < 4 {
return false
}
// u_short msglen, u_char version, u_char type
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
return false
}
switch msg[3] {
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
return true
default:
return false
}
}
-244
View File
@@ -1,244 +0,0 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"os"
"testing"
"time"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"golang.org/x/sys/unix"
)
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
func routeMsg(msgType uint8, extra int) []byte {
msg := make([]byte, 4+extra)
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
msg[2] = unix.RTM_VERSION
msg[3] = msgType
return msg
}
func TestIsNetworkChange(t *testing.T) {
// The three that mean our addressing may have moved
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
// Route churn is not something a rebind helps with
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
// Garbage must not be mistaken for a change
assert.False(t, isNetworkChange(nil), "empty")
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
wrongVersion[2] = unix.RTM_VERSION + 1
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
lying := routeMsg(unix.RTM_NEWADDR, 0)
binary.NativeEndian.PutUint16(lying[0:2], 512)
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
}
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
// stands in for the kernel so the watch loop can be driven with synthetic messages.
func socketPair(t *testing.T) (*os.File, int) {
t.Helper()
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
require.NoError(t, err)
require.NoError(t, unix.SetNonblock(fds[0], true))
f := os.NewFile(uintptr(fds[0]), "route")
t.Cleanup(func() {
_ = f.Close()
_ = unix.Close(fds[1])
})
return f, fds[1]
}
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
// exactly one report rather than one apiece.
for range 5 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
}
// Uninteresting messages in the middle of a burst must not add a report of their own either.
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a burst should have reported a change")
}
// Nothing more from that burst
select {
case <-changes:
t.Fatal("a burst should report exactly once")
case <-time.After(netChangeSettleWindow):
}
// A change after the window has closed is a separate event and gets its own report.
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a later change should report again")
}
// Closing the socket is how the real thing shuts down
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
_, err := unix.Write(kernel, routeMsg(msgType, 8))
require.NoError(t, err)
}
select {
case <-changes:
t.Fatal("route churn alone must not report a change")
case <-time.After(netChangeSettleWindow * 2):
}
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
for range 3 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
}
// The watcher must still be alive and responsive to a close
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket wedged on a full channel")
}
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
}
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
// cancelling the context closes the channel so a ranging caller falls out of its loop.
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
changes, err := watchNetworkChanges(ctx, test.NewLogger())
require.NoError(t, err)
require.NotNil(t, changes, "darwin should support watching")
drained := make(chan struct{})
go func() {
for range changes {
}
close(drained)
}()
cancel()
select {
case <-drained:
case <-time.After(time.Second * 5):
t.Fatal("cancelling the context should close the changes channel")
}
}
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
// must return it and clean up the watch goroutines.
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
// leaks itself.
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
ctx, cancel := context.WithCancel(context.Background())
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
m := NewNetworkChangeMonitor(ctx, l, c)
done := make(chan struct{})
go func() {
m.Start(func() {})
close(done)
}()
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
// watch or no-op'd, both of which we want to catch.
select {
case <-done:
t.Fatal("Start returned instead of watching")
case <-time.After(time.Millisecond * 250):
}
cancel()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start did not return after the context was cancelled")
}
// Starting again after the context is dead must not open anything.
m.Start(func() {})
}
-22
View File
@@ -1,22 +0,0 @@
//go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp
import (
"context"
"log/slog"
)
// watchNetworkChanges is a no-op outside of darwin.
//
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
// and two things racing to rebind the same socket is worse than one.
//
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
// selecting on it, since a receive from a nil channel blocks forever.
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
return nil, nil
}
-39
View File
@@ -1,39 +0,0 @@
package udp
import (
"context"
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
t.Helper()
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString(cfg))
return NewNetworkChangeMonitor(ctx, l, c)
}
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
// Says nothing about rebinding, so this covers the default.
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
assert.True(t, m.enabled, "should default to on")
}
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
require.False(t, m.enabled)
// Must return without opening a socket. If it watched anything this would block.
m.Start(func() {})
}
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
// Nothing to rebind, so there is no point watching, on any platform.
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
m.Start(nil)
}
+62
View File
@@ -0,0 +1,62 @@
//go:build !android && !e2e_testing
// +build !android,!e2e_testing
package udp
import (
"net"
"syscall"
"unsafe"
"golang.org/x/sys/unix"
)
// rawSendmmsg performs sendmmsg(2) over a syscall.RawConn without
// allocating a closure per call. The struct holds preallocated in/out
// scratch (chunk/sent/errno) and a method-value bound at construction so
// rawConn.Write receives a stable function pointer instead of a fresh
// closure on every send.
type rawSendmmsg struct {
msgs []rawMessage
chunk int
sent int
errno syscall.Errno
callback func(fd uintptr) bool
}
// bind wires r.callback to r.run. Must be called once after r.msgs is set;
// subsequent send calls invoke r.callback without rebinding.
func (r *rawSendmmsg) bind() { r.callback = r.run }
// run is the preallocated callback rawConn.Write invokes. It reads its
// input (r.chunk) and writes its outputs (r.sent, r.errno) through the
// rawSendmmsg fields so the method value does not capture per-call locals
// and therefore does not heap-allocate.
func (r *rawSendmmsg) run(fd uintptr) bool {
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
uintptr(unsafe.Pointer(&r.msgs[0])), uintptr(r.chunk),
0, 0, 0,
)
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
return false
}
r.sent = int(r1)
r.errno = errno
return true
}
// send issues sendmmsg over rc against the first n entries of r.msgs.
// Returns the number of entries the kernel processed and any error;
// matches the original sendmmsg helper's contract.
func (r *rawSendmmsg) send(rc syscall.RawConn, n int) (int, error) {
r.chunk = n
r.sent = 0
r.errno = 0
if err := rc.Write(r.callback); err != nil {
return r.sent, err
}
if r.errno != 0 {
return r.sent, &net.OpError{Op: "sendmmsg", Err: r.errno}
}
return r.sent, nil
}
+10 -16
View File
@@ -140,20 +140,13 @@ 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
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
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)
if err := u.WriteTo(b, addrs[i]); err != nil {
return err
}
}
return written, nil
return nil
}
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
@@ -195,7 +188,7 @@ func (u *StdConn) ListenOut(r EncReader, flush func()) error {
continue
}
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
flush()
}
}
@@ -204,9 +197,6 @@ func (u *StdConn) SupportsMultipleReaders() bool {
return false
}
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
// this way on its own, which is what strands us after the underlying network changes.
func (u *StdConn) Rebind() error {
var err error
if u.isV4 {
@@ -215,5 +205,9 @@ func (u *StdConn) Rebind() error {
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
}
return err
if err != nil {
u.l.Error("Failed to rebind udp socket", "error", err)
}
return nil
}
+61
View File
@@ -0,0 +1,61 @@
//go:build linux && !android && !e2e_testing
package udp
import (
"net/netip"
"testing"
)
// TestPlanRunBreaksOnECNChange confirms that two same-destination, same-size
// packets with different outer ECN end up in separate sendmmsg entries (the
// kernel stamps one outer codepoint per entry, so a run that straddled the
// boundary would silently lose information).
func TestPlanRunBreaksOnECNChange(t *testing.T) {
u := &StdConn{gsoSupported: true, maxGSOSegments: 63}
dst := netip.MustParseAddrPort("10.0.0.1:4242")
bufs := [][]byte{
make([]byte, 1200),
make([]byte, 1200),
make([]byte, 1200),
}
addrs := []netip.AddrPort{dst, dst, dst}
t.Run("uniform_ecn_runs_together", func(t *testing.T) {
ecns := []byte{0x02, 0x02, 0x02}
runLen, segSize := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 3 {
t.Errorf("runLen=%d want 3 (uniform ECT(0))", runLen)
}
if segSize != 1200 {
t.Errorf("segSize=%d want 1200", segSize)
}
})
t.Run("ecn_change_truncates_run", func(t *testing.T) {
// 0,0,3: first two run together, CE seeds a fresh entry.
ecns := []byte{0x00, 0x00, 0x03}
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 2 {
t.Errorf("runLen=%d want 2 (ECN changes at index 2)", runLen)
}
})
t.Run("nil_ecns_runs_full", func(t *testing.T) {
runLen, _ := u.planRun(bufs, addrs, nil, 0, 64)
if runLen != 3 {
t.Errorf("runLen=%d want 3 (nil ecns means no break)", runLen)
}
})
t.Run("first_ecn_is_singleton", func(t *testing.T) {
// Second packet has different ECN from the first → run halts at 1
// (the first packet alone forms the run).
ecns := []byte{0x00, 0x03, 0x03}
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 1 {
t.Errorf("runLen=%d want 1 (different ECN immediately)", runLen)
}
})
}
+5 -9
View File
@@ -44,17 +44,13 @@ 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
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
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)
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
return err
}
}
return written, nil
return nil
}
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
@@ -106,7 +102,7 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
continue
}
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
flush()
}
}
+716 -196
View File
File diff suppressed because it is too large Load Diff
+33
View File
@@ -30,6 +30,39 @@ type rawMessage struct {
Len uint32
}
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
msgs := make([]rawMessage, n)
buffers := make([][]byte, n)
names := make([][]byte, 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)
vs := []iovec{
{Base: &buffers[i][0], Len: uint32(len(buffers[i]))},
}
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]))
if cmsgSpace > 0 {
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
}
}
return msgs, buffers, names, cmsgs
}
func setIovLen(v *iovec, n int) {
v.Len = uint32(n)
}
+33
View File
@@ -33,6 +33,39 @@ type rawMessage struct {
Pad0 [4]byte
}
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
msgs := make([]rawMessage, n)
buffers := make([][]byte, n)
names := make([][]byte, 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)
vs := []iovec{
{Base: &buffers[i][0], Len: uint64(len(buffers[i]))},
}
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]))
if cmsgSpace > 0 {
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
}
}
return msgs, buffers, names, cmsgs
}
func setIovLen(v *iovec, n int) {
v.Len = uint64(n)
}
+100 -555
View File
@@ -3,11 +3,11 @@
package udp
import (
"fmt"
"encoding/binary"
"log/slog"
"net"
"net/netip"
"slices"
"syscall"
"testing"
"time"
"unsafe"
@@ -54,6 +54,39 @@ func buildCmsg(level, typ int32, data []byte) []byte {
return buf
}
// TestParseRecvCmsgOuterECNFamily is the RX half of the dual-stack ECN fix:
// parseRecvCmsg must read the outer ECN from whichever family the kernel
// delivered, not from the socket family. On the default `::` dual-stack bind
// a v4 peer's outer ECN arrives as an IP_TOS cmsg, which the old socket-family
// gate ignored entirely.
func TestParseRecvCmsgOuterECNFamily(t *testing.T) {
tc := make([]byte, 4)
binary.NativeEndian.PutUint32(tc, 0x02)
cases := []struct {
name string
ctrl []byte
want byte
}{
{"ip_tos_ce", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x03}), 0x03},
{"ip_tos_ect0", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x02}), 0x02},
{"ipv6_tclass_ect0", buildCmsg(int32(unix.IPPROTO_IPV6), int32(unix.IPV6_TCLASS), tc), 0x02},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
hdr := &msghdr{Control: &c.ctrl[0]}
setMsgControllen(hdr, len(c.ctrl))
gso, ecn := parseRecvCmsg(hdr, false, true)
if gso != 0 {
t.Errorf("gso = %d, want 0 (no UDP_GRO cmsg present)", gso)
}
if ecn != c.want {
t.Errorf("ecn = 0x%02x, want 0x%02x", ecn, c.want)
}
})
}
}
func testLogger() *slog.Logger {
return slog.New(slog.DiscardHandler)
}
@@ -90,13 +123,9 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
addrs := []netip.AddrPort{good, bad, good}
n, err := sender.WriteBatch(bufs, addrs)
if err != nil {
if err := sender.WriteBatch(bufs, addrs, nil); 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))
@@ -116,11 +145,12 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
}
}
// 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) {
// TestWriteBatchOuterTOSToV4Mapped is the TX half of the dual-stack ECN fix,
// verified against a live kernel: WriteBatch on the default `::` dual-stack
// socket, sending to a v4-mapped destination, must stamp the outer ECN via an
// IP_TOS cmsg (not IPV6_TCLASS, which the kernel's v4 path ignores) so a v4
// receiver actually sees it.
func TestWriteBatchOuterTOSToV4Mapped(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)
@@ -128,561 +158,76 @@ func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
defer rx.Close()
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 1)
// Ask the kernel to deliver the received outer TOS as ancillary data.
rxRaw, err := rx.SyscallConn()
if err != nil {
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
t.Fatalf("SyscallConn: %v", err)
}
var soErr error
if err := rxRaw.Control(func(fd uintptr) {
soErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
}); err != nil || soErr != nil {
t.Skipf("cannot enable IP_RECVTOS (sandbox/kernel?): ctrl=%v so=%v", err, soErr)
}
c, err := NewListener(testLogger(), netip.IPv6Unspecified(), 0, false, 1)
if err != nil {
t.Skipf("cannot open dual-stack 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")
if sender.isV4 {
t.Skipf("sender came up v4-only; need a dual-stack v6 socket for this test")
}
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)
}
}
}
// v4-mapped-in-v6 destination: routed through the kernel's IPv4 path.
dst := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
const wantECN = byte(0x02) // ECT(0)
// 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()}
w.prepareWriteMessages(MaxWriteBatch)
w.gsoSupported = true
w.maxGSOSegments = 63
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 {
if err := sender.WriteBatch([][]byte{[]byte("tos-probe")}, []netip.AddrPort{dst}, []byte{wantECN}); 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)
// Read the datagram plus its ancillary TOS.
rx.SetReadDeadline(time.Now().Add(3 * time.Second))
payload := make([]byte, 128)
oob := make([]byte, 512)
var n, oobn int
var rerr error
if err := rxRaw.Read(func(fd uintptr) bool {
n, oobn, _, _, rerr = unix.Recvmsg(int(fd), payload, oob, 0)
if rerr == syscall.EAGAIN || rerr == syscall.EWOULDBLOCK {
return false
}
return true
}); err != nil {
t.Fatalf("waiting for datagram failed (no delivery?): %v", err)
}
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))
if rerr != nil {
t.Fatalf("Recvmsg: %v", rerr)
}
for i, b := range wire {
if b[0] != wantTags[i] {
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
if string(payload[:n]) != "tos-probe" {
t.Fatalf("payload = %q, want %q", string(payload[:n]), "tos-probe")
}
cmsgs, err := unix.ParseSocketControlMessage(oob[:oobn])
if err != nil {
t.Fatalf("ParseSocketControlMessage: %v", err)
}
found := false
var gotTOS byte
for _, m := range cmsgs {
if m.Header.Level == unix.IPPROTO_IP && m.Header.Type == unix.IP_TOS && len(m.Data) >= 1 {
found = true
gotTOS = m.Data[0]
}
}
// 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()
uc, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 8)
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))
}
}
if !found {
t.Fatalf("no IP_TOS cmsg delivered to v4 receiver — outer ECN did not land (%d cmsgs)", len(cmsgs))
}
if gotTOS&0x03 != wantECN {
t.Errorf("received outer TOS = 0x%02x, want low-2-bits = 0x%02x", gotTOS, wantECN)
} else {
t.Logf("verified: v4 receiver saw outer TOS 0x%02x (ECN=0x%02x) from dual-stack sender", gotTOS, gotTOS&0x03)
}
}

Some files were not shown because too many files have changed in this diff Show More