mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 20:17:01 +02:00
Compare commits
125 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 43fb7bff60 | |||
| ff672a3a1f | |||
| 03269ffaa1 | |||
| e196a7a7ca | |||
| d14b97977f | |||
| e1d96b932a | |||
| d8d5ce344d | |||
| a3eef407b2 | |||
| 3b1004588d | |||
| 4cd433309b | |||
| 5ea48c1677 | |||
| cdfba18ea5 | |||
| e04473f30b | |||
| 42c937d86f | |||
| adfacc43c3 | |||
| 096a06238a | |||
| bc70bf47f8 | |||
| d7bcfb5d6b | |||
| ddb90ad4b7 | |||
| 549de9fd29 | |||
| 93946faf7a | |||
| d3779b6a39 | |||
| 6783c90e72 | |||
| 245eb61444 | |||
| 9d8e830e4e | |||
| da8a640a56 | |||
| b7bf32240b | |||
| b3002c2d13 | |||
| 575b97904d | |||
| 16878eec1c | |||
| e8d6be1dd9 | |||
| d13c53db43 | |||
| 7ee5e29758 | |||
| 9c68c60ba6 | |||
| 2190900107 | |||
| 69cf816f80 | |||
| 8044b40c82 | |||
| fbcdd83359 | |||
| e9870686c8 | |||
| 9ce0596f5c | |||
| 7455ac56eb | |||
| 0b0e45582a | |||
| d43bb81ae2 | |||
| 8c06802f5f | |||
| 5a8014585d | |||
| ffab005f9d | |||
| 38f11f6e3a | |||
| 095a421708 | |||
| b39bae57ec | |||
| 5c0f6e2b5f | |||
| df8955177e | |||
| 5c2a5607e5 | |||
| ed88422770 | |||
| 0b817c50b3 | |||
| db54e05bfe | |||
| 46a02b663e | |||
| b7d06615c0 | |||
| 3803146bc9 | |||
| f9eb86df9d | |||
| 66cb98e13a | |||
| 3fe2cb970e | |||
| 865dc9725c | |||
| 8cebecc087 | |||
| fb0d20f123 | |||
| 921ed4360a | |||
| 6fb10c1c1a | |||
| 8006b58758 | |||
| 56d2a3841d | |||
| 6c0305c7ae | |||
| 5c0a5ee4be | |||
| 35596c7708 | |||
| d74c5ac5c5 | |||
| c06bfb46be | |||
| be35d4f059 | |||
| eea3c81626 | |||
| 88872a8433 | |||
| 6bf424f749 | |||
| 9688d32f5b | |||
| 8c91fa2699 | |||
| 92d51c042e | |||
| 0a0b2404a2 | |||
| ef0e3015f9 | |||
| c6ebe71c08 | |||
| 1d768ac4e4 | |||
| 69fa9e4a2e | |||
| ed55cf40d5 | |||
| a6ae44ddb1 | |||
| 7cc37323a8 | |||
| 68e3fae870 | |||
| 720990ddcd | |||
| 05443523bd | |||
| 99bf613a2c | |||
| 9e7c783eb3 | |||
| 8b14a6ee56 | |||
| ae17513bbf | |||
| 1aca2f75ae | |||
| c7918d1096 | |||
| 44dd2e9ca4 | |||
| 733dc06192 | |||
| 7fee3a97b2 | |||
| 1e218737dc | |||
| 243c920f88 | |||
| 9bdab873f2 | |||
| a081fba023 | |||
| b50d6276e3 | |||
| 3b16f1adb6 | |||
| 410bac9688 | |||
| 22adae0b8c | |||
| f7cc437d88 | |||
| a66af843d1 | |||
| 67a742ddfb | |||
| 5681e510c4 | |||
| 9104dc4c34 | |||
| f31f6c5d1f | |||
| 264a25337b | |||
| 694414771c | |||
| ca84bcb38f | |||
| f04dd3bcc3 | |||
| 187afac7b5 | |||
| e7b121c82f | |||
| 72bf111209 | |||
| 1617897043 | |||
| f8775bb6ca | |||
| 15f0f0d5d0 | |||
| 7902ce674e |
@@ -14,7 +14,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -40,7 +40,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -80,7 +80,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Import certificates
|
||||
|
||||
@@ -34,7 +34,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -66,7 +66,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -92,7 +92,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||
|
||||
@@ -22,7 +22,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: build
|
||||
|
||||
@@ -22,7 +22,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Install goimports
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
- name: golangci-lint
|
||||
uses: golangci/golangci-lint-action@v9
|
||||
with:
|
||||
version: v2.5
|
||||
version: v2.12
|
||||
|
||||
test:
|
||||
name: Test ${{ matrix.name }}
|
||||
@@ -82,7 +82,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -127,7 +127,7 @@ jobs:
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build ${{ matrix.name }}
|
||||
|
||||
@@ -7,6 +7,88 @@ 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
|
||||
|
||||
@@ -161,6 +161,10 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
||||
bin-pkcs11: CGO_ENABLED = 1
|
||||
bin-pkcs11: bin
|
||||
|
||||
# Build with the pprof debug server (serves on :6060). See startPprofServer.
|
||||
debug: BUILD_ARGS += -tags debug
|
||||
debug: bin
|
||||
|
||||
bin:
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||
@@ -280,5 +284,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
|
||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||
|
||||
.FORCE:
|
||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||
.DEFAULT_GOAL := bin
|
||||
|
||||
@@ -44,11 +44,6 @@ type connectionManager struct {
|
||||
inactivityTimeout atomic.Int64
|
||||
dropInactive atomic.Bool
|
||||
|
||||
// Wake-from-sleep handling, sampled once per tick in Start
|
||||
wakeDetector *wakeDetector
|
||||
clearOnWake atomic.Bool
|
||||
wakeClearThreshold atomic.Int64
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
@@ -59,7 +54,6 @@ func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p
|
||||
punchy: p,
|
||||
relayUsed: make(map[uint32]struct{}),
|
||||
relayUsedLock: &sync.RWMutex{},
|
||||
wakeDetector: newWakeDetector(),
|
||||
}
|
||||
|
||||
cm.reload(c, true)
|
||||
@@ -104,38 +98,12 @@ func (cm *connectionManager) reload(c *config.C, initial bool) {
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if initial || c.HasChanged("tunnels.clear_on_wake") {
|
||||
old := cm.clearOnWake.Load()
|
||||
cm.clearOnWake.Store(c.GetBool("tunnels.clear_on_wake", true))
|
||||
if !initial {
|
||||
cm.l.Info("Clear on wake setting has changed",
|
||||
"oldBool", old,
|
||||
"newBool", cm.clearOnWake.Load(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if initial || c.HasChanged("tunnels.wake_clear_threshold") {
|
||||
old := cm.getWakeClearThreshold()
|
||||
cm.wakeClearThreshold.Store((int64)(c.GetDuration("tunnels.wake_clear_threshold", 30*time.Second)))
|
||||
if !initial {
|
||||
cm.l.Info("Wake clear threshold has changed",
|
||||
"oldDuration", old,
|
||||
"newDuration", cm.getWakeClearThreshold(),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (cm *connectionManager) getInactivityTimeout() time.Duration {
|
||||
return (time.Duration)(cm.inactivityTimeout.Load())
|
||||
}
|
||||
|
||||
func (cm *connectionManager) getWakeClearThreshold() time.Duration {
|
||||
return (time.Duration)(cm.wakeClearThreshold.Load())
|
||||
}
|
||||
|
||||
func (cm *connectionManager) In(h *HostInfo) {
|
||||
h.in.Store(true)
|
||||
}
|
||||
@@ -168,73 +136,6 @@ func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time)
|
||||
return in, out
|
||||
}
|
||||
|
||||
// checkWake runs once per tick and clears every tunnel when the machine has just returned from system sleep.
|
||||
// Tunnels rarely survive a suspend: our NAT mappings have expired and our address has usually changed, so every
|
||||
// established hostinfo is a corpse that will eat 15-20s of traffic checks before the wheel declares it dead.
|
||||
// Clearing now means the first packet after wake starts a fresh handshake immediately.
|
||||
//
|
||||
// The suspend itself costs nothing here: the ticker driving us is frozen with the rest of the process and this
|
||||
// fires within one tick of resume.
|
||||
func (cm *connectionManager) checkWake() {
|
||||
slept, ok := cm.wakeDetector.Sample()
|
||||
if !ok || slept == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// The clock pair is read non-atomically, so scheduling jitter shows up as tiny sub-millisecond "sleeps".
|
||||
// Keep the floor well above that so a zero/nonsense threshold can't clear tunnels on every tick.
|
||||
threshold := max(cm.getWakeClearThreshold(), time.Second)
|
||||
|
||||
if slept < threshold {
|
||||
// Short suspends (lid closed and quickly reopened) often come back before NAT state expires; those
|
||||
// tunnels may well be alive, leave them to the normal traffic checks.
|
||||
if slept >= time.Second {
|
||||
cm.l.Debug("Woke from sleep below the clear threshold, leaving tunnels alone",
|
||||
"sleptFor", slept,
|
||||
"threshold", threshold,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if !cm.clearOnWake.Load() {
|
||||
cm.l.Info("Woke from sleep, tunnels.clear_on_wake is disabled so tunnels are left to the normal traffic checks", "sleptFor", slept)
|
||||
return
|
||||
}
|
||||
|
||||
closed := cm.clearAllTunnels()
|
||||
cm.l.Info("Woke from sleep, cleared tunnels", "sleptFor", slept, "tunnelsCleared", closed)
|
||||
|
||||
// Our public address almost certainly changed; get it to the lighthouses as soon as possible so peers can
|
||||
// find us again. The update rides over a fresh lighthouse handshake. If the network isn't back up yet these
|
||||
// sends fail harmlessly and the periodic update worker retries within lighthouse.interval.
|
||||
cm.intf.lightHouse.TriggerUpdate()
|
||||
}
|
||||
|
||||
// clearAllTunnels closes every tunnel in the hostmap locally, without notifying the remotes. It is the wake-from-
|
||||
// sleep counterpart to Control.CloseAllTunnels: after a suspend the remotes stopped hearing from us long ago, and
|
||||
// close packets fired into a network that may not even be up yet are wasted, so we only tear down our own state
|
||||
// and let the next packet to each host start a fresh handshake.
|
||||
func (cm *connectionManager) clearAllTunnels() int {
|
||||
cm.hostMap.RLock()
|
||||
hostinfos := make([]*HostInfo, 0, len(cm.hostMap.Indexes))
|
||||
for _, h := range cm.hostMap.Indexes {
|
||||
hostinfos = append(hostinfos, h)
|
||||
}
|
||||
cm.hostMap.RUnlock()
|
||||
|
||||
for _, h := range hostinfos {
|
||||
cm.intf.closeTunnel(h)
|
||||
}
|
||||
|
||||
// With every tunnel gone no relay can be in use, drop the usage tracking wholesale.
|
||||
cm.relayUsedLock.Lock()
|
||||
clear(cm.relayUsed)
|
||||
cm.relayUsedLock.Unlock()
|
||||
|
||||
return len(hostinfos)
|
||||
}
|
||||
|
||||
func (cm *connectionManager) Start(ctx context.Context) {
|
||||
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
||||
defer clockSource.Stop()
|
||||
@@ -249,7 +150,6 @@ func (cm *connectionManager) Start(ctx context.Context) {
|
||||
return
|
||||
|
||||
case now := <-clockSource.C:
|
||||
cm.checkWake()
|
||||
cm.trafficTimer.Advance(now)
|
||||
for {
|
||||
localIndex, has := cm.trafficTimer.Purge()
|
||||
|
||||
@@ -25,6 +25,7 @@ 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)
|
||||
|
||||
@@ -501,85 +502,3 @@ func (d *dummyCert) MarshalJSON() ([]byte, error) {
|
||||
func (d *dummyCert) Copy() cert.Certificate {
|
||||
return d
|
||||
}
|
||||
|
||||
func TestConnectionManager_WakeClear(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
||||
vpnIp := netip.MustParseAddr("172.1.1.2")
|
||||
preferredRanges := []netip.Prefix{localrange}
|
||||
|
||||
// Very incomplete mock objects
|
||||
hostMap := newHostMap(l)
|
||||
hostMap.preferredRanges.Store(&preferredRanges)
|
||||
|
||||
cs := &CertState{
|
||||
initiatingVersion: cert.Version1,
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{version: cert.Version1},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
lh := newTestLighthouse()
|
||||
ifce := &Interface{
|
||||
hostMap: hostMap,
|
||||
inside: &overlaytest.NoopTun{},
|
||||
outside: &udp.NoopConn{},
|
||||
firewall: &Firewall{},
|
||||
lightHouse: lh,
|
||||
pki: &PKI{},
|
||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
||||
l: l,
|
||||
}
|
||||
ifce.pki.cs.Store(cs)
|
||||
|
||||
// Create manager
|
||||
conf := config.NewC(test.NewLogger())
|
||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||
nc.intf = ifce
|
||||
|
||||
// Drive the wake detector from a fake clock pair
|
||||
suspended := time.Duration(0)
|
||||
nc.wakeDetector = &wakeDetector{read: func() (time.Duration, bool) { return suspended, true }}
|
||||
nc.checkWake() // primes the baseline
|
||||
|
||||
addTunnel := func(localIndex uint32) *HostInfo {
|
||||
hostinfo := &HostInfo{
|
||||
vpnAddrs: []netip.Addr{vpnIp},
|
||||
localIndexId: localIndex,
|
||||
remoteIndexId: 9901,
|
||||
}
|
||||
hostinfo.ConnectionState = &ConnectionState{
|
||||
myCert: &dummyCert{version: cert.Version1},
|
||||
}
|
||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||
return hostinfo
|
||||
}
|
||||
|
||||
addTunnel(1099)
|
||||
nc.RelayUsed(5000)
|
||||
|
||||
// No suspend, nothing happens
|
||||
nc.checkWake()
|
||||
assert.Contains(t, nc.hostMap.Indexes, uint32(1099))
|
||||
|
||||
// A suspend below the threshold leaves tunnels alone
|
||||
suspended += 5 * time.Second
|
||||
nc.checkWake()
|
||||
assert.Contains(t, nc.hostMap.Indexes, uint32(1099))
|
||||
|
||||
// A suspend past the threshold clears everything, including relay usage tracking
|
||||
suspended += time.Hour
|
||||
nc.checkWake()
|
||||
assert.Empty(t, nc.hostMap.Indexes)
|
||||
assert.Empty(t, nc.hostMap.Hosts)
|
||||
assert.Empty(t, nc.relayUsed)
|
||||
|
||||
// With clear_on_wake disabled the tunnels survive a long suspend
|
||||
addTunnel(1100)
|
||||
nc.clearOnWake.Store(false)
|
||||
suspended += time.Hour
|
||||
nc.checkWake()
|
||||
assert.Contains(t, nc.hostMap.Indexes, uint32(1100))
|
||||
assert.Contains(t, nc.hostMap.Hosts, vpnIp)
|
||||
}
|
||||
|
||||
+19
-6
@@ -12,7 +12,15 @@ import (
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
)
|
||||
|
||||
const ReplayWindow = 1024
|
||||
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
|
||||
@@ -24,6 +32,8 @@ type ConnectionState struct {
|
||||
window *Bits
|
||||
decryptLock sync.Mutex
|
||||
writeLock sync.Mutex
|
||||
// epoch is this session's sessionEpoch ordinal. Immutable after creation.
|
||||
epoch uint64
|
||||
}
|
||||
|
||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
||||
@@ -38,6 +48,7 @@ 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++ {
|
||||
@@ -58,8 +69,7 @@ func (cs *ConnectionState) Curve() cert.Curve {
|
||||
return cs.myCert.Curve()
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) {
|
||||
var err error
|
||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
@@ -67,7 +77,7 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
|
||||
out, err = cs.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -81,7 +91,6 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// VerifyRelay verifies AEAD protected (but not encrypted) relay frames. packet must be length-checked by the caller.
|
||||
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
@@ -90,6 +99,11 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
// The entire body is sent as AD, not encrypted.
|
||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||
@@ -103,6 +117,5 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
+9
-1
@@ -53,6 +53,7 @@ type Control struct {
|
||||
statsStart func()
|
||||
dnsStart func()
|
||||
lighthouseStart func()
|
||||
networkChangeStart func(rebind func())
|
||||
connectionManagerStart func(context.Context)
|
||||
}
|
||||
|
||||
@@ -104,6 +105,9 @@ 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)
|
||||
}
|
||||
@@ -198,7 +202,11 @@ func (c *Control) RebindUDPServer() {
|
||||
return
|
||||
}
|
||||
|
||||
_ = c.f.outside.Rebind()
|
||||
// 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)
|
||||
}
|
||||
|
||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||
c.f.lightHouse.SendUpdate()
|
||||
|
||||
+40
-23
@@ -11,6 +11,8 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -30,9 +32,9 @@ func newFakeDevice() *fakeDevice {
|
||||
|
||||
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||
// the same way a closed device does
|
||||
func (d *fakeDevice) Read(p []byte) (int, error) {
|
||||
func (d *fakeDevice) Read() ([]tio.Packet, error) {
|
||||
<-d.closedCh
|
||||
return 0, io.EOF
|
||||
return nil, io.EOF
|
||||
}
|
||||
|
||||
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
||||
@@ -49,10 +51,8 @@ func (d *fakeDevice) Activate() error { return nil }
|
||||
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
||||
func (d *fakeDevice) Name() string { return "fake" }
|
||||
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||
func (d *fakeDevice) SupportsMultiqueue() bool { return false }
|
||||
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, errors.New("unsupported")
|
||||
}
|
||||
|
||||
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
|
||||
|
||||
// newReadyControl hand-builds the minimum Control that Main would have
|
||||
// produced right before Start, including the construction token NewInterface
|
||||
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
readers: make([]io.ReadWriteCloser, 1),
|
||||
batchers: make([]*batch.MultiCoalescer, 1),
|
||||
routines: 1,
|
||||
hostMap: newHostMap(l),
|
||||
lightHouse: lh,
|
||||
@@ -109,7 +109,8 @@ func TestControl_StopBeforeStart(t *testing.T) {
|
||||
require.NoError(t, c.Wait())
|
||||
|
||||
// A stopped control can never be started
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
err := c.Start()
|
||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||
|
||||
// A second Stop is a harmless no-op
|
||||
c.Stop()
|
||||
@@ -143,19 +144,29 @@ type fakeConn struct {
|
||||
rebinds int
|
||||
}
|
||||
|
||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||
return len(bufs), nil
|
||||
}
|
||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||
|
||||
type multiqueueDevice struct {
|
||||
*fakeDevice
|
||||
}
|
||||
|
||||
func (d *multiqueueDevice) SupportsMultiqueue() bool { return true }
|
||||
// Queues claims multiqueue support but fails to open the second queue,
|
||||
// exercising the activation error path.
|
||||
func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) {
|
||||
if n > 1 {
|
||||
return nil, errors.New("second queue failed to open")
|
||||
}
|
||||
return d.fakeDevice.Queues(n)
|
||||
}
|
||||
|
||||
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||
@@ -166,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
readers: make([]io.ReadWriteCloser, 2),
|
||||
batchers: make([]*batch.MultiCoalescer, 2),
|
||||
routines: 2,
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
@@ -181,7 +192,8 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
}
|
||||
|
||||
// The second reader fails to open, everything must be released
|
||||
require.Error(t, c.Start())
|
||||
err := c.Start()
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||
@@ -251,15 +263,18 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
||||
// panic and Wait must observe the final state
|
||||
require.NoError(t, c.Wait())
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
err := c.Start()
|
||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||
}
|
||||
|
||||
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||
c, dev, conn := newReadyControl(t)
|
||||
|
||||
require.NoError(t, c.Start())
|
||||
err := c.Start()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StateStarted, c.State())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
||||
err = c.Start()
|
||||
require.ErrorIs(t, err, ErrAlreadyStarted)
|
||||
|
||||
// Stop must unpark the reader blocked in the device and release everything
|
||||
c.Stop()
|
||||
@@ -270,7 +285,8 @@ func TestControl_StartStopLifecycle(t *testing.T) {
|
||||
|
||||
// The reader drained off a closed device, that is not a fatal error
|
||||
require.NoError(t, c.Wait())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
err = c.Start()
|
||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||
}
|
||||
|
||||
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||
@@ -280,7 +296,8 @@ func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||
c.RebindUDPServer()
|
||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||
|
||||
require.NoError(t, c.Start())
|
||||
err := c.Start()
|
||||
require.NoError(t, err)
|
||||
c.RebindUDPServer()
|
||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||
|
||||
|
||||
+13
-1
@@ -108,7 +108,19 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||
}
|
||||
|
||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||
return c.f.outside.(*udp.TesterConn).Addr
|
||||
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
|
||||
}
|
||||
|
||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
// Package cpupick chooses which CPUs the tun reader threads pin to when the
|
||||
// operator has not chosen for us (tun.cpu_affinity). The stock spread —
|
||||
// allowed[i] for routine i — has two failure modes this package exists to fix:
|
||||
//
|
||||
// - every co-located nebula starts its spread at allowed[0], so N instances
|
||||
// on one box stack their readers onto the same cores, and allowed[0] is
|
||||
// usually CPU 0, the core housekeeping and default IRQ affinity already
|
||||
// favor;
|
||||
// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact
|
||||
// cores) low IDs are not necessarily fast cores, and pinning an encrypt
|
||||
// thread to an efficiency core caps that queue's throughput.
|
||||
//
|
||||
// Default instead returns a preference-ordered pin list: the allowed set
|
||||
// filtered to performance cores (when the platform distinguishes them and
|
||||
// enough remain for every routine), confined to a single NUMA node and spread
|
||||
// across distinct physical cores when the topology permits, CPU 0's physical
|
||||
// core demoted to last resort, and the order rotated by a stable per-instance
|
||||
// key so co-located instances spread instead of stacking.
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
|
||||
// topology is the slice of machine layout arrange consults: the NUMA node
|
||||
// and the physical core behind each candidate CPU, plus which core CPU 0
|
||||
// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's
|
||||
// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate).
|
||||
// Probed from sysfs on Linux; flatTopology stands in when the platform can't
|
||||
// say, which turns every topology rule into a no-op rather than a wrong
|
||||
// answer.
|
||||
type topology struct {
|
||||
nodeOf map[int]int
|
||||
coreOf map[int]int
|
||||
zeroCore int
|
||||
}
|
||||
|
||||
// flatTopology places every CPU on node 0 and on a physical core of its own.
|
||||
func flatTopology(cpus []int) topology {
|
||||
t := topology{
|
||||
nodeOf: make(map[int]int, len(cpus)),
|
||||
coreOf: make(map[int]int, len(cpus)),
|
||||
zeroCore: -1,
|
||||
}
|
||||
for i, c := range cpus {
|
||||
t.nodeOf[c] = 0
|
||||
t.coreOf[c] = i
|
||||
if c == 0 {
|
||||
t.zeroCore = i
|
||||
}
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
// Default computes the pin order for `routines` tun readers. key is any
|
||||
// stable per-instance value; the bound UDP port is ideal — distinct across
|
||||
// co-located instances, stable across restarts so benchmark runs stay
|
||||
// comparable. Returns nil when there is nothing useful to say (no affinity
|
||||
// support on this platform, lookup failure); callers keep their existing
|
||||
// fallback spread.
|
||||
func Default(routines int, key uint64, l *slog.Logger) []int {
|
||||
allowed, err := util.AllowedCPUs()
|
||||
if err != nil || len(allowed) == 0 {
|
||||
return nil
|
||||
}
|
||||
perf, signal := perfCPUs(allowed)
|
||||
cands := pickCandidates(allowed, perf, routines)
|
||||
if len(cands) == 0 {
|
||||
return nil
|
||||
}
|
||||
if len(perf) < routines {
|
||||
signal = ""
|
||||
}
|
||||
cpus := arrange(cands, readTopology(cands), routines, splitmix64(key))
|
||||
if l != nil {
|
||||
l.Info("chose default pin CPUs for tun readers",
|
||||
"cpus", cpus[:min(routines, len(cpus))],
|
||||
"perfSignal", signal)
|
||||
}
|
||||
return cpus
|
||||
}
|
||||
|
||||
// pickCandidates applies the enough-for-everyone guard: a perf filter that
|
||||
// leaves fewer candidates than routines is discarded — giving every reader
|
||||
// its own (possibly slow) core beats stacking two readers on a fast one.
|
||||
func pickCandidates(allowed, perf []int, routines int) []int {
|
||||
if len(perf) < routines {
|
||||
return allowed
|
||||
}
|
||||
return perf
|
||||
}
|
||||
|
||||
// arrange turns the candidate set into the final pin order:
|
||||
//
|
||||
// 1. NUMA: when at least one node holds enough candidates for every
|
||||
// routine, confine to one such node, chosen by the instance hash. The
|
||||
// readers share hostmap and cipher state, so splitting one instance
|
||||
// across nodes taxes every packet — and co-located instances that hash
|
||||
// to different nodes stop competing entirely. When no node is big
|
||||
// enough, span nodes rather than stack readers.
|
||||
// 2. Rotate the preferred candidates by the hash so instances spread.
|
||||
// 3. SMT: emit one thread per physical core before any of their siblings —
|
||||
// two encrypt threads on one core split its execution units. Siblings
|
||||
// still follow for the routines > cores case.
|
||||
// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ
|
||||
// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the
|
||||
// sibling precedes CPU 0 itself, which only catches the bleed-through.
|
||||
//
|
||||
// The rotation happens before the SMT pass so each instance's one-per-core
|
||||
// walk also starts at a different core, and CPU 0's core is excluded from
|
||||
// the rotation so no hash value can put it back at the front.
|
||||
func arrange(cands []int, topo topology, routines int, h uint64) []int {
|
||||
byNode := map[int][]int{}
|
||||
var nodes []int
|
||||
for _, c := range cands {
|
||||
n := topo.nodeOf[c]
|
||||
if _, ok := byNode[n]; !ok {
|
||||
nodes = append(nodes, n)
|
||||
}
|
||||
byNode[n] = append(byNode[n], c)
|
||||
}
|
||||
var eligible []int
|
||||
for _, n := range nodes {
|
||||
if len(byNode[n]) >= routines {
|
||||
eligible = append(eligible, n)
|
||||
}
|
||||
}
|
||||
if len(eligible) > 0 {
|
||||
cands = byNode[eligible[int(h%uint64(len(eligible)))]]
|
||||
}
|
||||
|
||||
// Split off CPU 0's core: its siblings tail the list, CPU 0 tails them.
|
||||
preferred := make([]int, 0, len(cands))
|
||||
var zeroTail []int
|
||||
hasZero := false
|
||||
for _, c := range cands {
|
||||
switch {
|
||||
case c == 0:
|
||||
hasZero = true
|
||||
case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore:
|
||||
zeroTail = append(zeroTail, c)
|
||||
default:
|
||||
preferred = append(preferred, c)
|
||||
}
|
||||
}
|
||||
if hasZero {
|
||||
zeroTail = append(zeroTail, 0)
|
||||
}
|
||||
if len(preferred) == 0 {
|
||||
return zeroTail // CPU 0's core is all we have
|
||||
}
|
||||
|
||||
// The node pick consumed the low hash bits; rotate by the high ones so
|
||||
// the two choices stay independent.
|
||||
off := int((h >> 32) % uint64(len(preferred)))
|
||||
rot := make([]int, 0, len(preferred))
|
||||
rot = append(rot, preferred[off:]...)
|
||||
rot = append(rot, preferred[:off]...)
|
||||
|
||||
seenCore := make(map[int]bool, len(rot))
|
||||
out := make([]int, 0, len(cands))
|
||||
var siblings []int
|
||||
for _, c := range rot {
|
||||
g := topo.coreOf[c]
|
||||
if seenCore[g] {
|
||||
siblings = append(siblings, c)
|
||||
continue
|
||||
}
|
||||
seenCore[g] = true
|
||||
out = append(out, c)
|
||||
}
|
||||
out = append(out, siblings...)
|
||||
out = append(out, zeroTail...)
|
||||
return out
|
||||
}
|
||||
|
||||
// splitmix64 decorrelates instance keys before the selection modulos: ports
|
||||
// on one box often share spacing (4242/4243, or round steps like +1000) that
|
||||
// raw key%len arithmetic would fold onto the same offset.
|
||||
func splitmix64(x uint64) uint64 {
|
||||
x += 0x9e3779b97f4a7c15
|
||||
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9
|
||||
x = (x ^ (x >> 27)) * 0x94d049bb133111eb
|
||||
return x ^ (x >> 31)
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// pairTopo builds a topology where consecutive candidate pairs are SMT
|
||||
// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ...
|
||||
// All CPUs land on node 0.
|
||||
func pairTopo(cpus []int) topology {
|
||||
t := topology{
|
||||
nodeOf: make(map[int]int, len(cpus)),
|
||||
coreOf: make(map[int]int, len(cpus)),
|
||||
zeroCore: -1,
|
||||
}
|
||||
for i, c := range cpus {
|
||||
t.nodeOf[c] = 0
|
||||
t.coreOf[c] = i / 2
|
||||
if c == 0 {
|
||||
t.zeroCore = i / 2
|
||||
}
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
func TestArrangeDemotesZeroForEveryKey(t *testing.T) {
|
||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
for key := range uint64(64) {
|
||||
got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key))
|
||||
if len(got) != len(candidates) {
|
||||
t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates))
|
||||
}
|
||||
if got[0] == 0 {
|
||||
t.Errorf("key %d: CPU 0 at the front: %v", key, got)
|
||||
}
|
||||
if got[len(got)-1] != 0 {
|
||||
t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got)
|
||||
}
|
||||
sorted := slices.Clone(got)
|
||||
slices.Sort(sorted)
|
||||
if !slices.Equal(sorted, candidates) {
|
||||
t.Errorf("key %d: not a permutation: %v", key, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeDemotesZeroSiblings(t *testing.T) {
|
||||
// Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 —
|
||||
// must tail the list, sibling ahead of 0 itself.
|
||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
for key := range uint64(64) {
|
||||
got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key))
|
||||
n := len(got)
|
||||
if got[n-1] != 0 || got[n-2] != 1 {
|
||||
t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeZeroSiblingWithoutZero(t *testing.T) {
|
||||
// CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still
|
||||
// tails the list when the topology knows which core CPU 0 lives on.
|
||||
candidates := []int{1, 2, 3, 4, 5}
|
||||
topo := pairTopo([]int{0, 1, 2, 3, 4, 5})
|
||||
got := arrange(candidates, topo, 2, splitmix64(7))
|
||||
if got[len(got)-1] != 1 {
|
||||
t.Errorf("CPU 0's sibling not demoted: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeRotatesByKey(t *testing.T) {
|
||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
||||
seen := map[int]bool{}
|
||||
for key := range uint64(64) {
|
||||
seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true
|
||||
}
|
||||
// 64 hashed keys over 8 slots must hit more than one starting CPU, or
|
||||
// co-located instances would all stack again.
|
||||
if len(seen) < 2 {
|
||||
t.Errorf("rotation never varied across keys: %v", seen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeStableForSameKey(t *testing.T) {
|
||||
candidates := []int{0, 2, 4, 6}
|
||||
topo := flatTopology(candidates)
|
||||
a := arrange(candidates, topo, 2, splitmix64(4242))
|
||||
b := arrange(candidates, topo, 2, splitmix64(4242))
|
||||
if !slices.Equal(a, b) {
|
||||
t.Errorf("same key ordered differently: %v vs %v", a, b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeZeroOnly(t *testing.T) {
|
||||
if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) {
|
||||
t.Errorf("sole CPU 0 must survive: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeSMTSiblingsLast(t *testing.T) {
|
||||
// Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four
|
||||
// distinct physical cores before any sibling repeats.
|
||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
||||
topo := pairTopo(candidates)
|
||||
for key := range uint64(16) {
|
||||
got := arrange(candidates, topo, 4, splitmix64(key))
|
||||
seen := map[int]bool{}
|
||||
for _, c := range got[:4] {
|
||||
g := topo.coreOf[c]
|
||||
if seen[g] {
|
||||
t.Fatalf("key %d: sibling before all cores covered: %v", key, got)
|
||||
}
|
||||
seen[g] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeNUMAConfinesToOneNode(t *testing.T) {
|
||||
// Two nodes of four; both fit routines=3, so the result must sit
|
||||
// entirely inside one of them, and the hash must pick both across keys.
|
||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
||||
topo := flatTopology(candidates)
|
||||
for _, c := range []int{10, 11, 12, 13} {
|
||||
topo.nodeOf[c] = 1
|
||||
}
|
||||
nodesSeen := map[int]bool{}
|
||||
for key := range uint64(32) {
|
||||
got := arrange(candidates, topo, 3, splitmix64(key))
|
||||
if len(got) != 4 {
|
||||
t.Fatalf("key %d: not confined to one node: %v", key, got)
|
||||
}
|
||||
n := topo.nodeOf[got[0]]
|
||||
for _, c := range got {
|
||||
if topo.nodeOf[c] != n {
|
||||
t.Fatalf("key %d: spans nodes: %v", key, got)
|
||||
}
|
||||
}
|
||||
nodesSeen[n] = true
|
||||
}
|
||||
if len(nodesSeen) != 2 {
|
||||
t.Errorf("hash never spread instances across nodes: %v", nodesSeen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) {
|
||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
||||
topo := flatTopology(candidates)
|
||||
for _, c := range []int{10, 11, 12, 13} {
|
||||
topo.nodeOf[c] = 1
|
||||
}
|
||||
got := arrange(candidates, topo, 6, splitmix64(1))
|
||||
if len(got) != len(candidates) {
|
||||
t.Errorf("undersized nodes must span, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickCandidates(t *testing.T) {
|
||||
allowed := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
perf := []int{4, 5}
|
||||
|
||||
// Enough perf cores for every routine: only they are used.
|
||||
if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) {
|
||||
t.Errorf("perf filter not applied: %v", got)
|
||||
}
|
||||
// Perf filter too small for the routine count: discarded, everyone
|
||||
// gets their own core from the full allowed set.
|
||||
if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) {
|
||||
t.Errorf("undersized perf filter not discarded: %v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// capacityKeepPct is the cpu_capacity admission threshold, relative to the
|
||||
// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big
|
||||
// core's 1024 while mid cores sit at ~75%+, so half of max separates little
|
||||
// from the rest without splitting prime from mid on three-tier parts.
|
||||
const capacityKeepPct = 50
|
||||
|
||||
// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core
|
||||
// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE
|
||||
// cores, and AMD compact cores all sit >= 20% below their siblings' max.
|
||||
const freqKeepPct = 85
|
||||
|
||||
// perfCPUs partitions allowed into the subset that are "performance" cores,
|
||||
// consulting (in order of authority):
|
||||
//
|
||||
// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on
|
||||
// arm/arm64/riscv; the scheduler's own view of big vs LITTLE.
|
||||
// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present
|
||||
// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright.
|
||||
// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact
|
||||
// cores, which neither of the above covers.
|
||||
//
|
||||
// Returns allowed unchanged (signal "") when nothing distinguishes the
|
||||
// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable.
|
||||
func perfCPUs(allowed []int) ([]int, string) {
|
||||
return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed)
|
||||
}
|
||||
|
||||
func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) {
|
||||
if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok {
|
||||
return cpus, "cpu_capacity"
|
||||
}
|
||||
if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok {
|
||||
return cpus, "intel_core_pmu"
|
||||
}
|
||||
if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok {
|
||||
return cpus, "max_freq"
|
||||
}
|
||||
return allowed, ""
|
||||
}
|
||||
|
||||
// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least
|
||||
// keepPct percent of the maximum across allowed. Inconclusive (ok=false)
|
||||
// when any CPU is missing the file or when every value is equal.
|
||||
func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) {
|
||||
vals := make([]int, len(allowed))
|
||||
minV, maxV := 0, 0
|
||||
for i, cpu := range allowed {
|
||||
v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file))
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
vals[i] = v
|
||||
if i == 0 || v < minV {
|
||||
minV = v
|
||||
}
|
||||
if v > maxV {
|
||||
maxV = v
|
||||
}
|
||||
}
|
||||
if minV == maxV {
|
||||
return nil, false // homogeneous by this signal; try the next one
|
||||
}
|
||||
keep := make([]int, 0, len(allowed))
|
||||
for i, cpu := range allowed {
|
||||
if vals[i]*100 >= maxV*keepPct {
|
||||
keep = append(keep, cpu)
|
||||
}
|
||||
}
|
||||
return keep, true
|
||||
}
|
||||
|
||||
// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU
|
||||
// mask. Inconclusive when the file is absent (non-hybrid x86, other arches)
|
||||
// or no allowed CPU is in the mask (the process was deliberately confined
|
||||
// to E-cores; nothing useful to prefer within that).
|
||||
func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) {
|
||||
b, err := os.ReadFile(maskPath)
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
set, err := parseCPUList(strings.TrimSpace(string(b)))
|
||||
if err != nil || len(set) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
pcore := make(map[int]bool, len(set))
|
||||
for _, c := range set {
|
||||
pcore[c] = true
|
||||
}
|
||||
keep := make([]int, 0, len(allowed))
|
||||
for _, cpu := range allowed {
|
||||
if pcore[cpu] {
|
||||
keep = append(keep, cpu)
|
||||
}
|
||||
}
|
||||
if len(keep) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
return keep, true
|
||||
}
|
||||
|
||||
// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into
|
||||
// individual CPU IDs. Empty input yields an empty list.
|
||||
func parseCPUList(s string) ([]int, error) {
|
||||
if s == "" {
|
||||
return nil, nil
|
||||
}
|
||||
var out []int
|
||||
for part := range strings.SplitSeq(s, ",") {
|
||||
part = strings.TrimSpace(part)
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
lo, hi, isRange := strings.Cut(part, "-")
|
||||
a, err := strconv.Atoi(lo)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
||||
}
|
||||
if !isRange {
|
||||
out = append(out, a)
|
||||
continue
|
||||
}
|
||||
b, err := strconv.Atoi(hi)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
||||
}
|
||||
if b < a || b-a > 8192 {
|
||||
return nil, fmt.Errorf("bad cpulist range %q", part)
|
||||
}
|
||||
for v := a; v <= b; v++ {
|
||||
out = append(out, v)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func readIntFile(path string) (int, error) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return strconv.Atoi(strings.TrimSpace(string(b)))
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeSysfs builds a cpuDir tree with the given per-CPU file values.
|
||||
// A nil map for a file means "file absent on every CPU".
|
||||
func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
write := func(cpu int, rel string, v int) {
|
||||
p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel)
|
||||
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for cpu, v := range capacity {
|
||||
write(cpu, "cpu_capacity", v)
|
||||
}
|
||||
for cpu, v := range maxFreq {
|
||||
write(cpu, "cpufreq/cpuinfo_max_freq", v)
|
||||
}
|
||||
return dir
|
||||
}
|
||||
|
||||
func writeCoreMask(t *testing.T, mask string) string {
|
||||
t.Helper()
|
||||
p := filepath.Join(t.TempDir(), "cpus")
|
||||
if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func TestPerfCPUsBigLittleCapacity(t *testing.T) {
|
||||
// 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM.
|
||||
dir := fakeSysfs(t, map[int]int{
|
||||
0: 1024, 1: 1024, 2: 1024, 3: 1024,
|
||||
4: 290, 5: 290, 6: 290, 7: 290,
|
||||
}, nil)
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
||||
if signal != "cpu_capacity" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsThreeTierKeepsMid(t *testing.T) {
|
||||
// prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid.
|
||||
dir := fakeSysfs(t, map[int]int{
|
||||
0: 280, 1: 280, 2: 280, 3: 280,
|
||||
4: 780, 5: 780, 6: 780,
|
||||
7: 1024,
|
||||
}, nil)
|
||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
||||
if !slices.Equal(got, []int{4, 5, 6, 7}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsIntelHybridMask(t *testing.T) {
|
||||
// No cpu_capacity on x86; the P-core PMU mask decides.
|
||||
dir := fakeSysfs(t, nil, nil)
|
||||
mask := writeCoreMask(t, "0-7")
|
||||
got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11})
|
||||
if signal != "intel_core_pmu" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) {
|
||||
// Confined to E-cores only: the mask can't help, and equal freqs below
|
||||
// mean nothing else distinguishes them either -> allowed unchanged.
|
||||
dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000})
|
||||
mask := writeCoreMask(t, "0-7")
|
||||
got, signal := perfCPUsFrom(dir, mask, []int{8, 9})
|
||||
if signal != "" || !slices.Equal(got, []int{8, 9}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsMaxFreqCompactCores(t *testing.T) {
|
||||
// AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz.
|
||||
dir := fakeSysfs(t, nil, map[int]int{
|
||||
0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000,
|
||||
})
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
||||
if signal != "max_freq" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) {
|
||||
// Turbo Boost Max favored cores run a few percent hot; they must not
|
||||
// shrink the candidate set to one or two cores.
|
||||
dir := fakeSysfs(t, nil, map[int]int{
|
||||
0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000,
|
||||
})
|
||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("favored-core skew filtered CPUs: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsHomogeneousInconclusive(t *testing.T) {
|
||||
dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000})
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1})
|
||||
if signal != "" || !slices.Equal(got, []int{0, 1}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsNoSysfs(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2})
|
||||
if signal != "" || !slices.Equal(got, []int{0, 1, 2}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCPUList(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
want []int
|
||||
wantErr bool
|
||||
}{
|
||||
{"0-3", []int{0, 1, 2, 3}, false},
|
||||
{"0-1,16-17", []int{0, 1, 16, 17}, false},
|
||||
{"5", []int{5}, false},
|
||||
{"", nil, false},
|
||||
{"3-1", nil, true},
|
||||
{"a-b", nil, true},
|
||||
{"1,x", nil, true},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := parseCPUList(c.in)
|
||||
if (err != nil) != c.wantErr {
|
||||
t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr)
|
||||
continue
|
||||
}
|
||||
if !c.wantErr && !slices.Equal(got, c.want) {
|
||||
t.Errorf("%q: got %v want %v", c.in, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build !linux
|
||||
|
||||
package cpupick
|
||||
|
||||
// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction".
|
||||
// Default already returns nil off-Linux (util.AllowedCPUs has no answer
|
||||
// there), so this exists to keep the package compiling everywhere.
|
||||
func perfCPUs(allowed []int) ([]int, string) {
|
||||
return allowed, ""
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// readTopology probes the NUMA node and physical-core layout of cpus from
|
||||
// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown
|
||||
// node becomes node 0, an unknown core becomes a core of its own — either
|
||||
// way the corresponding arrange rule becomes a no-op instead of a wrong
|
||||
// answer.
|
||||
func readTopology(cpus []int) topology {
|
||||
return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus)
|
||||
}
|
||||
|
||||
func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology {
|
||||
coreOf, zeroCore := coreGroups(cpuDir, cpus)
|
||||
return topology{
|
||||
nodeOf: numaNodes(nodeDir, cpus),
|
||||
coreOf: coreOf,
|
||||
zeroCore: zeroCore,
|
||||
}
|
||||
}
|
||||
|
||||
// numaNodes maps each cpu to its NUMA node via
|
||||
// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node
|
||||
// dirs at all: VMs, non-NUMA kernels) land on node 0.
|
||||
func numaNodes(nodeDir string, cpus []int) map[int]int {
|
||||
out := make(map[int]int, len(cpus))
|
||||
for _, c := range cpus {
|
||||
out[c] = 0
|
||||
}
|
||||
entries, err := os.ReadDir(nodeDir)
|
||||
if err != nil {
|
||||
return out
|
||||
}
|
||||
want := make(map[int]bool, len(cpus))
|
||||
for _, c := range cpus {
|
||||
want[c] = true
|
||||
}
|
||||
for _, e := range entries {
|
||||
id, ok := strings.CutPrefix(e.Name(), "node")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
n, err := strconv.Atoi(id)
|
||||
if err != nil {
|
||||
continue // has_cpu, possible, ... share the prefix
|
||||
}
|
||||
b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist"))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
list, err := parseCPUList(strings.TrimSpace(string(b)))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, c := range list {
|
||||
if want[c] {
|
||||
out[c] = n
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// coreGroups maps each cpu to a dense physical-core id derived from its
|
||||
// (physical_package_id, core_id) pair — core_id alone repeats across
|
||||
// sockets. CPUs whose topology files are unreadable get a core of their own.
|
||||
// The second return is the group id of the core CPU 0 lives on, or -1 when
|
||||
// that can't be determined; CPU 0's own files are consulted even when 0 is
|
||||
// not a candidate, so its SMT siblings are recognized under cpusets that
|
||||
// exclude CPU 0 itself.
|
||||
func coreGroups(cpuDir string, cpus []int) (map[int]int, int) {
|
||||
type pkgCore struct{ pkg, core int }
|
||||
pairOf := func(cpu int) (pkgCore, bool) {
|
||||
topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
||||
pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id"))
|
||||
core, err2 := readIntFile(filepath.Join(topoDir, "core_id"))
|
||||
if err1 != nil || err2 != nil {
|
||||
return pkgCore{}, false
|
||||
}
|
||||
return pkgCore{pkg, core}, true
|
||||
}
|
||||
|
||||
ids := map[pkgCore]int{}
|
||||
out := make(map[int]int, len(cpus))
|
||||
next := 0
|
||||
for _, cpu := range cpus {
|
||||
k, ok := pairOf(cpu)
|
||||
if !ok {
|
||||
out[cpu] = next
|
||||
next++
|
||||
continue
|
||||
}
|
||||
id, ok := ids[k]
|
||||
if !ok {
|
||||
id = next
|
||||
next++
|
||||
ids[k] = id
|
||||
}
|
||||
out[cpu] = id
|
||||
}
|
||||
|
||||
zeroCore := -1
|
||||
if k, ok := pairOf(0); ok {
|
||||
if id, ok := ids[k]; ok {
|
||||
zeroCore = id
|
||||
}
|
||||
}
|
||||
return out, zeroCore
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist
|
||||
// string; cores maps cpu -> (package, core) pair.
|
||||
func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) {
|
||||
t.Helper()
|
||||
base := t.TempDir()
|
||||
nodeDir := filepath.Join(base, "node")
|
||||
cpuDir := filepath.Join(base, "cpu")
|
||||
for n, list := range nodes {
|
||||
d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n))
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for cpu, pc := range cores {
|
||||
d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return nodeDir, cpuDir
|
||||
}
|
||||
|
||||
func TestReadTopology(t *testing.T) {
|
||||
// Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1.
|
||||
// core_id repeats across packages on purpose: the pair must disambiguate.
|
||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
||||
map[int]string{0: "0-1,4-5", 1: "2-3,6-7"},
|
||||
map[int][2]int{
|
||||
0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1},
|
||||
2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1},
|
||||
})
|
||||
cpus := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
topo := readTopologyFrom(nodeDir, cpuDir, cpus)
|
||||
|
||||
for _, c := range []int{0, 1, 4, 5} {
|
||||
if topo.nodeOf[c] != 0 {
|
||||
t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c])
|
||||
}
|
||||
}
|
||||
for _, c := range []int{2, 3, 6, 7} {
|
||||
if topo.nodeOf[c] != 1 {
|
||||
t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c])
|
||||
}
|
||||
}
|
||||
pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}}
|
||||
for _, p := range pairs {
|
||||
if topo.coreOf[p[0]] != topo.coreOf[p[1]] {
|
||||
t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]])
|
||||
}
|
||||
}
|
||||
if topo.coreOf[0] == topo.coreOf[2] {
|
||||
t.Error("cross-package cores with equal core_id must not merge")
|
||||
}
|
||||
if topo.zeroCore != topo.coreOf[0] {
|
||||
t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) {
|
||||
// CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is:
|
||||
// zeroCore must still identify their shared core.
|
||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
||||
map[int]string{0: "0-7"},
|
||||
map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}})
|
||||
topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5})
|
||||
if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore {
|
||||
t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4])
|
||||
}
|
||||
if topo.coreOf[1] == topo.zeroCore {
|
||||
t.Error("cpu 1 wrongly grouped with CPU 0's core")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadTopologyMissingSysfs(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
cpus := []int{0, 1, 2}
|
||||
topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus)
|
||||
seen := map[int]bool{}
|
||||
for _, c := range cpus {
|
||||
if topo.nodeOf[c] != 0 {
|
||||
t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c])
|
||||
}
|
||||
if seen[topo.coreOf[c]] {
|
||||
t.Errorf("cpu %d shares a fallback core group", c)
|
||||
}
|
||||
seen[topo.coreOf[c]] = true
|
||||
}
|
||||
if topo.zeroCore != -1 {
|
||||
t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build !linux
|
||||
|
||||
package cpupick
|
||||
|
||||
// readTopology has no sysfs to consult off Linux; the flat stand-in makes
|
||||
// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux
|
||||
// (util.AllowedCPUs has no answer there) — this keeps the package compiling.
|
||||
func readTopology(cpus []int) topology {
|
||||
return flatTopology(cpus)
|
||||
}
|
||||
+2
-4
@@ -4,15 +4,13 @@
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"log/slog"
|
||||
|
||||
"dario.cat/mergo"
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
@@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
||||
func NewTestLogger() *slog.Logger {
|
||||
v := os.Getenv("TEST_LOGS")
|
||||
if v == "" {
|
||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
return slog.New(slog.DiscardHandler)
|
||||
}
|
||||
|
||||
level := slog.LevelInfo
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
//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()
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
//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
|
||||
}
|
||||
+148
-29
@@ -114,6 +114,28 @@ 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() {
|
||||
@@ -131,6 +153,9 @@ 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
|
||||
@@ -141,7 +166,9 @@ 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())
|
||||
|
||||
if err := os.MkdirAll("mermaid", 0755); err != nil {
|
||||
// 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 {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
@@ -152,7 +179,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
outNat: make(map[outNatKey]netip.AddrPort),
|
||||
flow: []flowEntry{},
|
||||
ignoreFlows: []ignoreFlow{},
|
||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||
fn: fn,
|
||||
t: t,
|
||||
cancelRender: cancel,
|
||||
}
|
||||
@@ -249,7 +276,7 @@ func (r *R) renderFlow() {
|
||||
continue
|
||||
}
|
||||
|
||||
addr := e.packet.from.GetUDPAddr()
|
||||
addr := e.packet.fromAddr()
|
||||
if _, ok := participants[addr]; ok {
|
||||
continue
|
||||
}
|
||||
@@ -268,7 +295,6 @@ 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)
|
||||
@@ -280,21 +306,22 @@ 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 = "->>"
|
||||
}
|
||||
|
||||
fmt.Fprintf(f,
|
||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
||||
normalizeName(p.from.GetUDPAddr().String()),
|
||||
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()),
|
||||
line,
|
||||
normalizeName(p.to.GetUDPAddr().String()),
|
||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
||||
normalizeName(p.toAddr().String()),
|
||||
detail,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -408,29 +435,34 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
||||
|
||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||
|
||||
if len(r.ignoreFlows) > 0 {
|
||||
var h header.H
|
||||
err := h.Parse(p.Data)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
var h header.H
|
||||
var parseErr error
|
||||
if !tun {
|
||||
parseErr = h.Parse(p.Data)
|
||||
}
|
||||
|
||||
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 {
|
||||
// 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
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
}
|
||||
|
||||
fp := &packet{
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
h: h,
|
||||
parseErr: parseErr,
|
||||
}
|
||||
|
||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||
@@ -660,6 +692,10 @@ 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)
|
||||
@@ -690,6 +726,85 @@ 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 {
|
||||
@@ -782,6 +897,10 @@ 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)
|
||||
|
||||
+30
-14
@@ -131,6 +131,9 @@ listen:
|
||||
port: 4242
|
||||
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
|
||||
# default is 64, does not support reload
|
||||
# Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced
|
||||
# superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the
|
||||
# default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts.
|
||||
#batch: 64
|
||||
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
|
||||
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
|
||||
@@ -146,6 +149,14 @@ 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.
|
||||
@@ -254,6 +265,25 @@ tun:
|
||||
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
||||
mtu: 1300
|
||||
|
||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
||||
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
||||
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
|
||||
#pin_threads: true
|
||||
|
||||
# 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.
|
||||
#cpu_affinity:
|
||||
# - 2
|
||||
# - 4
|
||||
|
||||
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
||||
routes:
|
||||
#- mtu: 8800
|
||||
@@ -390,20 +420,6 @@ logging:
|
||||
# This setting is reloadable
|
||||
#inactivity_timeout: 10m
|
||||
|
||||
# clear_on_wake controls whether all tunnels are immediately torn down (locally, without notifying the remotes)
|
||||
# when the machine detects it has just woken from system sleep. Tunnels rarely survive a suspend: NAT mappings
|
||||
# expire and the machine's address usually changes, so waiting for the normal liveness checks costs 15-20 seconds
|
||||
# of black-holed traffic per tunnel after wake. Clearing them means the first packet after wake starts a fresh
|
||||
# handshake right away.
|
||||
# This setting is reloadable
|
||||
#clear_on_wake: true
|
||||
|
||||
# wake_clear_threshold is the minimum time the machine must have been suspended for clear_on_wake to act.
|
||||
# Suspends shorter than this often come back before NAT state expires, so those tunnels may still be alive and
|
||||
# are left to the normal liveness checks. Values below 1s are treated as 1s.
|
||||
# This setting is reloadable
|
||||
#wake_clear_threshold: 30s
|
||||
|
||||
# Nebula security group configuration
|
||||
firewall:
|
||||
# Action to take when a packet is not allowed by the firewall rules.
|
||||
|
||||
+4
-2
@@ -5,6 +5,8 @@ import (
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||
@@ -56,8 +58,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||
c.cacheV = tick
|
||||
if ll := len(c.cache); ll > 0 {
|
||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
||||
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||
}
|
||||
c.cache = make(ConntrackCache, ll)
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
@@ -30,27 +31,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
||||
|
||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 3)
|
||||
c.Get()
|
||||
|
||||
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
||||
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 2)
|
||||
c.Get()
|
||||
|
||||
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
||||
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||
|
||||
c := newFixedTicker(t, l, 5)
|
||||
c.Get()
|
||||
@@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
||||
|
||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 0)
|
||||
c.Get()
|
||||
|
||||
@@ -65,3 +65,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
"Fragment": fp.Fragment,
|
||||
})
|
||||
}
|
||||
|
||||
// ParsedPacket is a Packet plus the parse byproducts the RX path reuses
|
||||
type ParsedPacket struct {
|
||||
Packet
|
||||
IPHdrLen int
|
||||
// FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6.
|
||||
// Distinct from Packet.Fragment, which is true only for NON-FIRST fragments
|
||||
FragAny bool
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
module github.com/slackhq/nebula
|
||||
|
||||
go 1.25.0
|
||||
go 1.26.0
|
||||
|
||||
require (
|
||||
dario.cat/mergo v1.0.2
|
||||
|
||||
+11
-2
@@ -295,7 +295,13 @@ 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 {
|
||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||
// 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",
|
||||
"udpAddr", addr,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
@@ -969,6 +975,9 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
nb := make([]byte, 12, 12)
|
||||
out := make([]byte, mtu)
|
||||
for _, cp := range hh.packetStore {
|
||||
// TODO: use a SendBatch here. Each callback lands in
|
||||
// sendNoMetrics -> WriteTo: one syscall per cached packet,
|
||||
// where one sendmmsg could flush the whole store.
|
||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||
}
|
||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||
@@ -1080,7 +1089,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
||||
// We received a valid handshake on this relay, so make sure the relay
|
||||
// state reflects that, in case it had been marked Disestablished.
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0)
|
||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
|
||||
return
|
||||
}
|
||||
|
||||
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
|
||||
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
+11
-6
@@ -190,13 +190,18 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
||||
}
|
||||
|
||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
switch t {
|
||||
case Message:
|
||||
return s == MessageNone || s == MessageRelay
|
||||
case Handshake:
|
||||
return s == HandshakeIXPSK0
|
||||
case Test:
|
||||
return s == TestReply || s == TestRequest
|
||||
case Control, CloseTunnel, RecvError, LightHouse:
|
||||
return s == 0
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// NewHeader turns bytes into a header
|
||||
|
||||
@@ -102,6 +102,57 @@ func TestTypeMap(t *testing.T) {
|
||||
}, subTypeMap)
|
||||
}
|
||||
|
||||
// mapIsValidSubType is the pre-refactor, map-driven definition of a valid
|
||||
// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps
|
||||
// the original behavior around so we can prove the switch is equivalent to it.
|
||||
func mapIsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func TestIsValidSubType(t *testing.T) {
|
||||
// Explicit intent table: documents exactly which subtypes are valid so the
|
||||
// test stays meaningful even if both the switch and subTypeMap change.
|
||||
assert.True(t, IsValidSubType(Message, MessageNone))
|
||||
assert.True(t, IsValidSubType(Message, MessageRelay))
|
||||
assert.False(t, IsValidSubType(Message, 2))
|
||||
|
||||
assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0))
|
||||
// HandshakeXXPSK0 is defined but not a wire-valid subtype.
|
||||
assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0))
|
||||
|
||||
assert.True(t, IsValidSubType(Test, TestRequest))
|
||||
assert.True(t, IsValidSubType(Test, TestReply))
|
||||
assert.False(t, IsValidSubType(Test, 2))
|
||||
|
||||
// These types only ever carry subtype 0.
|
||||
for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} {
|
||||
assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt)
|
||||
assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt)
|
||||
}
|
||||
|
||||
// Unknown/unassigned types are never valid.
|
||||
assert.False(t, IsValidSubType(99, 0))
|
||||
|
||||
// Exhaustive proof of equivalence with the original map-driven logic across
|
||||
// the entire (type, subtype) input space.
|
||||
for ti := 0; ti <= 0xff; ti++ {
|
||||
for si := 0; si <= 0xff; si++ {
|
||||
mt, mst := MessageType(ti), MessageSubType(si)
|
||||
assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst),
|
||||
"IsValidSubType(%d, %d) diverged from map-driven definition", ti, si)
|
||||
}
|
||||
}
|
||||
|
||||
// H method must delegate to the package function.
|
||||
assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType())
|
||||
assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType())
|
||||
}
|
||||
|
||||
func TestHeader_String(t *testing.T) {
|
||||
assert.Equal(
|
||||
t,
|
||||
|
||||
+11
@@ -543,6 +543,17 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||
return final
|
||||
}
|
||||
|
||||
func (hm *HostMap) QueryIndexCached(index uint32, cache map[uint32]*HostInfo) *HostInfo {
|
||||
if out, ok := cache[index]; ok {
|
||||
return out
|
||||
}
|
||||
out := hm.QueryIndex(index)
|
||||
if out != nil {
|
||||
cache[index] = out
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||
hm.RLock()
|
||||
if h, ok := hm.Indexes[index]; ok {
|
||||
|
||||
@@ -2,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
|
||||
@@ -9,10 +10,24 @@ import (
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||
// only valid until the next Read on that queue. Every consumer below
|
||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||
// synchronously; do not retain pkt outside this call. If a future
|
||||
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
||||
// the borrow.
|
||||
//
|
||||
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
|
||||
// superpacket. In both cases the L3+L4 headers at the start describe
|
||||
// the same 5-tuple every segment will share, so a single newPacket /
|
||||
// firewall check covers the whole superpacket.
|
||||
packet := pkt.Bytes
|
||||
err := newPacket(packet, false, fwPacket)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
@@ -37,7 +52,14 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||
// TUN device.
|
||||
if immediatelyForwardToSelf {
|
||||
_, err := f.readers[q].Write(packet)
|
||||
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
||||
// A self-forwarded superpacket would be re-handed to the
|
||||
// kernel as one giant blob; segment first so the loopback
|
||||
// path sees one IP datagram per Write.
|
||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||
_, werr := f.queues[q].Write(seg)
|
||||
return werr
|
||||
})
|
||||
if err != nil {
|
||||
f.l.Error("Failed to forward to tun", "error", err)
|
||||
}
|
||||
@@ -52,12 +74,24 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||
// so retaining segments past the loop is safe.
|
||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||
return nil
|
||||
})
|
||||
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||
"error", err,
|
||||
"vpnAddr", fwPacket.RemoteAddr,
|
||||
)
|
||||
}
|
||||
})
|
||||
|
||||
if hostinfo == nil {
|
||||
f.rejectInside(packet, out, q)
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||
"vpnAddr", fwPacket.RemoteAddr,
|
||||
@@ -71,12 +105,11 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
return
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
if dropReason == nil {
|
||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
||||
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||
} else {
|
||||
f.rejectInside(packet, out, q)
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
@@ -86,6 +119,125 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
ci.writeLock.Lock()
|
||||
}
|
||||
c := ci.messageCounter.Add(1)
|
||||
|
||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||
|
||||
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
ci.writeLock.Unlock()
|
||||
}
|
||||
if encErr != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||
"error", encErr,
|
||||
"udpAddr", hostinfo.GetRemote(),
|
||||
"counter", c,
|
||||
)
|
||||
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
|
||||
return nil
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
||||
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
||||
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||
// kernel-supplied superpacket bytes never get written into a separate
|
||||
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
|
||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
|
||||
ci := hostinfo.ConnectionState
|
||||
if ci.eKey == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// One traffic-out mark covers every segment of the superpacket; doing it
|
||||
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
|
||||
// times per TSO packet, inside writeLock under boring crypto.
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
remote := hostinfo.GetRemote()
|
||||
if hostinfo.lastRebindCount != f.rebindCount {
|
||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||
hostinfo.lastRebindCount = f.rebindCount
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if !remote.IsValid() { //the relay path
|
||||
//first, find our relay hostinfo:
|
||||
var relayHostInfo *HostInfo
|
||||
var relay *Relay
|
||||
var err error
|
||||
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||
if err != nil {
|
||||
hostinfo.relayState.DeleteRelay(relayIP)
|
||||
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||
"relay", relayIP,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if relayHostInfo == nil || relay == nil {
|
||||
//failure already logged
|
||||
return
|
||||
}
|
||||
|
||||
err = tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
|
||||
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
|
||||
|
||||
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
|
||||
if innerPacket == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
//now we need to do a relay-encrypt:
|
||||
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
|
||||
if err != nil {
|
||||
//already logged
|
||||
return nil
|
||||
}
|
||||
|
||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
||||
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
||||
|
||||
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
||||
if out == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
sendBatch.Commit(out, remote)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
if !f.firewall.OutboundSendReject {
|
||||
return
|
||||
@@ -96,33 +248,36 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
return
|
||||
}
|
||||
|
||||
_, err := f.readers[q].Write(out)
|
||||
_, err := f.queues[q].Write(out)
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
|
||||
if !f.firewall.InboundSendReject {
|
||||
return
|
||||
}
|
||||
|
||||
out = iputil.CreateRejectPacket(packet, out)
|
||||
// split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything
|
||||
// we can't re-use packet, if we need to send an icmp reject, it won't be long enough.
|
||||
half := len(rejectBuf) / 2
|
||||
encryptBuf := rejectBuf[0:0:half] //the first half of rejectBuf's capacity, len set to 0
|
||||
buildBuf := rejectBuf[half:]
|
||||
|
||||
out := iputil.CreateRejectPacket(packet, buildBuf)
|
||||
if len(out) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if len(out) > iputil.MaxRejectPacketSize {
|
||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||
f.l.Info("rejectOutside: packet too big, not sending",
|
||||
"packet", packet,
|
||||
"outPacket", out,
|
||||
)
|
||||
f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, q)
|
||||
}
|
||||
|
||||
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
||||
@@ -216,7 +371,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
||||
}
|
||||
|
||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||
fp := &firewall.Packet{}
|
||||
fp := &firewall.ParsedPacket{}
|
||||
err := newPacket(p, false, fp)
|
||||
if err != nil {
|
||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||
@@ -224,7 +379,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
||||
}
|
||||
|
||||
// check if packet is in outbound fw rules
|
||||
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
if dropReason != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping cached packet",
|
||||
@@ -275,21 +430,13 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||
}
|
||||
|
||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||
// to the payload for the ultimate target host, making this a useful method for sending
|
||||
// handshake messages to peers through relay tunnels.
|
||||
// via is the HostInfo through which the message is relayed.
|
||||
// ad is the plaintext data to authenticate, but not encrypt
|
||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||
// out is a buffer used to store the result of the Encrypt operation
|
||||
// q indicates which writer to use to send the packet.
|
||||
func (f *Interface) SendVia(via *HostInfo,
|
||||
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
) {
|
||||
) ([]byte, error) {
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||
via.ConnectionState.writeLock.Lock()
|
||||
@@ -311,7 +458,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
||||
"headerLen", len(out),
|
||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||
)
|
||||
return
|
||||
return nil, io.ErrShortBuffer
|
||||
}
|
||||
|
||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||
@@ -331,13 +478,31 @@ func (f *Interface) SendVia(via *HostInfo,
|
||||
}
|
||||
if err != nil {
|
||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||
// to the payload for the ultimate target host, making this a useful method for sending
|
||||
// handshake messages to peers through relay tunnels.
|
||||
// via is the HostInfo through which the message is relayed.
|
||||
// ad is the plaintext data to authenticate, but not encrypt
|
||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||
// out is a buffer used to store the result of the Encrypt operation
|
||||
// q indicates which writer to use to send the packet.
|
||||
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||
if err != nil {
|
||||
// already logged by prepareSendVia
|
||||
return
|
||||
}
|
||||
err = f.writers[0].WriteTo(out, via.GetRemote())
|
||||
|
||||
err = f.writers[q].WriteTo(toSend, via.GetRemote())
|
||||
if err != nil {
|
||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||
}
|
||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||
}
|
||||
|
||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||
@@ -408,7 +573,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", remote,
|
||||
"udpAddr", hr,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
@@ -423,7 +588,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
)
|
||||
continue
|
||||
}
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
+165
-44
@@ -4,9 +4,9 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"slices"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -14,12 +14,15 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/util"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
@@ -49,7 +52,19 @@ type InterfaceConfig struct {
|
||||
reQueryWait time.Duration
|
||||
|
||||
ConntrackCacheTimeout time.Duration
|
||||
l *slog.Logger
|
||||
|
||||
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
|
||||
// shorter lists than `routines` cycle. Empty list keeps the default
|
||||
// pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true.
|
||||
CpuAffinity []int
|
||||
// PinThreads controls whether each TUN reader OS thread is pinned to a
|
||||
// single CPU (via tun.pin_threads, default true). Pinning keeps each
|
||||
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
|
||||
// packets stay ordered on the wire.
|
||||
PinThreads bool
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type Interface struct {
|
||||
@@ -73,7 +88,16 @@ type Interface struct {
|
||||
routines int
|
||||
disconnectInvalid atomic.Bool
|
||||
closed atomic.Bool
|
||||
relayManager *relayManager
|
||||
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
|
||||
// Empty falls back to the default pin-to-(allowed CPU) behavior.
|
||||
// Only consulted when pinThreads is true.
|
||||
cpuAffinity []int
|
||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||
// left free to migrate as on stock nebula.
|
||||
pinThreads bool
|
||||
relayManager *relayManager
|
||||
|
||||
tryPromoteEvery atomic.Uint32
|
||||
reQueryEvery atomic.Uint32
|
||||
@@ -90,8 +114,14 @@ type Interface struct {
|
||||
|
||||
ctx context.Context
|
||||
writers []udp.Conn
|
||||
readers []io.ReadWriteCloser
|
||||
wg sync.WaitGroup
|
||||
queues []tio.Queue
|
||||
// batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
|
||||
// commits plaintext into the batcher; the plaintext is decrypted
|
||||
// in place inside the UDP receive buffers, so listenOut must call Flush
|
||||
// at the end of each UDP recvmmsg batch, before those buffers are
|
||||
// reused (every udp.Conn ListenOut guarantees that ordering).
|
||||
batchers []*batch.MultiCoalescer
|
||||
wg sync.WaitGroup
|
||||
|
||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||
// nil means "no fatal error" (yet)
|
||||
@@ -102,18 +132,13 @@ type Interface struct {
|
||||
metricHandshakes metrics.Histogram
|
||||
messageMetrics *MessageMetrics
|
||||
cachedPacketMetrics *cachedPacketMetrics
|
||||
metricTxDropped metrics.Counter
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type EncWriter interface {
|
||||
SendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
)
|
||||
SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
|
||||
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
||||
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
||||
Handshake(vpnAddr netip.Addr)
|
||||
@@ -172,6 +197,10 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
return nil, errors.New("no connection manager")
|
||||
}
|
||||
|
||||
if c.routines <= 1 {
|
||||
c.PinThreads = false //pinning is not useful unless there's more than one tun reader
|
||||
}
|
||||
|
||||
cs := c.pki.getCertState()
|
||||
ifce := &Interface{
|
||||
ctx: ctx,
|
||||
@@ -189,7 +218,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
routines: c.routines,
|
||||
version: c.version,
|
||||
writers: make([]udp.Conn, c.routines),
|
||||
readers: make([]io.ReadWriteCloser, c.routines),
|
||||
batchers: make([]*batch.MultiCoalescer, c.routines),
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrs: cs.myVpnAddrs,
|
||||
@@ -198,8 +227,11 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
relayManager: c.relayManager,
|
||||
connectionManager: c.connectionManager,
|
||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||
cpuAffinity: c.CpuAffinity,
|
||||
pinThreads: c.PinThreads,
|
||||
|
||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
|
||||
messageMetrics: c.MessageMetrics,
|
||||
cachedPacketMetrics: &cachedPacketMetrics{
|
||||
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
||||
@@ -240,25 +272,36 @@ func (f *Interface) activate() error {
|
||||
"boringcrypto", boringEnabled(),
|
||||
)
|
||||
|
||||
if f.routines > 1 {
|
||||
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
||||
f.routines = 1
|
||||
f.l.Warn("routines is not supported on this platform, falling back to a single routine")
|
||||
}
|
||||
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
||||
f.routines = 1
|
||||
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
|
||||
}
|
||||
|
||||
// Prepare the tun queues. A device that can't open that many hands back
|
||||
// fewer (a single queue on platforms without multiqueue support) and we
|
||||
// size the reader routines to what we actually got.
|
||||
queues, err := f.inside.Queues(f.routines)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(queues) < f.routines {
|
||||
// TODO: this clamp is only safe because it is unreachable when the
|
||||
// udp side has multiple readers (linux Queues opens exactly n or
|
||||
// errors; every other platform already clamped routines to 1 above).
|
||||
// If a platform ever returns fewer queues than routines with
|
||||
// SO_REUSEPORT sockets already bound, the surplus sockets get no
|
||||
// listenOut and the kernel blackholes every flow it hashes to them —
|
||||
// fail loudly or close the extra sockets instead.
|
||||
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
|
||||
"requested", f.routines, "opened", len(queues))
|
||||
f.routines = len(queues)
|
||||
}
|
||||
f.queues = queues
|
||||
|
||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||
|
||||
// Prepare n tun queues
|
||||
var reader io.ReadWriteCloser = f.inside
|
||||
for i := 0; i < f.routines; i++ {
|
||||
if i > 0 {
|
||||
reader, err = f.inside.NewMultiQueueReader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
f.readers[i] = reader
|
||||
for i := range f.queues {
|
||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
|
||||
}
|
||||
|
||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||
@@ -281,7 +324,7 @@ func (f *Interface) run() {
|
||||
// Launch n queues to read packets from tun dev
|
||||
for i := 0; i < f.routines; i++ {
|
||||
f.wg.Go(func() {
|
||||
f.listenIn(f.readers[i], i)
|
||||
f.listenIn(f.queues[i], i)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -306,6 +349,31 @@ func (f *Interface) onFatal(err error) {
|
||||
}
|
||||
}
|
||||
|
||||
type rxContext struct {
|
||||
q int
|
||||
scratch []byte
|
||||
// nb is a re-usable nonce buffer for decrypt calls to use
|
||||
nb []byte
|
||||
h *header.H
|
||||
fwPacket *firewall.ParsedPacket
|
||||
hostmapCache map[uint32]*HostInfo
|
||||
lhh *LightHouseHandler
|
||||
ctCache *firewall.ConntrackCacheTicker
|
||||
}
|
||||
|
||||
func newRxContext(f *Interface, q int) *rxContext {
|
||||
return &rxContext{
|
||||
q: q,
|
||||
scratch: make([]byte, mtu),
|
||||
nb: make([]byte, 12, 12),
|
||||
h: &header.H{},
|
||||
fwPacket: &firewall.ParsedPacket{},
|
||||
hostmapCache: map[uint32]*HostInfo{},
|
||||
lhh: f.lightHouse.NewRequestHandler(),
|
||||
ctCache: firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout),
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) listenOut(i int) {
|
||||
var li udp.Conn
|
||||
if i > 0 {
|
||||
@@ -314,16 +382,20 @@ func (f *Interface) listenOut(i int) {
|
||||
li = f.outside
|
||||
}
|
||||
|
||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
lhh := f.lightHouse.NewRequestHandler()
|
||||
plaintext := make([]byte, udp.MTU)
|
||||
h := &header.H{}
|
||||
fwPacket := &firewall.Packet{}
|
||||
nb := make([]byte, 12, 12)
|
||||
rxc := newRxContext(f, i)
|
||||
|
||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||
})
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
|
||||
}
|
||||
|
||||
flusher := func() {
|
||||
if err := f.batchers[i].Flush(); err != nil {
|
||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||
}
|
||||
clear(rxc.hostmapCache)
|
||||
}
|
||||
|
||||
err := li.ListenOut(listener, flusher)
|
||||
|
||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||
@@ -336,16 +408,42 @@ func (f *Interface) listenOut(i int) {
|
||||
f.l.Debug("underlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||
packet := make([]byte, mtu)
|
||||
out := make([]byte, mtu)
|
||||
fwPacket := &firewall.Packet{}
|
||||
func (f *Interface) pinThisThread(i int) {
|
||||
var cpu int
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||
// validated the entries against the allowed CPU set.
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||
// Default: spread queues across the CPUs we're actually allowed to
|
||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||
cpu = allowed[i%len(allowed)]
|
||||
} else {
|
||||
cpu = i % runtime.NumCPU()
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
||||
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
|
||||
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
||||
if f.pinThreads {
|
||||
f.pinThisThread(i)
|
||||
}
|
||||
|
||||
rejectBuf := make([]byte, mtu)
|
||||
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||
fwPacket := &firewall.ParsedPacket{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
|
||||
for {
|
||||
n, err := reader.Read(packet)
|
||||
pkts, err := queue.Read()
|
||||
if err != nil {
|
||||
// Same shutdown noise handling as listenOut
|
||||
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||
@@ -355,12 +453,35 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||
break
|
||||
}
|
||||
|
||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
||||
for _, pkt := range pkts {
|
||||
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||
// Flush incrementally once a full sendmmsg batch has
|
||||
// accumulated so the first packets of a deep read drain
|
||||
// hit the wire while the rest are still being encrypted.
|
||||
if sb.Len() >= batch.SendBatchCap {
|
||||
f.flushSendBatch(sb, i)
|
||||
}
|
||||
}
|
||||
f.flushSendBatch(sb, i)
|
||||
}
|
||||
|
||||
f.l.Debug("overlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
// flushSendBatch drains sb to the underlay and accounts for anything it could not deliver. A shortfall means
|
||||
// specific destinations were undeliverable (a stale remote, a reject rule), which the backend logs per peer at
|
||||
// debug; here it is only a counter, so one unreachable peer cannot spam a log line per batch.
|
||||
func (f *Interface) flushSendBatch(sb *batch.SendBatch, q int) {
|
||||
queued := sb.Len()
|
||||
written, err := sb.Flush()
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||
}
|
||||
if dropped := queued - written; dropped > 0 {
|
||||
f.metricTxDropped.Inc(int64(dropped))
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||
c.RegisterReloadCallback(f.reloadFirewall)
|
||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||
|
||||
+27
-3
@@ -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,11 +333,34 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
||||
return out
|
||||
}
|
||||
|
||||
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||
// 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
|
||||
}
|
||||
nextHeader = packet[6]
|
||||
offset = ipv6.HeaderLen
|
||||
|
||||
for {
|
||||
for range maxIPv6ExtHeaders {
|
||||
switch nextHeader {
|
||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||
if len(packet) < offset+2 {
|
||||
@@ -367,6 +390,7 @@ func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm
|
||||
return nextHeader, offset, isFragment
|
||||
}
|
||||
}
|
||||
return nextHeader, offset, isFragment
|
||||
}
|
||||
|
||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package iputil
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"testing"
|
||||
@@ -179,6 +180,46 @@ func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
|
||||
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
|
||||
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
|
||||
src := net.ParseIP("fd00::1")
|
||||
dst := net.ParseIP("fd00::2")
|
||||
|
||||
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
|
||||
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
|
||||
// than the inner packet length.
|
||||
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
|
||||
|
||||
// The ciphertext scratch reused as the reject buffer is the received
|
||||
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
|
||||
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
|
||||
const nebulaOverhead = 32
|
||||
segLen := len(inner) + nebulaOverhead
|
||||
|
||||
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
|
||||
const neighborHdr = 16
|
||||
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
|
||||
|
||||
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
|
||||
// the overrun that silently drops the neighbor packet.
|
||||
backing := make([]byte, segLen+neighborHdr)
|
||||
copy(backing[segLen:], sentinel)
|
||||
reject := CreateRejectPacket(inner, backing[:segLen])
|
||||
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
|
||||
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||
"without the cap the oversized reject overruns into the neighbor segment")
|
||||
|
||||
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
|
||||
// reject does not fit, so it is refused rather than corrupting the neighbor.
|
||||
backing = make([]byte, segLen+neighborHdr)
|
||||
copy(backing[segLen:], sentinel)
|
||||
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
|
||||
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
|
||||
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||
"capped segment must leave the neighbor untouched")
|
||||
}
|
||||
|
||||
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||
b[0] = ipv6.Version << 4
|
||||
@@ -474,3 +515,121 @@ 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)
|
||||
})
|
||||
}
|
||||
|
||||
+9
-1
@@ -36,6 +36,10 @@ 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
|
||||
@@ -107,6 +111,10 @@ 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{})
|
||||
@@ -918,7 +926,7 @@ func (lh *LightHouse) SendUpdate() {
|
||||
}
|
||||
|
||||
lal := lh.GetLocalAllowList()
|
||||
for _, e := range localAddrs(lh.l, lal) {
|
||||
for _, e := range lh.localAddrsFn(lal) {
|
||||
if lh.myVpnNetworksTable.Contains(e) {
|
||||
continue
|
||||
}
|
||||
|
||||
+1
-1
@@ -498,7 +498,7 @@ type testEncWriter struct {
|
||||
protocolVersion cert.Version
|
||||
}
|
||||
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
}
|
||||
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||
}
|
||||
|
||||
@@ -6,11 +6,14 @@ import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime/debug"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/cpupick"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -33,6 +36,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
buildVersion = moduleVersion()
|
||||
}
|
||||
|
||||
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise.
|
||||
startPprofServer(ctx, l)
|
||||
|
||||
// Print the config if in test, the exit comes later
|
||||
if configTest {
|
||||
b, err := yaml.Marshal(c.Settings)
|
||||
@@ -161,7 +167,13 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
for i := 0; i < routines; i++ {
|
||||
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
||||
}
|
||||
@@ -210,6 +222,21 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
l.Warn("Failed to start DNS responder", "error", err)
|
||||
}
|
||||
|
||||
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
|
||||
// The operator didn't choose pin CPUs, so pick a default set that
|
||||
// prefers performance cores and doesn't stack co-located instances
|
||||
// onto allowed[0]. The bound UDP port keys the per-instance spread:
|
||||
// distinct across instances sharing a box, stable across restarts.
|
||||
// A nil result keeps listenIn's stock allowed[i] fallback.
|
||||
key := uint64(os.Getpid())
|
||||
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
|
||||
key = uint64(ap.Port())
|
||||
}
|
||||
cpuAffinity = cpupick.Default(routines, key, l)
|
||||
}
|
||||
|
||||
ifConfig := &InterfaceConfig{
|
||||
HostMap: hostMap,
|
||||
Inside: tun,
|
||||
@@ -231,6 +258,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||
punchy: punchy,
|
||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||
CpuAffinity: cpuAffinity,
|
||||
PinThreads: pinThreads,
|
||||
l: l,
|
||||
}
|
||||
|
||||
@@ -268,6 +297,8 @@ 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,
|
||||
@@ -278,10 +309,75 @@ 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
|
||||
}
|
||||
|
||||
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
|
||||
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
|
||||
// (listenIn falls back to spreading queues across the allowed CPU set).
|
||||
// Length mismatch with `routines` is a warning, not an error: shorter lists
|
||||
// are modulo-cycled across queues, longer lists' tail is ignored. Invalid
|
||||
// entries (non-integer, or a CPU ID we're not allowed to run on) are also a
|
||||
// warning and disable the override entirely so we don't silently pin to the
|
||||
// wrong CPU. Entries are validated against the process's current affinity
|
||||
// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or
|
||||
// taskset the runnable IDs are frequently not that contiguous range, and
|
||||
// pinning to an unrunnable ID always fails. If the allowed set can't be
|
||||
// determined we fall back to a plain non-negative check.
|
||||
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
||||
raw := c.Get("tun.cpu_affinity")
|
||||
if raw == nil {
|
||||
return nil
|
||||
}
|
||||
rv, ok := raw.([]any)
|
||||
if !ok {
|
||||
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
|
||||
return nil
|
||||
}
|
||||
// allowed is the set of CPU IDs we're actually permitted to run on. A nil
|
||||
// slice (unsupported platform or lookup error) means "can't tell", so we
|
||||
// only apply the weaker non-negative check in that case.
|
||||
allowed, err := util.AllowedCPUs()
|
||||
if err != nil {
|
||||
l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err)
|
||||
allowed = nil
|
||||
}
|
||||
cpus := make([]int, 0, len(rv))
|
||||
for i, e := range rv {
|
||||
var cpu int
|
||||
switch v := e.(type) {
|
||||
case int:
|
||||
cpu = v
|
||||
case int64:
|
||||
cpu = int(v)
|
||||
case float64:
|
||||
cpu = int(v)
|
||||
default:
|
||||
l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity",
|
||||
"index", i, "value", e)
|
||||
return nil
|
||||
}
|
||||
if cpu < 0 {
|
||||
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
|
||||
"index", i, "cpu", cpu)
|
||||
return nil
|
||||
}
|
||||
if len(allowed) > 0 && !slices.Contains(allowed, cpu) {
|
||||
l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity",
|
||||
"index", i, "cpu", cpu, "allowed", allowed)
|
||||
return nil
|
||||
}
|
||||
cpus = append(cpus, cpu)
|
||||
}
|
||||
if len(cpus) != routines {
|
||||
l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list",
|
||||
"affinity_len", len(cpus), "routines", routines)
|
||||
}
|
||||
return cpus
|
||||
}
|
||||
|
||||
func moduleVersion() string {
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestParseCpuAffinity(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
|
||||
// newConfig returns a config.C with tun.cpu_affinity set to v. A nil v
|
||||
// leaves the key unset.
|
||||
newConfig := func(v any) *config.C {
|
||||
c := config.NewC(l)
|
||||
if v != nil {
|
||||
c.Settings["tun"] = map[string]any{"cpu_affinity": v}
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// unset -> nil (listenIn falls back to spreading across the allowed set)
|
||||
assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1))
|
||||
|
||||
// Pick a CPU we're actually allowed to run on so a valid list survives
|
||||
// validation regardless of the host's affinity mask.
|
||||
allowed, _ := util.AllowedCPUs()
|
||||
validCPU := 0
|
||||
if len(allowed) > 0 {
|
||||
validCPU = allowed[0]
|
||||
}
|
||||
|
||||
// valid list -> parsed through unchanged
|
||||
assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2))
|
||||
|
||||
// a negative entry is out of range on every platform -> disables the override
|
||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2))
|
||||
|
||||
// a non-integer entry -> disables the override
|
||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2))
|
||||
|
||||
// a CPU id outside the allowed set -> disables the override. Only assertable
|
||||
// where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond
|
||||
// any representable CPU id so it can never be in the mask.
|
||||
if len(allowed) > 0 {
|
||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1))
|
||||
}
|
||||
}
|
||||
@@ -164,3 +164,48 @@ func TestCipherStateNilSafety(t *testing.T) {
|
||||
assert.Empty(t, out)
|
||||
assert.Equal(t, 0, cc.Overhead())
|
||||
}
|
||||
|
||||
func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
}
|
||||
|
||||
func inPlaceDecrypt(t *testing.T, enc, dec CipherState) {
|
||||
t.Helper()
|
||||
const hdrLen = 16
|
||||
plaintext := []byte("in-place decrypt should replace the ciphertext bytes")
|
||||
nb := make([]byte, 12)
|
||||
|
||||
// packet = [16-byte header | ciphertext+tag], like a nebula Message.
|
||||
packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead())
|
||||
for i := range packet {
|
||||
packet[i] = byte(i)
|
||||
}
|
||||
packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Simulate a GRO row: [packet | next segment]. A failed auth on packet
|
||||
// may zero packet's plaintext region but must not touch the header, the
|
||||
// tag, or the neighboring segment.
|
||||
neighbor := []byte("next coalesced segment, must stay intact")
|
||||
row := append(append([]byte(nil), packet...), neighbor...)
|
||||
tampered := row[:len(packet)]
|
||||
tampered[hdrLen] ^= 0x01
|
||||
_, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb)
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header")
|
||||
assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():],
|
||||
"failed auth must not touch the tag")
|
||||
assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment")
|
||||
|
||||
out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, plaintext, out)
|
||||
// The plaintext must be IN the packet buffer, not a fresh allocation.
|
||||
assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer")
|
||||
}
|
||||
|
||||
+56
-34
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"golang.org/x/net/ipv4"
|
||||
)
|
||||
|
||||
@@ -22,7 +23,11 @@ const (
|
||||
|
||||
var ErrOutOfWindow = errors.New("out of window packet")
|
||||
|
||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// readOutsidePackets processes one received underlay packet.
|
||||
// Message payloads are decrypted IN PLACE, so packet must stay untouched
|
||||
// by the caller until the batcher for queue q has been flushed
|
||||
func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxContext) {
|
||||
h := rxc.h
|
||||
err := h.Parse(packet)
|
||||
if err != nil {
|
||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||
@@ -90,7 +95,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
if isMessageRelay {
|
||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||
} else {
|
||||
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||
hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
|
||||
}
|
||||
|
||||
// At this point we should have a valid existing tunnel, verify and send
|
||||
@@ -113,17 +118,18 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
// All remaining packets are encrypted
|
||||
if isMessageRelay {
|
||||
// Relay packets are special, this branch should always early-return
|
||||
if err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, nb); err != nil {
|
||||
err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
||||
return
|
||||
}
|
||||
|
||||
out, err = hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, out, packet, nb)
|
||||
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||
@@ -139,7 +145,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
case header.Message:
|
||||
switch h.Subtype {
|
||||
case header.MessageNone:
|
||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
||||
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -147,15 +153,23 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
|
||||
case header.LightHouse:
|
||||
//TODO: assert via is not relayed
|
||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
|
||||
case header.Test:
|
||||
switch h.Subtype {
|
||||
case header.TestReply:
|
||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||
case header.TestRequest:
|
||||
//recycle the input packet ciphertext as our output buffer
|
||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, nb, packet)
|
||||
const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
|
||||
const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead
|
||||
if maxOverhead+len(out) > len(rxc.scratch) {
|
||||
// A reply that cannot fit in scratch is dropped no matter the log level.
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping oversized test request", "payloadLen", len(out), "from", via)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, rxc.nb, rxc.scratch[:0])
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -173,7 +187,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
|
||||
h := rxc.h
|
||||
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
@@ -186,9 +201,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if !ok {
|
||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||
// its internal mapping. This should never happen.
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||
"relayRemoteIndex", h.RemoteIndex,
|
||||
)
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -202,7 +215,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||
f.readOutsidePackets(via, signedPayload, rxc)
|
||||
case ForwardingType:
|
||||
// Find the target HostInfo relay object
|
||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||
@@ -221,8 +234,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||
fwdBuf := packet[:0:len(packet)] // Cap to len(packet) to protect memory from a larger parent buffer
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true)
|
||||
fwdBuf := packet[:0]
|
||||
//todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
return
|
||||
@@ -303,7 +317,11 @@ var (
|
||||
)
|
||||
|
||||
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// fp is reused across packets; reset the parse byproducts so an early-error return cannot
|
||||
// leak the previous packet's offsets.
|
||||
fp.IPHdrLen = 0
|
||||
fp.FragAny = false
|
||||
if len(data) < 1 {
|
||||
return ErrPacketTooShort
|
||||
}
|
||||
@@ -318,7 +336,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return ErrUnknownIPVersion
|
||||
}
|
||||
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
dataLen := len(data)
|
||||
if dataLen < ipv6.HeaderLen {
|
||||
return ErrIPv6PacketTooShort
|
||||
@@ -344,6 +362,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
switch proto {
|
||||
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
||||
fp.Protocol = uint8(proto)
|
||||
fp.IPHdrLen = offset
|
||||
fp.RemotePort = 0
|
||||
fp.LocalPort = 0
|
||||
fp.Fragment = false
|
||||
@@ -354,6 +373,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) 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 {
|
||||
@@ -371,6 +391,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) 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])
|
||||
@@ -388,6 +411,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) 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 {
|
||||
@@ -429,7 +455,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return ErrIPv6CouldNotFindPayload
|
||||
}
|
||||
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// Do we at least have an ipv4 header worth of data?
|
||||
if len(data) < ipv4.HeaderLen {
|
||||
return ErrIPv4PacketTooShort
|
||||
@@ -446,6 +472,10 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
// Check if this is the second or further fragment of a fragmented packet.
|
||||
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
||||
// Any fragmentation at all (MF or offset): first fragments have readable ports for the
|
||||
// firewall but must never be coalesced.
|
||||
fp.FragAny = (flagsfrags & 0x3fff) != 0
|
||||
fp.IPHdrLen = ihl
|
||||
|
||||
// Firewall handles protocol checks
|
||||
fp.Protocol = data[9]
|
||||
@@ -489,31 +519,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||
err := newPacket(out, true, fwPacket)
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
||||
err := newPacket(out, true, rxc.fwPacket)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||
"error", err,
|
||||
"packet", out,
|
||||
)
|
||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
|
||||
return
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
|
||||
if dropReason != nil {
|
||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
||||
// This gives us a buffer to build the reject packet in
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
"reason", dropReason,
|
||||
)
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
_, err = f.readers[q].Write(out)
|
||||
err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
|
||||
+90
-5
@@ -17,7 +17,7 @@ import (
|
||||
)
|
||||
|
||||
func Test_newPacket(t *testing.T) {
|
||||
p := &firewall.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
// 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.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
// 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.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
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.Packet{}
|
||||
fp := &firewall.ParsedPacket{}
|
||||
|
||||
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.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
const (
|
||||
hdrLen = 40 // IPv6 header
|
||||
@@ -675,3 +675,88 @@ 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)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math/rand"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4
|
||||
// checksum field is pre-loaded with the folded (not inverted) pseudo-header
|
||||
// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed
|
||||
// produces packets every receiver silently drops, with nothing failing on
|
||||
// our side — so these tests check the helpers against an independent
|
||||
// RFC 1071 reference built from explicit pseudo-header bytes, never against
|
||||
// the production checksum code.
|
||||
|
||||
// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded)
|
||||
// into a wide one's-complement accumulator.
|
||||
func refSum(b []byte) uint64 {
|
||||
var s uint64
|
||||
for i := 0; i+1 < len(b); i += 2 {
|
||||
s += uint64(b[i])<<8 | uint64(b[i+1])
|
||||
}
|
||||
if len(b)%2 == 1 {
|
||||
s += uint64(b[len(b)-1]) << 8
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// refFold folds a wide one's-complement accumulator to 16 bits.
|
||||
func refFold(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
func TestFoldOnceNoInvertEdgeCases(t *testing.T) {
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // fold yields 0x1fffd: needs a second fold
|
||||
0xffffffff, // worst case
|
||||
0x00010001, // simple two-word
|
||||
}
|
||||
for _, c := range cases {
|
||||
want := refFold(uint64(c))
|
||||
if got := foldOnceNoInvert(c); got != want {
|
||||
t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
// Folding a folded value must be a no-op.
|
||||
if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) {
|
||||
t.Errorf("foldOnceNoInvert not idempotent at %#x", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv4MatchesReference(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [4]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20},
|
||||
{"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0},
|
||||
{"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535},
|
||||
{"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535},
|
||||
{"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2).
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], c.src[:])
|
||||
copy(ph[4:8], c.dst[:])
|
||||
ph[9] = c.proto
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len))
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv6MatchesReference(t *testing.T) {
|
||||
ones := func(b byte) (a [16]byte) {
|
||||
for i := range a {
|
||||
a[i] = b
|
||||
}
|
||||
return
|
||||
}
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [16]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20},
|
||||
{"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0},
|
||||
{"max-u16-len", ones(0xff), ones(0xfe), 6, 65535},
|
||||
{"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1).
|
||||
ph := make([]byte, 40)
|
||||
copy(ph[0:16], c.src[:])
|
||||
copy(ph[16:32], c.dst[:])
|
||||
binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len))
|
||||
ph[39] = c.proto
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1791))
|
||||
for _, hdrLen := range []int{20, 24, 40, 60} {
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
hdr := make([]byte, hdrLen)
|
||||
rng.Read(hdr)
|
||||
hdr[0] = 0x40 | byte(hdrLen/4)
|
||||
hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires
|
||||
|
||||
want := ^refFold(refSum(hdr))
|
||||
got := ipv4HdrChecksum(hdr)
|
||||
if got != want {
|
||||
t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want)
|
||||
}
|
||||
|
||||
// Receiver-side property: with the checksum stored, the full
|
||||
// header must sum to all-ones.
|
||||
binary.BigEndian.PutUint16(hdr[10:12], got)
|
||||
if v := refFold(refSum(hdr)); v != 0xffff {
|
||||
t.Fatalf("stored checksum does not validate: full-header fold = %#x", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers
|
||||
// exist for: seed the TCP checksum field with fold(pseudoSum), do what the
|
||||
// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4
|
||||
// bytes including the seed, then invert, then store), and verify the result
|
||||
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
|
||||
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1826))
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
payLen := rng.Intn(1500)
|
||||
l4 := make([]byte, 20+payLen)
|
||||
rng.Read(l4)
|
||||
|
||||
// Seed exactly as flushSlot does.
|
||||
seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4)))
|
||||
binary.BigEndian.PutUint16(l4[16:18], seed)
|
||||
|
||||
// Kernel NEEDS_CSUM completion: sum the L4 region (seed included,
|
||||
// which is equivalent to summing with the field zeroed and folding
|
||||
// the seed in), invert, store.
|
||||
final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:]))
|
||||
binary.BigEndian.PutUint16(l4[16:18], final)
|
||||
|
||||
// Receiver validation.
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], src[:])
|
||||
copy(ph[4:8], dst[:])
|
||||
ph[9] = 6
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4)))
|
||||
if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff {
|
||||
t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)",
|
||||
trial, v, seed, final, payLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
)
|
||||
|
||||
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
||||
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
||||
type flowKey struct {
|
||||
src, dst [16]byte
|
||||
sport, dport uint16
|
||||
isV6 bool
|
||||
}
|
||||
|
||||
// initialSlots is the starting capacity of the slot pool.
|
||||
// One flow per packet is the worst case,
|
||||
// so this matches a typical carrier-side recvmmsg batch on the UDP socket.
|
||||
const initialSlots = 64
|
||||
|
||||
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
|
||||
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
|
||||
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
|
||||
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at
|
||||
// byte 40.
|
||||
//
|
||||
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
|
||||
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
|
||||
// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these
|
||||
// functions are too big to inline, and returning structs by value put five 64-byte copies on the
|
||||
// per-packet path.
|
||||
func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) {
|
||||
if len(pkt) < 20 {
|
||||
return nil, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
if ipHdrLen != 20 {
|
||||
return nil, false
|
||||
}
|
||||
return fk.parseIPv4Prologue(pkt)
|
||||
case 6:
|
||||
if ipHdrLen != 40 || len(pkt) < 40 {
|
||||
return nil, false
|
||||
}
|
||||
return fk.parseIPv6Prologue(pkt)
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
|
||||
// len(pkt) >= 20 and the version.
|
||||
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl != 20 {
|
||||
return nil, false
|
||||
}
|
||||
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
|
||||
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
|
||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||
return nil, false
|
||||
}
|
||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||
if totalLen > len(pkt) || totalLen < ihl {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = false
|
||||
copy(fk.src[:4], pkt[12:16])
|
||||
copy(fk.dst[:4], pkt[16:20])
|
||||
return pkt[:totalLen], true
|
||||
}
|
||||
|
||||
// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version,
|
||||
// and that the L4 header sits at byte 40.
|
||||
func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) {
|
||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||
if 40+payloadLen > len(pkt) {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = true
|
||||
copy(fk.src[:], pkt[8:24])
|
||||
copy(fk.dst[:], pkt[24:40])
|
||||
return pkt[:40+payloadLen], true
|
||||
}
|
||||
|
||||
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||
// byte-for-byte equality on every field that must be identical across coalesced segments.
|
||||
// Size/IPID/IPCsum are masked out.
|
||||
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
|
||||
// segments with differing ECN codepoints must not coalesce,
|
||||
// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable.
|
||||
//
|
||||
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||
if isV6 {
|
||||
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
|
||||
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
|
||||
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
|
||||
}
|
||||
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
|
||||
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
|
||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
|
||||
}
|
||||
|
||||
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
|
||||
const ipv4FlagDF = 0x40
|
||||
|
||||
// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
|
||||
// nextHdr may join a chain whose seed header is seedHdr as segment index seg
|
||||
// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as
|
||||
// seed_id+n, so coalescing is only transparent when that re-stamp is either
|
||||
// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or
|
||||
// reproduces the original IDs exactly (DF clear + IDs already sequential —
|
||||
// the same admission rule kernel GRO applies). Without this, a DF=0 sender
|
||||
// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs
|
||||
// rewritten into ranges that collide across superpackets, corrupting
|
||||
// reassembly if the packets are fragmented after the TUN write.
|
||||
//
|
||||
// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6
|
||||
// is inside its compared range), so checking the seed's copy suffices.
|
||||
func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool {
|
||||
if seedHdr[6]&ipv4FlagDF != 0 {
|
||||
return true
|
||||
}
|
||||
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
|
||||
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
|
||||
}
|
||||
|
||||
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||
// slices via Reserve and releases them in bulk via Reset.
|
||||
type Arena struct {
|
||||
buf []byte
|
||||
}
|
||||
|
||||
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
|
||||
func NewArena(capacity int) *Arena {
|
||||
return &Arena{buf: make([]byte, 0, capacity)}
|
||||
}
|
||||
|
||||
// Reserve hands out a non-overlapping sz-byte slice from the arena.
|
||||
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
|
||||
// Already-borrowed slices reference the old backing and remain valid until Reset.
|
||||
func (a *Arena) Reserve(sz int) []byte {
|
||||
if len(a.buf)+sz > cap(a.buf) {
|
||||
newCap := max(cap(a.buf)*2, sz)
|
||||
a.buf = make([]byte, 0, newCap)
|
||||
}
|
||||
start := len(a.buf)
|
||||
a.buf = a.buf[:start+sz]
|
||||
return a.buf[start : start+sz : start+sz]
|
||||
}
|
||||
|
||||
// Reset releases every slice handed out since the last Reset.
|
||||
// Callers must not use any previously-borrowed slice after this returns.
|
||||
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
||||
func (a *Arena) Reset() {
|
||||
a.buf = a.buf[:0]
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks
|
||||
// bypass staging and the sort entirely.
|
||||
func stagePackets(pkts [][]byte) []stagedPacket {
|
||||
staged := make([]stagedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pp := testPP(p)
|
||||
staged[i] = stagedPacket{
|
||||
pkt: p,
|
||||
key: SortKey{Epoch: 1, Counter: uint64(i + 1)},
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
}
|
||||
}
|
||||
return staged
|
||||
}
|
||||
|
||||
func flushLanes(b *testing.B, m *MultiCoalescer) {
|
||||
b.Helper()
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if m.udp != nil {
|
||||
if err := m.udp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the
|
||||
// batcher, which is where the production profile concentrates.
|
||||
func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
staged := stagePackets(pkts)
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if err := m.dispatch(staged[i%len(staged)]); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
flushLanes(b, m)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
flushLanes(b, m)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends.
|
||||
func BenchmarkDispatchSingleFlow(b *testing.B) {
|
||||
runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the
|
||||
// lastSlot cache on every packet.
|
||||
func BenchmarkDispatchInterleaved16(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a
|
||||
// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction).
|
||||
func BenchmarkDispatchAckHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
var pkts [][]byte
|
||||
seq := uint32(1000)
|
||||
for range tcpCoalesceMaxSegs / 2 {
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, pay))
|
||||
seq += uint32(len(pay))
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, nil))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape.
|
||||
func BenchmarkDispatchUDPFlow(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, udpCoalesceMaxSegs)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(2000, 443, pay)
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes
|
||||
// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including
|
||||
// the parsedTCP-to-slot field transfer) can cost.
|
||||
func BenchmarkDispatchSeedHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, tcpCoalesceMaxSegs)
|
||||
seq := uint32(1000)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildTCPv4(seq, tcpAckPsh, pay)
|
||||
seq += uint32(len(pay))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package batch
|
||||
|
||||
//TODO refactor this away
|
||||
// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind
|
||||
// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At
|
||||
// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top
|
||||
// of a local parse, so tests and benches can drive one lane with nothing but a packet.
|
||||
|
||||
// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP,
|
||||
// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4
|
||||
// offset; fk must be zero on entry and is filled in place.
|
||||
func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) {
|
||||
if len(pkt) < 20 {
|
||||
return nil, 0, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
if pkt[9] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv4Prologue(pkt)
|
||||
return trimmed, 20, ok
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return nil, 0, false
|
||||
}
|
||||
if pkt[6] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv6Prologue(pkt)
|
||||
return trimmed, 40, ok
|
||||
}
|
||||
return nil, 0, false
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for
|
||||
// coalescing or not. Returns false for non-TCP or malformed input.
|
||||
func (p *parsedTCP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||
func (p *parsedUDP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedTCP
|
||||
if !info.parseBase(pkt) {
|
||||
// Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it.
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedUDP
|
||||
if !info.parseBase(pkt) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
)
|
||||
|
||||
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
|
||||
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
|
||||
//
|
||||
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
|
||||
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
|
||||
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
|
||||
// lanes carry no reorder-repair machinery.
|
||||
//
|
||||
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
|
||||
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
|
||||
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
|
||||
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
|
||||
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
|
||||
// to the later-flushed pt lane.
|
||||
//
|
||||
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
|
||||
type MultiCoalescer struct {
|
||||
tcp *TCPCoalescer
|
||||
udp *UDPCoalescer
|
||||
pt *Passthrough
|
||||
|
||||
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
||||
// each pkt alive until Flush returns.
|
||||
staged []stagedPacket
|
||||
}
|
||||
|
||||
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
|
||||
// value: pp is reused by the caller per packet and must not be retained past Commit.
|
||||
type stagedPacket struct {
|
||||
pkt []byte
|
||||
key SortKey
|
||||
proto byte
|
||||
fragAny bool
|
||||
ipHdrLen uint16
|
||||
}
|
||||
|
||||
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
|
||||
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
|
||||
// transmission-order repair.
|
||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
|
||||
m := &MultiCoalescer{
|
||||
pt: NewPassthrough(w),
|
||||
staged: make([]stagedPacket, 0, initialSlots),
|
||||
}
|
||||
m.tcp = NewTCPCoalescer(w, l)
|
||||
m.udp = NewUDPCoalescer(w)
|
||||
return m
|
||||
}
|
||||
|
||||
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
|
||||
// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed:
|
||||
// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a
|
||||
// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only
|
||||
// for this call, so the fields dispatch needs are copied here.
|
||||
func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error {
|
||||
m.staged = append(m.staged, stagedPacket{
|
||||
pkt: pkt,
|
||||
key: key,
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// compareStaged orders staged packets by (epoch, counter)
|
||||
func compareStaged(a, b stagedPacket) int {
|
||||
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
|
||||
return c
|
||||
}
|
||||
return cmp.Compare(a.key.Counter, b.key.Counter)
|
||||
}
|
||||
|
||||
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
|
||||
// passthrough when the lane has no GSO support.
|
||||
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
|
||||
switch sp.proto {
|
||||
case ipProtoTCP:
|
||||
if m.tcp != nil {
|
||||
return m.tcp.commitStaged(sp)
|
||||
}
|
||||
case ipProtoUDP:
|
||||
if m.udp != nil {
|
||||
return m.udp.commitStaged(sp)
|
||||
}
|
||||
}
|
||||
return m.pt.enqueue(sp.pkt)
|
||||
}
|
||||
|
||||
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
|
||||
// Drains everything and returns the joined errors; one bad packet does not hold up the rest.
|
||||
// After Flush returns, committed payload slices may be recycled.
|
||||
func (m *MultiCoalescer) Flush() error {
|
||||
// Arrival order is already almost sorted (reorder is the exception), which pdqsort detects
|
||||
// and handles in near-linear time.
|
||||
slices.SortFunc(m.staged, compareStaged)
|
||||
|
||||
var errs []error
|
||||
for _, sp := range m.staged {
|
||||
if err := m.dispatch(sp); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
clear(m.staged) // drop borrowed pkt refs
|
||||
m.staged = m.staged[:0]
|
||||
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
if m.udp != nil {
|
||||
if err := m.udp.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,437 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// keySeq hands out SortKeys with ascending counters in a fixed epoch, for
|
||||
// tests where commit order IS transmission order.
|
||||
type keySeq struct {
|
||||
epoch, counter uint64
|
||||
}
|
||||
|
||||
func (k *keySeq) next() SortKey {
|
||||
k.counter++
|
||||
return SortKey{Epoch: k.epoch, Counter: k.counter}
|
||||
}
|
||||
|
||||
// newTestMultiCoalescer builds a batcher over w.
|
||||
func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer {
|
||||
tb.Helper()
|
||||
return NewMultiCoalescer(w, test.NewLogger())
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||
// else (ICMP here) falls through to plain Write.
|
||||
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
tcpPay := make([]byte, 1200)
|
||||
udpPay := make([]byte, 1200)
|
||||
icmp := make([]byte, 28)
|
||||
icmp[0] = 0x45
|
||||
icmp[2] = 0
|
||||
icmp[3] = 28
|
||||
icmp[9] = 1
|
||||
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
|
||||
// property: packets committed out of counter order (wire reorder inside one
|
||||
// flush batch) are replayed into the lanes in transmission order, so the
|
||||
// reorder never fragments the coalesce chain — one superpacket, in seq
|
||||
// order, exactly as if the wire had never reordered. The retransmit shape
|
||||
// falls out of the same key: a retransmit carries a lower seq but a HIGHER
|
||||
// counter (it was encrypted later), so it emits after the data it trails.
|
||||
func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3).
|
||||
// Arrival order: 3400, 1000, 2200.
|
||||
if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||
t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if len(g.pays) != 3 {
|
||||
t.Fatalf("segs=%d want 3", len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000", seedSeq)
|
||||
}
|
||||
|
||||
// Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3).
|
||||
w.writes, w.gsoWrites, w.order = nil, nil, nil
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
first := binary.BigEndian.Uint32(w.writes[0][24:28])
|
||||
second := binary.BigEndian.Uint32(w.writes[1][24:28])
|
||||
if first != 4600 || second != 1000 {
|
||||
t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows;
|
||||
// the staging sort must repair each flow into one superpacket without any
|
||||
// cross-flow contamination.
|
||||
func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4).
|
||||
// Arrival: A.1300, B.1700, A.100, B.500.
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 2 {
|
||||
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
|
||||
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
|
||||
switch sport {
|
||||
case 1000:
|
||||
if seedSeq != 100 {
|
||||
t.Errorf("flow A seed seq=%d want 100", seedSeq)
|
||||
}
|
||||
case 3000:
|
||||
if seedSeq != 500 {
|
||||
t.Errorf("flow B seed seq=%d want 500", seedSeq)
|
||||
}
|
||||
default:
|
||||
t.Errorf("unexpected sport %d", sport)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces
|
||||
// the tunnel, and the replacement's counter space starts near zero — raw
|
||||
// counter order would emit the new tunnel's packets first while the old
|
||||
// tunnel's backlog is still arriving. The epoch key must dominate:
|
||||
// everything from the old tunnel emits before anything from the new one.
|
||||
func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// New session's first data arrives before the old session's last data.
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Same flow, contiguous seq, identical headers: after the epoch sort the
|
||||
// two segments append into one superpacket seeded by the OLD session's
|
||||
// packet.
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO
|
||||
// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP
|
||||
// packets still reach the kernel via verbatim rather than being lost.
|
||||
func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noUSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.udp != nil {
|
||||
t.Fatal("UDP lane must not come up without USO")
|
||||
}
|
||||
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
|
||||
// anything. Both lane constructors refuse, so every packet rides the
|
||||
// verbatim lane — but the staging sort still applies, so emission follows
|
||||
// transmission order even without GSO.
|
||||
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: false}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
if m.tcp != nil || m.udp != nil {
|
||||
t.Fatal("no lane may come up without offloads")
|
||||
}
|
||||
pkts := [][]byte{
|
||||
buildTCPv4(1000, tcpAck, make([]byte, 1200)),
|
||||
buildUDPv4(1000, 53, make([]byte, 800)),
|
||||
buildTCPv4(2200, tcpAck, make([]byte, 1200)),
|
||||
}
|
||||
// Committed in reverse transmission order; keys carry the truth.
|
||||
for i := len(pkts) - 1; i >= 0; i-- {
|
||||
if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != len(pkts) {
|
||||
t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes))
|
||||
}
|
||||
// One lane for everything means the sorted order survives end to end.
|
||||
for i, want := range pkts {
|
||||
if !bytes.Equal(w.writes[i], want) {
|
||||
t.Errorf("write %d out of order or corrupt", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a
|
||||
// single fragment header (NH=44) naming UDP as the terminal protocol —
|
||||
// a first fragment (offset 0, MF set) carrying the UDP header and a
|
||||
// partial payload.
|
||||
func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 40
|
||||
const fragHdrLen = 8
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen))
|
||||
pkt[6] = 44 // fragment extension header
|
||||
pkt[7] = 64
|
||||
pkt[8] = 0xfe
|
||||
pkt[9] = 0x80
|
||||
pkt[23] = 1
|
||||
pkt[24] = 0xfe
|
||||
pkt[25] = 0x80
|
||||
pkt[39] = 2
|
||||
|
||||
pkt[40] = ipProtoUDP // fragment's next header
|
||||
binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set
|
||||
binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[48:50], sport)
|
||||
binary.BigEndian.PutUint16(pkt[50:52], dport)
|
||||
binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload)))
|
||||
copy(pkt[56:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header
|
||||
// routing: a fragment whose chain terminates in UDP must ride the UDP lane
|
||||
// as an in-lane verbatim — emitted ahead of later same-flow datagrams —
|
||||
// not the verbatim lane, which flushes after every coalescer lane and
|
||||
// would reorder it behind data that arrived after it.
|
||||
func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites))
|
||||
}
|
||||
// Transmission order was fragment-then-data; same-lane routing must keep it.
|
||||
if w.order[0] != "write" {
|
||||
t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram
|
||||
// (fragment) seals every open UDP chain, so datagrams from before and after
|
||||
// it land in separate superpackets and the fragment holds its transmission-
|
||||
// order position between them.
|
||||
func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
want := []string{"gso", "write", "gso"}
|
||||
if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] {
|
||||
t.Fatalf("emission order = %v, want %v", w.order, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case.
|
||||
func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noTSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.tcp != nil {
|
||||
t.Fatal("TCP lane must not come up without TSO")
|
||||
}
|
||||
|
||||
pay := make([]byte, 1200)
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// testPP derives the ParsedPacket newPacket would produce for the packet
|
||||
// shapes the tests build: plain v4/v6, v4 with options or fragment bits set,
|
||||
// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything
|
||||
// unrecognizable stays zero (proto 0 routes to the passthrough lane).
|
||||
func testPP(pkt []byte) *firewall.ParsedPacket {
|
||||
pp := &firewall.ParsedPacket{}
|
||||
if len(pkt) < 20 {
|
||||
return pp
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
pp.Protocol = pkt[9]
|
||||
pp.IPHdrLen = int(pkt[0]&0x0f) * 4
|
||||
pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0
|
||||
case 6:
|
||||
pp.Protocol = pkt[6]
|
||||
pp.IPHdrLen = 40
|
||||
if pp.Protocol == 44 { // fragment extension header
|
||||
pp.Protocol = pkt[40]
|
||||
pp.IPHdrLen = 48
|
||||
pp.FragAny = true
|
||||
}
|
||||
}
|
||||
return pp
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
|
||||
// order enqueued.
|
||||
type Passthrough struct {
|
||||
out io.Writer
|
||||
slots [][]byte
|
||||
}
|
||||
|
||||
func NewPassthrough(w io.Writer) *Passthrough {
|
||||
return &Passthrough{
|
||||
out: w,
|
||||
slots: make([][]byte, 0, 128),
|
||||
}
|
||||
}
|
||||
|
||||
// enqueue accepts one packet, already sorted into transmission order by dispatch.
|
||||
func (p *Passthrough) enqueue(pkt []byte) error {
|
||||
p.slots = append(p.slots, pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Passthrough) Flush() error {
|
||||
var firstErr error
|
||||
for _, s := range p.slots {
|
||||
_, err := p.out.Write(s)
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
clear(p.slots)
|
||||
p.slots = p.slots[:0]
|
||||
return firstErr
|
||||
}
|
||||
@@ -0,0 +1,472 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
|
||||
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out.
|
||||
const ipProtoTCP = 6
|
||||
|
||||
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
||||
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
||||
const tcpCoalesceBufSize = 65535
|
||||
|
||||
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||
const tcpCoalesceMaxSegs = 64
|
||||
|
||||
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
|
||||
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
|
||||
// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the
|
||||
// caller's plaintext buffers; the caller must keep them alive until Flush.
|
||||
type coalesceSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce
|
||||
// slots. A slot that never grows past one segment is emitted from rawPkt so its original
|
||||
// (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
|
||||
fk flowKey
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int
|
||||
numSeg int
|
||||
totalPay int
|
||||
nextSeq uint32
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
|
||||
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
|
||||
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
|
||||
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
|
||||
// commitParsed. Owns no locks; one coalescer per TUN write queue.
|
||||
type TCPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
|
||||
// slots is the ordered event queue. Flush walks it once and emits each
|
||||
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
||||
slots []*coalesceSlot
|
||||
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
|
||||
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
|
||||
// non-admissible packet for the flow, or in Flush.
|
||||
openSlots map[flowKey]*coalesceSlot
|
||||
// lastSlot caches the most recently touched open slot. Bulk traffic
|
||||
// arrives in same-flow runs (single-flow steady state, or GRO bursts
|
||||
// under multi-flow), so comparing the incoming key against the cached
|
||||
// slot's own fk lets the hot path skip the map lookup (and the aeshash
|
||||
// of a 38-byte key) for the length of each run.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||
// at is removed.
|
||||
lastSlot *coalesceSlot
|
||||
pool []*coalesceSlot // free list for reuse
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
|
||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &TCPCoalescer{
|
||||
w: gw,
|
||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||
l: l,
|
||||
}
|
||||
}
|
||||
|
||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedTCP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
|
||||
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
|
||||
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+20 {
|
||||
return false
|
||||
}
|
||||
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
|
||||
if tcpOff < 20 || tcpOff > 60 {
|
||||
return false
|
||||
}
|
||||
if len(pkt) < ipHdrLen+tcpOff {
|
||||
return false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + tcpOff
|
||||
p.payLen = len(pkt) - p.hdrLen
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8])
|
||||
p.flags = pkt[ipHdrLen+13]
|
||||
return true
|
||||
}
|
||||
|
||||
// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named;
|
||||
// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed.
|
||||
const (
|
||||
tcpFlagPsh = 0x08
|
||||
tcpFlagAck = 0x10
|
||||
tcpFlagEce = 0x40
|
||||
)
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *TCPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues).
|
||||
func (c *TCPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedTCP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
|
||||
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
|
||||
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
|
||||
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
|
||||
if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||
// SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later
|
||||
// in-flow packets cannot extend it and emit ahead of this verbatim.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
if info.payLen == 0 {
|
||||
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
|
||||
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
|
||||
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
|
||||
// kernel GRO. This is the only place emission deviates from transmission order.
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
|
||||
// many flows: wire-side GRO delivers runs of same-flow packets
|
||||
// (deliverSegments splits a superdatagram into up to 64), so the cache
|
||||
// hits for the length of each run and a miss costs one fk compare
|
||||
// before the map lookup carries the weight.
|
||||
var open *coalesceSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (PSH or short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend (seq gap from upstream loss, header change, or a full
|
||||
// chain): evict it from openSlots and fall through to seed a fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) Flush() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to its seed packet; ship the original so
|
||||
// its valid checksum rides the DATA_VALID path instead of a kernel software csum.
|
||||
// rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is
|
||||
// pristine here.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
if err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
c.release(s)
|
||||
}
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
|
||||
return first
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
||||
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
s.fk = info.fk
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
if info.flags&tcpFlagPsh == 0 {
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
} else {
|
||||
// PSH on the seed closes the chain immediately; it is never registered as open.
|
||||
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
|
||||
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
|
||||
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
|
||||
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
|
||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
if info.seq != s.nextSeq {
|
||||
return false
|
||||
}
|
||||
if s.numSeg >= tcpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if info.payLen > s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// ECE state must be stable across a burst.
|
||||
// Receivers expect the flag set on every segment of a CE-echoing window or none.
|
||||
seedFlags := s.rawPkt[s.ipHdrLen+13]
|
||||
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||
return false
|
||||
}
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: the
|
||||
// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH.
|
||||
// The caller must deregister a closed slot from openSlots.
|
||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
if info.flags&tcpFlagPsh != 0 {
|
||||
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
|
||||
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
|
||||
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||
if n := len(c.pool); n > 0 {
|
||||
s := c.pool[n-1]
|
||||
c.pool[n-1] = nil
|
||||
c.pool = c.pool[:n-1]
|
||||
return s
|
||||
}
|
||||
return &coalesceSlot{}
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||
clear(s.payIovs)
|
||||
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header
|
||||
// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after,
|
||||
// so nothing re-reads the patched header. Does not remove the slot from c.slots.
|
||||
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||
total := s.hdrLen + s.totalPay
|
||||
l4Len := total - s.ipHdrLen
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||
hdr[10] = 0
|
||||
hdr[11] = 0
|
||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||
}
|
||||
|
||||
var psum uint32
|
||||
if s.isV6 {
|
||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
||||
} else {
|
||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
||||
}
|
||||
tcsum := s.ipHdrLen + 16
|
||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||
}
|
||||
|
||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||
// equality on every field that must be identical across coalesced
|
||||
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
||||
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
||||
// [18:tcpHdrLen] options (incl. urgent).
|
||||
tcp := ipHdrLen
|
||||
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||
// already have its checksum field zeroed) and returns the folded/inverted
|
||||
// 16-bit value to store.
|
||||
func ipv4HdrChecksum(hdr []byte) uint16 {
|
||||
var sum uint32
|
||||
for i := 0; i+1 < len(hdr); i += 2 {
|
||||
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
||||
}
|
||||
if len(hdr)%2 == 1 {
|
||||
sum += uint32(hdr[len(hdr)-1]) << 8
|
||||
}
|
||||
for sum>>16 != 0 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
}
|
||||
return ^uint16(sum)
|
||||
}
|
||||
|
||||
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
||||
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
||||
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
||||
// reuses these helpers.
|
||||
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
var sum uint32
|
||||
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
||||
sum += uint32(proto)
|
||||
sum += uint32(l4Len)
|
||||
return sum
|
||||
}
|
||||
|
||||
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
var sum uint32
|
||||
for i := 0; i < 16; i += 2 {
|
||||
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
||||
}
|
||||
sum += uint32(l4Len >> 16)
|
||||
sum += uint32(l4Len & 0xffff)
|
||||
sum += uint32(proto)
|
||||
return sum
|
||||
}
|
||||
|
||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
|
||||
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
|
||||
func foldOnceNoInvert(sum uint32) uint16 {
|
||||
for sum>>16 != 0 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
}
|
||||
return uint16(sum)
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
|
||||
// everything but satisfies the interface the coalescer detects.
|
||||
type nopTunWriter struct{}
|
||||
|
||||
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
|
||||
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
||||
return nil
|
||||
}
|
||||
func (nopTunWriter) Capabilities() tio.Capabilities {
|
||||
return tio.Capabilities{TSO: true, USO: true}
|
||||
}
|
||||
|
||||
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
|
||||
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
|
||||
// contiguous so every packet is coalesceable onto the previous one.
|
||||
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||
pkts := make([][]byte, n)
|
||||
pay := make([]byte, payloadLen)
|
||||
seq := uint32(1000)
|
||||
for i := range n {
|
||||
pkts[i] = buildTCPv4(seq, tcpAck, pay)
|
||||
seq += uint32(payloadLen)
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
|
||||
// seq continuity but round-robin across flows — worst case for any
|
||||
// "last-slot" cache.
|
||||
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
seqs := make([]uint32, nFlows)
|
||||
for i := range seqs {
|
||||
seqs[i] = uint32(1000 + i*1000000)
|
||||
}
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for range perFlow {
|
||||
for f := range nFlows {
|
||||
sport := uint16(10000 + f)
|
||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||
seqs[f] += uint32(payloadLen)
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in
|
||||
// runs of runLen per flow — the arrival pattern wire-side GRO actually
|
||||
// produces (deliverSegments splits each superdatagram into up to 64
|
||||
// same-flow packets back to back). Contrast with buildTCPv4Interleaved's
|
||||
// per-packet round-robin, the adversarial worst case for a last-slot cache.
|
||||
func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
seqs := make([]uint32, nFlows)
|
||||
for i := range seqs {
|
||||
seqs[i] = uint32(1000 + i*1000000)
|
||||
}
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(10000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||
seqs[f] += uint32(payloadLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim
|
||||
// branch in Commit.
|
||||
func buildICMPv4() []byte {
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||
pkt[9] = 1 // ICMP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
return pkt
|
||||
}
|
||||
|
||||
// runCommitBench drives Commit over pkts batchSize at a time, flushing
|
||||
// between batches, and reports per-packet cost.
|
||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestTCPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := c.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
// Drain any trailing partial batch so slot state doesn't leak across runs.
|
||||
_ = c.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
|
||||
// contiguous seq, 1200-byte payloads. Every packet past the seed should
|
||||
// append onto the open slot. This is the case we most care about.
|
||||
func BenchmarkCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
|
||||
// A single-entry fast-path cache will miss on every packet; an N-way
|
||||
// cache or map lookup carries the weight.
|
||||
func BenchmarkCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
|
||||
func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
|
||||
// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot
|
||||
// cache hits for the length of each run; the per-packet round-robin
|
||||
// benches above are its worst case.
|
||||
func BenchmarkCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase
|
||||
// bails early and addVerbatim is the only work.
|
||||
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||
pkt := buildICMPv4()
|
||||
pkts := make([][]byte, 64)
|
||||
for i := range pkts {
|
||||
pkts[i] = pkt
|
||||
}
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||
// Each packet takes the "TCP but not admissible" branch which does a
|
||||
// map delete + verbatim. Measures the seal-without-slot cost.
|
||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
pay := make([]byte, 0)
|
||||
pkts := make([][]byte, 64)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
|
||||
}
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
|
||||
// it includes the staging sort's already-sorted fast path plus the
|
||||
// dispatch-time parse — the full steady-state cost of the batcher. The
|
||||
// ParsedPackets are precomputed: in production they fall out of the
|
||||
// firewall's newPacket, which this bench does not model.
|
||||
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
pps := make([]*firewall.ParsedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pps[i] = testPP(p)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
j := i % len(pkts)
|
||||
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := m.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = m.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
|
||||
// BenchmarkCommitSingleFlow — same workload but routed through the
|
||||
// dispatcher. The delta vs the single-lane bench measures dispatcher
|
||||
// overhead.
|
||||
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
|
||||
// through the dispatcher.
|
||||
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,59 @@
|
||||
package batch
|
||||
|
||||
import "net/netip"
|
||||
|
||||
const SendBatchCap = 128
|
||||
|
||||
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||
type batchWriter interface {
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
}
|
||||
|
||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||
// Slots are backed by an Arena (see its docs)
|
||||
type SendBatch struct {
|
||||
out batchWriter
|
||||
bufs [][]byte
|
||||
dsts []netip.AddrPort
|
||||
arena *Arena
|
||||
}
|
||||
|
||||
// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots
|
||||
func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
||||
return &SendBatch{
|
||||
out: out,
|
||||
bufs: make([][]byte, 0, batchCap),
|
||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||
arena: NewArena(arenaSize),
|
||||
}
|
||||
}
|
||||
|
||||
func (b *SendBatch) Reserve(sz int) []byte {
|
||||
return b.arena.Reserve(sz)
|
||||
}
|
||||
|
||||
// Len reports how many packets are queued for the next Flush. Callers use
|
||||
// it to flush incrementally once a full sendmmsg batch has accumulated,
|
||||
// bounding how long the first packet of a large read batch waits.
|
||||
func (b *SendBatch) Len() int { return len(b.bufs) }
|
||||
|
||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
||||
b.bufs = append(b.bufs, pkt)
|
||||
b.dsts = append(b.dsts, dst)
|
||||
}
|
||||
|
||||
// Flush writes every queued packet and reports how many actually went out. A short count means some destinations
|
||||
// were undeliverable; the batch is drained either way.
|
||||
func (b *SendBatch) Flush() (int, error) {
|
||||
var err error
|
||||
written := 0
|
||||
if len(b.bufs) > 0 {
|
||||
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
||||
}
|
||||
clear(b.bufs)
|
||||
b.bufs = b.bufs[:0]
|
||||
b.dsts = b.dsts[:0]
|
||||
b.arena.Reset()
|
||||
return written, err
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fakeBatchWriter struct {
|
||||
bufs [][]byte
|
||||
addrs []netip.AddrPort
|
||||
}
|
||||
|
||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||
// returns, so tests must capture data before that happens.
|
||||
w.bufs = make([][]byte, len(bufs))
|
||||
for i, b := range bufs {
|
||||
cp := make([]byte, len(b))
|
||||
copy(cp, b)
|
||||
w.bufs[i] = cp
|
||||
}
|
||||
w.addrs = append(w.addrs[:0], addrs...)
|
||||
return len(bufs), nil
|
||||
}
|
||||
|
||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
b := NewSendBatch(fw, 4, 32)
|
||||
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||
for i := 0; i < 4; i++ {
|
||||
slot := b.Reserve(32)
|
||||
if cap(slot) != 32 {
|
||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||
}
|
||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 4 {
|
||||
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
||||
}
|
||||
for i, buf := range fw.bufs {
|
||||
if len(buf) != 3 || buf[0] != byte(i) {
|
||||
t.Errorf("buf %d: %x", i, buf)
|
||||
}
|
||||
if fw.addrs[i] != ap {
|
||||
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
||||
}
|
||||
}
|
||||
|
||||
// Flush again with nothing committed — should be a no-op.
|
||||
fw.bufs = nil
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("empty Flush: %v", err)
|
||||
}
|
||||
if fw.bufs != nil {
|
||||
t.Fatalf("empty Flush triggered WriteBatch")
|
||||
}
|
||||
|
||||
// Reuse after Flush.
|
||||
slot := b.Reserve(32)
|
||||
if cap(slot) != 32 {
|
||||
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
b := NewSendBatch(fw, 3, 8)
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
s := b.Reserve(8)
|
||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
|
||||
for i, buf := range fw.bufs {
|
||||
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
// Tiny initial backing forces a grow on the second Reserve.
|
||||
b := NewSendBatch(fw, 1, 4)
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||
|
||||
s1 := b.Reserve(4)
|
||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||
b.Commit(pkt1, ap)
|
||||
|
||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||
b.Commit(pkt2, ap)
|
||||
|
||||
// pkt1 must still be intact even though backing reallocated.
|
||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||
}
|
||||
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 2 {
|
||||
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
||||
}
|
||||
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
||||
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
||||
}
|
||||
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
||||
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
|
||||
// ipProtoUDP is the IANA protocol number for UDP.
|
||||
const ipProtoUDP = 17
|
||||
|
||||
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
|
||||
// kernel's gso_max_size; payloads beyond this are emitted as-is.
|
||||
const udpCoalesceBufSize = 65535
|
||||
|
||||
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
|
||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||
const udpCoalesceMaxSegs = 64
|
||||
|
||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||
type udpSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed
|
||||
// packet for coalesce slots. A coalesce slot that never grows past one
|
||||
// segment is emitted from rawPkt so its original (already valid) L4
|
||||
// checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
|
||||
fk flowKey
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int // per-segment UDP payload length
|
||||
numSeg int
|
||||
totalPay int
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
|
||||
// Preserves the in-flow order of packets as they are Commit-ed
|
||||
//
|
||||
// Owns no locks; one coalescer per TUN write queue.
|
||||
type UDPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
slots []*udpSlot
|
||||
openSlots map[flowKey]*udpSlot
|
||||
// lastSlot caches the most recently touched open slot; see the
|
||||
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
|
||||
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
|
||||
// the fk compare beats the map's 38-byte key hash on most packets.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
|
||||
// is removed.
|
||||
lastSlot *udpSlot
|
||||
pool []*udpSlot
|
||||
}
|
||||
|
||||
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &UDPCoalescer{
|
||||
w: gw,
|
||||
slots: make([]*udpSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||
pool: make([]*udpSlot, 0, initialSlots),
|
||||
}
|
||||
}
|
||||
|
||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedUDP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int // ipHdrLen + 8
|
||||
payLen int
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
|
||||
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place. Returns false for malformed input or any shape that must not
|
||||
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+8 {
|
||||
return false
|
||||
}
|
||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
|
||||
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
|
||||
return false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + 8
|
||||
p.payLen = udpLen - 8
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
return true
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open.
|
||||
func (c *UDPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedUDP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
|
||||
// A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be
|
||||
// coalesced.
|
||||
if info.payLen == 0 {
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
// Cached-slot fast path; see the TCPCoalescer equivalent.
|
||||
var open *udpSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend: evict it from openSlots and fall through to seed a
|
||||
// fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) Flush() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to the packet it was
|
||||
// seeded from; ship the original so its valid checksum rides the
|
||||
// DATA_VALID path instead of paying a kernel software csum.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
if err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
c.release(s)
|
||||
}
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
return first
|
||||
}
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *UDPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) addVerbatim(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
||||
if info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
s.fk = info.fk
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed.
|
||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
if s.numSeg >= udpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if info.payLen > s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches
|
||||
// here; closing removes the slot from openSlots, the only path in.
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel
|
||||
// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be
|
||||
// the final one. The caller must deregister a closed slot from openSlots.
|
||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
return info.payLen < s.gsoSize
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) take() *udpSlot {
|
||||
if n := len(c.pool); n > 0 {
|
||||
s := c.pool[n-1]
|
||||
c.pool[n-1] = nil
|
||||
c.pool = c.pool[:n-1]
|
||||
return s
|
||||
}
|
||||
return &udpSlot{}
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||
// Reset every field, identity ones included; see TCPCoalescer.release.
|
||||
clear(s.payIovs)
|
||||
*s = udpSlot{payIovs: s.payIovs[:0]}
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
||||
// the UDP length to the *total* across all coalesced segments, then seeds
|
||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
|
||||
// slot is released right after, so nothing re-reads the patched header.
|
||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||
hdr[10] = 0
|
||||
hdr[11] = 0
|
||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||
}
|
||||
|
||||
// UDP length field (offset 4 inside the UDP header) = total UDP size.
|
||||
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
|
||||
|
||||
var psum uint32
|
||||
if s.isV6 {
|
||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
|
||||
} else {
|
||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
|
||||
}
|
||||
udpCsumOff := s.ipHdrLen + 6
|
||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||
}
|
||||
|
||||
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||
// every field that must be identical across coalesced segments
|
||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
|
||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
||||
udp := ipHdrLen
|
||||
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the
|
||||
// steady state for single-flow QUIC bulk, the workload USO exists for.
|
||||
func buildUDPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, n)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(40000, 443, pay)
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow
|
||||
// datagrams arriving in GRO-burst runs of runLen per flow.
|
||||
func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(40000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildUDPv4(sport, 443, pay))
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a
|
||||
// time, flushing between batches, and reports per-packet cost.
|
||||
func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestUDPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := c.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = c.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state.
|
||||
func BenchmarkUDPCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200)
|
||||
runUDPCommitBench(b, pkts, udpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin.
|
||||
func BenchmarkUDPCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16.
|
||||
func BenchmarkUDPCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -0,0 +1,536 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
|
||||
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 20
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x45
|
||||
pkt[1] = 0x00
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
||||
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
||||
pkt[8] = 64
|
||||
pkt[9] = ipProtoUDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
||||
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
|
||||
binary.BigEndian.PutUint16(pkt[26:28], 0)
|
||||
|
||||
copy(pkt[28:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// buildUDPv6 builds a minimal IPv6+UDP packet.
|
||||
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 40
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
|
||||
pkt[6] = ipProtoUDP
|
||||
pkt[7] = 64
|
||||
pkt[8] = 0xfe
|
||||
pkt[9] = 0x80
|
||||
pkt[23] = 1
|
||||
pkt[24] = 0xfe
|
||||
pkt[25] = 0x80
|
||||
pkt[39] = 2
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[40:42], sport)
|
||||
binary.BigEndian.PutUint16(pkt[42:44], dport)
|
||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
|
||||
binary.BigEndian.PutUint16(pkt[46:48], 0)
|
||||
|
||||
copy(pkt[48:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
|
||||
// do USO. See newTestTCPCoalescer.
|
||||
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
|
||||
tb.Helper()
|
||||
c := NewUDPCoalescer(w)
|
||||
if c == nil {
|
||||
tb.Fatal("NewUDPCoalescer: writer does not support USO")
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition:
|
||||
// no USO, no coalescer.
|
||||
func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) {
|
||||
if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil {
|
||||
t.Fatalf("want nil for a non-USO writer, got %v", c)
|
||||
}
|
||||
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
|
||||
t.Fatalf("want nil for a plain writer, got %v", c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
// ICMP packet
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||
pkt[9] = 1
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A slot that never grew past one datagram flushes as a plain Write of
|
||||
// the original packet bytes: the original (already valid) checksum
|
||||
// ships via the DATA_VALID path, so the kernel does no csum work.
|
||||
// WriteGSO is reserved for slots that actually coalesced (>=2 segs).
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if !bytes.Equal(w.writes[0], pkt) {
|
||||
t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if g.gsoSize != 1200 {
|
||||
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
||||
}
|
||||
if len(g.pays) != 3 {
|
||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
||||
}
|
||||
if g.csumStart != 20 {
|
||||
t.Errorf("csumStart=%d want 20", g.csumStart)
|
||||
}
|
||||
// IP totalLen and UDP length must be the TOTAL across all segments —
|
||||
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
|
||||
// value would silently drop everything but the first segment. Total =
|
||||
// IP(20) + UDP(8) + 3*1200 = 3628.
|
||||
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
|
||||
if gotTotalLen != 3628 {
|
||||
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
|
||||
}
|
||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
|
||||
if gotUDPLen != 8+3*1200 {
|
||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||
}
|
||||
}
|
||||
|
||||
// Last segment may be shorter, sealing the chain.
|
||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 1200)
|
||||
tail := make([]byte, 600)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A 4th packet, even same-sized, must NOT join — chain is sealed.
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
|
||||
// single-segment and flushes as a plain write of the original packet.
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 1 {
|
||||
t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
if len(w.gsoWrites[0].pays) != 3 {
|
||||
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
}
|
||||
if got, want := len(w.writes[0]), 20+8+1200; got != want {
|
||||
t.Errorf("re-seed plain write len=%d want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes in arrival order.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Different 5-tuples must not coalesce.
|
||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
||||
}
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 2 {
|
||||
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Caps at udpCoalesceMaxSegs.
|
||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 100)
|
||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
|
||||
// reseeds a new one.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
|
||||
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
|
||||
}
|
||||
if len(w.gsoWrites[1].pays) != 5 {
|
||||
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
|
||||
}
|
||||
}
|
||||
|
||||
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
||||
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
||||
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
|
||||
// reseeds again. All three stay single-segment, so each ships as a plain
|
||||
// write of its original bytes, keeping its own codepoint.
|
||||
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
pkt1[1] = 0x03 // CE
|
||||
pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again
|
||||
for _, p := range [][]byte{pkt0, pkt1, pkt2} {
|
||||
if err := c.Commit(p); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
wantECN := []byte{0x00, 0x03, 0x00}
|
||||
for i, p := range w.writes {
|
||||
if got := p[1] & 0x03; got != wantECN[i] {
|
||||
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 path: same flow, equal-sized → coalesced.
|
||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if !g.isV6 {
|
||||
t.Errorf("expected v6 write")
|
||||
}
|
||||
if g.csumStart != 40 {
|
||||
t.Errorf("csumStart=%d want 40", g.csumStart)
|
||||
}
|
||||
// IPv6 payload_len and UDP length must be TOTAL — kernel's
|
||||
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
|
||||
// 3*1200 = 3608.
|
||||
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
|
||||
if gotPlen != 8+3*1200 {
|
||||
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
|
||||
}
|
||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
|
||||
if gotUDPLen != 8+3*1200 {
|
||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||
}
|
||||
}
|
||||
|
||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
pkt1[1] = 0xb8 // EF DSCP, ECN=0
|
||||
if err := c.Commit(pkt0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(pkt1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes, no gso.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// Fragmented IPv4 must not be coalesced.
|
||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// A zero-length UDP datagram (UDP length == 8, no payload) is legal and
|
||||
// must be delivered as a plain single datagram — never coalesced. Seeding
|
||||
// it into a GSO slot stores an empty payload iovec that panics WriteGSO
|
||||
// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it
|
||||
// reach the GSO path. Regression: must not panic and must be written.
|
||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes[0]) != len(pkt) {
|
||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes[0]) != len(pkt) {
|
||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// A zero-length datagram arriving mid-flow must seal the open chain so the
|
||||
// datagram after it seeds a fresh superpacket *after* the empty one on the
|
||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The empty datagram sealed the first slot, so the trailing full packet
|
||||
// can't join it. All three emit as plain writes (the two full datagrams
|
||||
// stayed single-segment; the empty one is verbatim) in per-flow
|
||||
// arrival order: full, empty, full.
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IPv4 with options is not admissible (we require IHL=5).
|
||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF
|
||||
// clear is fine as long as the IDs already run seed+1 per datagram, so
|
||||
// kernel USO's re-stamp reproduces them.
|
||||
func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
for i := range 2 {
|
||||
pkt := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(pkt, uint16(40+i), false)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 {
|
||||
t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks
|
||||
// the chain; each datagram stays a single-segment slot and flushes as a
|
||||
// plain write that keeps its own (meaningful) ID.
|
||||
func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
p1 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p1, 40, false)
|
||||
p2 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p2, 50, false)
|
||||
|
||||
if err := c.Commit(p1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(p2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []uint16{40, 50} {
|
||||
if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want {
|
||||
t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package checksum
|
||||
|
||||
import (
|
||||
"golang.org/x/sys/cpu"
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
//go:noescape
|
||||
func checksumAVX2(buf []byte, initial uint16) uint16
|
||||
|
||||
var hasAVX2 = cpu.X86.HasAVX2
|
||||
|
||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
||||
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
||||
// falling back to gvisor's pure-Go implementation otherwise. The result
|
||||
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
if hasAVX2 {
|
||||
return checksumAVX2(buf, initial)
|
||||
}
|
||||
return gvisorchecksum.Checksum(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
#include "textflag.h"
|
||||
|
||||
// func checksumAVX2(buf []byte, initial uint16) uint16
|
||||
//
|
||||
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
||||
//
|
||||
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
||||
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
||||
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
||||
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
||||
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
||||
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
||||
// equivalent to summing as 16-bit big-endian words.
|
||||
//
|
||||
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
||||
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
||||
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
||||
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
||||
// chunk to break the dep chain. After the vector loop the lane sums are
|
||||
// horizontally reduced and merged with a scalar accumulator that handles
|
||||
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
||||
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
||||
MOVQ buf_base+0(FP), SI
|
||||
MOVQ buf_len+8(FP), CX
|
||||
MOVWQZX initial+24(FP), AX
|
||||
|
||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||
// with the rest of the accumulator. The final fold's bswap16 will undo
|
||||
// this and convert the whole result back to BE.
|
||||
XCHGB AH, AL
|
||||
|
||||
CMPQ CX, $32
|
||||
JLT scalar_tail
|
||||
|
||||
VPXOR Y4, Y4, Y4
|
||||
VPXOR Y5, Y5, Y5
|
||||
VPXOR Y6, Y6, Y6
|
||||
VPXOR Y7, Y7, Y7
|
||||
|
||||
CMPQ CX, $64
|
||||
JLT loop32
|
||||
|
||||
loop64:
|
||||
VPMOVZXDQ (SI), Y0
|
||||
VPMOVZXDQ 16(SI), Y1
|
||||
VPMOVZXDQ 32(SI), Y2
|
||||
VPMOVZXDQ 48(SI), Y3
|
||||
VPADDQ Y0, Y4, Y4
|
||||
VPADDQ Y1, Y5, Y5
|
||||
VPADDQ Y2, Y6, Y6
|
||||
VPADDQ Y3, Y7, Y7
|
||||
ADDQ $64, SI
|
||||
SUBQ $64, CX
|
||||
CMPQ CX, $64
|
||||
JGE loop64
|
||||
|
||||
loop32:
|
||||
CMPQ CX, $32
|
||||
JLT reduce_vec
|
||||
VPMOVZXDQ (SI), Y0
|
||||
VPMOVZXDQ 16(SI), Y1
|
||||
VPADDQ Y0, Y4, Y4
|
||||
VPADDQ Y1, Y5, Y5
|
||||
ADDQ $32, SI
|
||||
SUBQ $32, CX
|
||||
JMP loop32
|
||||
|
||||
reduce_vec:
|
||||
// Combine the four ymm accumulators into Y4.
|
||||
VPADDQ Y5, Y4, Y4
|
||||
VPADDQ Y7, Y6, Y6
|
||||
VPADDQ Y6, Y4, Y4
|
||||
|
||||
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
||||
VEXTRACTI128 $1, Y4, X5
|
||||
VPADDQ X5, X4, X4
|
||||
VPSHUFD $0x4e, X4, X5
|
||||
VPADDQ X5, X4, X4
|
||||
VMOVQ X4, R8
|
||||
VZEROUPPER
|
||||
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
|
||||
scalar_tail:
|
||||
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
||||
CMPQ CX, $8
|
||||
JLT tail4
|
||||
|
||||
loop8:
|
||||
ADDQ (SI), AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $8, SI
|
||||
SUBQ $8, CX
|
||||
CMPQ CX, $8
|
||||
JGE loop8
|
||||
|
||||
tail4:
|
||||
CMPQ CX, $4
|
||||
JLT tail2
|
||||
MOVL (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $4, SI
|
||||
SUBQ $4, CX
|
||||
|
||||
tail2:
|
||||
CMPQ CX, $2
|
||||
JLT tail1
|
||||
MOVWQZX (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $2, SI
|
||||
SUBQ $2, CX
|
||||
|
||||
tail1:
|
||||
TESTQ CX, CX
|
||||
JZ fold
|
||||
MOVBQZX (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
|
||||
fold:
|
||||
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
||||
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
||||
// the truncation steps absorb the single bit that may be left over
|
||||
// after each add so the next round's bound holds.
|
||||
|
||||
// 64 → 33 bits.
|
||||
MOVQ AX, R8
|
||||
SHRQ $32, R8
|
||||
MOVL AX, AX
|
||||
ADDQ R8, AX
|
||||
|
||||
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
||||
MOVQ AX, R8
|
||||
SHRQ $32, R8
|
||||
ADDQ R8, AX
|
||||
MOVL AX, AX
|
||||
|
||||
// 32 → 17 bits.
|
||||
MOVQ AX, R8
|
||||
SHRQ $16, R8
|
||||
MOVWQZX AX, AX
|
||||
ADDQ R8, AX
|
||||
|
||||
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
||||
MOVQ AX, R8
|
||||
SHRQ $16, R8
|
||||
ADDQ R8, AX
|
||||
|
||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||
// to big-endian to match the gvisor API contract.
|
||||
XCHGB AH, AL
|
||||
|
||||
MOVW AX, ret+32(FP)
|
||||
RET
|
||||
@@ -0,0 +1,12 @@
|
||||
package checksum
|
||||
|
||||
//go:noescape
|
||||
func checksumNEON(buf []byte, initial uint16) uint16
|
||||
|
||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
||||
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
||||
// armv8 so no feature check is needed.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
return checksumNEON(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
#include "textflag.h"
|
||||
|
||||
// func checksumNEON(buf []byte, initial uint16) uint16
|
||||
//
|
||||
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
||||
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
||||
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
||||
// at the very end to recover the on-wire (big-endian) result.
|
||||
//
|
||||
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
||||
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
||||
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
||||
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
||||
// updated twice per iter to break the dep chain. Tail bytes go through a
|
||||
// scalar ADCS chain seeded with the byte-swapped initial.
|
||||
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
||||
MOVD buf_base+0(FP), R0
|
||||
MOVD buf_len+8(FP), R1
|
||||
MOVHU initial+24(FP), R2
|
||||
|
||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||
// with the rest of the accumulator.
|
||||
REV16W R2, R2
|
||||
|
||||
MOVD ZR, R3 // scalar accumulator
|
||||
|
||||
CMP $32, R1
|
||||
BLT scalar_tail
|
||||
|
||||
VEOR V8.B16, V8.B16, V8.B16
|
||||
VEOR V9.B16, V9.B16, V9.B16
|
||||
VEOR V10.B16, V10.B16, V10.B16
|
||||
VEOR V11.B16, V11.B16, V11.B16
|
||||
|
||||
CMP $64, R1
|
||||
BLT loop16_init
|
||||
|
||||
loop64:
|
||||
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
||||
VUADDW V0.S2, V8.D2, V8.D2
|
||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||
VUADDW V1.S2, V10.D2, V10.D2
|
||||
VUADDW2 V1.S4, V11.D2, V11.D2
|
||||
VUADDW V2.S2, V8.D2, V8.D2
|
||||
VUADDW2 V2.S4, V9.D2, V9.D2
|
||||
VUADDW V3.S2, V10.D2, V10.D2
|
||||
VUADDW2 V3.S4, V11.D2, V11.D2
|
||||
SUB $64, R1, R1
|
||||
CMP $64, R1
|
||||
BGE loop64
|
||||
|
||||
loop16_init:
|
||||
CMP $16, R1
|
||||
BLT reduce_vec
|
||||
|
||||
loop16:
|
||||
VLD1.P 16(R0), [V0.B16]
|
||||
VUADDW V0.S2, V8.D2, V8.D2
|
||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||
SUB $16, R1, R1
|
||||
CMP $16, R1
|
||||
BGE loop16
|
||||
|
||||
reduce_vec:
|
||||
// Combine the four accumulators into V8.
|
||||
VADD V9.D2, V8.D2, V8.D2
|
||||
VADD V11.D2, V10.D2, V10.D2
|
||||
VADD V10.D2, V8.D2, V8.D2
|
||||
|
||||
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
||||
VADDP V8.D2, V8.D2, V8.D2
|
||||
VMOV V8.D[0], R8
|
||||
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
scalar_tail:
|
||||
CMP $8, R1
|
||||
BLT tail4
|
||||
|
||||
loop8:
|
||||
MOVD.P 8(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $8, R1, R1
|
||||
CMP $8, R1
|
||||
BGE loop8
|
||||
|
||||
tail4:
|
||||
CMP $4, R1
|
||||
BLT tail2
|
||||
MOVWU.P 4(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $4, R1, R1
|
||||
|
||||
tail2:
|
||||
CMP $2, R1
|
||||
BLT tail1
|
||||
MOVHU.P 2(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $2, R1, R1
|
||||
|
||||
tail1:
|
||||
CBZ R1, fold
|
||||
MOVBU (R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
fold:
|
||||
// Merge the byte-swapped initial into our LE-form accumulator.
|
||||
ADDS R2, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
// 64 → 33 bits.
|
||||
LSR $32, R3, R8
|
||||
AND $0xffffffff, R3, R3
|
||||
ADD R8, R3, R3
|
||||
|
||||
// 33 → 32 (truncate after adding bit 32 back).
|
||||
LSR $32, R3, R8
|
||||
ADD R8, R3, R3
|
||||
AND $0xffffffff, R3, R3
|
||||
|
||||
// 32 → 17.
|
||||
LSR $16, R3, R8
|
||||
AND $0xffff, R3, R3
|
||||
ADD R8, R3, R3
|
||||
|
||||
// 17 → 16 (truncation absorbs bit 16 below).
|
||||
LSR $16, R3, R8
|
||||
ADD R8, R3, R3
|
||||
|
||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
||||
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
||||
// 16-bit byte-swap on the live low 16.
|
||||
REV16W R3, R3
|
||||
AND $0xffff, R3, R3
|
||||
|
||||
MOVH R3, ret+32(FP)
|
||||
RET
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build !amd64 && !arm64
|
||||
|
||||
package checksum
|
||||
|
||||
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
|
||||
// Checksum delegates to gvisor on architectures without a hand-written body.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
return gvisorchecksum.Checksum(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package checksum
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand/v2"
|
||||
"testing"
|
||||
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
// archImpl names one checksum function under test. The per-arch
|
||||
// export_*_test.go files enumerate the hand-written implementations so the
|
||||
// suite compares each one against gvisor directly, regardless of which one
|
||||
// the public Checksum dispatches to on the running CPU. Testing only the
|
||||
// dispatcher was tautological wherever it resolved to the gvisor fallback
|
||||
// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself,
|
||||
// assembly untested, suite green.
|
||||
type archImpl struct {
|
||||
name string
|
||||
fn func([]byte, uint16) uint16
|
||||
available bool
|
||||
}
|
||||
|
||||
// implsUnderTest is the public dispatcher plus every arch implementation.
|
||||
func implsUnderTest() []archImpl {
|
||||
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
|
||||
}
|
||||
|
||||
// requireAvailable skips loudly when the running CPU can't execute an
|
||||
// implementation — visible in test output, unlike the old silent tautology.
|
||||
func requireAvailable(t *testing.T, impl archImpl) {
|
||||
t.Helper()
|
||||
if !impl.available {
|
||||
t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||
// seeds and a handful of starting alignments, asserting that each local
|
||||
// implementation matches gvisor's reference bit-for-bit.
|
||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(1, 2))
|
||||
const padFront = 16
|
||||
|
||||
// Random pool large enough for the longest case + alignment slop.
|
||||
pool := make([]byte, 4096+padFront)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||
|
||||
for length := 0; length <= 4096; length++ {
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||
length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumPatternedBuffers exercises specific byte patterns that have
|
||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||
// alternating, and ascending sequences.
|
||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
for length := 0; length <= 256; length++ {
|
||||
patterns := map[string][]byte{
|
||||
"zeros": make([]byte, length),
|
||||
"ones": bytes(length, 0xff),
|
||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||
"ascending": ascending(length),
|
||||
}
|
||||
for name, buf := range patterns {
|
||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||
name, length, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func bytes(n int, v byte) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = v
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func pattern(n int, p []byte) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = p[i%len(p)]
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func ascending(n int) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = byte(i)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
|
||||
// trailing tail bytes) the asm handlers walk through. The tail handlers
|
||||
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
|
||||
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
|
||||
// representative spread of k values, including k=0 (no main loop, all tail)
|
||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
||||
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
||||
func TestChecksumTailPaths(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(42, 17))
|
||||
const padFront = 16
|
||||
const maxK = 8
|
||||
|
||||
pool := make([]byte, 64*maxK+padFront+64)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||
|
||||
for k := 0; k <= maxK; k++ {
|
||||
for tail := 0; tail < 64; tail++ {
|
||||
length := 64*k + tail
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||
k, tail, length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
|
||||
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
|
||||
// handler is meaningfully on the hot path. Sizes are picked to either exercise
|
||||
// every tail branch (tiny lengths) or sit slightly off realistic packet
|
||||
// boundaries (e.g. 1499 = MTU − 1).
|
||||
func BenchmarkChecksumTailSizes(b *testing.B) {
|
||||
sizes := []int{
|
||||
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
|
||||
33, 35, 47, 63, // one loop32 + assorted tails
|
||||
65, 95, 127, // one loop64 + assorted tails
|
||||
1447, 1471, 1499, 1501, // around MTU
|
||||
8191, 8193, // around USO
|
||||
65531, 65533, // near the kernel max
|
||||
}
|
||||
for _, size := range sizes {
|
||||
buf := make([]byte, size)
|
||||
for i := range buf {
|
||||
buf[i] = byte(i)
|
||||
}
|
||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = gvisorchecksum.Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
|
||||
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
|
||||
// USO size (8192), and the kernel's max GSO superpacket (65535).
|
||||
func BenchmarkChecksum(b *testing.B) {
|
||||
for _, size := range []int{60, 1448, 8192, 65535} {
|
||||
buf := make([]byte, size)
|
||||
for i := range buf {
|
||||
buf[i] = byte(i)
|
||||
}
|
||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = gvisorchecksum.Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// so the tests exercise them directly, independent of what the public
|
||||
// Checksum dispatches to on the running CPU. Without this, running the
|
||||
// suite on a non-AVX2 machine compared gvisor against itself and left the
|
||||
// assembly untested — silently. available=false makes the test skip loudly
|
||||
// instead.
|
||||
var archImpls = []archImpl{
|
||||
{name: "avx2", fn: checksumAVX2, available: hasAVX2},
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// for direct testing; see export_amd64_test.go for the rationale. NEON is
|
||||
// mandatory in armv8, so it is always available.
|
||||
var archImpls = []archImpl{
|
||||
{name: "neon", fn: checksumNEON, available: true},
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
//go:build !amd64 && !arm64
|
||||
|
||||
package checksum
|
||||
|
||||
// No hand-written implementations on this architecture; the dispatcher is
|
||||
// pure gvisor and there is nothing separate to test.
|
||||
var archImpls []archImpl
|
||||
+13
-3
@@ -4,15 +4,25 @@ import (
|
||||
"io"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||
const defaultBatchBufSize = 65535
|
||||
|
||||
type Device interface {
|
||||
io.ReadWriteCloser
|
||||
io.Closer
|
||||
Activate() error
|
||||
Networks() []netip.Prefix
|
||||
Name() string
|
||||
RoutesFor(netip.Addr) routing.Gateways
|
||||
SupportsMultiqueue() bool
|
||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
||||
// Queues returns the device's packet queues, opening additional ones as
|
||||
// needed until there are n. Platforms without multiqueue support return
|
||||
// their single queue regardless of n, so callers must size reader loops
|
||||
// to len(result), not n; implementations never return more than n. An
|
||||
// error means a queue that should have opened could not; the caller owns
|
||||
// cleanup via Close. Called once, during interface activation.
|
||||
Queues(n int) ([]tio.Queue, error)
|
||||
}
|
||||
|
||||
@@ -3,10 +3,9 @@
|
||||
package overlaytest
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
@@ -31,20 +30,16 @@ func (NoopTun) Name() string {
|
||||
return "noop"
|
||||
}
|
||||
|
||||
func (NoopTun) Read([]byte) (int, error) {
|
||||
return 0, nil
|
||||
func (NoopTun) Read() ([]tio.Packet, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (NoopTun) Write([]byte) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func (NoopTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, errors.New("unsupported")
|
||||
func (NoopTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{NoopTun{}}, nil
|
||||
}
|
||||
|
||||
func (NoopTun) Close() error {
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
||||
// (events is POLLIN for reads, POLLOUT for writes)
|
||||
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
||||
//
|
||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
||||
func blockOn(fd, shutdownFd int32, events int16) error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
pfds := [2]unix.PollFd{
|
||||
{Fd: fd, Events: events},
|
||||
{Fd: shutdownFd, Events: unix.POLLIN},
|
||||
}
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(pfds[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
tunEvents := pfds[0].Revents
|
||||
shutdownEvents := pfds[1].Revents
|
||||
// Check err before trusting the potentially bogus bits we just got.
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type offloadQueueSet struct {
|
||||
pq []*Offload
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
// l is handed to each queue for its bad-vnet-header drop logging.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &offloadQueueSet{
|
||||
pq: []*Offload{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
l: l,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type pollQueueSet struct {
|
||||
pq []*Poll
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func NewPollQueueSet() (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &pollQueueSet{
|
||||
pq: []*Poll{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newPoll(fd, c.shutdownFd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import "testing"
|
||||
|
||||
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
||||
// of pointer-capturing closure that sendInsideMessage builds.
|
||||
type fakeBatch struct{ buf [65536]byte }
|
||||
|
||||
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
||||
func (b *fakeBatch) Commit([]byte) {}
|
||||
|
||||
type fakeHostInfo struct {
|
||||
remoteIndexId uint32
|
||||
counter uint64
|
||||
}
|
||||
type fakeIface struct {
|
||||
rebindCount uint8
|
||||
hi *fakeHostInfo
|
||||
}
|
||||
|
||||
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
||||
// SegmentSuperpacket call when a closure captures pointer-bearing
|
||||
// receivers — the realistic shape of sendInsideMessage's closure.
|
||||
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
||||
const mss = 1400
|
||||
const numSeg = 32
|
||||
pkt := buildTSOv6(mss*numSeg, mss)
|
||||
gso := GSOInfo{
|
||||
Size: mss,
|
||||
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
||||
CsumStart: 40,
|
||||
Proto: GSOProtoTCP,
|
||||
}
|
||||
p := Packet{Bytes: pkt, GSO: gso}
|
||||
|
||||
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
||||
f := &fakeIface{rebindCount: 7, hi: hi}
|
||||
fb := &fakeBatch{}
|
||||
|
||||
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
||||
// copy each iter (matches the production pattern where every TUN read
|
||||
// hands the segmenter a fresh kernel-supplied buffer).
|
||||
master := append([]byte(nil), pkt...)
|
||||
work := make([]byte, len(pkt))
|
||||
p.Bytes = work
|
||||
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
copy(work, master)
|
||||
err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||
out := fb.Reserve(16 + len(seg) + 16)
|
||||
out[0] = byte(f.rebindCount)
|
||||
out[1] = byte(hi.counter)
|
||||
hi.counter++
|
||||
fb.Commit(out)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
b.Fatalf("SegmentSuperpacket: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package tio
|
||||
|
||||
import "fmt"
|
||||
|
||||
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||
return 0, fmt.Errorf("GSO unsupported")
|
||||
}
|
||||
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if pkt.GSO.IsSuperpacket() {
|
||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||
}
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package tio
|
||||
|
||||
import "io"
|
||||
|
||||
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
||||
// Read fills a private scratch buffer and returns exactly one Packet whose
|
||||
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
|
||||
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
|
||||
type singleQueue struct {
|
||||
rw io.ReadWriter
|
||||
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
||||
buf []byte
|
||||
ret [1]Packet
|
||||
}
|
||||
|
||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
|
||||
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
|
||||
// Close closes rwc.
|
||||
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
||||
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else,
|
||||
// e.g. several queues sharing one device. Close on the returned Queue is a
|
||||
// no-op so one queue can't tear the shared source out from under its
|
||||
// siblings; the owner remains responsible for closing the source itself.
|
||||
func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue {
|
||||
return &singleQueue{rw: rw, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
func (q *singleQueue) Read() ([]Packet, error) {
|
||||
n, err := q.rw.Read(q.buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q.ret[0] = Packet{Bytes: q.buf[:n]}
|
||||
return q.ret[:], nil
|
||||
}
|
||||
|
||||
func (q *singleQueue) Write(p []byte) (int, error) {
|
||||
return q.rw.Write(p)
|
||||
}
|
||||
|
||||
func (q *singleQueue) Close() error {
|
||||
if q.closer == nil {
|
||||
return nil
|
||||
}
|
||||
return q.closer.Close()
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package tio
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
||||
type QueueSet interface {
|
||||
io.Closer
|
||||
Queues() []Queue
|
||||
|
||||
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
||||
Add(fd int) error
|
||||
}
|
||||
|
||||
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||
type Capabilities struct {
|
||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
|
||||
// and WriteGSO with GSOProtoTCP is safe.
|
||||
TSO bool
|
||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
|
||||
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||
USO bool
|
||||
}
|
||||
|
||||
// Queue is a readable/writable Poll queue.
|
||||
// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers;
|
||||
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
|
||||
//
|
||||
// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd
|
||||
// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close,
|
||||
// which first signals a shared shutdown eventfd every reader polls alongside its own fd.
|
||||
// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns
|
||||
// os.ErrClosed from Read, so it cannot be used to stop a single Queue.
|
||||
type Queue interface {
|
||||
io.Closer
|
||||
|
||||
// Read returns one or more packets.
|
||||
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
|
||||
// until the next Read or Close on this Queue.
|
||||
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
|
||||
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
||||
Read() ([]Packet, error)
|
||||
|
||||
// Write emits a single packet on the plaintext (outside→inside) delivery path.
|
||||
// Safe for concurrent use.
|
||||
Write(p []byte) (int, error)
|
||||
}
|
||||
|
||||
// Packet is the unit Queue.Read returns.
|
||||
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
|
||||
// GSO is the zero value for an already-segmented IP datagram;
|
||||
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
|
||||
type Packet struct {
|
||||
Bytes []byte
|
||||
GSO GSOInfo
|
||||
}
|
||||
|
||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
|
||||
type GSOInfo struct {
|
||||
// Size is the GSO segment size: max payload bytes per segment
|
||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
|
||||
Size uint16
|
||||
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
|
||||
HdrLen uint16
|
||||
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
|
||||
CsumStart uint16
|
||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
|
||||
Proto GSOProto
|
||||
}
|
||||
|
||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
||||
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
|
||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
||||
|
||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
||||
// safe to retain past the next Read or Close on the originating Queue.
|
||||
// GSO metadata is copied verbatim.
|
||||
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
|
||||
func (p Packet) Clone() Packet {
|
||||
if p.Bytes == nil {
|
||||
return p
|
||||
}
|
||||
cp := make([]byte, len(p.Bytes))
|
||||
copy(cp, p.Bytes)
|
||||
return Packet{Bytes: cp, GSO: p.GSO}
|
||||
}
|
||||
|
||||
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
|
||||
// Callers pick a write-path coalescer based on the result.
|
||||
// Queues that don't implement it are treated as having no offload capability.
|
||||
type CapsProvider interface {
|
||||
Capabilities() Capabilities
|
||||
}
|
||||
|
||||
// GSOProto selects the L4 protocol for a GSO superpacket.
|
||||
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||
// inside the transport header virtio NEEDS_CSUM expects.
|
||||
type GSOProto uint8
|
||||
|
||||
const (
|
||||
GSOProtoUnknown GSOProto = iota
|
||||
GSOProtoTCP
|
||||
GSOProtoUDP
|
||||
)
|
||||
|
||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||
// assembled from a header prefix plus one or more borrowed payload fragments,
|
||||
// in a single vectored write (writev with a leading virtio_net_hdr).
|
||||
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
|
||||
// Backends without GSO support do not implement this interface and coalescing is skipped.
|
||||
//
|
||||
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
|
||||
// transportHdr is the TCP or UDP header
|
||||
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
|
||||
// They are read-only from the writer's perspective and must remain valid until the call returns.
|
||||
// Every segment in pays except possibly the last must be exactly the same size.
|
||||
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
|
||||
//
|
||||
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
|
||||
// USO may not have been negotiated even when TSO was.
|
||||
type GSOWriter interface {
|
||||
io.Writer
|
||||
CapsProvider
|
||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||
}
|
||||
|
||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||
// queue advertises the negotiated capability for `want`.
|
||||
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
|
||||
gw, ok := w.(GSOWriter)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
caps := gw.Capabilities()
|
||||
switch want {
|
||||
case GSOProtoTCP:
|
||||
return gw, caps.TSO
|
||||
case GSOProtoUDP:
|
||||
return gw, caps.USO
|
||||
default:
|
||||
return gw, false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,417 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
const maxSuperpacketLen = 65535
|
||||
|
||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB.
|
||||
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||
// We round up to give margin for the drain headroom check below.
|
||||
const tunRxBufSize = 64 * 1024
|
||||
|
||||
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
|
||||
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||
// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load,
|
||||
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
|
||||
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
|
||||
const tunRxBufCap = tunRxBufSize * 8
|
||||
|
||||
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
|
||||
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
|
||||
const tunDrainCap = 64
|
||||
|
||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
|
||||
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
|
||||
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
|
||||
// so realistic coalesced bursts never touch the limit.
|
||||
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
|
||||
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
|
||||
const gsoMaxIovs = 256
|
||||
|
||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
|
||||
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
|
||||
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
|
||||
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
|
||||
// L4 checksum anyway. What matters here is what the header does NOT say:
|
||||
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
|
||||
// All packets that reach the plain Write paths already carry a valid L4 checksum.
|
||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||
|
||||
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
//
|
||||
// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill
|
||||
// the first cache line, and the state the reader mutates per packet (rxOff, pending,
|
||||
// readIovs) all sits after it, so per-packet reader stores never invalidate the line
|
||||
// concurrent Write callers load fd from.
|
||||
type Offload struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
||||
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
|
||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
||||
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
||||
// so non-GSO Writes can ship that constant directly while WriteGSO
|
||||
// rewrites this scratch on every call.
|
||||
gsoHdrBuf [virtio.Size]byte
|
||||
// gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to
|
||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
||||
// (and drops the call) if a caller hands it more fragments than fit.
|
||||
gsoIovs []unix.Iovec
|
||||
|
||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||
rxOff int // cursor into rxBuf for the current Read drain
|
||||
pending []Packet // packets returned from the most recent Read
|
||||
|
||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||
// lets us read the body directly into rxBuf at the current rxOff with
|
||||
// no userspace copy on the GSO_NONE fast path.
|
||||
readVnetScratch [virtio.Size]byte
|
||||
// readIovs is the readv(2) iovec scratch wired once at construction,
|
||||
// iovec[0] points at readVnetScratch
|
||||
// iovec[1].Base/Len is updated per read to address the current rxBuf slot.
|
||||
readIovs [2]unix.Iovec
|
||||
|
||||
// l is only consulted on the rare bad-vnet-header drop path; it lives
|
||||
// after the hot state on purpose. May be nil (tests); drops go unlogged then.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Offload{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
closed: atomic.Bool{},
|
||||
l: l,
|
||||
|
||||
rxBuf: make([]byte, tunRxBufCap),
|
||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
||||
}
|
||||
|
||||
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
||||
out.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
// readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we
|
||||
// only repoint readIovs[1] at the next rxBuf slot (see readPacket).
|
||||
out.readIovs[0].Base = &out.readVnetScratch[0]
|
||||
out.readIovs[0].SetLen(virtio.Size)
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnRead() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnWrite() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
|
||||
// and reading the packet body directly into rxBuf at the current rxOff.
|
||||
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
|
||||
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
|
||||
func (r *Offload) readPacket(block bool) (int, error) {
|
||||
for {
|
||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||
r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff)
|
||||
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if !block {
|
||||
return 0, errno
|
||||
}
|
||||
if err := r.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Read returns one or more packets from the tun.
|
||||
// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket).
|
||||
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
|
||||
// - the kernel queue is empty (EAGAIN)
|
||||
// - we've collected tunDrainCap packets,
|
||||
// - or we're out of rxBuf headroom.
|
||||
//
|
||||
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
|
||||
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
|
||||
func (r *Offload) Read() ([]Packet, error) {
|
||||
r.pending = r.pending[:0]
|
||||
r.rxOff = 0
|
||||
|
||||
// Initial (blocking) read.
|
||||
// Retry on decode errors so a single bad packet does not stall the reader.
|
||||
for {
|
||||
n, err := r.readPacket(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop and read again. A bad packet should not kill the reader,
|
||||
// but a systematic decode failure must not be invisible either.
|
||||
r.logDroppedRead(err)
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
||||
// cap is reached, or rxBuf no longer has room for another worst-case
|
||||
// kernel-supplied packet (tunRxBufSize).
|
||||
for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize {
|
||||
n, err := r.readPacket(false)
|
||||
if err != nil {
|
||||
// EAGAIN / EINTR / anything else: stop draining. We already
|
||||
// have a valid batch from the first read.
|
||||
break
|
||||
}
|
||||
if n <= 0 {
|
||||
break
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop this packet and stop the drain; we'd rather hand off
|
||||
// what we have than keep spinning here.
|
||||
r.logDroppedRead(err)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return r.pending, nil
|
||||
}
|
||||
|
||||
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
|
||||
// header. Debug-gated so the happy path never pays for attribute assembly.
|
||||
func (r *Offload) logDroppedRead(err error) {
|
||||
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen).
|
||||
// The bytes stay in rxBuf:
|
||||
// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||
// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time.
|
||||
//
|
||||
// rxOff advances by pktLen on success
|
||||
func (r *Offload) decodeRead(pktLen int) error {
|
||||
if pktLen <= 0 {
|
||||
return fmt.Errorf("short tun read: %d", pktLen)
|
||||
}
|
||||
var hdr virtio.Hdr
|
||||
hdr.Decode(r.readVnetScratch[:])
|
||||
|
||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
r.pending = append(r.pending, Packet{Bytes: body})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.pending = append(r.pending, Packet{
|
||||
Bytes: body,
|
||||
GSO: GSOInfo{
|
||||
Size: hdr.GSOSize,
|
||||
HdrLen: hdr.HdrLen,
|
||||
CsumStart: hdr.CsumStart,
|
||||
Proto: proto,
|
||||
},
|
||||
})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Offload) Write(buf []byte) (int, error) {
|
||||
if len(buf) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
iovs := [2]unix.Iovec{
|
||||
{Base: &validVnetHdr[0]},
|
||||
{Base: &buf[0]},
|
||||
}
|
||||
iovs[0].SetLen(virtio.Size)
|
||||
iovs[1].SetLen(len(buf))
|
||||
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
|
||||
}
|
||||
|
||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||
for {
|
||||
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, io.ErrShortWrite
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if err := r.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
|
||||
func (r *Offload) Capabilities() Capabilities {
|
||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
||||
}
|
||||
|
||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
||||
if len(pays) == 0 {
|
||||
// There are no payload fragments. There is nothing to send.
|
||||
return nil
|
||||
}
|
||||
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
|
||||
switch proto {
|
||||
case GSOProtoUDP:
|
||||
csumOff = 6
|
||||
case GSOProtoTCP:
|
||||
csumOff = 16
|
||||
default:
|
||||
return fmt.Errorf("unknown GSO proto: %d", proto)
|
||||
}
|
||||
// Incorrect geometry must cause an error, not a silent drop.
|
||||
// No sane packet should ever make it inside this branch.
|
||||
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
|
||||
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
|
||||
}
|
||||
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
|
||||
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
|
||||
need := 3 + len(pays)
|
||||
if need > cap(r.gsoIovs) {
|
||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||
}
|
||||
r.gsoIovs = r.gsoIovs[:need]
|
||||
r.gsoIovs[1].Base = &hdr[0]
|
||||
r.gsoIovs[1].SetLen(len(hdr))
|
||||
r.gsoIovs[2].Base = &transportHdr[0]
|
||||
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||
|
||||
segSize := len(pays[0])
|
||||
total := len(hdr) + len(transportHdr)
|
||||
for i, p := range pays {
|
||||
if len(p) == 0 {
|
||||
// The coalescers route zero-payload packets down the non-GSO path,
|
||||
// so an empty fragment means the caller's accounting is broken.
|
||||
return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays))
|
||||
} else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) {
|
||||
// all segments must be the same size, except for the last one
|
||||
return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize)
|
||||
}
|
||||
total += len(p)
|
||||
r.gsoIovs[3+i].Base = &p[0]
|
||||
r.gsoIovs[3+i].SetLen(len(p))
|
||||
}
|
||||
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
|
||||
if total > maxSuperpacketLen {
|
||||
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
|
||||
}
|
||||
|
||||
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
|
||||
// Multiple segments carry the real GSO type and segSize, which the loop
|
||||
// above verified is the size of every fragment except possibly the last.
|
||||
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||
if len(pays) > 1 {
|
||||
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
|
||||
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
|
||||
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
|
||||
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
|
||||
}
|
||||
}
|
||||
var gsoSize uint16
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
gsoSize = uint16(segSize)
|
||||
}
|
||||
virtio.EncodeHeader(
|
||||
r.gsoHdrBuf[:],
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
gsoType, /*gsoType*/
|
||||
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
|
||||
gsoSize, /*gsoSize*/
|
||||
uint16(len(hdr)), /*csumStart*/
|
||||
csumOff, /*csumOffset*/
|
||||
)
|
||||
|
||||
_, err := r.rawWrite(r.gsoIovs)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *Offload) Close() error {
|
||||
if r.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type Poll struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
|
||||
readBuf []byte
|
||||
batchRet [1]Packet
|
||||
}
|
||||
|
||||
// newPoll wraps an existing tun fd.
|
||||
// On failure it does NOT close fd: the caller owns fd and is the sole closer
|
||||
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
|
||||
// This matches the newOffload convention and keeps closes at exactly one on every path.
|
||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Poll{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
readBuf: make([]byte, 65535), // largest possible size Linux permits
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
|
||||
// Returns os.ErrClosed if Close was called.
|
||||
func (t *Poll) blockOnRead() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (t *Poll) blockOnWrite() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// TODO: port Offload's post-wake drain loop here so one poll wake amortizes
|
||||
// over a burst (up to tunDrainCap packets) instead of paying a syscall and a
|
||||
// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd
|
||||
// config currently lose that batching. blockOn and the EAGAIN plumbing are
|
||||
// already shared; kept one-packet-per-Read for now to preserve behavior.
|
||||
func (t *Poll) Read() ([]Packet, error) {
|
||||
n, err := t.readOne(t.readBuf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t.batchRet[0] = Packet{Bytes: t.readBuf[:n]}
|
||||
return t.batchRet[:], nil
|
||||
}
|
||||
|
||||
func (t *Poll) readOne(to []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Read(t.fd, to)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Write is safe for concurrent use
|
||||
func (t *Poll) Write(from []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Write(t.fd, from)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Poll) Close() error {
|
||||
if t.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(t.fd)
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||
// The caller takes ownership of the read fd (pass it into a QueueSet).
|
||||
func newReadPipe(t *testing.T) int {
|
||||
t.Helper()
|
||||
var fds [2]int
|
||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||
t.Fatalf("pipe2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||
return fds[0]
|
||||
}
|
||||
|
||||
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||
pipe1 := newReadPipe(t)
|
||||
pipe2 := newReadPipe(t)
|
||||
parent, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, parent.Add(pipe1))
|
||||
require.NoError(t, parent.Add(pipe2))
|
||||
t.Cleanup(func() {
|
||||
_ = unix.Close(pipe1)
|
||||
_ = unix.Close(pipe2)
|
||||
})
|
||||
|
||||
readers := parent.Queues()
|
||||
errs := make([]error, len(readers))
|
||||
var wg sync.WaitGroup
|
||||
for i, r := range readers {
|
||||
wg.Add(1)
|
||||
go func(i int, r Queue) {
|
||||
defer wg.Done()
|
||||
_, errs[i] = r.Read()
|
||||
}(i, r)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if err := parent.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("readers did not wake")
|
||||
}
|
||||
|
||||
for i, err := range errs {
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer
|
||||
// goroutines while a reader drains the other end of the pipe. The writers
|
||||
// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same
|
||||
// time — the exact scenario that raced on the old shared writePoll member
|
||||
// array. Run under -race; a shared-array regression trips the detector here.
|
||||
func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
|
||||
var fds [2]int
|
||||
require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC))
|
||||
readFd, writeFd := fds[0], fds[1]
|
||||
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
|
||||
|
||||
p, err := newPoll(writeFd, shutdownFd)
|
||||
require.NoError(t, err)
|
||||
|
||||
const writers = 2
|
||||
const perWriter = 4000
|
||||
payload := make([]byte, 100)
|
||||
total := writers * perWriter * len(payload)
|
||||
|
||||
// Reader: drain the read end (blocking) until every writer's bytes are
|
||||
// consumed, so the writers keep making progress rather than wedging on a
|
||||
// permanently full pipe.
|
||||
readDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(readDone)
|
||||
buf := make([]byte, 4096)
|
||||
got := 0
|
||||
for got < total {
|
||||
n, rerr := unix.Read(readFd, buf)
|
||||
got += n
|
||||
if rerr != nil {
|
||||
if rerr == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
return
|
||||
}
|
||||
if n == 0 { // EOF
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for w := 0; w < writers; w++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < perWriter; i++ {
|
||||
if _, werr := p.Write(payload); werr != nil {
|
||||
t.Errorf("write: %v", werr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
select {
|
||||
case <-readDone:
|
||||
case <-time.After(10 * time.Second):
|
||||
t.Fatal("reader did not drain")
|
||||
}
|
||||
|
||||
require.NoError(t, p.Close())
|
||||
_ = unix.Close(readFd)
|
||||
}
|
||||
|
||||
// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when
|
||||
// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's
|
||||
// callers in tun_linux.go) is the sole closer. If newPoll also closed fd,
|
||||
// the poll path would double-close on Add error. We force the failure with
|
||||
// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not
|
||||
// permitted on O_PATH fds and fails with EBADF, while the fd itself stays
|
||||
// open so we can observe that newPoll left it alone.
|
||||
func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) {
|
||||
fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||
|
||||
p, err := newPoll(fd, 1)
|
||||
require.Error(t, err, "SetNonblock on an O_PATH fd should fail")
|
||||
require.Nil(t, p)
|
||||
|
||||
// If newPoll had closed fd, F_GETFD would report it closed. It staying
|
||||
// open proves newPoll left the fd for the caller to close exactly once.
|
||||
require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer")
|
||||
}
|
||||
|
||||
func TestPoll_Close_Idempotent(t *testing.T) {
|
||||
tf, err := newPoll(newReadPipe(t), 1)
|
||||
require.NoError(t, err)
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("first Close: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// fdOpen reports whether fd currently refers to an open file description.
|
||||
// A closed (or never-allocated) fd makes F_GETFD fail with EBADF.
|
||||
func fdOpen(t *testing.T, fd int) bool {
|
||||
t.Helper()
|
||||
_, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0)
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
if errors.Is(err, unix.EBADF) {
|
||||
return false
|
||||
}
|
||||
t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err)
|
||||
return false
|
||||
}
|
||||
|
||||
// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the
|
||||
// leaked shutdown eventfd: the container that owns shutdownFd must close it in
|
||||
// Close, and a second Close must be a safe no-op.
|
||||
func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*pollQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
|
||||
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
||||
// for the GSO/offload queueset.
|
||||
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*offloadQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
|
||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
||||
// value. The caller should only invoke this on a confirmed superpacket.
|
||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
return GSOProtoTCP, nil
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
return GSOProtoUDP, nil
|
||||
default:
|
||||
return 0, fmt.Errorf("unsupported virtio gso type: %d", t)
|
||||
}
|
||||
}
|
||||
|
||||
// gsoTypeFromProto is the reverse of protoFromGSOType
|
||||
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
|
||||
switch {
|
||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||
case ipVer == 6:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||
case ipVer == 4:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||
default:
|
||||
return unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
}
|
||||
}
|
||||
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt.
|
||||
// For non-GSO pkts fn is called once with pkt.Bytes.
|
||||
// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext
|
||||
// (a freshly-patched L3+L4 header sliced in front of the original payload chunk).
|
||||
// This slicing is destructive: pkt is consumed by this call.
|
||||
// Aborts and returns the first error from fn or from per-segment construction.
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if !pkt.GSO.IsSuperpacket() {
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
switch pkt.GSO.Proto {
|
||||
case GSOProtoTCP:
|
||||
return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
case GSOProtoUDP:
|
||||
return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
default:
|
||||
return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,75 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||
// not set).
|
||||
const Size = 10
|
||||
|
||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||
type Hdr struct {
|
||||
Flags uint8
|
||||
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
|
||||
HdrLen uint16
|
||||
GSOSize uint16
|
||||
CsumStart uint16
|
||||
CsumOffset uint16
|
||||
}
|
||||
|
||||
func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr {
|
||||
return Hdr{
|
||||
Flags: flags,
|
||||
gsoType: gsoType,
|
||||
HdrLen: hdrLen,
|
||||
GSOSize: gsoSize,
|
||||
CsumStart: csumStart,
|
||||
CsumOffset: csumOffset,
|
||||
}
|
||||
}
|
||||
|
||||
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||
func (h *Hdr) Decode(b []byte) {
|
||||
h.Flags = b[0]
|
||||
h.gsoType = b[1]
|
||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
||||
}
|
||||
|
||||
func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) {
|
||||
b[0] = flags
|
||||
b[1] = gsoType
|
||||
binary.NativeEndian.PutUint16(b[2:4], hdrLen)
|
||||
binary.NativeEndian.PutUint16(b[4:6], gsoSize)
|
||||
binary.NativeEndian.PutUint16(b[6:8], csumStart)
|
||||
binary.NativeEndian.PutUint16(b[8:10], csumOffset)
|
||||
}
|
||||
|
||||
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
|
||||
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
||||
func (h *Hdr) Encode(b []byte) {
|
||||
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
|
||||
}
|
||||
|
||||
// GSOType returns gsoType with the ECN-flag masked out
|
||||
func (h *Hdr) GSOType() uint8 {
|
||||
return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN
|
||||
}
|
||||
|
||||
func (h *Hdr) HasECNFlag() bool {
|
||||
return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0
|
||||
}
|
||||
|
||||
func (h *Hdr) SetGSOType(x uint8) {
|
||||
h.gsoType = x
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package virtio
|
||||
@@ -0,0 +1,441 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
// Package virtio implements the pure validation, header-correction, and
|
||||
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
|
||||
// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte
|
||||
// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes,
|
||||
// so it can be unit-tested in isolation from the tio Queue runtime.
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
||||
const (
|
||||
ipv4HeaderMinLen = 20 // IHL=5, no options
|
||||
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
||||
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
||||
tcpHeaderMinLen = 20 // data-offset=5, no options
|
||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||
)
|
||||
|
||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
|
||||
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
|
||||
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||
|
||||
// Byte offsets inside an IPv4 header.
|
||||
const (
|
||||
ipv4TotalLenOff = 2
|
||||
ipv4IDOff = 4
|
||||
ipv4ChecksumOff = 10
|
||||
ipv4SrcOff = 12
|
||||
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
||||
)
|
||||
|
||||
// Byte offsets inside an IPv6 header.
|
||||
const (
|
||||
ipv6PayloadLenOff = 4
|
||||
ipv6SrcOff = 8
|
||||
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
||||
)
|
||||
|
||||
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
||||
const (
|
||||
tcpSeqOff = 4
|
||||
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
||||
tcpFlagsOff = 13
|
||||
tcpChecksumOff = 16
|
||||
)
|
||||
|
||||
// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}.
|
||||
const (
|
||||
udpHeaderLen = 8
|
||||
udpLengthOff = 4
|
||||
udpChecksumOff = 6
|
||||
)
|
||||
|
||||
var errPacketTooShort = errors.New("packet too short")
|
||||
|
||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||
|
||||
// tcpCwrFlag is cleared on every segment except the first.
|
||||
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
|
||||
// and must appear on the first segment of a TSO burst only.
|
||||
const tcpCwrFlag = 0x80
|
||||
|
||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
||||
// the GSO type must agree with the IP version nibble.
|
||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||
}
|
||||
if len(pkt) < ipv4HeaderMinLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
ipVersion := pkt[0] >> 4
|
||||
if ipVersion == 6 && len(pkt) < ipv6FixedLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
|
||||
gsoType := hdr.GSOType()
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 {
|
||||
// A GSO type with no segment size would dodge IsSuperpacket() downstream and
|
||||
// travel as a plain jumbo datagram with an unfinished checksum.
|
||||
return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType)
|
||||
}
|
||||
if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) {
|
||||
return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType)
|
||||
}
|
||||
switch gsoType {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||
if ipVersion != 4 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
if ipVersion != 6 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
||||
if !(ipVersion == 4 || ipVersion == 6) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
default:
|
||||
if !(ipVersion == 6 || ipVersion == 4) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
|
||||
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
|
||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
// Thank you wireguard-go for documenting these edge-cases
|
||||
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
||||
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
|
||||
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
hdr.HdrLen = hdr.CsumStart + 8
|
||||
} else {
|
||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
||||
return errors.New("packet is too short")
|
||||
}
|
||||
|
||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
||||
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
|
||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||
}
|
||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||
}
|
||||
|
||||
if len(pkt) < int(hdr.HdrLen) {
|
||||
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
|
||||
}
|
||||
|
||||
if hdr.HdrLen < hdr.CsumStart {
|
||||
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
|
||||
}
|
||||
cSumAt := int(hdr.CsumStart + hdr.CsumOffset)
|
||||
if cSumAt+1 >= len(pkt) {
|
||||
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
|
||||
// with a floor of one so a header-only superpacket still yields a single segment.
|
||||
func segCount(payLen, gsoSize int) int {
|
||||
n := (payLen + gsoSize - 1) / gsoSize
|
||||
if n == 0 {
|
||||
return 1
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
|
||||
// for every segment: the source and destination addresses plus the protocol
|
||||
// number. The per-segment L4 length is added by the caller inside the loop.
|
||||
func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 {
|
||||
if isV4 {
|
||||
return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto
|
||||
}
|
||||
return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto
|
||||
}
|
||||
|
||||
// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments.
|
||||
// csumStart is the L3 header length, which bounds a valid IHL.
|
||||
func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
// total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment.
|
||||
sum := uint32(checksum.Checksum(pkt[:ihl], 0))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum, nil
|
||||
}
|
||||
|
||||
// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite
|
||||
func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
flags := uint16(pkt[csumStart+tcpFlagsOff])
|
||||
|
||||
sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0))
|
||||
sum += uint32(^uint16(seq >> 16))
|
||||
sum += uint32(^uint16(seq))
|
||||
sum += uint32(^flags)
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum
|
||||
}
|
||||
|
||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt.
|
||||
// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize,
|
||||
// where it sits immediately before that segment's payload chunk in the original buffer.
|
||||
// pkt is consumed by this call and must not be inspected by the caller after the final yield.
|
||||
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
|
||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
|
||||
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// TSO bumps the ID per segment, so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once. '
|
||||
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
|
||||
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
|
||||
if i > 0 {
|
||||
// Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
segSeq := origSeq + uint32(segStart)
|
||||
segFlags := origFlags
|
||||
if i != 0 {
|
||||
segFlags &^= tcpCwrFlag
|
||||
}
|
||||
if i != numSeg-1 {
|
||||
segFlags &^= tcpFinPshMask
|
||||
}
|
||||
totalLen := segLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
||||
seg[csumStart+tcpFlagsOff] = segFlags
|
||||
|
||||
tcpLen := tcpHdrLen + segPayLen
|
||||
// Payload bytes still live at their original offset in pkt.
|
||||
// The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed)
|
||||
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
|
||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
|
||||
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
|
||||
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
|
||||
// pkt is consumed destructively.
|
||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
if headerLen-csumStart != udpHeaderLen {
|
||||
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
|
||||
}
|
||||
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// Software UDP GSO bumps the ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once and stamp every segment from it
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
if i > 0 {
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
totalLen := segLen
|
||||
udpLen := udpHeaderLen + segPayLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||
|
||||
// Sum the UDP header (length just written, checksum zeroed) together with
|
||||
// this segment's payload in one pass, seeded with the pseudo-header sum.
|
||||
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
|
||||
pseudo := baseProtoSum + uint32(udpLen)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
|
||||
if csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum)
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
|
||||
// CsumStart / CsumOffset point at the 16-bit checksum field.
|
||||
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
|
||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||
cs := int(hdr.CsumStart)
|
||||
co := int(hdr.CsumOffset)
|
||||
if cs+co+2 > len(seg) {
|
||||
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
||||
}
|
||||
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
||||
// L4 region starting at cs, folding the prior partial in as the seed.
|
||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||
seg[cs+co] = 0
|
||||
seg[cs+co+1] = 0
|
||||
csum := ^checksum.Checksum(seg[cs:], partial)
|
||||
// RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value.
|
||||
if co == udpChecksumOff && csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum)
|
||||
return nil
|
||||
}
|
||||
|
||||
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
||||
// complements it, yielding the on-wire Internet checksum value.
|
||||
func foldComplement(sum uint32) uint16 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return ^uint16(sum)
|
||||
}
|
||||
@@ -0,0 +1,602 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// verifyChecksum confirms that the one's-complement sum across b, seeded with
|
||||
// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum).
|
||||
// A corrupted header stamped into a segment makes this fail even when the
|
||||
// checksum field itself was computed from the (pristine) base sums, because
|
||||
// the bytes the receiver would sum no longer match what was checksummed.
|
||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||
return checksum.Checksum(b, pseudo) == 0xffff
|
||||
}
|
||||
|
||||
// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own
|
||||
// address and length fields, used to independently verify its L4 checksum.
|
||||
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
|
||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||
s += uint32(proto) + uint32(l4Len)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a
|
||||
// payload of payLen bytes and returns it alongside the header fields the
|
||||
// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP).
|
||||
func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const tcpLen = 20
|
||||
pkt = make([]byte, ipLen+tcpLen+payLen)
|
||||
|
||||
// IPv4 header.
|
||||
pkt[0] = 0x45 // version 4, IHL 5
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID
|
||||
pkt[8] = 64 // TTL
|
||||
pkt[9] = unix.IPPROTO_TCP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
|
||||
|
||||
// TCP header.
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
|
||||
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
|
||||
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
|
||||
pkt[32] = 0x50 // data offset 5 words
|
||||
pkt[33] = 0x18 // ACK | PSH
|
||||
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + tcpLen, ipLen
|
||||
}
|
||||
|
||||
// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a
|
||||
// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP).
|
||||
func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const udpLen = 8
|
||||
pkt = make([]byte, ipLen+udpLen+payLen)
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + udpLen, ipLen
|
||||
}
|
||||
|
||||
// collectTCP segments a fresh copy of pkt and returns each segment as an
|
||||
// independent slice so assertions can run after segmentation completes.
|
||||
func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentTCP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentUDP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place
|
||||
// header-slide bug: when gsoSize < headerLen the old code stamped each
|
||||
// segment's header from pkt[:headerLen], which had already been overwritten
|
||||
// by the previous segment's overlapping stamp, so segments 2..n carried a
|
||||
// corrupted header (garbage src/dst/ports/seq). Every segment must instead
|
||||
// carry the ORIGINAL constant header fields with correct per-segment seq.
|
||||
func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
|
||||
const origSeq = 10000
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
// gsoSize (8) < headerLen (40): the bug's trigger. Even split.
|
||||
{"small-gso-even", 40, 8},
|
||||
// gsoSize (8) < headerLen (40) with a short final segment.
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
// gsoSize (100) >= headerLen (40): the normal path, must still work.
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
// Constant header fields must be identical to the original in
|
||||
// EVERY segment. These are exactly the bytes the old code
|
||||
// corrupted in segments 2..n.
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_TCP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 {
|
||||
t.Errorf("seg %d: dport=%d want 80", i, dport)
|
||||
}
|
||||
if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 {
|
||||
t.Errorf("seg %d: ack=%d want 20000", i, ack)
|
||||
}
|
||||
if seg[32] != 0x50 {
|
||||
t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32])
|
||||
}
|
||||
|
||||
// Per-segment seq must advance by the payload offset.
|
||||
segStart := i * gso
|
||||
if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) {
|
||||
t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart)
|
||||
}
|
||||
|
||||
// Payload bytes must be the original contiguous slice.
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
// End-to-end: the stamped header must checksum-verify. A
|
||||
// corrupted header fails here because the written checksum was
|
||||
// derived from the pristine header.
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad TCP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in
|
||||
// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check
|
||||
// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a
|
||||
// regression that doubled CsumStart and thus over-tightened the bound (since
|
||||
// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus
|
||||
// bound spuriously rejected valid small USO superpackets in decodeRead.
|
||||
func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
||||
// A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments
|
||||
// (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP
|
||||
// checksum field lives at bytes 26..27, comfortably inside the 40-byte
|
||||
// packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and
|
||||
// rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset
|
||||
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
||||
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
||||
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize: two 6-byte segments*/
|
||||
csumStart, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||
}
|
||||
if hdr.HdrLen != csumStart+udpHeaderLen {
|
||||
t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen)
|
||||
}
|
||||
})
|
||||
|
||||
// A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the
|
||||
// checksum field would end at byte 27, but the packet is only 25 bytes
|
||||
// (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it.
|
||||
t.Run("too-short-rejected", func(t *testing.T) {
|
||||
pkt := make([]byte, 25)
|
||||
pkt[0] = 0x45 // IPv4, IHL 5
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize*/
|
||||
20, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs
|
||||
// the same header stamp and must be correct when gsoSize < headerLen.
|
||||
func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
{"small-gso-even", 40, 8},
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_UDP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
||||
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
||||
}
|
||||
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only).
|
||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||
}
|
||||
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) {
|
||||
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen)
|
||||
}
|
||||
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad UDP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
|
||||
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
|
||||
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
|
||||
const ipLen, udpLen = 20, 8
|
||||
pkt = make([]byte, ipLen+udpLen+len(payload))
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
|
||||
copy(pkt[ipLen+udpLen:], payload)
|
||||
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
|
||||
|
||||
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
|
||||
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
|
||||
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
|
||||
var payload []byte
|
||||
for i := 0; i < 0x10000; i++ {
|
||||
p := []byte{byte(i >> 8), byte(i)}
|
||||
pkt, hdr := buildUDPv4Single(p)
|
||||
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
|
||||
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
|
||||
pkt[cs+co], pkt[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
|
||||
payload = p
|
||||
break
|
||||
}
|
||||
}
|
||||
if payload == nil {
|
||||
t.Fatal("no 2-byte payload produced a zero checksum")
|
||||
}
|
||||
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
|
||||
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
|
||||
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
|
||||
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
|
||||
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
|
||||
const cs, co = 20, tcpChecksumOff
|
||||
|
||||
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
|
||||
seg := make([]byte, cs+co+2)
|
||||
for i := range seg[cs:] {
|
||||
seg[cs+i] = byte(i * 7)
|
||||
}
|
||||
var partial uint16
|
||||
for i := 0; i <= 0xffff; i++ {
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
|
||||
probe := append([]byte(nil), seg...)
|
||||
probe[cs+co], probe[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
|
||||
partial = uint16(i)
|
||||
break
|
||||
}
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
|
||||
|
||||
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
|
||||
if err := FinishChecksum(seg, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
|
||||
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
|
||||
func TestFinishChecksumUDPValidates(t *testing.T) {
|
||||
payload := []byte("the definitive tun offloads branch")
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
|
||||
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
|
||||
t.Fatal("completed UDP checksum does not validate")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
|
||||
// into gso_type for TSO superpackets with CWR set. CheckValid must
|
||||
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
|
||||
// fell into the default case and skipped the IP-version agreement check.
|
||||
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
|
||||
func TestCheckValidMasksGSOECN(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
v6pkt := make([]byte, len(v4pkt))
|
||||
copy(v6pkt, v4pkt)
|
||||
v6pkt[0] = 0x60 // claim IPv6
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
gsoType uint8
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
|
||||
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
|
||||
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
|
||||
if tc.wantErr && err == nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
|
||||
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
|
||||
// dodging both segmentation and FinishChecksum on its way downstream.
|
||||
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
|
||||
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
|
||||
// against an independent RFC 1071 reference fold, hitting the carry edge
|
||||
// cases (values whose first fold produces another carry).
|
||||
func TestFoldComplementMatchesReference(t *testing.T) {
|
||||
refFold := func(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // first fold yields another carry
|
||||
0xffffffff, // worst case
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
|
||||
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
|
||||
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
|
||||
// into scratch, zero the fields the segment loop rewrites, sum. The production
|
||||
// versions instead sum in place and subtract those fields via one's-complement
|
||||
// arithmetic, which is faster but far less obvious — particularly for the TCP
|
||||
// flags byte, which is only half of a 16-bit word. These references exist so
|
||||
// that trade is checked rather than asserted.
|
||||
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
tcpLen := headerLen - csumStart
|
||||
var tmp [tcpHeaderMaxLen]byte
|
||||
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
|
||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||
tmp[tcpFlagsOff] = 0
|
||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
|
||||
}
|
||||
|
||||
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
|
||||
// what the file already has and reproduces identically on every run.
|
||||
func randByte(state *uint32) byte {
|
||||
*state = *state*1664525 + 1013904223
|
||||
return byte(*state >> 24)
|
||||
}
|
||||
|
||||
func TestBaseSumsMatchZeroingReference(t *testing.T) {
|
||||
state := uint32(12345)
|
||||
|
||||
t.Run("ipv4", func(t *testing.T) {
|
||||
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, ihl)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = byte(0x40 | (ihl / 4))
|
||||
|
||||
want := referenceBaseIPv4HdrSum(pkt, ihl)
|
||||
got, err := baseIPv4HdrSum(pkt, ihl)
|
||||
if err != nil {
|
||||
t.Fatalf("ihl=%d: %v", ihl, err)
|
||||
}
|
||||
// Compare the value that reaches the wire: the raw partial
|
||||
// sums may legally differ by one's-complement -0 vs +0.
|
||||
for _, tl := range []uint32{20, 1500, 65535} {
|
||||
for _, id := range []uint32{0, 0x4242, 0xffff} {
|
||||
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
|
||||
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp", func(t *testing.T) {
|
||||
const csumStart = 20
|
||||
for dataOff := 5; dataOff <= 15; dataOff++ {
|
||||
tcpLen := dataOff * 4
|
||||
headerLen := csumStart + tcpLen
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, headerLen+64)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = 0x45
|
||||
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
|
||||
|
||||
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
got := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
|
||||
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
|
||||
for _, l4 := range []uint32{20, 1460, 65535} {
|
||||
a := foldComplement(want + seq + fl + l4)
|
||||
b := foldComplement(got + seq + fl + l4)
|
||||
if a != b {
|
||||
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
|
||||
dataOff, seq, fl, l4, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
@@ -63,7 +64,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||
return r
|
||||
}
|
||||
|
||||
func (t tun) Activate() error {
|
||||
func (t *tun) Activate() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -96,10 +97,6 @@ func (t *tun) Name() string {
|
||||
return "android"
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -16,6 +15,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -550,7 +550,9 @@ func (t *tun) Read(to []byte) (int, error) {
|
||||
return n - 4, nil
|
||||
}
|
||||
|
||||
// Write pushes one IP packet onto the utun device.
|
||||
// Write pushes one IP packet onto the utun device. Safe for concurrent use:
|
||||
// the AF prefix and iovecs are per-call stack state, and the fd write itself
|
||||
// serializes on the runtime's fd mutex (see the Queue contract in tio.go).
|
||||
func (t *tun) Write(from []byte) (int, error) {
|
||||
if len(from) == 0 {
|
||||
return 0, syscall.EIO
|
||||
@@ -606,10 +608,6 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
+26
-24
@@ -10,6 +10,7 @@ import (
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
@@ -23,6 +24,23 @@ type disabledTun struct {
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// Read hands the next queued packet to a reader, copying it into b. Reads
|
||||
// from concurrent queues are safe: the channel receive serializes them and
|
||||
// each queue copies into its own private scratch buffer.
|
||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
||||
r, ok := <-t.read
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
t.tx.Inc(1)
|
||||
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
t.l.Debug("Write payload", "raw", prettyPacket(r))
|
||||
}
|
||||
|
||||
return copy(b, r), nil
|
||||
}
|
||||
|
||||
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
||||
tun := &disabledTun{
|
||||
vpnNetworks: vpnNetworks,
|
||||
@@ -57,24 +75,6 @@ func (*disabledTun) Name() string {
|
||||
return "disabled"
|
||||
}
|
||||
|
||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
||||
r, ok := <-t.read
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
if len(r) > len(b) {
|
||||
return 0, fmt.Errorf("packet larger than mtu: %d > %d bytes", len(r), len(b))
|
||||
}
|
||||
|
||||
t.tx.Inc(1)
|
||||
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
t.l.Debug("Write payload", "raw", prettyPacket(r))
|
||||
}
|
||||
|
||||
return copy(b, r), nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
||||
out := make([]byte, len(b))
|
||||
out = iputil.CreateICMPEchoResponse(b, out)
|
||||
@@ -106,12 +106,14 @@ func (t *disabledTun) Write(b []byte) (int, error) {
|
||||
return len(b), nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) SupportsMultiqueue() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return t, nil
|
||||
func (t *disabledTun) Queues(n int) ([]tio.Queue, error) {
|
||||
out := make([]tio.Queue, n)
|
||||
for i := range out {
|
||||
// NoClose: the shared channel and metrics are owned by the
|
||||
// disabledTun; Close on the device tears them down once for everybody.
|
||||
out[i] = tio.NewSingleQueueNoClose(t, defaultBatchBufSize)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) Close() error {
|
||||
|
||||
@@ -1,120 +0,0 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||
// The caller takes ownership of the read fd (pass it to newTunFd / newFriend).
|
||||
func newReadPipe(t *testing.T) int {
|
||||
t.Helper()
|
||||
var fds [2]int
|
||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||
t.Fatalf("pipe2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||
return fds[0]
|
||||
}
|
||||
|
||||
func TestTunFile_WakeForShutdown_UnblocksRead(t *testing.T) {
|
||||
tf, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = tf.Close() })
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := tf.Read(make([]byte, 64))
|
||||
done <- err
|
||||
}()
|
||||
|
||||
// Verify Read is actually blocked in poll.
|
||||
select {
|
||||
case err := <-done:
|
||||
t.Fatalf("Read returned before shutdown signal: %v", err)
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
|
||||
if err := tf.wakeForShutdown(); err != nil {
|
||||
t.Fatalf("wakeForShutdown: %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Fatalf("expected os.ErrClosed, got %v", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Read did not wake on shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunFile_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||
parent, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
friend, err := parent.newFriend(newReadPipe(t))
|
||||
if err != nil {
|
||||
_ = parent.Close()
|
||||
t.Fatalf("newFriend: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = friend.Close()
|
||||
_ = parent.Close()
|
||||
})
|
||||
|
||||
readers := []*tunFile{parent, friend}
|
||||
errs := make([]error, len(readers))
|
||||
var wg sync.WaitGroup
|
||||
for i, r := range readers {
|
||||
wg.Add(1)
|
||||
go func(i int, r *tunFile) {
|
||||
defer wg.Done()
|
||||
_, errs[i] = r.Read(make([]byte, 64))
|
||||
}(i, r)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if err := parent.wakeForShutdown(); err != nil {
|
||||
t.Fatalf("wakeForShutdown: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("readers did not wake")
|
||||
}
|
||||
|
||||
for i, err := range errs {
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunFile_Close_Idempotent(t *testing.T) {
|
||||
tf, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("first Close: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
@@ -20,7 +19,7 @@ import (
|
||||
"github.com/gaissmai/bart"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -561,12 +560,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
+3
-6
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -159,10 +160,6 @@ func (t *tun) Name() string {
|
||||
return "iOS"
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
+163
-256
@@ -4,9 +4,7 @@
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
@@ -19,188 +17,25 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"github.com/vishvananda/netlink"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// tunFile wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
type tunFile struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
lastOne bool
|
||||
readPoll [2]unix.PollFd
|
||||
writePoll [2]unix.PollFd
|
||||
closed bool
|
||||
}
|
||||
|
||||
// newFriend makes a tunFile for a MultiQueueReader that copies the shutdown eventfd from the parent tun
|
||||
func (r *tunFile) newFriend(fd int) (*tunFile, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
return &tunFile{
|
||||
fd: fd,
|
||||
shutdownFd: r.shutdownFd,
|
||||
readPoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLIN},
|
||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
writePoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newTunFd(fd int) (*tunFile, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &tunFile{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
lastOne: true,
|
||||
readPoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLIN},
|
||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
writePoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *tunFile) blockOnRead() error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(r.readPoll[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
//always reset these!
|
||||
tunEvents := r.readPoll[0].Revents
|
||||
shutdownEvents := r.readPoll[1].Revents
|
||||
r.readPoll[0].Revents = 0
|
||||
r.readPoll[1].Revents = 0
|
||||
//do the err check before trusting the potentially bogus bits we just got
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
} else if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *tunFile) blockOnWrite() error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(r.writePoll[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
//always reset these!
|
||||
tunEvents := r.writePoll[0].Revents
|
||||
shutdownEvents := r.writePoll[1].Revents
|
||||
r.writePoll[0].Revents = 0
|
||||
r.writePoll[1].Revents = 0
|
||||
//do the err check before trusting the potentially bogus bits we just got
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
} else if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *tunFile) Read(buf []byte) (int, error) {
|
||||
for {
|
||||
if n, err := unix.Read(r.fd, buf); err == nil {
|
||||
return n, nil
|
||||
} else if err == unix.EAGAIN {
|
||||
if err = r.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
} else if err == unix.EINTR {
|
||||
continue
|
||||
} else if err == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
} else {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *tunFile) Write(buf []byte) (int, error) {
|
||||
for {
|
||||
if n, err := unix.Write(r.fd, buf); err == nil {
|
||||
return n, nil
|
||||
} else if err == unix.EAGAIN {
|
||||
if err = r.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
} else if err == unix.EINTR {
|
||||
continue
|
||||
} else if err == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
} else {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *tunFile) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(int(r.readPoll[1].Fd), buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *tunFile) Close() error {
|
||||
if r.closed { // avoid closing more than once. Technically a fd could get re-used, which would be a problem
|
||||
return nil
|
||||
}
|
||||
r.closed = true
|
||||
if r.lastOne {
|
||||
_ = unix.Close(r.shutdownFd)
|
||||
}
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
|
||||
type tun struct {
|
||||
*tunFile
|
||||
readers []*tunFile
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
readers tio.QueueSet
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
vnetHdr bool
|
||||
offloadFlags uint
|
||||
|
||||
Routes atomic.Pointer[[]Route]
|
||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||
@@ -239,56 +74,110 @@ type ifreqQLEN struct {
|
||||
}
|
||||
|
||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
// We don't know what flags the caller opened this fd with and can't turn
|
||||
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||
return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
|
||||
}
|
||||
|
||||
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||
// missing (docker containers occasionally omit it).
|
||||
func openTunDev() (int, error) {
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err == nil {
|
||||
return fd, nil
|
||||
}
|
||||
if !os.IsNotExist(err) {
|
||||
return -1, err
|
||||
}
|
||||
if err = os.MkdirAll("/dev/net", 0755); err != nil {
|
||||
return -1, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||
}
|
||||
if err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))); err != nil {
|
||||
return -1, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||
}
|
||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||
}
|
||||
return fd, nil
|
||||
}
|
||||
|
||||
t.Device = "tun0"
|
||||
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
|
||||
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||
var req ifReq
|
||||
req.Flags = flags
|
||||
copy(req.Name[:], name)
|
||||
if err := ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||
}
|
||||
|
||||
return t, nil
|
||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
|
||||
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
||||
|
||||
// usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags.
|
||||
// Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO
|
||||
const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||
|
||||
func offloadUSOEnabled(offloadFlags uint) bool {
|
||||
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
|
||||
}
|
||||
|
||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
||||
if os.IsNotExist(err) {
|
||||
err = os.MkdirAll("/dev/net", 0755)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||
}
|
||||
err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200)))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||
}
|
||||
|
||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||
}
|
||||
} else {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
var req ifReq
|
||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||
// IFF_TUN_EXCL prevents us from attaching to an already-running tun
|
||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL)
|
||||
if multiqueue {
|
||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||
}
|
||||
nameStr := c.GetString("tun.dev", "")
|
||||
copy(req.Name[:], nameStr)
|
||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
|
||||
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
|
||||
var offloadFlags uint
|
||||
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, &NameError{
|
||||
Name: nameStr,
|
||||
Underlying: err,
|
||||
vnetHdr = false
|
||||
} else {
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
|
||||
offloadFlags = usoAndTSOOffloadFlags
|
||||
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
||||
offloadFlags = tsoOffloadFlags
|
||||
} else {
|
||||
l.Warn("Failed to enable TUN offload (TSO); proceeding without virtio headers", "error", err)
|
||||
_ = unix.Close(fd)
|
||||
vnetHdr = false
|
||||
}
|
||||
}
|
||||
name := strings.Trim(string(req.Name[:]), "\x00")
|
||||
|
||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
||||
if !vnetHdr {
|
||||
fd, err = openTunDev()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
name, err = tunSetIff(fd, nameStr, baseFlags)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, &NameError{Name: nameStr, Underlying: err}
|
||||
}
|
||||
}
|
||||
|
||||
if vnetHdr {
|
||||
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
|
||||
}
|
||||
|
||||
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -298,17 +187,37 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
tfd, err := newTunFd(fd)
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths.
|
||||
// It will close your files on error.
|
||||
// offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false)
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) {
|
||||
var qs tio.QueueSet
|
||||
var err error
|
||||
if vnetHdr {
|
||||
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
|
||||
} else {
|
||||
qs, err = tio.NewPollQueueSet()
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
}
|
||||
err = qs.Add(fd)
|
||||
if err != nil {
|
||||
// Add only appends on success, so closing the set here can't
|
||||
// double-close fd; it releases the set's shutdown eventfd.
|
||||
_ = unix.Close(fd)
|
||||
_ = qs.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
t := &tun{
|
||||
tunFile: tfd,
|
||||
readers: []*tunFile{tfd},
|
||||
Device: name,
|
||||
readers: qs,
|
||||
closeLock: sync.Mutex{},
|
||||
vnetHdr: vnetHdr,
|
||||
offloadFlags: offloadFlags,
|
||||
vpnNetworks: vpnNetworks,
|
||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
||||
@@ -406,36 +315,49 @@ func (t *tun) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return true
|
||||
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
|
||||
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||
for len(t.readers.Queues()) < n {
|
||||
if err := t.addQueue(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return t.readers.Queues(), nil
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
|
||||
func (t *tun) addQueue() error {
|
||||
t.closeLock.Lock()
|
||||
defer t.closeLock.Unlock()
|
||||
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
var req ifReq
|
||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||
copy(req.Name[:], t.Device)
|
||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||
if t.vnetHdr {
|
||||
flags |= unix.IFF_VNET_HDR
|
||||
}
|
||||
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
out, err := t.tunFile.newFriend(fd)
|
||||
if t.vnetHdr {
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = t.readers.Add(fd)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
t.readers = append(t.readers, out)
|
||||
|
||||
return out, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||
@@ -603,6 +525,13 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
||||
Table: unix.RT_TABLE_MAIN,
|
||||
Type: unix.RTN_UNICAST,
|
||||
}
|
||||
// Match the metric the kernel uses for its auto-installed connected route,
|
||||
// so RouteReplace overwrites it in place instead of adding a second route at a worse metric.
|
||||
// IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0.
|
||||
// Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6.
|
||||
if cidr.Addr().Is6() {
|
||||
nr.Priority = 256
|
||||
}
|
||||
err := netlink.RouteReplace(&nr)
|
||||
if err != nil {
|
||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||
@@ -878,32 +807,10 @@ func (t *tun) Close() error {
|
||||
t.routeChan = nil
|
||||
}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit
|
||||
_ = t.tunFile.wakeForShutdown()
|
||||
|
||||
if t.ioctlFd > 0 {
|
||||
_ = unix.Close(int(t.ioctlFd))
|
||||
t.ioctlFd = 0
|
||||
}
|
||||
|
||||
for i := range t.readers {
|
||||
if i == 0 {
|
||||
continue //we want to close the zeroth reader last
|
||||
}
|
||||
err := t.readers[i].Close()
|
||||
if err != nil {
|
||||
t.l.Error("error closing tun reader", "reader", i, "error", err)
|
||||
} else {
|
||||
t.l.Info("closed tun reader", "reader", i)
|
||||
}
|
||||
}
|
||||
|
||||
//this is t.readers[0] too
|
||||
err := t.tunFile.Close()
|
||||
if err != nil {
|
||||
t.l.Error("error closing tun reader", "reader", 0, "error", err)
|
||||
} else {
|
||||
t.l.Info("closed tun reader", "reader", 0)
|
||||
}
|
||||
return err
|
||||
return t.readers.Close()
|
||||
}
|
||||
|
||||
@@ -3,7 +3,9 @@
|
||||
|
||||
package overlay
|
||||
|
||||
import "testing"
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
var runAdvMSSTests = []struct {
|
||||
name string
|
||||
@@ -32,3 +34,65 @@ func TestTunAdvMSS(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestOffloadUSOEnabled pins the single source of truth for the per-queue USO
|
||||
// capability: it is derived from the negotiated offload mask, so the mask
|
||||
// stored on the tun and the capability reported to coalescers cannot drift.
|
||||
func TestOffloadUSOEnabled(t *testing.T) {
|
||||
// usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
||||
// addQueue) would not actually be a downgrade.
|
||||
if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags)
|
||||
}
|
||||
if usoAndTSOOffloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
offloadFlags uint
|
||||
wantUSO bool
|
||||
}{
|
||||
{"uso-negotiated", usoAndTSOOffloadFlags, true},
|
||||
{"tso-fallback", tsoOffloadFlags, false},
|
||||
{"no-vnet-hdr", 0, false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := offloadUSOEnabled(tc.offloadFlags); got != tc.wantUSO {
|
||||
t.Fatalf("offloadUSOEnabled(%#x) = %v, want %v", tc.offloadFlags, got, tc.wantUSO)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddQueueReplaysNegotiatedMask guards the device-wide TUNSETOFFLOAD
|
||||
// downgrade bug: addQueue must issue the exact mask newTun negotiated
|
||||
// (t.offloadFlags), not a hardcoded TSO-only mask. Because TUNSETOFFLOAD is
|
||||
// per-netdev, a narrower mask on an added queue silently disables USO for
|
||||
// every queue on a USO-capable kernel while the queues keep advertising it.
|
||||
//
|
||||
// A full multi-queue exercise needs /dev/net/tun and CAP_NET_ADMIN, which are
|
||||
// not available in CI/sandbox, so this asserts on the struct field that the
|
||||
// TUNSETOFFLOAD argument is read from.
|
||||
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
||||
t.Run("uso-negotiated", func(t *testing.T) {
|
||||
tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
|
||||
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
||||
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
||||
// mask (the original bug).
|
||||
if tn.offloadFlags != usoAndTSOOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags)
|
||||
}
|
||||
if tn.offloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
||||
}
|
||||
})
|
||||
t.Run("tso-fallback", func(t *testing.T) {
|
||||
tn := &tun{vnetHdr: true, offloadFlags: tsoOffloadFlags}
|
||||
if tn.offloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, tsoOffloadFlags)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -17,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -390,12 +390,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -17,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -138,8 +138,8 @@ func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||
//go:noescape
|
||||
func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||
|
||||
// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from the
|
||||
// packet so the payload lands directly in to.
|
||||
// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from
|
||||
// the packet so the payload lands directly in to.
|
||||
func (t *tun) Read(to []byte) (int, error) {
|
||||
var head [4]byte
|
||||
|
||||
@@ -369,12 +369,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
@@ -177,10 +178,6 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (t *TestTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
||||
func (t *TestTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, udp.MTU)}, nil
|
||||
}
|
||||
|
||||
+7
-11
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"crypto"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -18,6 +17,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"github.com/slackhq/nebula/wintun"
|
||||
@@ -47,6 +47,10 @@ type winTun struct {
|
||||
tun *wintun.NativeTun
|
||||
}
|
||||
|
||||
func (t *winTun) Read(b []byte) (int, error) {
|
||||
return t.tun.Read(b, 0)
|
||||
}
|
||||
|
||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
||||
}
|
||||
@@ -255,20 +259,12 @@ func (t *winTun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *winTun) Read(b []byte) (int, error) {
|
||||
return t.tun.Read(b, 0)
|
||||
}
|
||||
|
||||
func (t *winTun) Write(b []byte) (int, error) {
|
||||
return t.tun.Write(b, 0)
|
||||
}
|
||||
|
||||
func (t *winTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
||||
func (t *winTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *winTun) Close() error {
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user