mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 11:27:02 +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
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -40,7 +40,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -80,7 +80,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Import certificates
|
- name: Import certificates
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: add hashicorp source
|
- name: add hashicorp source
|
||||||
@@ -66,7 +66,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: add hashicorp source
|
- name: add hashicorp source
|
||||||
@@ -92,7 +92,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: build
|
- name: build
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Install goimports
|
- name: Install goimports
|
||||||
@@ -42,7 +42,7 @@ jobs:
|
|||||||
- name: golangci-lint
|
- name: golangci-lint
|
||||||
uses: golangci/golangci-lint-action@v9
|
uses: golangci/golangci-lint-action@v9
|
||||||
with:
|
with:
|
||||||
version: v2.5
|
version: v2.12
|
||||||
|
|
||||||
test:
|
test:
|
||||||
name: Test ${{ matrix.name }}
|
name: Test ${{ matrix.name }}
|
||||||
@@ -82,7 +82,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -127,7 +127,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build ${{ matrix.name }}
|
- name: Build ${{ matrix.name }}
|
||||||
|
|||||||
@@ -7,6 +7,88 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [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
|
## [1.10.3] - 2026-02-06
|
||||||
|
|
||||||
### Security
|
### Security
|
||||||
|
|||||||
@@ -161,6 +161,10 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
|||||||
bin-pkcs11: CGO_ENABLED = 1
|
bin-pkcs11: CGO_ENABLED = 1
|
||||||
bin-pkcs11: bin
|
bin-pkcs11: bin
|
||||||
|
|
||||||
|
# Build with the pprof debug server (serves on :6060). See startPprofServer.
|
||||||
|
debug: BUILD_ARGS += -tags debug
|
||||||
|
debug: bin
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
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
|
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 $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.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
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ func newTestLighthouse() *LightHouse {
|
|||||||
lighthouses := []netip.Addr{}
|
lighthouses := []netip.Addr{}
|
||||||
staticList := map[netip.Addr]struct{}{}
|
staticList := map[netip.Addr]struct{}{}
|
||||||
|
|
||||||
|
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
|
||||||
lh.lighthouses.Store(&lighthouses)
|
lh.lighthouses.Store(&lighthouses)
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
|
|
||||||
|
|||||||
+19
-6
@@ -12,7 +12,15 @@ import (
|
|||||||
"github.com/slackhq/nebula/noiseutil"
|
"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 {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey noiseutil.CipherState
|
||||||
@@ -24,6 +32,8 @@ type ConnectionState struct {
|
|||||||
window *Bits
|
window *Bits
|
||||||
decryptLock sync.Mutex
|
decryptLock sync.Mutex
|
||||||
writeLock 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
|
// 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),
|
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
|
epoch: sessionEpoch.Add(1),
|
||||||
}
|
}
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
ci.messageCounter.Add(r.MessageIndex)
|
||||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||||
@@ -58,8 +69,7 @@ func (cs *ConnectionState) Curve() cert.Curve {
|
|||||||
return cs.myCert.Curve()
|
return cs.myCert.Curve()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) {
|
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
|
||||||
var err error
|
|
||||||
cs.decryptLock.Lock()
|
cs.decryptLock.Lock()
|
||||||
result := cs.window.Check(l, messageCounter)
|
result := cs.window.Check(l, messageCounter)
|
||||||
cs.decryptLock.Unlock()
|
cs.decryptLock.Unlock()
|
||||||
@@ -67,7 +77,7 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []
|
|||||||
return nil, ErrAlreadySeen
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -81,7 +91,6 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []
|
|||||||
return out, nil
|
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 {
|
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||||
cs.decryptLock.Lock()
|
cs.decryptLock.Lock()
|
||||||
result := cs.window.Check(l, messageCounter)
|
result := cs.window.Check(l, messageCounter)
|
||||||
@@ -90,6 +99,11 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa
|
|||||||
return ErrAlreadySeen
|
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()]
|
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
_, 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 {
|
if !result {
|
||||||
return ErrAlreadySeen
|
return ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
+9
-1
@@ -53,6 +53,7 @@ type Control struct {
|
|||||||
statsStart func()
|
statsStart func()
|
||||||
dnsStart func()
|
dnsStart func()
|
||||||
lighthouseStart func()
|
lighthouseStart func()
|
||||||
|
networkChangeStart func(rebind func())
|
||||||
connectionManagerStart func(context.Context)
|
connectionManagerStart func(context.Context)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -104,6 +105,9 @@ func (c *Control) Start() error {
|
|||||||
if c.dnsStart != nil {
|
if c.dnsStart != nil {
|
||||||
go c.dnsStart()
|
go c.dnsStart()
|
||||||
}
|
}
|
||||||
|
if c.networkChangeStart != nil {
|
||||||
|
go c.networkChangeStart(c.RebindUDPServer)
|
||||||
|
}
|
||||||
if c.connectionManagerStart != nil {
|
if c.connectionManagerStart != nil {
|
||||||
go c.connectionManagerStart(c.ctx)
|
go c.connectionManagerStart(c.ctx)
|
||||||
}
|
}
|
||||||
@@ -198,7 +202,11 @@ func (c *Control) RebindUDPServer() {
|
|||||||
return
|
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
|
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||||
c.f.lightHouse.SendUpdate()
|
c.f.lightHouse.SendUpdate()
|
||||||
|
|||||||
+34
-17
@@ -11,6 +11,8 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"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/routing"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/slackhq/nebula/udp"
|
"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
|
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||||
// the same way a closed device does
|
// the same way a closed device does
|
||||||
func (d *fakeDevice) Read(p []byte) (int, error) {
|
func (d *fakeDevice) Read() ([]tio.Packet, error) {
|
||||||
<-d.closedCh
|
<-d.closedCh
|
||||||
return 0, io.EOF
|
return nil, io.EOF
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
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) Networks() []netip.Prefix { return nil }
|
||||||
func (d *fakeDevice) Name() string { return "fake" }
|
func (d *fakeDevice) Name() string { return "fake" }
|
||||||
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||||
func (d *fakeDevice) SupportsMultiqueue() bool { return false }
|
|
||||||
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
|
||||||
return nil, errors.New("unsupported")
|
|
||||||
}
|
|
||||||
|
|
||||||
// newReadyControl hand-builds the minimum Control that Main would have
|
// newReadyControl hand-builds the minimum Control that Main would have
|
||||||
// produced right before Start, including the construction token NewInterface
|
// produced right before Start, including the construction token NewInterface
|
||||||
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
readers: make([]io.ReadWriteCloser, 1),
|
batchers: make([]*batch.MultiCoalescer, 1),
|
||||||
routines: 1,
|
routines: 1,
|
||||||
hostMap: newHostMap(l),
|
hostMap: newHostMap(l),
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -109,7 +109,8 @@ func TestControl_StopBeforeStart(t *testing.T) {
|
|||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
|
|
||||||
// A stopped control can never be started
|
// 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
|
// A second Stop is a harmless no-op
|
||||||
c.Stop()
|
c.Stop()
|
||||||
@@ -145,8 +146,11 @@ type fakeConn struct {
|
|||||||
|
|
||||||
func (c *fakeConn) Rebind() error { c.rebinds++; 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) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) 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) ReloadConfig(_ *config.C) {}
|
||||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||||
@@ -155,7 +159,14 @@ type multiqueueDevice struct {
|
|||||||
*fakeDevice
|
*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) {
|
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||||
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||||
@@ -166,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
readers: make([]io.ReadWriteCloser, 2),
|
batchers: make([]*batch.MultiCoalescer, 2),
|
||||||
routines: 2,
|
routines: 2,
|
||||||
l: test.NewLogger(),
|
l: test.NewLogger(),
|
||||||
}
|
}
|
||||||
@@ -181,7 +192,8 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// The second reader fails to open, everything must be released
|
// 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.Equal(t, StateStopped, c.State())
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
assert.True(t, conn.closed, "the udp socket 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
|
// panic and Wait must observe the final state
|
||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
assert.Equal(t, StateStopped, c.State())
|
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) {
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
c, dev, conn := newReadyControl(t)
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
require.NoError(t, c.Start())
|
err := c.Start()
|
||||||
|
require.NoError(t, err)
|
||||||
assert.Equal(t, StateStarted, c.State())
|
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
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
c.Stop()
|
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
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
require.NoError(t, c.Wait())
|
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) {
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
@@ -280,7 +296,8 @@ func TestControl_RebindIsGatedByState(t *testing.T) {
|
|||||||
c.RebindUDPServer()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
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()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
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 {
|
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 {
|
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
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
@@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *slog.Logger {
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
return slog.New(slog.DiscardHandler)
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
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
|
||||||
|
}
|
||||||
+140
-21
@@ -114,6 +114,28 @@ type packet struct {
|
|||||||
packet *udp.Packet
|
packet *udp.Packet
|
||||||
tun bool // a packet pulled off a tun device
|
tun bool // a packet pulled off a tun device
|
||||||
rx bool // the packet was received by a udp 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() {
|
func (p *packet) WasReceived() {
|
||||||
@@ -131,6 +153,9 @@ const (
|
|||||||
ExitNow ExitType = 1
|
ExitNow ExitType = 1
|
||||||
// RouteAndExit routes this packet and exits immediately afterwards
|
// RouteAndExit routes this packet and exits immediately afterwards
|
||||||
RouteAndExit ExitType = 2
|
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
|
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 {
|
func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
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)
|
panic(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -152,7 +179,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
outNat: make(map[outNatKey]netip.AddrPort),
|
outNat: make(map[outNatKey]netip.AddrPort),
|
||||||
flow: []flowEntry{},
|
flow: []flowEntry{},
|
||||||
ignoreFlows: []ignoreFlow{},
|
ignoreFlows: []ignoreFlow{},
|
||||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
fn: fn,
|
||||||
t: t,
|
t: t,
|
||||||
cancelRender: cancel,
|
cancelRender: cancel,
|
||||||
}
|
}
|
||||||
@@ -249,7 +276,7 @@ func (r *R) renderFlow() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
addr := e.packet.from.GetUDPAddr()
|
addr := e.packet.fromAddr()
|
||||||
if _, ok := participants[addr]; ok {
|
if _, ok := participants[addr]; ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -268,7 +295,6 @@ func (r *R) renderFlow() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Print packets
|
// Print packets
|
||||||
h := &header.H{}
|
|
||||||
for _, e := range r.flow {
|
for _, e := range r.flow {
|
||||||
if e.packet == nil {
|
if e.packet == nil {
|
||||||
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
//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))
|
fmt.Fprintln(f, r.formatUdpPacket(p))
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
if err := h.Parse(p.packet.Data); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
line := "--x"
|
line := "--x"
|
||||||
if p.rx {
|
if p.rx {
|
||||||
line = "->>"
|
line = "->>"
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Fprintf(f,
|
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
||||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
||||||
normalizeName(p.from.GetUDPAddr().String()),
|
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,
|
line,
|
||||||
normalizeName(p.to.GetUDPAddr().String()),
|
normalizeName(p.toAddr().String()),
|
||||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
detail,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -408,21 +435,24 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
|||||||
|
|
||||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||||
|
|
||||||
if len(r.ignoreFlows) > 0 {
|
|
||||||
var h header.H
|
var h header.H
|
||||||
err := h.Parse(p.Data)
|
var parseErr error
|
||||||
if err != nil {
|
if !tun {
|
||||||
panic(err)
|
parseErr = h.Parse(p.Data)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||||
for _, i := range r.ignoreFlows {
|
for _, i := range r.ignoreFlows {
|
||||||
if !tun {
|
if tun {
|
||||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
if i.tun.HasValue && i.tun.IsTrue {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
continue
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -431,6 +461,8 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
|||||||
to: to,
|
to: to,
|
||||||
packet: p.Copy(),
|
packet: p.Copy(),
|
||||||
tun: tun,
|
tun: tun,
|
||||||
|
h: h,
|
||||||
|
parseErr: parseErr,
|
||||||
}
|
}
|
||||||
|
|
||||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||||
@@ -660,6 +692,10 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
p.Release()
|
p.Release()
|
||||||
return
|
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:
|
case KeepRouting:
|
||||||
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
||||||
receiver.InjectUDPPacket(p)
|
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) {
|
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
||||||
@@ -782,6 +897,10 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
p.Release()
|
p.Release()
|
||||||
return
|
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:
|
case KeepRouting:
|
||||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
|
|||||||
@@ -131,6 +131,9 @@ listen:
|
|||||||
port: 4242
|
port: 4242
|
||||||
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
|
# 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
|
# 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
|
#batch: 64
|
||||||
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
|
# 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)
|
# 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.
|
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
||||||
#windows_bypass_wdf: true
|
#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
|
# 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
|
# 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.
|
# 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
|
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
||||||
mtu: 1300
|
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
|
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
||||||
routes:
|
routes:
|
||||||
#- mtu: 8800
|
#- mtu: 8800
|
||||||
|
|||||||
+4
-2
@@ -5,6 +5,8 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// 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 {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"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) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
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) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
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{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
@@ -65,3 +65,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
|||||||
"Fragment": fp.Fragment,
|
"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
|
module github.com/slackhq/nebula
|
||||||
|
|
||||||
go 1.25.0
|
go 1.26.0
|
||||||
|
|
||||||
require (
|
require (
|
||||||
dario.cat/mergo v1.0.2
|
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)
|
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||||
err := hm.outside.WriteTo(stage0, addr)
|
err := hm.outside.WriteTo(stage0, addr)
|
||||||
if err != nil {
|
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,
|
"udpAddr", addr,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", hsFields,
|
||||||
@@ -969,6 +975,9 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
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)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
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
|
// We received a valid handshake on this relay, so make sure the relay
|
||||||
// state reflects that, in case it had been marked Disestablished.
|
// state reflects that, in case it had been marked Disestablished.
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
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])...)
|
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
|
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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+11
-6
@@ -190,14 +190,19 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||||
if n, ok := subTypeMap[t]; ok {
|
switch t {
|
||||||
if _, ok := (*n)[s]; ok {
|
case Message:
|
||||||
return true
|
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
|
// NewHeader turns bytes into a header
|
||||||
func NewHeader(b []byte) (*H, error) {
|
func NewHeader(b []byte) (*H, error) {
|
||||||
|
|||||||
@@ -102,6 +102,57 @@ func TestTypeMap(t *testing.T) {
|
|||||||
}, subTypeMap)
|
}, 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) {
|
func TestHeader_String(t *testing.T) {
|
||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
|
|||||||
+11
@@ -543,6 +543,17 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|||||||
return final
|
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 {
|
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||||
hm.RLock()
|
hm.RLock()
|
||||||
if h, ok := hm.Indexes[index]; ok {
|
if h, ok := hm.Indexes[index]; ok {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,10 +10,24 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"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)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
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
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
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 {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -52,12 +74,24 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// 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 {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -71,12 +105,11 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
return
|
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 {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"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) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.OutboundSendReject {
|
if !f.firewall.OutboundSendReject {
|
||||||
return
|
return
|
||||||
@@ -96,33 +248,36 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := f.readers[q].Write(out)
|
_, err := f.queues[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
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 {
|
if !f.firewall.InboundSendReject {
|
||||||
return
|
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 {
|
if len(out) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(out) > iputil.MaxRejectPacketSize {
|
if len(out) > iputil.MaxRejectPacketSize {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||||
f.l.Info("rejectOutside: packet too big, not sending",
|
f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out)
|
||||||
"packet", packet,
|
|
||||||
"outPacket", out,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
return
|
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
|
// 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) {
|
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)
|
err := newPacket(p, false, fp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
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
|
// 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 dropReason != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping cached packet",
|
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)
|
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
|
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||||
// 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,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) {
|
) ([]byte, error) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -311,7 +458,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"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.
|
// 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 {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
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
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.GetRemote())
|
|
||||||
|
err = f.writers[q].WriteTo(toSend, via.GetRemote())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
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) {
|
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 {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", remote,
|
"udpAddr", hr,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
@@ -423,7 +588,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
)
|
)
|
||||||
continue
|
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
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+159
-38
@@ -4,9 +4,9 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"runtime"
|
||||||
"slices"
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -14,12 +14,15 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -49,6 +52,18 @@ type InterfaceConfig struct {
|
|||||||
reQueryWait time.Duration
|
reQueryWait time.Duration
|
||||||
|
|
||||||
ConntrackCacheTimeout time.Duration
|
ConntrackCacheTimeout time.Duration
|
||||||
|
|
||||||
|
// 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
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -73,6 +88,15 @@ type Interface struct {
|
|||||||
routines int
|
routines int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
|
// 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
|
relayManager *relayManager
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
@@ -90,7 +114,13 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []io.ReadWriteCloser
|
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
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
@@ -102,18 +132,13 @@ type Interface struct {
|
|||||||
metricHandshakes metrics.Histogram
|
metricHandshakes metrics.Histogram
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
cachedPacketMetrics *cachedPacketMetrics
|
cachedPacketMetrics *cachedPacketMetrics
|
||||||
|
metricTxDropped metrics.Counter
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
type EncWriter interface {
|
type EncWriter interface {
|
||||||
SendVia(via *HostInfo,
|
SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
|
||||||
relay *Relay,
|
|
||||||
ad,
|
|
||||||
nb,
|
|
||||||
out []byte,
|
|
||||||
nocopy bool,
|
|
||||||
)
|
|
||||||
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
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)
|
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
||||||
Handshake(vpnAddr netip.Addr)
|
Handshake(vpnAddr netip.Addr)
|
||||||
@@ -172,6 +197,10 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
return nil, errors.New("no connection manager")
|
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()
|
cs := c.pki.getCertState()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
@@ -189,7 +218,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]io.ReadWriteCloser, c.routines),
|
batchers: make([]*batch.MultiCoalescer, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -198,8 +227,11 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
relayManager: c.relayManager,
|
relayManager: c.relayManager,
|
||||||
connectionManager: c.connectionManager,
|
connectionManager: c.connectionManager,
|
||||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||||
|
cpuAffinity: c.CpuAffinity,
|
||||||
|
pinThreads: c.PinThreads,
|
||||||
|
|
||||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||||
|
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
|
||||||
messageMetrics: c.MessageMetrics,
|
messageMetrics: c.MessageMetrics,
|
||||||
cachedPacketMetrics: &cachedPacketMetrics{
|
cachedPacketMetrics: &cachedPacketMetrics{
|
||||||
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
||||||
@@ -240,25 +272,36 @@ func (f *Interface) activate() error {
|
|||||||
"boringcrypto", boringEnabled(),
|
"boringcrypto", boringEnabled(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if f.routines > 1 {
|
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
||||||
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
|
||||||
f.routines = 1
|
f.routines = 1
|
||||||
f.l.Warn("routines is not supported on this platform, falling back to a single routine")
|
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
// 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
|
||||||
// Prepare n tun queues
|
// size the reader routines to what we actually got.
|
||||||
var reader io.ReadWriteCloser = f.inside
|
queues, err := f.inside.Queues(f.routines)
|
||||||
for i := 0; i < f.routines; i++ {
|
|
||||||
if i > 0 {
|
|
||||||
reader, err = f.inside.NewMultiQueueReader()
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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.readers[i] = reader
|
f.queues = queues
|
||||||
|
|
||||||
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
|
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
|
// 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
|
// Launch n queues to read packets from tun dev
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
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) {
|
func (f *Interface) listenOut(i int) {
|
||||||
var li udp.Conn
|
var li udp.Conn
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
@@ -314,16 +382,20 @@ func (f *Interface) listenOut(i int) {
|
|||||||
li = f.outside
|
li = f.outside
|
||||||
}
|
}
|
||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
rxc := newRxContext(f, i)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
h := &header.H{}
|
|
||||||
fwPacket := &firewall.Packet{}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
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
|
// 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
|
// 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)
|
f.l.Debug("underlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
func (f *Interface) pinThisThread(i int) {
|
||||||
packet := make([]byte, mtu)
|
var cpu int
|
||||||
out := make([]byte, mtu)
|
if n := len(f.cpuAffinity); n > 0 {
|
||||||
fwPacket := &firewall.Packet{}
|
// 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)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
pkts, err := queue.Read()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Same shutdown noise handling as listenOut
|
// Same shutdown noise handling as listenOut
|
||||||
if !f.closed.Load() && f.ctx.Err() == nil {
|
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
@@ -355,12 +453,35 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
|||||||
break
|
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)
|
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) {
|
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||||
c.RegisterReloadCallback(f.reloadFirewall)
|
c.RegisterReloadCallback(f.reloadFirewall)
|
||||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||||
|
|||||||
+27
-3
@@ -199,7 +199,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
proto, offset, isFragment := IPv6FindUpperProtocol(packet)
|
||||||
if isFragment {
|
if isFragment {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -333,11 +333,34 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
|||||||
return out
|
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]
|
nextHeader = packet[6]
|
||||||
offset = ipv6.HeaderLen
|
offset = ipv6.HeaderLen
|
||||||
|
|
||||||
for {
|
for range maxIPv6ExtHeaders {
|
||||||
switch nextHeader {
|
switch nextHeader {
|
||||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||||
if len(packet) < offset+2 {
|
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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return nextHeader, offset, isFragment
|
||||||
}
|
}
|
||||||
|
|
||||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"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 {
|
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||||
b[0] = ipv6.Version << 4
|
b[0] = ipv6.Version << 4
|
||||||
@@ -474,3 +515,121 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
|||||||
result := CreateICMPEchoResponse(packet, out)
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
assert.Nil(t, result)
|
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
|
myVpnNetworksTable *bart.Lite
|
||||||
punchy *Punchy
|
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
|
// Local cache of answers from light houses
|
||||||
// map of vpn addr to answers
|
// map of vpn addr to answers
|
||||||
addrMap map[netip.Addr]*RemoteList
|
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)),
|
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
||||||
|
return localAddrs(h.l, al)
|
||||||
|
}
|
||||||
|
|
||||||
lighthouses := make([]netip.Addr, 0)
|
lighthouses := make([]netip.Addr, 0)
|
||||||
h.lighthouses.Store(&lighthouses)
|
h.lighthouses.Store(&lighthouses)
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
@@ -918,7 +926,7 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
lal := lh.GetLocalAllowList()
|
lal := lh.GetLocalAllowList()
|
||||||
for _, e := range localAddrs(lh.l, lal) {
|
for _, e := range lh.localAddrsFn(lal) {
|
||||||
if lh.myVpnNetworksTable.Contains(e) {
|
if lh.myVpnNetworksTable.Contains(e) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -498,7 +498,7 @@ type testEncWriter struct {
|
|||||||
protocolVersion cert.Version
|
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) {
|
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,11 +6,14 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/cpupick"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/sshd"
|
"github.com/slackhq/nebula/sshd"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -33,6 +36,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
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
|
// Print the config if in test, the exit comes later
|
||||||
if configTest {
|
if configTest {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
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++ {
|
for i := 0; i < routines; i++ {
|
||||||
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
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 {
|
if err != nil {
|
||||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
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)
|
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{
|
ifConfig := &InterfaceConfig{
|
||||||
HostMap: hostMap,
|
HostMap: hostMap,
|
||||||
Inside: tun,
|
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),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
|
CpuAffinity: cpuAffinity,
|
||||||
|
PinThreads: pinThreads,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -268,6 +297,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
attachCommands(l, c, ssh, ifce)
|
attachCommands(l, c, ssh, ifce)
|
||||||
|
|
||||||
|
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||||
|
|
||||||
return &Control{
|
return &Control{
|
||||||
state: StateReady,
|
state: StateReady,
|
||||||
f: ifce,
|
f: ifce,
|
||||||
@@ -278,10 +309,75 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
statsStart: stats.Start,
|
statsStart: stats.Start,
|
||||||
dnsStart: ds.Start,
|
dnsStart: ds.Start,
|
||||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||||
|
networkChangeStart: networkChanges.Start,
|
||||||
connectionManagerStart: connManager.Start,
|
connectionManagerStart: connManager.Start,
|
||||||
}, nil
|
}, 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 {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
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.Empty(t, out)
|
||||||
assert.Equal(t, 0, cc.Overhead())
|
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/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -22,7 +23,11 @@ const (
|
|||||||
|
|
||||||
var ErrOutOfWindow = errors.New("out of window packet")
|
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)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// 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 {
|
if isMessageRelay {
|
||||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||||
} else {
|
} 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
|
// 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
|
// All remaining packets are encrypted
|
||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
// Relay packets are special, this branch should always early-return
|
// 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) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
||||||
return
|
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 err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||||
@@ -139,7 +145,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
case header.Message:
|
case header.Message:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
|
||||||
default:
|
default:
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
@@ -147,15 +153,23 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
|
|
||||||
case header.LightHouse:
|
case header.LightHouse:
|
||||||
//TODO: assert via is not relayed
|
//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:
|
case header.Test:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.TestReply:
|
case header.TestReply:
|
||||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||||
case header.TestRequest:
|
case header.TestRequest:
|
||||||
//recycle the input packet ciphertext as our output buffer
|
const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
|
||||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, nb, packet)
|
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:
|
default:
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||||
return
|
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
|
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
// Pull the Roaming parts up here, and return in all call paths.
|
// 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 {
|
if !ok {
|
||||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
// 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.
|
// its internal mapping. This should never happen.
|
||||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
|
||||||
"relayRemoteIndex", h.RemoteIndex,
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -202,7 +215,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
f.readOutsidePackets(via, signedPayload, rxc)
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
@@ -221,8 +234,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel, rebuilding it in place.
|
// 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
|
// 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
|
fwdBuf := packet[:0]
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true)
|
//todo it would potentially be nice to batch these
|
||||||
|
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
return
|
return
|
||||||
@@ -303,7 +317,11 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
// 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 {
|
if len(data) < 1 {
|
||||||
return ErrPacketTooShort
|
return ErrPacketTooShort
|
||||||
}
|
}
|
||||||
@@ -318,7 +336,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return ErrUnknownIPVersion
|
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)
|
dataLen := len(data)
|
||||||
if dataLen < ipv6.HeaderLen {
|
if dataLen < ipv6.HeaderLen {
|
||||||
return ErrIPv6PacketTooShort
|
return ErrIPv6PacketTooShort
|
||||||
@@ -344,6 +362,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
switch proto {
|
switch proto {
|
||||||
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
||||||
fp.Protocol = uint8(proto)
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.IPHdrLen = offset
|
||||||
fp.RemotePort = 0
|
fp.RemotePort = 0
|
||||||
fp.LocalPort = 0
|
fp.LocalPort = 0
|
||||||
fp.Fragment = false
|
fp.Fragment = false
|
||||||
@@ -354,6 +373,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return ErrIPv6PacketTooShort
|
return ErrIPv6PacketTooShort
|
||||||
}
|
}
|
||||||
fp.Protocol = uint8(proto)
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.IPHdrLen = offset
|
||||||
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||||
icmptype := data[offset+1]
|
icmptype := data[offset+1]
|
||||||
switch icmptype {
|
switch icmptype {
|
||||||
@@ -371,6 +391,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fp.Protocol = uint8(proto)
|
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 {
|
if incoming {
|
||||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||||
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
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
|
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
|
// 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
|
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
|
||||||
if fragmentOffset != 0 {
|
if fragmentOffset != 0 {
|
||||||
@@ -429,7 +455,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return ErrIPv6CouldNotFindPayload
|
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?
|
// Do we at least have an ipv4 header worth of data?
|
||||||
if len(data) < ipv4.HeaderLen {
|
if len(data) < ipv4.HeaderLen {
|
||||||
return ErrIPv4PacketTooShort
|
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.
|
// Check if this is the second or further fragment of a fragmented packet.
|
||||||
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||||
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
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
|
// Firewall handles protocol checks
|
||||||
fp.Protocol = data[9]
|
fp.Protocol = data[9]
|
||||||
@@ -489,31 +519,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
||||||
err := newPacket(out, true, fwPacket)
|
err := newPacket(out, true, rxc.fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
|
||||||
"error", err,
|
|
||||||
"packet", out,
|
|
||||||
)
|
|
||||||
return
|
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 {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
|
||||||
// This gives us a buffer to build the reject packet in
|
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
|
||||||
"fwPacket", fwPacket,
|
|
||||||
"reason", dropReason,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
return
|
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 {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+90
-5
@@ -17,7 +17,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func Test_newPacket(t *testing.T) {
|
func Test_newPacket(t *testing.T) {
|
||||||
p := &firewall.Packet{}
|
p := &firewall.ParsedPacket{}
|
||||||
|
|
||||||
// length fails
|
// length fails
|
||||||
err := newPacket([]byte{}, true, p)
|
err := newPacket([]byte{}, true, p)
|
||||||
@@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Test_newPacket_v6(t *testing.T) {
|
func Test_newPacket_v6(t *testing.T) {
|
||||||
p := &firewall.Packet{}
|
p := &firewall.ParsedPacket{}
|
||||||
|
|
||||||
// invalid ipv6
|
// invalid ipv6
|
||||||
ip := layers.IPv6{
|
ip := layers.IPv6{
|
||||||
@@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||||
p := &firewall.Packet{}
|
p := &firewall.ParsedPacket{}
|
||||||
|
|
||||||
ip := &layers.IPv6{
|
ip := &layers.IPv6{
|
||||||
Version: 6,
|
Version: 6,
|
||||||
@@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
secondFrag = append(secondFrag, fragHeader...)
|
secondFrag = append(secondFrag, fragHeader...)
|
||||||
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||||
|
|
||||||
fp := &firewall.Packet{}
|
fp := &firewall.ParsedPacket{}
|
||||||
|
|
||||||
b.Run("Normal", func(b *testing.B) {
|
b.Run("Normal", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
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
|
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
||||||
// on the same offset the host does.
|
// on the same offset the host does.
|
||||||
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||||
p := &firewall.Packet{}
|
p := &firewall.ParsedPacket{}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
hdrLen = 40 // IPv6 header
|
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.
|
// 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")
|
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"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"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 {
|
type Device interface {
|
||||||
io.ReadWriteCloser
|
io.Closer
|
||||||
Activate() error
|
Activate() error
|
||||||
Networks() []netip.Prefix
|
Networks() []netip.Prefix
|
||||||
Name() string
|
Name() string
|
||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
SupportsMultiqueue() bool
|
// Queues returns the device's packet queues, opening additional ones as
|
||||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
// 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
|
package overlaytest
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -31,20 +30,16 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read([]byte) (int, error) {
|
func (NoopTun) Read() ([]tio.Packet, error) {
|
||||||
return 0, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Write([]byte) (int, error) {
|
func (NoopTun) Write([]byte) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) SupportsMultiqueue() bool {
|
func (NoopTun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{NoopTun{}}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, errors.New("unsupported")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Close() error {
|
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/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
@@ -63,7 +64,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t tun) Activate() error {
|
func (t *tun) Activate() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,10 +97,6 @@ func (t *tun) Name() string {
|
|||||||
return "android"
|
return "android"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -16,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -550,7 +550,9 @@ func (t *tun) Read(to []byte) (int, error) {
|
|||||||
return n - 4, nil
|
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) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
if len(from) == 0 {
|
if len(from) == 0 {
|
||||||
return 0, syscall.EIO
|
return 0, syscall.EIO
|
||||||
@@ -606,10 +608,6 @@ func (t *tun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
|
||||||
}
|
}
|
||||||
|
|||||||
+25
-23
@@ -10,6 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -23,6 +24,23 @@ type disabledTun struct {
|
|||||||
l *slog.Logger
|
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 {
|
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
||||||
tun := &disabledTun{
|
tun := &disabledTun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
@@ -57,24 +75,6 @@ func (*disabledTun) Name() string {
|
|||||||
return "disabled"
|
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 {
|
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
||||||
out := make([]byte, len(b))
|
out := make([]byte, len(b))
|
||||||
out = iputil.CreateICMPEchoResponse(b, out)
|
out = iputil.CreateICMPEchoResponse(b, out)
|
||||||
@@ -106,12 +106,14 @@ func (t *disabledTun) Write(b []byte) (int, error) {
|
|||||||
return len(b), nil
|
return len(b), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) SupportsMultiqueue() bool {
|
func (t *disabledTun) Queues(n int) ([]tio.Queue, error) {
|
||||||
return true
|
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) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return t, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) Close() error {
|
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"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -20,7 +19,7 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -561,12 +560,8 @@ func (t *tun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
|
|||||||
+3
-6
@@ -16,6 +16,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
@@ -159,10 +160,6 @@ func (t *tun) Name() string {
|
|||||||
return "iOS"
|
return "iOS"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
|
||||||
}
|
}
|
||||||
|
|||||||
+158
-251
@@ -4,9 +4,7 @@
|
|||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -19,180 +17,15 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/vishvananda/netlink"
|
"github.com/vishvananda/netlink"
|
||||||
"golang.org/x/sys/unix"
|
"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 {
|
type tun struct {
|
||||||
*tunFile
|
readers tio.QueueSet
|
||||||
readers []*tunFile
|
|
||||||
closeLock sync.Mutex
|
closeLock sync.Mutex
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
@@ -201,6 +34,8 @@ type tun struct {
|
|||||||
TXQueueLen int
|
TXQueueLen int
|
||||||
deviceIndex int
|
deviceIndex int
|
||||||
ioctlFd uintptr
|
ioctlFd uintptr
|
||||||
|
vnetHdr bool
|
||||||
|
offloadFlags uint
|
||||||
|
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
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) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
// We don't know what flags the caller opened this fd with and can't turn
|
||||||
if err != nil {
|
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||||
return nil, err
|
return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Device = "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
|
||||||
|
}
|
||||||
|
|
||||||
return t, nil
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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) {
|
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)
|
// IFF_TUN_EXCL prevents us from attaching to an already-running tun
|
||||||
if err != nil {
|
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL)
|
||||||
// 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)
|
|
||||||
if multiqueue {
|
if multiqueue {
|
||||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||||
}
|
}
|
||||||
nameStr := c.GetString("tun.dev", "")
|
nameStr := c.GetString("tun.dev", "")
|
||||||
copy(req.Name[:], nameStr)
|
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, &NameError{
|
|
||||||
Name: nameStr,
|
|
||||||
Underlying: err,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
name := strings.Trim(string(req.Name[:]), "\x00")
|
|
||||||
|
|
||||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
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)
|
||||||
|
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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -298,17 +187,37 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
// newTunGeneric does all the stuff common to different tun initialization paths.
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
// It will close your files on error.
|
||||||
tfd, err := newTunFd(fd)
|
// 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 {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
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{
|
t := &tun{
|
||||||
tunFile: tfd,
|
Device: name,
|
||||||
readers: []*tunFile{tfd},
|
readers: qs,
|
||||||
closeLock: sync.Mutex{},
|
closeLock: sync.Mutex{},
|
||||||
|
vnetHdr: vnetHdr,
|
||||||
|
offloadFlags: offloadFlags,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
|
||||||
return true
|
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()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
|
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var req ifReq
|
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
if t.vnetHdr {
|
||||||
copy(req.Name[:], t.Device)
|
flags |= unix.IFF_VNET_HDR
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
}
|
||||||
|
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = 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 {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
t.readers = append(t.readers, out)
|
return nil
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
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,
|
Table: unix.RT_TABLE_MAIN,
|
||||||
Type: unix.RTN_UNICAST,
|
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)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
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
|
t.routeChan = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit
|
|
||||||
_ = t.tunFile.wakeForShutdown()
|
|
||||||
|
|
||||||
if t.ioctlFd > 0 {
|
if t.ioctlFd > 0 {
|
||||||
_ = unix.Close(int(t.ioctlFd))
|
_ = unix.Close(int(t.ioctlFd))
|
||||||
t.ioctlFd = 0
|
t.ioctlFd = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := range t.readers {
|
return t.readers.Close()
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,7 +3,9 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
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 (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,6 +16,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -390,12 +390,8 @@ func (t *tun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,6 +16,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -138,8 +138,8 @@ func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error)
|
|||||||
//go:noescape
|
//go:noescape
|
||||||
func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error)
|
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
|
// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from
|
||||||
// packet so the payload lands directly in to.
|
// the packet so the payload lands directly in to.
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
var head [4]byte
|
var head [4]byte
|
||||||
|
|
||||||
@@ -369,12 +369,8 @@ func (t *tun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) SupportsMultiqueue() bool {
|
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
@@ -177,10 +178,6 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
|||||||
return n, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) SupportsMultiqueue() bool {
|
func (t *TestTun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, udp.MTU)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
|
||||||
}
|
}
|
||||||
|
|||||||
+7
-11
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"crypto"
|
"crypto"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -18,6 +17,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/slackhq/nebula/wintun"
|
"github.com/slackhq/nebula/wintun"
|
||||||
@@ -47,6 +47,10 @@ type winTun struct {
|
|||||||
tun *wintun.NativeTun
|
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) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||||
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
||||||
}
|
}
|
||||||
@@ -255,20 +259,12 @@ func (t *winTun) Name() string {
|
|||||||
return t.Device
|
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) {
|
func (t *winTun) Write(b []byte) (int, error) {
|
||||||
return t.tun.Write(b, 0)
|
return t.tun.Write(b, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) SupportsMultiqueue() bool {
|
func (t *winTun) Queues(int) ([]tio.Queue, error) {
|
||||||
return false
|
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||||
}
|
|
||||||
|
|
||||||
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) Close() error {
|
func (t *winTun) Close() error {
|
||||||
|
|||||||
+12
-5
@@ -6,6 +6,7 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -46,12 +47,16 @@ func (d *UserDevice) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return routing.Gateways{routing.NewGateway(ip, 1)}
|
return routing.Gateways{routing.NewGateway(ip, 1)}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) SupportsMultiqueue() bool {
|
func (d *UserDevice) Queues(n int) ([]tio.Queue, error) {
|
||||||
return true
|
out := make([]tio.Queue, n)
|
||||||
|
for i := range out {
|
||||||
|
// All queues share the underlying pipes (the io.Pipe serializes
|
||||||
|
// concurrent callers) but each owns a private scratch buffer so
|
||||||
|
// concurrent Reads across queues never alias. NoClose: the pipes are
|
||||||
|
// owned by the UserDevice and torn down once by UserDevice.Close.
|
||||||
|
out[i] = tio.NewSingleQueueNoClose(d, defaultBatchBufSize)
|
||||||
}
|
}
|
||||||
|
return out, nil
|
||||||
func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|
||||||
return d, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
||||||
@@ -61,9 +66,11 @@ func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
|||||||
func (d *UserDevice) Read(p []byte) (n int, err error) {
|
func (d *UserDevice) Read(p []byte) (n int, err error) {
|
||||||
return d.outboundReader.Read(p)
|
return d.outboundReader.Read(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
||||||
return d.inboundWriter.Write(p)
|
return d.inboundWriter.Write(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Close() error {
|
func (d *UserDevice) Close() error {
|
||||||
d.inboundWriter.Close()
|
d.inboundWriter.Close()
|
||||||
d.outboundWriter.Close()
|
d.outboundWriter.Close()
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user