mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 21:46:59 +02:00
Compare commits
133 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 | |||
| c2fbe215e6 | |||
| 94ac6db4ca | |||
| a60350e34e | |||
| 58f3b6fda7 | |||
| a99699e370 | |||
| 3615a79b8b | |||
| 147c202c27 | |||
| e290a6892f |
@@ -12,9 +12,9 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
@@ -38,9 +38,9 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
@@ -78,9 +78,9 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
|
|||||||
@@ -32,9 +32,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
@@ -64,9 +64,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
@@ -90,9 +90,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
|
|||||||
@@ -20,9 +20,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
|
|||||||
@@ -20,9 +20,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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 }}
|
||||||
@@ -80,9 +80,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
@@ -125,9 +125,9 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- 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
|
||||||
|
|||||||
@@ -0,0 +1,96 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
cert_test "github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
|
||||||
|
// a library, and on a config update dnclient calls Stop() in-process to tear the
|
||||||
|
// old instance down before starting a new one. This boots a real nebula (real
|
||||||
|
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
|
||||||
|
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
|
||||||
|
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
|
||||||
|
// dump instead of relying on a process signal to unstick them.
|
||||||
|
func TestControlStopClosesOnTimer(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
before := time.Now().Add(-time.Hour)
|
||||||
|
after := time.Now().Add(time.Hour)
|
||||||
|
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
|
||||||
|
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
||||||
|
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
|
||||||
|
|
||||||
|
caPath := filepath.Join(dir, "ca.pem")
|
||||||
|
certPath := filepath.Join(dir, "cert.pem")
|
||||||
|
keyPath := filepath.Join(dir, "key.pem")
|
||||||
|
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
|
||||||
|
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
|
||||||
|
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
|
||||||
|
|
||||||
|
// tun disabled so no device/root is needed; routines: 2 so we exercise the
|
||||||
|
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
|
||||||
|
configBody := fmt.Sprintf(`
|
||||||
|
pki:
|
||||||
|
ca: %s
|
||||||
|
cert: %s
|
||||||
|
key: %s
|
||||||
|
listen:
|
||||||
|
host: 127.0.0.1
|
||||||
|
port: 0
|
||||||
|
tun:
|
||||||
|
disabled: true
|
||||||
|
firewall:
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
routines: 2
|
||||||
|
`, caPath, certPath, keyPath)
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
|
||||||
|
|
||||||
|
c := config.NewC(l)
|
||||||
|
require.NoError(t, c.Load(dir))
|
||||||
|
|
||||||
|
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, ctrl.Start())
|
||||||
|
|
||||||
|
// Run like a live nebula, then close on a timer, exactly as dnclient does.
|
||||||
|
<-time.NewTimer(5 * time.Second).C
|
||||||
|
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
|
||||||
|
ctrl.Wait() // blocks until every reader goroutine has returned
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-stopped:
|
||||||
|
t.Log("nebula closed cleanly on timer")
|
||||||
|
case <-time.After(10 * time.Second):
|
||||||
|
buf := make([]byte, 1<<20)
|
||||||
|
n := runtime.Stack(buf, true)
|
||||||
|
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -2,16 +2,26 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"log/slog"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
"github.com/slackhq/nebula/handshake"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 8192
|
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
|
||||||
dKey noiseutil.CipherState
|
dKey noiseutil.CipherState
|
||||||
@@ -20,7 +30,10 @@ type ConnectionState struct {
|
|||||||
initiator bool
|
initiator bool
|
||||||
messageCounter atomic.Uint64
|
messageCounter atomic.Uint64
|
||||||
window *Bits
|
window *Bits
|
||||||
|
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
|
||||||
@@ -35,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++ {
|
||||||
@@ -54,3 +68,54 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
|||||||
func (cs *ConnectionState) Curve() cert.Curve {
|
func (cs *ConnectionState) Curve() cert.Curve {
|
||||||
return cs.myCert.Curve()
|
return cs.myCert.Curve()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result := cs.window.Check(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return nil, ErrAlreadySeen
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result = cs.window.Update(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return nil, ErrAlreadySeen
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result := cs.window.Check(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return ErrAlreadySeen
|
||||||
|
}
|
||||||
|
|
||||||
|
// The entire body is sent as AD, not encrypted.
|
||||||
|
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||||
|
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||||
|
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||||
|
// which will gracefully fail in the DecryptDanger call.
|
||||||
|
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||||
|
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||||
|
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result = cs.window.Update(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return ErrAlreadySeen
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
+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()
|
||||||
|
|||||||
@@ -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},
|
||||||
batchers: make([]batch.RxBatcher, 1),
|
batchers: make([]*batch.MultiCoalescer, 1),
|
||||||
routines: 1,
|
routines: 1,
|
||||||
hostMap: newHostMap(l),
|
hostMap: newHostMap(l),
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -148,8 +148,8 @@ func (c *fakeConn) Rebind() error { c.rebinds++; ret
|
|||||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) 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(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||||
return nil
|
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 }
|
||||||
@@ -177,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},
|
||||||
batchers: make([]batch.RxBatcher, 2),
|
batchers: make([]*batch.MultiCoalescer, 2),
|
||||||
routines: 2,
|
routines: 2,
|
||||||
l: test.NewLogger(),
|
l: test.NewLogger(),
|
||||||
}
|
}
|
||||||
|
|||||||
+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)
|
||||||
|
}
|
||||||
+16
-7
@@ -97,8 +97,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
newAddr := getDnsServerAddr(c)
|
newAddr := getDnsServerAddr(c)
|
||||||
|
|
||||||
d.serverMu.Lock()
|
d.serverMu.Lock()
|
||||||
running := d.server
|
running := d.server != nil
|
||||||
runningStarted := d.started
|
|
||||||
sameAddr := d.addr == newAddr
|
sameAddr := d.addr == newAddr
|
||||||
d.addr = newAddr
|
d.addr = newAddr
|
||||||
d.enabled.Store(enabled)
|
d.enabled.Store(enabled)
|
||||||
@@ -112,7 +111,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !enabled {
|
if !enabled {
|
||||||
if running != nil {
|
if running {
|
||||||
d.Stop()
|
d.Stop()
|
||||||
}
|
}
|
||||||
// Drop any records that accumulated while enabled; a later re-enable
|
// Drop any records that accumulated while enabled; a later re-enable
|
||||||
@@ -121,12 +120,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if running == nil {
|
if !running {
|
||||||
// Was disabled (or never started); bring it up now.
|
// Was disabled (or never started); bring it up now.
|
||||||
go d.Start()
|
go d.Start()
|
||||||
} else if !sameAddr {
|
} else if !sameAddr {
|
||||||
d.shutdownServer(running, runningStarted, "reload")
|
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
||||||
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
d.Stop()
|
||||||
go d.Start()
|
go d.Start()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -162,7 +161,9 @@ func (d *dnsServer) Start() {
|
|||||||
|
|
||||||
started := make(chan struct{})
|
started := make(chan struct{})
|
||||||
d.serverMu.Lock()
|
d.serverMu.Lock()
|
||||||
if d.ctx.Err() != nil {
|
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
||||||
|
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
|
||||||
|
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
|
||||||
d.serverMu.Unlock()
|
d.serverMu.Unlock()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -200,6 +201,14 @@ func (d *dnsServer) Start() {
|
|||||||
close(started)
|
close(started)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
||||||
|
d.serverMu.Lock()
|
||||||
|
if d.server == server {
|
||||||
|
d.server = nil
|
||||||
|
d.started = nil
|
||||||
|
}
|
||||||
|
d.serverMu.Unlock()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+206
-4
@@ -194,14 +194,51 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
ds, c := newTestDnsServer(t)
|
ds, c := newTestDnsServer(t)
|
||||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
|
||||||
require.NoError(t, ds.reload(c, true))
|
require.NoError(t, ds.reload(c, true))
|
||||||
// No server running yet, no addr change. Reload should not spawn anything.
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
before := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, before)
|
||||||
|
|
||||||
|
// Same address, so the running listener must be left alone rather than rebuilt under live queries
|
||||||
require.NoError(t, ds.reload(c, false))
|
require.NoError(t, ds.reload(c, false))
|
||||||
assert.True(t, ds.enabled.Load())
|
assert.True(t, ds.enabled.Load())
|
||||||
assert.Nil(t, ds.server)
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
after := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
assert.Same(t, before, after, "a same-address reload must not restart the listener")
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
|
||||||
|
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
|
||||||
|
// initial only records config, it never starts anything
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "the initial reload must not start a listener")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||||
@@ -427,3 +464,168 @@ func waitFor(t *testing.T, cond func() bool) {
|
|||||||
}
|
}
|
||||||
t.Fatal("timed out waiting for condition")
|
t.Fatal("timed out waiting for condition")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
|
||||||
|
func TestDnsServer_Start_isIdempotent(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
first := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, first)
|
||||||
|
|
||||||
|
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ds.Start()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("second Start never returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
second := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
assert.Same(t, first, second, "a second Start must not replace the running server")
|
||||||
|
|
||||||
|
// The real proof, after Stop the port must actually be free
|
||||||
|
ds.Stop()
|
||||||
|
waitFor(t, func() bool {
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = pc.Close()
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
|
||||||
|
// installed, so reload has to clear the slot before shutting the old one down.
|
||||||
|
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
|
||||||
|
first := freeUDPPort(t)
|
||||||
|
second := freeUDPPort(t)
|
||||||
|
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", first, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
|
||||||
|
for i := range 8 {
|
||||||
|
want := second
|
||||||
|
if i%2 == 1 {
|
||||||
|
want = first
|
||||||
|
}
|
||||||
|
setDnsConfig(c, "127.0.0.1", want, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
srv := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
|
||||||
|
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Land back on second so the port assertions below are meaningful
|
||||||
|
setDnsConfig(c, "127.0.0.1", second, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
// The old port must be released and the new one actually held
|
||||||
|
waitFor(t, func() bool {
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = pc.Close()
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
|
||||||
|
require.Error(t, err, "the new address should be bound by the DNS responder")
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
|
||||||
|
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
ds.Start() // returns once the bind fails
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
// With the slot released, a reload can retry once the port frees up
|
||||||
|
require.NoError(t, blocker.Close())
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
|
||||||
|
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
require.True(t, ds.enabled.Load())
|
||||||
|
|
||||||
|
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ds.Start()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
t.Fatal("Start returned early, the test never exercised the window")
|
||||||
|
case <-time.After(time.Millisecond * 100):
|
||||||
|
}
|
||||||
|
|
||||||
|
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
|
||||||
|
ds.enabled.Store(false)
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("Start never returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
require.NoError(t, err, "an orphaned listener is still holding the port")
|
||||||
|
_ = pc.Close()
|
||||||
|
}
|
||||||
|
|||||||
@@ -725,6 +725,70 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
// If them tears down the tunnel while me keeps Established relay state, me's next
|
||||||
|
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
|
||||||
|
// them's Disestablished terminal relay entry. them must re-establish that entry, or
|
||||||
|
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
|
||||||
|
// them can receive but every send is silently dropped.
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||||
|
|
||||||
|
// Teach my how to get to the relay and that their can be reached via the relay
|
||||||
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
|
||||||
|
// Build a router so we don't have to reason who gets which packet
|
||||||
|
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
// Start the servers
|
||||||
|
myControl.Start()
|
||||||
|
relayControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
|
||||||
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
|
||||||
|
|
||||||
|
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
|
||||||
|
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
|
t.Log("Re-handshake from me, riding the still-Established relay state")
|
||||||
|
myControl.ReHandshake(theirVpnIpNet[0].Addr())
|
||||||
|
for {
|
||||||
|
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||||
|
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
|
||||||
|
return router.RouteAndExit
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||||
|
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
|
||||||
|
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
|
||||||
|
|
||||||
|
t.Log("Send from them to me; their only relay entry must survive the transmit")
|
||||||
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
require.Never(t, func() bool {
|
||||||
|
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||||
|
return h == nil || len(h.CurrentRelaysToMe) == 0
|
||||||
|
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
|
||||||
|
|
||||||
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||||
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays(t *testing.T) {
|
func TestStage1RaceRelays(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
|
||||||
|
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
|
||||||
|
t.Helper()
|
||||||
|
cm := lh.QueryLighthouse(vpnAddr)
|
||||||
|
if cm == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var out []netip.AddrPort
|
||||||
|
for _, c := range *cm {
|
||||||
|
out = append(out, c.Reported...)
|
||||||
|
out = append(out, c.Learned...)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
|
||||||
|
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
|
||||||
|
t.Helper()
|
||||||
|
h := &header.H{}
|
||||||
|
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
|
if c != lh {
|
||||||
|
return router.KeepRouting
|
||||||
|
}
|
||||||
|
// Punches are a single byte and never parse, they are just not what we are after
|
||||||
|
if err := h.Parse(p.Data); err != nil {
|
||||||
|
return router.KeepRouting
|
||||||
|
}
|
||||||
|
if h.Type == header.LightHouse {
|
||||||
|
return router.RouteAndExit
|
||||||
|
}
|
||||||
|
return router.KeepRouting
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
|
||||||
|
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
|
||||||
|
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
|
||||||
|
// so we call RebindUDPServer directly, which is the same thing the monitor does.
|
||||||
|
func TestRebindSendsLighthouseUpdate(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
|
||||||
|
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
|
||||||
|
// Let the startup registration finish, then clear everything it left behind
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
// Nothing should be talking to the lighthouse on its own now
|
||||||
|
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
|
||||||
|
"nothing should reach the lighthouse before the rebind")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
|
||||||
|
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
|
||||||
|
"a rebind should push an update to the lighthouse rather than waiting out the interval")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
|
||||||
|
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
|
||||||
|
// whose remote NAT state died while we were on a different network.
|
||||||
|
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
lhCfg := m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
// Without this the peers advertise this machine's real addresses and then try to punch at them,
|
||||||
|
// which the router has no route for.
|
||||||
|
"local_allow_list": m{
|
||||||
|
"10.0.0.0/24": true,
|
||||||
|
"::/0": false,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
|
||||||
|
r.RouteFor(time.Second)
|
||||||
|
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
|
||||||
|
r.RouteFor(time.Millisecond * 300)
|
||||||
|
|
||||||
|
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
|
||||||
|
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
|
||||||
|
// so this cannot be satisfied by the update the rebind itself pushes.
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
|
||||||
|
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
|
||||||
|
"an ordinary send should not requery the lighthouse")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
|
||||||
|
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
|
||||||
|
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
|
||||||
|
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
|
||||||
|
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
|
||||||
|
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
|
||||||
|
// is picked up.
|
||||||
|
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
|
||||||
|
return []netip.Addr{myControl.GetUDPAddr().Addr()}
|
||||||
|
})
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
|
||||||
|
"the lighthouse should know the address we started on")
|
||||||
|
|
||||||
|
// Wake up somewhere else
|
||||||
|
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
|
||||||
|
myControl.SetUDPAddr(newAddr)
|
||||||
|
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
|
||||||
|
|
||||||
|
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||||
|
"the lighthouse should still be handing out the old address before the rebind")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||||
|
"after the rebind the lighthouse should hand peers our new address")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
}
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
|
||||||
|
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
|
||||||
|
// we are now, so everything we send disappears.
|
||||||
|
//
|
||||||
|
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
|
||||||
|
// since that is when its NAT opens and the tunnel is usable again.
|
||||||
|
//
|
||||||
|
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
|
||||||
|
func TestRecoveryTiming(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
rebind bool
|
||||||
|
}{
|
||||||
|
{"no trigger", false},
|
||||||
|
{"rebind counter", true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
d, lost := measureRecovery(t, tc.rebind)
|
||||||
|
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
|
||||||
|
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
|
||||||
|
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
|
||||||
|
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
|
||||||
|
t.Helper()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
peerCfg := m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
"local_allow_list": m{
|
||||||
|
"10.0.0.0/24": true,
|
||||||
|
"::/0": false,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
defer func() {
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
}()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
||||||
|
r.RouteFor(time.Second)
|
||||||
|
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
|
||||||
|
t.Fatal("failed to establish the tunnel we are measuring")
|
||||||
|
}
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
// From here the peer's NAT has no state for us, everything we send it disappears
|
||||||
|
start := time.Now()
|
||||||
|
blackholed := 0
|
||||||
|
var recovered time.Duration
|
||||||
|
|
||||||
|
if rebind {
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Keep the tun busy the way someone retrying a stalled connection would
|
||||||
|
stop := make(chan struct{})
|
||||||
|
defer close(stop)
|
||||||
|
go func() {
|
||||||
|
tick := time.NewTicker(time.Millisecond * 200)
|
||||||
|
defer tick.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-stop:
|
||||||
|
return
|
||||||
|
case <-tick.C:
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(
|
||||||
|
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
|
if c == theirControl && p.From == myControl.GetUDPAddr() {
|
||||||
|
blackholed++
|
||||||
|
return router.Drop
|
||||||
|
}
|
||||||
|
|
||||||
|
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
|
||||||
|
if c == myControl && p.From == theirUdpAddr {
|
||||||
|
recovered = time.Since(start)
|
||||||
|
return router.RouteAndExit
|
||||||
|
}
|
||||||
|
|
||||||
|
return router.KeepRouting
|
||||||
|
})
|
||||||
|
|
||||||
|
if recovered == 0 {
|
||||||
|
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
|
||||||
|
}
|
||||||
|
return recovered, blackholed
|
||||||
|
}
|
||||||
+148
-29
@@ -114,6 +114,28 @@ type packet struct {
|
|||||||
packet *udp.Packet
|
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,29 +435,34 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
|||||||
|
|
||||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||||
|
|
||||||
if len(r.ignoreFlows) > 0 {
|
var h header.H
|
||||||
var h header.H
|
var parseErr error
|
||||||
err := h.Parse(p.Data)
|
if !tun {
|
||||||
if err != nil {
|
parseErr = h.Parse(p.Data)
|
||||||
panic(err)
|
}
|
||||||
}
|
|
||||||
|
|
||||||
for _, i := range r.ignoreFlows {
|
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||||
if !tun {
|
for _, i := range r.ignoreFlows {
|
||||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
if tun {
|
||||||
return nil
|
if i.tun.HasValue && i.tun.IsTrue {
|
||||||
}
|
|
||||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// A packet we could not parse has no type to match against, so no rule can ignore it
|
||||||
|
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fp := &packet{
|
fp := &packet{
|
||||||
from: from,
|
from: from,
|
||||||
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)
|
||||||
|
|||||||
@@ -1,188 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"log/slog"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"golang.org/x/net/ipv4"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestInnerECN(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
pkt []byte
|
|
||||||
want byte
|
|
||||||
}{
|
|
||||||
{"empty", nil, 0},
|
|
||||||
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
|
||||||
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
|
||||||
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
|
||||||
{"v4_CE", v4WithToS(0x03), 0x03},
|
|
||||||
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
|
||||||
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
|
||||||
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
|
||||||
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
|
||||||
{"v6_CE", v6WithTC(0x03), 0x03},
|
|
||||||
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
|
||||||
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
t.Run(c.name, func(t *testing.T) {
|
|
||||||
got := innerECN(c.pkt)
|
|
||||||
if got != c.want {
|
|
||||||
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
|
||||||
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
|
||||||
// the DSCP and ECN portions through the byte 1 mask.
|
|
||||||
func v4WithToS(tos byte) []byte {
|
|
||||||
return []byte{0x45, tos}
|
|
||||||
}
|
|
||||||
|
|
||||||
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
|
||||||
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
|
||||||
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
|
||||||
func v6WithTC(tc byte) []byte {
|
|
||||||
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestApplyOuterECN(t *testing.T) {
|
|
||||||
silent := slog.New(slog.DiscardHandler)
|
|
||||||
hi := &HostInfo{}
|
|
||||||
|
|
||||||
// Build a v4 packet helper with a given inner ECN field.
|
|
||||||
v4 := func(innerECN byte) []byte {
|
|
||||||
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
|
||||||
return []byte{
|
|
||||||
0x45, innerECN, 0, 28,
|
|
||||||
0, 0, 0x40, 0,
|
|
||||||
64, 6, 0, 0,
|
|
||||||
10, 0, 0, 1,
|
|
||||||
10, 0, 0, 2,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
|
||||||
// TC[1:0] which sit at byte 1 mask 0x30.
|
|
||||||
v6 := func(innerECN byte) []byte {
|
|
||||||
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
|
||||||
pkt := make([]byte, 40)
|
|
||||||
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
|
||||||
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
type cell struct {
|
|
||||||
outer byte
|
|
||||||
inner byte
|
|
||||||
wantECN byte
|
|
||||||
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
|
||||||
}
|
|
||||||
|
|
||||||
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
|
||||||
table := []cell{
|
|
||||||
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
|
||||||
{ecnNotECT, ecnECT0, ecnECT0, true},
|
|
||||||
{ecnNotECT, ecnECT1, ecnECT1, true},
|
|
||||||
{ecnNotECT, ecnCE, ecnCE, true},
|
|
||||||
|
|
||||||
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
|
||||||
{ecnECT0, ecnECT0, ecnECT0, true},
|
|
||||||
{ecnECT0, ecnECT1, ecnECT1, true},
|
|
||||||
{ecnECT0, ecnCE, ecnCE, true},
|
|
||||||
|
|
||||||
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
|
||||||
{ecnECT1, ecnECT0, ecnECT0, true},
|
|
||||||
{ecnECT1, ecnECT1, ecnECT1, true},
|
|
||||||
{ecnECT1, ecnCE, ecnCE, true},
|
|
||||||
|
|
||||||
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
|
||||||
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
|
||||||
{ecnCE, ecnECT1, ecnCE, false},
|
|
||||||
{ecnCE, ecnCE, ecnCE, true},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, c := range table {
|
|
||||||
t.Run("v4", func(t *testing.T) {
|
|
||||||
pkt := v4(c.inner)
|
|
||||||
applyOuterECN(pkt, c.outer, hi, silent)
|
|
||||||
got := pkt[1] & 0x03
|
|
||||||
if got != c.wantECN {
|
|
||||||
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
t.Run("v6", func(t *testing.T) {
|
|
||||||
pkt := v6(c.inner)
|
|
||||||
applyOuterECN(pkt, c.outer, hi, silent)
|
|
||||||
got := (pkt[1] >> 4) & 0x03
|
|
||||||
if got != c.wantECN {
|
|
||||||
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
|
|
||||||
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
|
|
||||||
// The passthrough emit paths write the packet verbatim, so a stale checksum
|
|
||||||
// turns an underlay congestion mark into packet loss at the receiver.
|
|
||||||
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
|
|
||||||
silent := slog.New(slog.DiscardHandler)
|
|
||||||
hi := &HostInfo{}
|
|
||||||
|
|
||||||
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
|
|
||||||
// flips only the low two bits of the ToS byte while leaving DSCP intact.
|
|
||||||
pkt := []byte{
|
|
||||||
0x45, 0x88 | ecnECT0, 0, 40,
|
|
||||||
0x1c, 0x46, 0x40, 0x00,
|
|
||||||
64, 6, 0, 0,
|
|
||||||
10, 0, 0, 1,
|
|
||||||
10, 0, 0, 2,
|
|
||||||
}
|
|
||||||
// Stamp a correct header checksum before the fold.
|
|
||||||
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
|
|
||||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
|
||||||
t.Fatal("test setup: initial header checksum invalid")
|
|
||||||
}
|
|
||||||
|
|
||||||
applyOuterECN(pkt, ecnCE, hi, silent)
|
|
||||||
|
|
||||||
// CE folded in, DSCP preserved.
|
|
||||||
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
|
|
||||||
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
|
|
||||||
}
|
|
||||||
// The incremental RFC 1624 update must leave the checksum valid and equal
|
|
||||||
// to a full recompute over the mutated header.
|
|
||||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
|
||||||
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
|
|
||||||
}
|
|
||||||
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
|
|
||||||
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
|
|
||||||
// treating the checksum field (bytes 10:12) as zero.
|
|
||||||
func ipv4HeaderChecksum(hdr []byte) uint16 {
|
|
||||||
var sum uint32
|
|
||||||
for i := 0; i+1 < len(hdr); i += 2 {
|
|
||||||
if i == 10 {
|
|
||||||
continue // checksum field
|
|
||||||
}
|
|
||||||
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
|
|
||||||
}
|
|
||||||
for sum > 0xffff {
|
|
||||||
sum = (sum >> 16) + (sum & 0xffff)
|
|
||||||
}
|
|
||||||
return ^uint16(sum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
|
|
||||||
// computation over the header.
|
|
||||||
func ipv4HeaderChecksumValid(hdr []byte) bool {
|
|
||||||
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
|
|
||||||
}
|
|
||||||
+18
-20
@@ -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.
|
||||||
@@ -256,22 +267,19 @@ tun:
|
|||||||
|
|
||||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
# 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
|
# 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.
|
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
|
||||||
#
|
|
||||||
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
|
|
||||||
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
|
|
||||||
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
|
|
||||||
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
|
|
||||||
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
|
|
||||||
# equal N`) or set cpu_affinity explicitly to benefit.
|
|
||||||
#pin_threads: true
|
#pin_threads: true
|
||||||
|
|
||||||
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
# 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
|
# (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;
|
# 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
|
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
|
||||||
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
|
# CPUs. Only meaningful while pin_threads is true. Not reloadable.
|
||||||
# service your underlay NIC's RX queue IRQs. 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:
|
#cpu_affinity:
|
||||||
# - 2
|
# - 2
|
||||||
# - 4
|
# - 4
|
||||||
@@ -412,16 +420,6 @@ logging:
|
|||||||
# This setting is reloadable
|
# This setting is reloadable
|
||||||
#inactivity_timeout: 10m
|
#inactivity_timeout: 10m
|
||||||
|
|
||||||
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
|
|
||||||
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
|
|
||||||
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
|
|
||||||
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
|
|
||||||
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
|
|
||||||
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
|
|
||||||
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
|
|
||||||
# the route half of this setting to take effect.
|
|
||||||
#ecn: true
|
|
||||||
|
|
||||||
# Nebula security group configuration
|
# Nebula security group configuration
|
||||||
firewall:
|
firewall:
|
||||||
# Action to take when a packet is not allowed by the firewall rules.
|
# Action to take when a packet is not allowed by the firewall rules.
|
||||||
|
|||||||
@@ -8,6 +8,15 @@ Before=sshd.service
|
|||||||
Type=notify
|
Type=notify
|
||||||
NotifyAccess=main
|
NotifyAccess=main
|
||||||
SyslogIdentifier=nebula
|
SyslogIdentifier=nebula
|
||||||
|
|
||||||
|
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
|
||||||
|
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
|
||||||
|
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
|
||||||
|
#User=nebula
|
||||||
|
#Group=nebula
|
||||||
|
#CapabilityBoundingSet=CAP_NET_ADMIN
|
||||||
|
#AmbientCapabilities=CAP_NET_ADMIN
|
||||||
|
|
||||||
ExecReload=/bin/kill -HUP $MAINPID
|
ExecReload=/bin/kill -HUP $MAINPID
|
||||||
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
||||||
Restart=always
|
Restart=always
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -24,12 +24,12 @@ require (
|
|||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
go.uber.org/goleak v1.3.0
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.53.0
|
golang.org/x/crypto v0.54.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
golang.org/x/net v0.56.0
|
golang.org/x/net v0.57.0
|
||||||
golang.org/x/sync v0.21.0
|
golang.org/x/sync v0.22.0
|
||||||
golang.org/x/sys v0.46.0
|
golang.org/x/sys v0.47.0
|
||||||
golang.org/x/term v0.44.0
|
golang.org/x/term v0.45.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||||
|
|||||||
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
|||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
||||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||||
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||||
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||||
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
|
|||||||
+15
-5
@@ -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,
|
||||||
@@ -529,7 +535,9 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
|||||||
|
|
||||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
delete(hm.vpnIps, addr)
|
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
||||||
|
delete(hm.vpnIps, addr)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(hm.vpnIps) == 0 {
|
if len(hm.vpnIps) == 0 {
|
||||||
@@ -967,7 +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 sendbatcher
|
// 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)))
|
||||||
@@ -1078,8 +1088,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
|||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
// 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.remoteIdx, 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,13 +190,18 @@ 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
|
||||||
|
|||||||
@@ -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
-1
@@ -287,7 +287,6 @@ type HostInfo struct {
|
|||||||
type ViaSender struct {
|
type ViaSender struct {
|
||||||
UdpAddr netip.AddrPort
|
UdpAddr netip.AddrPort
|
||||||
relayHI *HostInfo // relayHI is the host info object of the relay
|
relayHI *HostInfo // relayHI is the host info object of the relay
|
||||||
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
|
|
||||||
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
||||||
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
||||||
}
|
}
|
||||||
@@ -544,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 {
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []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
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
// only valid until the next Read on that queue. Every consumer below
|
// only valid until the next Read on that queue. Every consumer below
|
||||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||||
@@ -74,7 +74,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
||||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
// so retaining segments past the loop is safe.
|
// so retaining segments past the loop is safe.
|
||||||
@@ -105,9 +105,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
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.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -126,7 +126,6 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
|||||||
c := ci.messageCounter.Add(1)
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
f.connectionManager.Out(hostinfo)
|
|
||||||
|
|
||||||
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
@@ -138,8 +137,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
|||||||
"udpAddr", hostinfo.GetRemote(),
|
"udpAddr", hostinfo.GetRemote(),
|
||||||
"counter", c,
|
"counter", c,
|
||||||
)
|
)
|
||||||
// Skip this segment; the rest of the superpacket can still
|
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
|
||||||
// go out — TCP will retransmit anything we drop here.
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -151,16 +149,19 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
|||||||
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||||
// kernel-supplied superpacket bytes never get written into a separate
|
// kernel-supplied superpacket bytes never get written into a separate
|
||||||
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
|
||||||
// SendBatch slot.
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
|
||||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
|
||||||
ci := hostinfo.ConnectionState
|
ci := hostinfo.ConnectionState
|
||||||
if ci.eKey == nil {
|
if ci.eKey == nil {
|
||||||
return
|
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()
|
remote := hostinfo.GetRemote()
|
||||||
ecnEnabled := f.ecnEnabled.Load()
|
|
||||||
if hostinfo.lastRebindCount != f.rebindCount {
|
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
|
//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.
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
@@ -211,11 +212,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var ecn byte
|
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
|
||||||
if ecnEnabled {
|
|
||||||
ecn = innerECN(seg)
|
|
||||||
}
|
|
||||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
|
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -233,36 +230,14 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var ecn byte
|
sendBatch.Commit(out, remote)
|
||||||
if ecnEnabled {
|
|
||||||
ecn = innerECN(seg)
|
|
||||||
}
|
|
||||||
sendBatch.Commit(out, remote, ecn)
|
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
|
||||||
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
|
||||||
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
|
||||||
func innerECN(pkt []byte) byte {
|
|
||||||
if len(pkt) < 2 {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
return pkt[1] & 0x03
|
|
||||||
case 6:
|
|
||||||
return (pkt[1] >> 4) & 0x03
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.OutboundSendReject {
|
if !f.firewall.OutboundSendReject {
|
||||||
return
|
return
|
||||||
@@ -279,27 +254,30 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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
|
||||||
@@ -393,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)
|
||||||
@@ -401,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",
|
||||||
@@ -514,20 +492,14 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
|||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
// 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
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
// q indicates which writer to use to send the packet.
|
// q indicates which writer to use to send the packet.
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||||
relay *Relay,
|
|
||||||
ad,
|
|
||||||
nb,
|
|
||||||
out []byte,
|
|
||||||
nocopy bool,
|
|
||||||
) {
|
|
||||||
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// already logged by prepareSendVia
|
// already logged by prepareSendVia
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
err = f.writers[0].WriteTo(toSend, 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)
|
||||||
}
|
}
|
||||||
@@ -601,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 {
|
||||||
@@ -616,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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+89
-79
@@ -96,12 +96,7 @@ type Interface struct {
|
|||||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||||
// left free to migrate as on stock nebula.
|
// left free to migrate as on stock nebula.
|
||||||
pinThreads bool
|
pinThreads bool
|
||||||
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
|
||||||
// inside.go copies the inner ECN onto the outer carrier on encap and
|
|
||||||
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
|
||||||
// via tunnels.ecn (default true).
|
|
||||||
ecnEnabled atomic.Bool
|
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
@@ -120,10 +115,12 @@ type Interface struct {
|
|||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
queues []tio.Queue
|
queues []tio.Queue
|
||||||
// batchers is one per tun queue, wrapping queues[i].
|
// batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
|
||||||
// decryptToTun sends plaintext into the batch.RxBatcher;
|
// commits plaintext into the batcher; the plaintext is decrypted
|
||||||
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
// in place inside the UDP receive buffers, so listenOut must call Flush
|
||||||
batchers []batch.RxBatcher
|
// 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.
|
||||||
@@ -135,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)
|
||||||
@@ -205,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,
|
||||||
@@ -222,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),
|
||||||
batchers: make([]batch.RxBatcher, 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,
|
||||||
@@ -235,6 +231,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
pinThreads: c.PinThreads,
|
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),
|
||||||
@@ -288,6 +285,13 @@ func (f *Interface) activate() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if len(queues) < f.routines {
|
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",
|
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
|
||||||
"requested", f.routines, "opened", len(queues))
|
"requested", f.routines, "opened", len(queues))
|
||||||
f.routines = len(queues)
|
f.routines = len(queues)
|
||||||
@@ -297,18 +301,7 @@ func (f *Interface) activate() error {
|
|||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
for i := range f.queues {
|
for i := range f.queues {
|
||||||
caps := tio.QueueCapabilities(f.queues[i])
|
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
|
||||||
if caps.TSO || caps.USO {
|
|
||||||
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
|
||||||
// is on, everything else (and either lane disabled) falls
|
|
||||||
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
|
||||||
// reaches the TUN.
|
|
||||||
arena := batch.NewArena(batch.DefaultMultiArenaCap)
|
|
||||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
|
|
||||||
} else {
|
|
||||||
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
|
|
||||||
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||||
@@ -356,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 {
|
||||||
@@ -364,21 +382,17 @@ 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()
|
|
||||||
h := &header.H{}
|
|
||||||
fwPacket := &firewall.Packet{}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
|
|
||||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
plaintext := f.batchers[i].Reserve(len(payload))
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
flusher := func() {
|
flusher := func() {
|
||||||
if err := f.batchers[i].Flush(); err != nil {
|
if err := f.batchers[i].Flush(); err != nil {
|
||||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
}
|
}
|
||||||
|
clear(rxc.hostmapCache)
|
||||||
}
|
}
|
||||||
|
|
||||||
err := li.ListenOut(listener, flusher)
|
err := li.ListenOut(listener, flusher)
|
||||||
@@ -394,32 +408,36 @@ 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) pinThisThread(i int) {
|
||||||
|
var cpu int
|
||||||
|
if n := len(f.cpuAffinity); n > 0 {
|
||||||
|
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||||
|
// validated the entries against the allowed CPU set.
|
||||||
|
cpu = f.cpuAffinity[i%n]
|
||||||
|
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||||
|
// Default: spread queues across the CPUs we're actually allowed to
|
||||||
|
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||||
|
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||||
|
cpu = allowed[i%len(allowed)]
|
||||||
|
} else {
|
||||||
|
cpu = i % runtime.NumCPU()
|
||||||
|
}
|
||||||
|
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||||
|
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
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
|
// 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.
|
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
||||||
if f.pinThreads {
|
if f.pinThreads {
|
||||||
var cpu int
|
f.pinThisThread(i)
|
||||||
if n := len(f.cpuAffinity); n > 0 {
|
|
||||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
|
||||||
// validated the entries against the allowed CPU set.
|
|
||||||
cpu = f.cpuAffinity[i%n]
|
|
||||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
|
||||||
// Default: spread queues across the CPUs we're actually allowed to
|
|
||||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
|
||||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
|
||||||
cpu = allowed[i%len(allowed)]
|
|
||||||
} else {
|
|
||||||
cpu = i % runtime.NumCPU()
|
|
||||||
}
|
|
||||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
|
||||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
rejectBuf := make([]byte, mtu)
|
rejectBuf := make([]byte, mtu)
|
||||||
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||||
fwPacket := &firewall.Packet{}
|
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)
|
||||||
@@ -441,26 +459,35 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|||||||
// accumulated so the first packets of a deep read drain
|
// accumulated so the first packets of a deep read drain
|
||||||
// hit the wire while the rest are still being encrypted.
|
// hit the wire while the rest are still being encrypted.
|
||||||
if sb.Len() >= batch.SendBatchCap {
|
if sb.Len() >= batch.SendBatchCap {
|
||||||
if err := sb.Flush(); err != nil {
|
f.flushSendBatch(sb, i)
|
||||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := sb.Flush(); err != nil {
|
f.flushSendBatch(sb, i)
|
||||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", 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)
|
||||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||||
c.RegisterReloadCallback(f.reloadMisc)
|
c.RegisterReloadCallback(f.reloadMisc)
|
||||||
c.RegisterReloadCallback(f.reloadEcn)
|
|
||||||
|
|
||||||
for _, udpConn := range f.writers {
|
for _, udpConn := range f.writers {
|
||||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||||
@@ -593,23 +620,6 @@ func (f *Interface) reloadMisc(c *config.C) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
|
||||||
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
|
||||||
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
|
||||||
func (f *Interface) reloadEcn(c *config.C) {
|
|
||||||
initial := c.InitialLoad()
|
|
||||||
if initial || c.HasChanged("tunnels.ecn") {
|
|
||||||
v := c.GetBool("tunnels.ecn", true)
|
|
||||||
changed := f.ecnEnabled.Swap(v) != v
|
|
||||||
if !initial {
|
|
||||||
f.l.Info("tunnels.ecn changed", "enabled", v)
|
|
||||||
if changed {
|
|
||||||
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||||
ticker := time.NewTicker(i)
|
ticker := time.NewTicker(i)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
|
|||||||
+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 {
|
||||||
|
|||||||
@@ -515,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,12 +6,14 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"slices"
|
"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"
|
||||||
@@ -165,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)
|
||||||
}
|
}
|
||||||
@@ -216,8 +224,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
pinThreads := c.GetBool("tun.pin_threads", true)
|
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||||
cpuAffinity := parseCpuAffinity(c, l, routines)
|
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||||
if pinThreads && len(cpuAffinity) == 0 && !configTest {
|
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
|
||||||
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
|
// 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{
|
||||||
@@ -260,7 +277,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
ifce.reloadSendRecvError(c)
|
ifce.reloadSendRecvError(c)
|
||||||
ifce.reloadAcceptRecvError(c)
|
ifce.reloadAcceptRecvError(c)
|
||||||
ifce.reloadEcn(c)
|
|
||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
@@ -281,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,
|
||||||
@@ -291,6 +309,7 @@ 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
|
||||||
}
|
}
|
||||||
@@ -359,57 +378,6 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
|||||||
return cpus
|
return cpus
|
||||||
}
|
}
|
||||||
|
|
||||||
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
|
|
||||||
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
|
|
||||||
// any physical NIC's interrupts. The stock allowed[i] spread pins the
|
|
||||||
// encrypt threads onto exactly the cores most drivers affine their first RX
|
|
||||||
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
|
|
||||||
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
|
|
||||||
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
|
|
||||||
//
|
|
||||||
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
|
|
||||||
// info is unavailable or when there aren't enough IRQ-free CPUs to give
|
|
||||||
// every routine its own core: silently doubling readers up on fewer cores
|
|
||||||
// is worse than the occasional IRQ collision. NICs whose vectors blanket
|
|
||||||
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
|
|
||||||
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
|
|
||||||
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
|
|
||||||
// effective.
|
|
||||||
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
|
|
||||||
irq, err := util.NICIRQCPUs()
|
|
||||||
if err != nil || len(irq) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
allowed, err := util.AllowedCPUs()
|
|
||||||
if err != nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
|
|
||||||
if cpus == nil {
|
|
||||||
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
|
|
||||||
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
|
|
||||||
return cpus
|
|
||||||
}
|
|
||||||
|
|
||||||
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
|
|
||||||
// irq, or nil if fewer than `routines` qualify.
|
|
||||||
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
|
|
||||||
free := make([]int, 0, routines)
|
|
||||||
for _, cpu := range allowed {
|
|
||||||
if irq[cpu] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
free = append(free, cpu)
|
|
||||||
if len(free) == routines {
|
|
||||||
return free
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func moduleVersion() string {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -9,26 +9,6 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestChooseIRQFreeCPUs(t *testing.T) {
|
|
||||||
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
|
|
||||||
|
|
||||||
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
|
|
||||||
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
|
|
||||||
|
|
||||||
// Exactly enough.
|
|
||||||
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
|
|
||||||
|
|
||||||
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
|
|
||||||
// than doubling readers up on shared cores.
|
|
||||||
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
|
|
||||||
|
|
||||||
// No IRQ info at all behaves like a plain prefix of allowed.
|
|
||||||
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
|
|
||||||
|
|
||||||
// Non-contiguous allowed set (cgroup cpuset) with holes.
|
|
||||||
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseCpuAffinity(t *testing.T) {
|
func TestParseCpuAffinity(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
|||||||
+75
-154
@@ -13,7 +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/udp"
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -23,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, meta udp.RxMeta) {
|
// 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
|
||||||
@@ -91,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
|
||||||
@@ -103,26 +107,32 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
|
||||||
|
f.messageMetrics.RxInvalid(1)
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("packet too small", "from", via, "length", len(packet))
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// All remaining packets are encrypted
|
// All remaining packets are encrypted
|
||||||
ci := hostinfo.ConnectionState
|
|
||||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Relay packets are special
|
|
||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
|
// Relay packets are special, this branch should always early-return
|
||||||
|
err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
|
||||||
|
if err != nil {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, 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",
|
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"header", h,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -135,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, meta)
|
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
|
||||||
@@ -143,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, ci, 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
|
||||||
@@ -169,28 +187,10 @@ 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, meta udp.RxMeta) {
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
|
||||||
// The entire body is sent as AD, not encrypted.
|
h := rxc.h
|
||||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||||
// 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
|
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
|
||||||
// which will gracefully fail in the DecryptDanger call.
|
|
||||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
|
||||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
|
||||||
var err error
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
// Advance the replay window now that the frame is authenticated
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
// Successfully validated the thing. Get rid of the Relay header.
|
|
||||||
signedPayload = signedPayload[header.Len:]
|
|
||||||
// Pull the Roaming parts up here, and return in all call paths.
|
// Pull the Roaming parts up here, and return in all call paths.
|
||||||
f.handleHostRoaming(hostinfo, via)
|
f.handleHostRoaming(hostinfo, via)
|
||||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||||
@@ -201,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -214,11 +212,10 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
via = ViaSender{
|
via = ViaSender{
|
||||||
UdpAddr: via.UdpAddr,
|
UdpAddr: via.UdpAddr,
|
||||||
relayHI: hostinfo,
|
relayHI: hostinfo,
|
||||||
remoteIdx: relay.RemoteIndex,
|
|
||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta)
|
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)
|
||||||
@@ -235,9 +232,11 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
if targetRelay.State == Established {
|
if targetRelay.State == Established {
|
||||||
switch targetRelay.Type {
|
switch targetRelay.Type {
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||||
// Find the target HostInfo //todo it would potentially be nice to batch these
|
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
fwdBuf := packet[:0]
|
||||||
|
//todo it would potentially be nice to batch these
|
||||||
|
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
return
|
return
|
||||||
@@ -318,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
|
||||||
}
|
}
|
||||||
@@ -333,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
|
||||||
@@ -359,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
|
||||||
@@ -369,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 {
|
||||||
@@ -386,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])
|
||||||
@@ -403,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 {
|
||||||
@@ -444,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
|
||||||
@@ -461,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]
|
||||||
@@ -504,117 +519,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
||||||
var err error
|
err := newPacket(out, true, rxc.fwPacket)
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
|
||||||
}
|
|
||||||
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
|
||||||
return nil, ErrOutOfWindow
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
|
||||||
const (
|
|
||||||
ecnNotECT = 0x00
|
|
||||||
ecnECT1 = 0x01
|
|
||||||
ecnECT0 = 0x02
|
|
||||||
ecnCE = 0x03
|
|
||||||
)
|
|
||||||
|
|
||||||
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
|
||||||
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
|
||||||
// codepoints are advisory only and leave the inner unchanged.
|
|
||||||
//
|
|
||||||
// Merge cases (outer × inner → action):
|
|
||||||
//
|
|
||||||
// outer != CE : no-op (inner is authoritative)
|
|
||||||
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
|
||||||
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
|
||||||
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
|
||||||
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
switch pkt[1] & 0x03 {
|
|
||||||
case ecnNotECT:
|
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
|
||||||
}
|
|
||||||
case ecnCE:
|
|
||||||
// Already CE.
|
|
||||||
default:
|
|
||||||
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
|
|
||||||
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
|
|
||||||
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
|
|
||||||
// checksum lives at pkt[10:12]. A header too short to carry a
|
|
||||||
// checksum can't be fixed up here, so leave it for newPacket to
|
|
||||||
// reject rather than emit a mangled packet.
|
|
||||||
if len(pkt) < ipv4.HeaderLen {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
m := binary.BigEndian.Uint16(pkt[0:2])
|
|
||||||
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
|
||||||
mNew := binary.BigEndian.Uint16(pkt[0:2])
|
|
||||||
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
|
|
||||||
for sum > 0xffff {
|
|
||||||
sum = (sum >> 16) + (sum & 0xffff)
|
|
||||||
}
|
|
||||||
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
|
|
||||||
}
|
|
||||||
case 6:
|
|
||||||
switch (pkt[1] >> 4) & 0x03 {
|
|
||||||
case ecnNotECT:
|
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
|
||||||
}
|
|
||||||
case ecnCE:
|
|
||||||
// Already CE.
|
|
||||||
default:
|
|
||||||
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
|
||||||
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
|
|
||||||
// underlay into the inner header before firewall + TUN write. Other
|
|
||||||
// outer codepoints are advisory only — we keep the inner unchanged.
|
|
||||||
if f.ecnEnabled.Load() {
|
|
||||||
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
|
||||||
}
|
|
||||||
|
|
||||||
err := newPacket(out, true, fwPacket)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
|
||||||
"error", err,
|
|
||||||
"packet", out,
|
|
||||||
)
|
|
||||||
return
|
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. With UDP GRO this is a single segment of a shared
|
|
||||||
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
|
|
||||||
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
|
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
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.batchers[q].Commit(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)
|
||||||
|
}
|
||||||
|
|||||||
+8
-25
@@ -1,28 +1,11 @@
|
|||||||
package batch
|
package batch
|
||||||
|
|
||||||
import "net/netip"
|
// 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:
|
||||||
type RxBatcher interface {
|
// a re-handshake replaces the tunnel outright and the replacement's epoch is higher,
|
||||||
// Reserve creates a pkt to borrow
|
// so the old tunnel's packets sort first during the cutover overlap.
|
||||||
Reserve(sz int) []byte
|
// Counter is the packet's AEAD message counter within that tunnel.
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
type SortKey struct {
|
||||||
Commit(pkt []byte) error
|
Epoch uint64
|
||||||
// Flush emits every queued packet in arrival order.
|
Counter uint64
|
||||||
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
|
|
||||||
// After Flush returns, borrowed payload slices may be recycled.
|
|
||||||
Flush() error
|
|
||||||
}
|
|
||||||
|
|
||||||
type TxBatcher interface {
|
|
||||||
// Reserve creates a pkt to borrow
|
|
||||||
Reserve(sz int) []byte
|
|
||||||
// Commit borrows pkt and records its destination plus the 2-bit
|
|
||||||
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
|
||||||
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
|
||||||
// to leave the outer ECN field unset.
|
|
||||||
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
|
||||||
// Flush emits every queued packet via the underlying batch writer in arrival order.
|
|
||||||
// Returns an errors.Join of one or more errors.
|
|
||||||
// After Flush returns, borrowed payload slices may be recycled.
|
|
||||||
Flush() error
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+106
-126
@@ -8,135 +8,125 @@ import (
|
|||||||
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||||
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
// 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
|
// 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
|
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
||||||
// never alias.
|
|
||||||
type flowKey struct {
|
type flowKey struct {
|
||||||
src, dst [16]byte
|
src, dst [16]byte
|
||||||
sport, dport uint16
|
sport, dport uint16
|
||||||
isV6 bool
|
isV6 bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// initialSlots is the starting capacity of the slot pool. One flow per
|
// initialSlots is the starting capacity of the slot pool.
|
||||||
// packet is the worst case so this matches a typical carrier-side
|
// One flow per packet is the worst case,
|
||||||
// recvmmsg batch on the encrypted UDP socket.
|
// so this matches a typical carrier-side recvmmsg batch on the UDP socket.
|
||||||
const initialSlots = 64
|
const initialSlots = 64
|
||||||
|
|
||||||
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
|
||||||
// L4-specific parsing (TCP / UDP) on top.
|
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
|
||||||
type parsedIP struct {
|
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
|
||||||
fk flowKey
|
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at
|
||||||
ipHdrLen int
|
// byte 40.
|
||||||
// pkt is the original buffer trimmed to the IP-declared total length.
|
//
|
||||||
// Anything below the IP layer (transport parsers) should slice into
|
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
|
||||||
// pkt rather than the unbounded original.
|
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
|
||||||
pkt []byte
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
|
||||||
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
// len(pkt) >= 20 and the version.
|
||||||
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
|
||||||
// or IPv6 with extension headers (all rejected by both coalescers in
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
// identical ways before this refactor).
|
if ihl != 20 {
|
||||||
//
|
return nil, false
|
||||||
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
|
||||||
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
|
||||||
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
|
||||||
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
|
||||||
var p parsedIP
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return p, false
|
|
||||||
}
|
}
|
||||||
v := pkt[0] >> 4
|
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
|
||||||
switch v {
|
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
|
||||||
case 4:
|
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
return nil, false
|
||||||
if ihl != 20 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if pkt[9] != wantProto {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
// Reject actual fragmentation (MF or non-zero frag offset).
|
|
||||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
|
||||||
if totalLen > len(pkt) || totalLen < ihl {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = 20
|
|
||||||
p.fk.isV6 = false
|
|
||||||
copy(p.fk.src[:4], pkt[12:16])
|
|
||||||
copy(p.fk.dst[:4], pkt[16:20])
|
|
||||||
p.pkt = pkt[:totalLen]
|
|
||||||
case 6:
|
|
||||||
if len(pkt) < 40 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if pkt[6] != wantProto {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
|
||||||
if 40+payloadLen > len(pkt) {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = 40
|
|
||||||
p.fk.isV6 = true
|
|
||||||
copy(p.fk.src[:], pkt[8:24])
|
|
||||||
copy(p.fk.dst[:], pkt[24:40])
|
|
||||||
p.pkt = pkt[:40+payloadLen]
|
|
||||||
default:
|
|
||||||
return p, false
|
|
||||||
}
|
}
|
||||||
return p, true
|
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
|
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||||
// byte-for-byte equality on every field that must be identical across
|
// byte-for-byte equality on every field that must be identical across coalesced segments.
|
||||||
// coalesced segments. Size/IPID/IPCsum are masked out. The full DSCP/ECN
|
// Size/IPID/IPCsum are masked out.
|
||||||
// byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel
|
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
|
||||||
// GRO: segments with differing ECN codepoints must not coalesce, otherwise
|
// segments with differing ECN codepoints must not coalesce,
|
||||||
// ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion)
|
// 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.
|
||||||
// mark or mark a Not-ECT flow as ECN-capable.
|
|
||||||
//
|
//
|
||||||
// The transport (L4) portion of the header is checked separately by the
|
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
||||||
// per-protocol matcher.
|
|
||||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||||
if isV6 {
|
if isV6 {
|
||||||
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
|
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
|
||||||
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
|
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
|
||||||
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
|
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
|
||||||
if a[0] != b[0] {
|
}
|
||||||
return false
|
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
|
||||||
}
|
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
|
||||||
if a[1] != b[1] {
|
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||||
return false
|
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
|
||||||
}
|
}
|
||||||
if !bytes.Equal(a[2:4], b[2:4]) {
|
|
||||||
return false
|
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
|
||||||
}
|
const ipv4FlagDF = 0x40
|
||||||
if !bytes.Equal(a[6:40], b[6:40]) {
|
|
||||||
return false
|
// 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
|
return true
|
||||||
}
|
}
|
||||||
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
|
||||||
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
|
||||||
// Compare byte 1 fully so ECN must match.
|
|
||||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
|
||||||
if a[0] != b[0] {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if a[1] != b[1] {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[6:10], b[6:10]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[12:20], b[12:20]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||||
@@ -145,17 +135,14 @@ type Arena struct {
|
|||||||
buf []byte
|
buf []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewArena returns an Arena with a pre-allocated backing of the given
|
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
|
||||||
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
|
||||||
// only feeds the coalescer pre-made []byte packets via Commit).
|
|
||||||
func NewArena(capacity int) *Arena {
|
func NewArena(capacity int) *Arena {
|
||||||
return &Arena{buf: make([]byte, 0, capacity)}
|
return &Arena{buf: make([]byte, 0, capacity)}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the
|
// Reserve hands out a non-overlapping sz-byte slice from the arena.
|
||||||
// request doesn't fit the current backing, a fresh, larger backing is
|
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
|
||||||
// allocated; already-borrowed slices reference the old backing and remain
|
// Already-borrowed slices reference the old backing and remain valid until Reset.
|
||||||
// valid until Reset.
|
|
||||||
func (a *Arena) Reserve(sz int) []byte {
|
func (a *Arena) Reserve(sz int) []byte {
|
||||||
if len(a.buf)+sz > cap(a.buf) {
|
if len(a.buf)+sz > cap(a.buf) {
|
||||||
newCap := max(cap(a.buf)*2, sz)
|
newCap := max(cap(a.buf)*2, sz)
|
||||||
@@ -166,16 +153,9 @@ func (a *Arena) Reserve(sz int) []byte {
|
|||||||
return a.buf[start : start+sz : start+sz]
|
return a.buf[start : start+sz : start+sz]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reset releases every slice handed out since the last Reset. Callers must
|
// Reset releases every slice handed out since the last Reset.
|
||||||
// not use any previously-borrowed slice after this returns. The underlying
|
// Callers must not use any previously-borrowed slice after this returns.
|
||||||
// backing array is retained so subsequent Reserves don't re-allocate.
|
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
||||||
func (a *Arena) Reset() {
|
func (a *Arena) Reset() {
|
||||||
a.buf = a.buf[:0]
|
a.buf = a.buf[:0]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reserver hands out an sz-byte slice valid until its Resetter runs.
|
|
||||||
type Reserver func(sz int) []byte
|
|
||||||
|
|
||||||
// Resetter clears all reservations held by a Reserver. Only the arena's
|
|
||||||
// owner holds one; lanes inside a MultiCoalescer get nil.
|
|
||||||
type Resetter func()
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -1,119 +1,121 @@
|
|||||||
package batch
|
package batch
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"cmp"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
)
|
)
|
||||||
|
|
||||||
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
|
||||||
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
|
||||||
// across lanes so the caller's allocation pattern is unchanged.
|
|
||||||
//
|
//
|
||||||
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
|
||||||
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
|
||||||
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
|
||||||
// ever lands in one lane and each lane preserves its own slot order.
|
// lanes carry no reorder-repair machinery.
|
||||||
//
|
//
|
||||||
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
|
||||||
// This is acceptable because the carrier-side recvmmsg path already
|
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
|
||||||
// stable-sorts by (peer, message counter) before delivering plaintext
|
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
|
||||||
// here, so replay-window invariants are unaffected, and apps observe
|
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
|
||||||
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
|
||||||
// Do not "fix" this by interleaving lane outputs at flush time; that
|
// to the later-flushed pt lane.
|
||||||
// negates the entire point of coalescing (each lane needs to see runs of
|
//
|
||||||
// adjacent same-flow packets to coalesce them).
|
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
|
||||||
type MultiCoalescer struct {
|
type MultiCoalescer struct {
|
||||||
tcp *TCPCoalescer
|
tcp *TCPCoalescer
|
||||||
udp *UDPCoalescer
|
udp *UDPCoalescer
|
||||||
pt *Passthrough
|
pt *Passthrough
|
||||||
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
|
|
||||||
// and Flush resets it exactly once after every lane has drained.
|
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
||||||
arena *Arena
|
// each pkt alive until Flush returns.
|
||||||
|
staged []stagedPacket
|
||||||
}
|
}
|
||||||
|
|
||||||
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
|
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
|
||||||
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
|
// value: pp is reused by the caller per packet and must not be retained past Commit.
|
||||||
// burst worth of MTU-sized packets without the arena growing.
|
type stagedPacket struct {
|
||||||
const DefaultMultiArenaCap = initialSlots * 65535
|
pkt []byte
|
||||||
|
key SortKey
|
||||||
|
proto byte
|
||||||
|
fragAny bool
|
||||||
|
ipHdrLen uint16
|
||||||
|
}
|
||||||
|
|
||||||
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
|
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
|
||||||
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
|
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
|
||||||
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
// transmission-order repair.
|
||||||
// Either lane disabled redirects its traffic into the passthrough lane.
|
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
|
||||||
// arena is the single backing slab shared across every lane; the caller
|
|
||||||
// pre-sizes it via NewArena so the hot path never allocates.
|
|
||||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
|
||||||
m := &MultiCoalescer{
|
m := &MultiCoalescer{
|
||||||
pt: NewPassthrough(w, arena.Reserve, nil),
|
pt: NewPassthrough(w),
|
||||||
arena: arena,
|
staged: make([]stagedPacket, 0, initialSlots),
|
||||||
}
|
|
||||||
if tcpEnabled {
|
|
||||||
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
|
|
||||||
}
|
|
||||||
if udpEnabled {
|
|
||||||
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
|
|
||||||
}
|
}
|
||||||
|
m.tcp = NewTCPCoalescer(w, l)
|
||||||
|
m.udp = NewUDPCoalescer(w)
|
||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
|
||||||
return m.arena.Reserve(sz)
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// Commit dispatches pkt to the appropriate lane based on IP version + L4
|
// compareStaged orders staged packets by (epoch, counter)
|
||||||
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
func compareStaged(a, b stagedPacket) int {
|
||||||
// pkt must remain valid until the next Flush.
|
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
|
||||||
//
|
return c
|
||||||
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
|
||||||
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
|
||||||
// re-walk the header.
|
|
||||||
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return m.pt.Commit(pkt)
|
|
||||||
}
|
}
|
||||||
v := pkt[0] >> 4
|
return cmp.Compare(a.key.Counter, b.key.Counter)
|
||||||
var proto byte
|
}
|
||||||
switch v {
|
|
||||||
case 4:
|
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
|
||||||
proto = pkt[9]
|
// passthrough when the lane has no GSO support.
|
||||||
case 6:
|
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
|
||||||
if len(pkt) < 40 {
|
switch sp.proto {
|
||||||
return m.pt.Commit(pkt)
|
|
||||||
}
|
|
||||||
proto = pkt[6]
|
|
||||||
default:
|
|
||||||
return m.pt.Commit(pkt)
|
|
||||||
}
|
|
||||||
switch proto {
|
|
||||||
case ipProtoTCP:
|
case ipProtoTCP:
|
||||||
if m.tcp != nil {
|
if m.tcp != nil {
|
||||||
info, ok := parseTCPBase(pkt)
|
return m.tcp.commitStaged(sp)
|
||||||
if !ok {
|
|
||||||
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
|
||||||
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
|
||||||
m.tcp.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return m.tcp.commitParsed(pkt, info)
|
|
||||||
}
|
}
|
||||||
case ipProtoUDP:
|
case ipProtoUDP:
|
||||||
if m.udp != nil {
|
if m.udp != nil {
|
||||||
info, ok := parseUDP(pkt)
|
return m.udp.commitStaged(sp)
|
||||||
if !ok {
|
|
||||||
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return m.udp.commitParsed(pkt, info)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return m.pt.Commit(pkt)
|
return m.pt.enqueue(sp.pkt)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush drains every lane in a fixed order, then resets the shared arena once.
|
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
|
||||||
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
|
// 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 {
|
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
|
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 m.tcp != nil {
|
||||||
if err := m.tcp.Flush(); err != nil {
|
if err := m.tcp.Flush(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
@@ -127,6 +129,5 @@ func (m *MultiCoalescer) Flush() error {
|
|||||||
if err := m.pt.Flush(); err != nil {
|
if err := m.pt.Flush(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
m.arena.Reset()
|
|
||||||
return errors.Join(errs...)
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,17 +1,39 @@
|
|||||||
package batch
|
package batch
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/test"
|
"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
|
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||||
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||||
// else (ICMP here) falls through to plain Write.
|
// else (ICMP here) falls through to plain Write.
|
||||||
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
|
m := newTestMultiCoalescer(t, w)
|
||||||
|
k := &keySeq{epoch: 1}
|
||||||
|
|
||||||
tcpPay := make([]byte, 1200)
|
tcpPay := make([]byte, 1200)
|
||||||
udpPay := make([]byte, 1200)
|
udpPay := make([]byte, 1200)
|
||||||
@@ -21,19 +43,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
|||||||
icmp[3] = 28
|
icmp[3] = 28
|
||||||
icmp[9] = 1
|
icmp[9] = 1
|
||||||
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil {
|
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil {
|
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(icmp); err != nil {
|
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Flush(); err != nil {
|
if err := m.Flush(); err != nil {
|
||||||
@@ -48,17 +70,162 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
|
||||||
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
// property: packets committed out of counter order (wire reorder inside one
|
||||||
// the kernel via the passthrough lane rather than being lost.
|
// flush batch) are replayed into the lanes in transmission order, so the
|
||||||
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
// 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}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
|
m := newTestMultiCoalescer(t, w)
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
|
||||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
// 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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Flush(); err != nil {
|
if err := m.Flush(); err != nil {
|
||||||
@@ -72,16 +239,164 @@ func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
|
||||||
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
// anything. Both lane constructors refuse, so every packet rides the
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
// verbatim lane — but the staging sort still applies, so emission follows
|
||||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
|
// transmission order even without GSO.
|
||||||
|
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
|
||||||
pay := make([]byte, 1200)
|
w := &fakeTunWriter{gsoEnabled: false}
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := m.Flush(); err != nil {
|
if err := m.Flush(); err != nil {
|
||||||
@@ -94,3 +409,29 @@ func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
|||||||
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,54 +2,29 @@ package batch
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"io"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/udp"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
|
||||||
|
// order enqueued.
|
||||||
type Passthrough struct {
|
type Passthrough struct {
|
||||||
out io.Writer
|
out io.Writer
|
||||||
slots [][]byte
|
slots [][]byte
|
||||||
reserver Reserver
|
|
||||||
resetter Resetter
|
|
||||||
cursor int
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const passthroughBaseNumSlots = 128
|
func NewPassthrough(w io.Writer) *Passthrough {
|
||||||
|
|
||||||
// DefaultPassthroughArenaCap is the recommended arena capacity for a
|
|
||||||
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
|
|
||||||
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
|
|
||||||
|
|
||||||
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
|
|
||||||
return &Passthrough{
|
return &Passthrough{
|
||||||
out: w,
|
out: w,
|
||||||
slots: make([][]byte, 0, passthroughBaseNumSlots),
|
slots: make([][]byte, 0, 128),
|
||||||
reserver: reserver,
|
|
||||||
resetter: resetter,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Passthrough) Reserve(sz int) []byte {
|
// enqueue accepts one packet, already sorted into transmission order by dispatch.
|
||||||
return p.reserver(sz)
|
func (p *Passthrough) enqueue(pkt []byte) error {
|
||||||
}
|
|
||||||
|
|
||||||
func (p *Passthrough) Commit(pkt []byte) error {
|
|
||||||
p.slots = append(p.slots, pkt)
|
p.slots = append(p.slots, pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush drains every queued packet and calls the configured Resetter
|
|
||||||
func (p *Passthrough) Flush() error {
|
func (p *Passthrough) Flush() error {
|
||||||
firstErr := p.drain()
|
|
||||||
if p.resetter != nil {
|
|
||||||
p.resetter()
|
|
||||||
}
|
|
||||||
return firstErr
|
|
||||||
}
|
|
||||||
|
|
||||||
// drain writes out every queued packet and clears the slot list.
|
|
||||||
func (p *Passthrough) drain() error {
|
|
||||||
var firstErr error
|
var firstErr error
|
||||||
for _, s := range p.slots {
|
for _, s := range p.slots {
|
||||||
_, err := p.out.Write(s)
|
_, err := p.out.Write(s)
|
||||||
|
|||||||
+182
-438
@@ -2,12 +2,9 @@ package batch
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
)
|
)
|
||||||
@@ -23,24 +20,19 @@ const tcpCoalesceBufSize = 65535
|
|||||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||||
const tcpCoalesceMaxSegs = 64
|
const tcpCoalesceMaxSegs = 64
|
||||||
|
|
||||||
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
|
||||||
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
|
||||||
const tcpCoalesceHdrCap = 100
|
// 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.
|
||||||
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
|
||||||
// passthrough is true the slot holds a single borrowed packet that must be
|
|
||||||
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
|
||||||
// passthrough is false the slot is an in-progress coalesced superpacket:
|
|
||||||
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
|
||||||
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
|
||||||
// slices from the caller's plaintext buffers — no payload is ever copied.
|
|
||||||
// The caller (listenOut) must keep those buffers alive until Flush.
|
|
||||||
type coalesceSlot struct {
|
type coalesceSlot struct {
|
||||||
passthrough bool
|
verbatim bool
|
||||||
rawPkt []byte // borrowed when passthrough
|
// 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
|
fk flowKey
|
||||||
hdrBuf [tcpCoalesceHdrCap]byte
|
|
||||||
hdrLen int
|
hdrLen int
|
||||||
ipHdrLen int
|
ipHdrLen int
|
||||||
isV6 bool
|
isV6 bool
|
||||||
@@ -48,216 +40,203 @@ type coalesceSlot struct {
|
|||||||
numSeg int
|
numSeg int
|
||||||
totalPay int
|
totalPay int
|
||||||
nextSeq uint32
|
nextSeq uint32
|
||||||
// psh closes the chain: set when the last-accepted segment had PSH or
|
payIovs [][]byte
|
||||||
// was sub-gsoSize. No further appends after that.
|
|
||||||
psh bool
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
|
||||||
// multiple concurrent flows and emits each flow's run as a single TSO
|
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
|
||||||
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
|
||||||
// deferred until Flush so arrival order is preserved on the wire. Owns
|
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
|
||||||
// no locks; one coalescer per TUN write queue.
|
// commitParsed. Owns no locks; one coalescer per TUN write queue.
|
||||||
type TCPCoalescer struct {
|
type TCPCoalescer struct {
|
||||||
plainW io.Writer
|
w tio.GSOWriter
|
||||||
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
|
||||||
|
|
||||||
// slots is the ordered event queue. Flush walks it once and emits each
|
// slots is the ordered event queue. Flush walks it once and emits each
|
||||||
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
||||||
slots []*coalesceSlot
|
slots []*coalesceSlot
|
||||||
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
|
||||||
// segments can extend an in-progress superpacket in O(1). Slots are
|
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
|
||||||
// removed from this map when they close (PSH or short-last-segment),
|
// non-admissible packet for the flow, or in Flush.
|
||||||
// when a non-admissible packet for that flow arrives, or in Flush.
|
|
||||||
openSlots map[flowKey]*coalesceSlot
|
openSlots map[flowKey]*coalesceSlot
|
||||||
// lastSlot caches the most recently touched open slot. Steady-state
|
// lastSlot caches the most recently touched open slot. Bulk traffic
|
||||||
// bulk traffic is dominated by a single flow, so comparing the
|
// arrives in same-flow runs (single-flow steady state, or GRO bursts
|
||||||
// incoming key against the cached slot's own fk lets the hot path
|
// under multi-flow), so comparing the incoming key against the cached
|
||||||
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
// 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
|
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||||
// at is removed/sealed.
|
// at is removed.
|
||||||
lastSlot *coalesceSlot
|
lastSlot *coalesceSlot
|
||||||
pool []*coalesceSlot // free list for reuse
|
pool []*coalesceSlot // free list for reuse
|
||||||
reserver Reserver
|
|
||||||
resetter Resetter
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
|
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
|
||||||
c := &TCPCoalescer{
|
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
|
||||||
plainW: w,
|
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &TCPCoalescer{
|
||||||
|
w: gw,
|
||||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||||
reserver: reserver,
|
|
||||||
resetter: resetter,
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
|
||||||
c.gsoW = gw
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
type parsedTCP struct {
|
type parsedTCP struct {
|
||||||
fk flowKey
|
fk flowKey
|
||||||
ipHdrLen int
|
ipHdrLen int
|
||||||
tcpHdrLen int
|
hdrLen int
|
||||||
hdrLen int
|
payLen int
|
||||||
payLen int
|
seq uint32
|
||||||
seq uint32
|
flags byte
|
||||||
flags byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
|
||||||
// regardless of whether it's admissible for coalescing. Returns ok=false
|
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||||
// for non-TCP or malformed input.
|
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
|
||||||
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
|
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||||
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||||
var p parsedTCP
|
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||||
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
|
||||||
if !ok {
|
if !ok {
|
||||||
return p, false
|
return false
|
||||||
}
|
}
|
||||||
pkt = ip.pkt
|
return p.parseTail(trimmed, ipHdrLen)
|
||||||
p.fk = ip.fk
|
|
||||||
p.ipHdrLen = ip.ipHdrLen
|
|
||||||
|
|
||||||
if len(pkt) < p.ipHdrLen+20 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
|
||||||
if tcpOff < 20 || tcpOff > 60 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if len(pkt) < p.ipHdrLen+tcpOff {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.tcpHdrLen = tcpOff
|
|
||||||
p.hdrLen = p.ipHdrLen + tcpOff
|
|
||||||
p.payLen = len(pkt) - p.hdrLen
|
|
||||||
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
|
||||||
p.flags = pkt[p.ipHdrLen+13]
|
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
|
||||||
return p, true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||||
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
// fk's addresses are already filled.
|
||||||
// negative mask in coalesceable, not by name.
|
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 (
|
const (
|
||||||
tcpFlagPsh = 0x08
|
tcpFlagPsh = 0x08
|
||||||
tcpFlagAck = 0x10
|
tcpFlagAck = 0x10
|
||||||
tcpFlagEce = 0x40
|
tcpFlagEce = 0x40
|
||||||
)
|
)
|
||||||
|
|
||||||
// coalesceable reports whether a parsed TCP segment is eligible for
|
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||||
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||||
// non-empty payload. CWR is excluded because it marks a one-shot
|
func (c *TCPCoalescer) sealAllOpen() {
|
||||||
// congestion-window-reduced transition the receiver must observe at a
|
clear(c.openSlots)
|
||||||
// segment boundary.
|
c.lastSlot = nil
|
||||||
func (p parsedTCP) coalesceable() bool {
|
|
||||||
if p.flags&tcpFlagAck == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.payLen > 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||||
return c.reserver(sz)
|
// 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
// coalesce (any fragmentation, unparseable header) seals every open chain
|
||||||
if c.gsoW == nil {
|
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||||
c.addPassthrough(pkt)
|
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
|
||||||
|
if sp.fragAny {
|
||||||
|
c.sealAllOpen()
|
||||||
|
c.addVerbatim(sp.pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
info, ok := parseTCPBase(pkt)
|
var info parsedTCP
|
||||||
if !ok {
|
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||||
c.addPassthrough(pkt)
|
c.sealAllOpen()
|
||||||
|
c.addVerbatim(sp.pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return c.commitParsed(pkt, info)
|
return c.commitParsed(sp.pkt, &info)
|
||||||
}
|
}
|
||||||
|
|
||||||
// commitParsed is the post-parse half of Commit. The caller must have
|
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
|
||||||
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
// valid parse so the header is not re-walked here.
|
||||||
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
|
||||||
// after the dispatcher has already done so.
|
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
|
||||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
|
||||||
if c.gsoW == nil {
|
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
|
||||||
c.addPassthrough(pkt)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
if !info.coalesceable() {
|
if info.payLen == 0 {
|
||||||
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
|
||||||
// Seal this flow's open slot so later in-flow packets don't extend
|
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
|
||||||
// it and accidentally reorder past this passthrough.
|
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
|
||||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
// kernel GRO. This is the only place emission deviates from transmission order.
|
||||||
c.lastSlot = nil
|
c.addVerbatim(pkt)
|
||||||
}
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Single-flow fast path: with only one open flow the cache hits every
|
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
|
||||||
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
// many flows: wire-side GRO delivers runs of same-flow packets
|
||||||
// when there are multiple flows in flight (where the hit rate would
|
// (deliverSegments splits a superdatagram into up to 64), so the cache
|
||||||
// be ~0 and the compare is pure overhead).
|
// hits for the length of each run and a miss costs one fk compare
|
||||||
|
// before the map lookup carries the weight.
|
||||||
var open *coalesceSlot
|
var open *coalesceSlot
|
||||||
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk {
|
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
open = last
|
open = last
|
||||||
} else {
|
} else {
|
||||||
open = c.openSlots[info.fk]
|
open = c.openSlots[info.fk]
|
||||||
}
|
}
|
||||||
if open != nil {
|
if open != nil {
|
||||||
if c.canAppend(open, pkt, info) {
|
if c.canAppend(open, pkt, info) {
|
||||||
c.appendPayload(open, pkt, info)
|
if c.appendPayload(open, pkt, info) {
|
||||||
if open.psh {
|
// Chain closed (PSH or short segment): stop extending it.
|
||||||
delete(c.openSlots, info.fk)
|
c.sealFlow(info.fk)
|
||||||
c.lastSlot = nil
|
|
||||||
} else {
|
} else {
|
||||||
c.lastSlot = open
|
c.lastSlot = open
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
// Can't extend (seq gap from upstream loss, header change, or a full
|
||||||
delete(c.openSlots, info.fk)
|
// chain): evict it from openSlots and fall through to seed a fresh slot.
|
||||||
if c.lastSlot == open {
|
c.sealFlow(info.fk)
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
c.seed(pkt, info)
|
c.seed(pkt, info)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush emits every queued event in (per-flow) seq order.
|
|
||||||
func (c *TCPCoalescer) Flush() error {
|
func (c *TCPCoalescer) Flush() error {
|
||||||
first := c.drain()
|
|
||||||
if c.resetter != nil {
|
|
||||||
c.resetter()
|
|
||||||
}
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
// drain emits every queued slot (reordering/merging coalesced runs first)
|
|
||||||
// and clears the slot state.
|
|
||||||
func (c *TCPCoalescer) drain() error {
|
|
||||||
c.reorderForFlush()
|
|
||||||
var first error
|
var first error
|
||||||
for _, s := range c.slots {
|
for _, s := range c.slots {
|
||||||
var err error
|
var err error
|
||||||
if s.passthrough {
|
if s.verbatim || s.numSeg == 1 {
|
||||||
_, err = c.plainW.Write(s.rawPkt)
|
// 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 {
|
} else {
|
||||||
err = c.flushSlot(s)
|
err = c.flushSlot(s)
|
||||||
}
|
}
|
||||||
@@ -274,23 +253,27 @@ func (c *TCPCoalescer) drain() error {
|
|||||||
return first
|
return first
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
||||||
s := c.take()
|
s := c.take()
|
||||||
s.passthrough = true
|
s.verbatim = true
|
||||||
s.rawPkt = pkt
|
s.rawPkt = pkt
|
||||||
c.slots = append(c.slots, s)
|
c.slots = append(c.slots, s)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
||||||
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||||
// Pathological shape — can't fit our scratch, emit as-is.
|
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||||
c.addPassthrough(pkt)
|
// 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
|
return
|
||||||
}
|
}
|
||||||
s := c.take()
|
s := c.take()
|
||||||
s.passthrough = false
|
s.verbatim = false
|
||||||
s.rawPkt = nil
|
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
// the superpacket header flushSlot patches in place.
|
||||||
|
s.rawPkt = pkt
|
||||||
s.hdrLen = info.hdrLen
|
s.hdrLen = info.hdrLen
|
||||||
s.ipHdrLen = info.ipHdrLen
|
s.ipHdrLen = info.ipHdrLen
|
||||||
s.isV6 = info.fk.isV6
|
s.isV6 = info.fk.isV6
|
||||||
@@ -299,26 +282,23 @@ func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
|||||||
s.numSeg = 1
|
s.numSeg = 1
|
||||||
s.totalPay = info.payLen
|
s.totalPay = info.payLen
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
s.psh = info.flags&tcpFlagPsh != 0
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
c.slots = append(c.slots, s)
|
c.slots = append(c.slots, s)
|
||||||
if !s.psh {
|
if info.flags&tcpFlagPsh == 0 {
|
||||||
c.openSlots[info.fk] = s
|
c.openSlots[info.fk] = s
|
||||||
c.lastSlot = s
|
c.lastSlot = s
|
||||||
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
} else {
|
||||||
// PSH-on-seed seals the slot immediately. Any prior cached open
|
// PSH on the seed closes the chain immediately; it is never registered as open.
|
||||||
// slot for this flow has just been sealed-and-replaced by this
|
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
|
||||||
// passthrough-shaped seed, so drop the cache too.
|
c.sealFlow(info.fk)
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed: same
|
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
|
||||||
// header shape and stable contents, adjacent seq, not oversized, chain not closed.
|
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
|
||||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
|
||||||
if s.psh {
|
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
|
||||||
return false
|
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||||
}
|
|
||||||
if info.hdrLen != s.hdrLen {
|
if info.hdrLen != s.hdrLen {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -334,31 +314,35 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bo
|
|||||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
// ECE state must be stable across a burst — receivers expect the
|
// ECE state must be stable across a burst.
|
||||||
// flag set on every segment of a CE-echoing window or none.
|
// Receivers expect the flag set on every segment of a CE-echoing window or none.
|
||||||
seedFlags := s.hdrBuf[s.ipHdrLen+13]
|
seedFlags := s.rawPkt[s.ipHdrLen+13]
|
||||||
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
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 false
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
// 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.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
s.numSeg++
|
s.numSeg++
|
||||||
s.totalPay += info.payLen
|
s.totalPay += info.payLen
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
if info.flags&tcpFlagPsh != 0 {
|
if info.flags&tcpFlagPsh != 0 {
|
||||||
// Propagate PSH into the seed header so kernel TSO sets it on the
|
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
|
||||||
// last segment. Without this the sender's push signal is dropped.
|
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
|
||||||
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
|
||||||
}
|
|
||||||
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
|
|
||||||
s.psh = true
|
|
||||||
}
|
}
|
||||||
|
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||||
@@ -372,21 +356,18 @@ func (c *TCPCoalescer) take() *coalesceSlot {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||||
s.passthrough = false
|
|
||||||
s.rawPkt = nil
|
|
||||||
clear(s.payIovs)
|
clear(s.payIovs)
|
||||||
s.payIovs = s.payIovs[:0]
|
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
||||||
s.numSeg = 0
|
|
||||||
s.totalPay = 0
|
|
||||||
s.psh = false
|
|
||||||
c.pool = append(c.pool, s)
|
c.pool = append(c.pool, s)
|
||||||
}
|
}
|
||||||
|
|
||||||
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots.
|
// 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 {
|
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||||
total := s.hdrLen + s.totalPay
|
total := s.hdrLen + s.totalPay
|
||||||
l4Len := total - s.ipHdrLen
|
l4Len := total - s.ipHdrLen
|
||||||
hdr := s.hdrBuf[:s.hdrLen]
|
hdr := s.rawPkt[:s.hdrLen]
|
||||||
|
|
||||||
if s.isV6 {
|
if s.isV6 {
|
||||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||||
@@ -406,7 +387,7 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
|||||||
tcsum := s.ipHdrLen + 16
|
tcsum := s.ipHdrLen + 16
|
||||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
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
|
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||||
@@ -437,242 +418,6 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
|
||||||
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
|
||||||
// this pass a small wire reorder — counter 250 arriving in batch K when
|
|
||||||
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
|
||||||
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
|
||||||
// receiver as a much larger reorder than the wire actually had.
|
|
||||||
//
|
|
||||||
// Two phases:
|
|
||||||
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
|
||||||
// Cross-flow ordering inside a segment isn't preserved (it never was
|
|
||||||
// and doesn't matter for any single flow's TCP correctness).
|
|
||||||
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
|
||||||
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
|
||||||
// matters because the kernel TSO splitter chops at gsoSize from the
|
|
||||||
// start of the merged payload — a short segment in the middle would
|
|
||||||
// desynchronize every later segment.
|
|
||||||
//
|
|
||||||
// Passthrough slots act as barriers: the merge check skips them on either
|
|
||||||
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
|
||||||
// data.
|
|
||||||
func (c *TCPCoalescer) reorderForFlush() {
|
|
||||||
if len(c.slots) <= 1 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
runStart := 0
|
|
||||||
for i := 0; i <= len(c.slots); i++ {
|
|
||||||
if i < len(c.slots) && !c.slots[i].passthrough {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
c.sortRun(c.slots[runStart:i])
|
|
||||||
runStart = i + 1
|
|
||||||
}
|
|
||||||
out := c.slots[:0]
|
|
||||||
logged := false
|
|
||||||
for _, s := range c.slots {
|
|
||||||
if n := len(out); n > 0 {
|
|
||||||
prev := out[n-1]
|
|
||||||
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
|
||||||
// Same-flow neighbors after sort. If they aren't seq-
|
|
||||||
// contiguous it's a real gap — packets the wire reordered
|
|
||||||
// across batches, or actual loss before nebula. Log it so
|
|
||||||
// the operator can quantify how often it happens; the data
|
|
||||||
// itself still emits in seq order, kernel TCP handles the
|
|
||||||
// gap via its OOO queue.
|
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
if prev.nextSeq != slotSeedSeq(s) {
|
|
||||||
logged = true
|
|
||||||
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
|
||||||
c.l.Debug("tcp coalesce: cross-slot seq gap",
|
|
||||||
"src", flowKeyAddr(s.fk, false),
|
|
||||||
"dst", flowKeyAddr(s.fk, true),
|
|
||||||
"sport", s.fk.sport,
|
|
||||||
"dport", s.fk.dport,
|
|
||||||
"prev_seed_seq", slotSeedSeq(prev),
|
|
||||||
"prev_next_seq", prev.nextSeq,
|
|
||||||
"this_seed_seq", slotSeedSeq(s),
|
|
||||||
"gap_bytes", gap,
|
|
||||||
"prev_seg_count", prev.numSeg,
|
|
||||||
"prev_total_pay", prev.totalPay,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if canMergeSlots(prev, s) {
|
|
||||||
mergeSlots(prev, s)
|
|
||||||
c.release(s)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
out = append(out, s)
|
|
||||||
}
|
|
||||||
if logged {
|
|
||||||
c.l.Warn("==== end of batch ====")
|
|
||||||
}
|
|
||||||
c.slots = out
|
|
||||||
}
|
|
||||||
|
|
||||||
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
|
||||||
// logging. Only used on the cold gap-log path so the netip allocation
|
|
||||||
// doesn't matter.
|
|
||||||
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
|
||||||
src := fk.src
|
|
||||||
if dst {
|
|
||||||
src = fk.dst
|
|
||||||
}
|
|
||||||
if fk.isV6 {
|
|
||||||
return netip.AddrFrom16(src)
|
|
||||||
}
|
|
||||||
var v4 [4]byte
|
|
||||||
copy(v4[:], src[:4])
|
|
||||||
return netip.AddrFrom4(v4)
|
|
||||||
}
|
|
||||||
|
|
||||||
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
|
||||||
// cluster together in seq order, ready for the merge sweep. Stable so
|
|
||||||
// equal-key slots keep their original relative position (defensive — a
|
|
||||||
// duplicate seedSeq would already mean something's wrong upstream).
|
|
||||||
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
|
||||||
if len(run) <= 1 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
|
||||||
// reflection + closure-escape allocations that sort.SliceStable forces.
|
|
||||||
slices.SortStableFunc(run, compareCoalesceSlots)
|
|
||||||
}
|
|
||||||
|
|
||||||
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
|
||||||
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
|
||||||
return cmp
|
|
||||||
}
|
|
||||||
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
|
||||||
if aSeq == bSeq {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
if tcpSeqLess(aSeq, bSeq) {
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
|
||||||
// nextSeq tracks the seq just past the last appended byte; subtracting
|
|
||||||
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
|
||||||
// arithmetic so no special-casing is needed.
|
|
||||||
func slotSeedSeq(s *coalesceSlot) uint32 {
|
|
||||||
return s.nextSeq - uint32(s.totalPay)
|
|
||||||
}
|
|
||||||
|
|
||||||
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
|
||||||
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
|
||||||
// into the right comparison even across the 2^32 wrap.
|
|
||||||
func tcpSeqLess(a, b uint32) bool {
|
|
||||||
return int32(a-b) < 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
|
||||||
// is irrelevant — only that same-flow slots cluster together so the
|
|
||||||
// post-sort sweep can merge contiguous pairs.
|
|
||||||
func flowKeyCompare(a, b flowKey) int {
|
|
||||||
// Cheap scalar fields first so most non-matching keys short-circuit
|
|
||||||
// without ever calling bytes.Compare. sport is the ephemeral port on
|
|
||||||
// egress flows and discriminates fastest. For matching keys (same
|
|
||||||
// flow), array equality on src/dst inlines to word-sized compares,
|
|
||||||
// so we only pay bytes.Compare when the arrays actually differ.
|
|
||||||
if a.sport != b.sport {
|
|
||||||
if a.sport < b.sport {
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
if a.dport != b.dport {
|
|
||||||
if a.dport < b.dport {
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
if a.dst != b.dst {
|
|
||||||
return bytes.Compare(a.dst[:], b.dst[:])
|
|
||||||
}
|
|
||||||
if a.src != b.src {
|
|
||||||
return bytes.Compare(a.src[:], b.src[:])
|
|
||||||
}
|
|
||||||
if a.isV6 != b.isV6 {
|
|
||||||
if !a.isV6 {
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
|
||||||
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
|
||||||
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
|
||||||
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
|
||||||
// would split the merged skb at gsoSize boundaries from the start, so a
|
|
||||||
// short segment in the middle would corrupt every later segment. PSH and
|
|
||||||
// ECE state must agree across both slots: PSH is a semantic delimiter
|
|
||||||
// (preserving the sender's push boundary) and ECE state must be uniform
|
|
||||||
// across a window (the same rule canAppend enforces for in-flow appends).
|
|
||||||
// The IP-level ECN codepoint must also match: this check calls headersMatch
|
|
||||||
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
|
|
||||||
// differing ECN marks stay separate superpackets, each keeping its own mark.
|
|
||||||
//
|
|
||||||
// Note: a slot sealed by reorder (canAppend returned false on seq
|
|
||||||
// mismatch) keeps psh=false, so this restriction does not block the
|
|
||||||
// reorder-fix merge — only legitimate PSH-set seals.
|
|
||||||
func canMergeSlots(prev, s *coalesceSlot) bool {
|
|
||||||
if prev.psh {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if prev.fk != s.fk {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if prev.gsoSize != s.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if prev.nextSeq != slotSeedSeq(s) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
|
||||||
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
|
||||||
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
|
||||||
// and totals updated, PSH OR'd into the seed header so the push signal is
|
|
||||||
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
|
|
||||||
// is responsible for releasing src (it's no longer in c.slots after this call).
|
|
||||||
func mergeSlots(dst, src *coalesceSlot) {
|
|
||||||
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
|
||||||
dst.numSeg += src.numSeg
|
|
||||||
dst.totalPay += src.totalPay
|
|
||||||
dst.nextSeq = src.nextSeq
|
|
||||||
if src.psh {
|
|
||||||
dst.psh = true
|
|
||||||
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||||
// already have its checksum field zeroed) and returns the folded/inverted
|
// already have its checksum field zeroed) and returns the folded/inverted
|
||||||
// 16-bit value to store.
|
// 16-bit value to store.
|
||||||
@@ -717,9 +462,8 @@ func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
|||||||
return sum
|
return sum
|
||||||
}
|
}
|
||||||
|
|
||||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
|
||||||
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
|
||||||
// the L4 checksum field — the kernel will add the payload sum and invert.
|
|
||||||
func foldOnceNoInvert(sum uint32) uint16 {
|
func foldOnceNoInvert(sum uint32) uint16 {
|
||||||
for sum>>16 != 0 {
|
for sum>>16 != 0 {
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
|||||||
@@ -2,9 +2,9 @@ package batch
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"runtime"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
)
|
)
|
||||||
@@ -55,7 +55,31 @@ func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
|||||||
return pkts
|
return pkts
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
// 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.
|
// branch in Commit.
|
||||||
func buildICMPv4() []byte {
|
func buildICMPv4() []byte {
|
||||||
pkt := make([]byte, 28)
|
pkt := make([]byte, 28)
|
||||||
@@ -71,8 +95,7 @@ func buildICMPv4() []byte {
|
|||||||
// between batches, and reports per-packet cost.
|
// between batches, and reports per-packet cost.
|
||||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
b.Helper()
|
b.Helper()
|
||||||
arena := NewArena(0)
|
c := newTestTCPCoalescer(b, nopTunWriter{})
|
||||||
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
|
|
||||||
b.ReportAllocs()
|
b.ReportAllocs()
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
@@ -113,8 +136,17 @@ func BenchmarkCommitInterleaved16(b *testing.B) {
|
|||||||
runCommitBench(b, pkts, len(pkts))
|
runCommitBench(b, pkts, len(pkts))
|
||||||
}
|
}
|
||||||
|
|
||||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
|
||||||
// bails early and addPassthrough is the only work.
|
// 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) {
|
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||||
pkt := buildICMPv4()
|
pkt := buildICMPv4()
|
||||||
pkts := make([][]byte, 64)
|
pkts := make([][]byte, 64)
|
||||||
@@ -126,7 +158,7 @@ func BenchmarkCommitPassthrough(b *testing.B) {
|
|||||||
|
|
||||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||||
// Each packet takes the "TCP but not admissible" branch which does a
|
// Each packet takes the "TCP but not admissible" branch which does a
|
||||||
// map delete + passthrough. Measures the seal-without-slot cost.
|
// map delete + verbatim. Measures the seal-without-slot cost.
|
||||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||||
pay := make([]byte, 0)
|
pay := make([]byte, 0)
|
||||||
pkts := make([][]byte, 64)
|
pkts := make([][]byte, 64)
|
||||||
@@ -136,18 +168,24 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
|||||||
runCommitBench(b, pkts, 64)
|
runCommitBench(b, pkts, 64)
|
||||||
}
|
}
|
||||||
|
|
||||||
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
|
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
|
||||||
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
// it includes the staging sort's already-sorted fast path plus the
|
||||||
// is the bench that shows the savings of skipping the lane's re-parse.
|
// 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) {
|
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
b.Helper()
|
b.Helper()
|
||||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true)
|
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||||
|
pps := make([]*firewall.ParsedPacket, len(pkts))
|
||||||
|
for i, p := range pkts {
|
||||||
|
pps[i] = testPP(p)
|
||||||
|
}
|
||||||
b.ReportAllocs()
|
b.ReportAllocs()
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
pkt := pkts[i%len(pkts)]
|
j := i % len(pkts)
|
||||||
if err := m.Commit(pkt); err != nil {
|
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
if (i+1)%batchSize == 0 {
|
if (i+1)%batchSize == 0 {
|
||||||
@@ -174,68 +212,3 @@ func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
|||||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||||
runMultiCommitBench(b, pkts, len(pkts))
|
runMultiCommitBench(b, pkts, len(pkts))
|
||||||
}
|
}
|
||||||
|
|
||||||
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
|
||||||
type flowKeyPair struct{ a, b flowKey }
|
|
||||||
|
|
||||||
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
|
||||||
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
|
||||||
var fk flowKey
|
|
||||||
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
|
||||||
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
|
||||||
fk.sport = sport
|
|
||||||
fk.dport = dport
|
|
||||||
return fk
|
|
||||||
}
|
|
||||||
|
|
||||||
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
|
||||||
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
|
||||||
// repeatedly when many segments share a flow).
|
|
||||||
// - sportDiffers: same src/dst/dport, different sport — the typical
|
|
||||||
// "sibling flows from one host to one server" pattern.
|
|
||||||
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
|
||||||
// servers from a fixed local port.
|
|
||||||
// - allDiffer: every field differs; worst case for short-circuiting.
|
|
||||||
func flowKeyCases() map[string][]flowKeyPair {
|
|
||||||
const n = 64
|
|
||||||
cases := map[string][]flowKeyPair{
|
|
||||||
"sameFlow": make([]flowKeyPair, n),
|
|
||||||
"sportDiffers": make([]flowKeyPair, n),
|
|
||||||
"dstDiffers": make([]flowKeyPair, n),
|
|
||||||
"allDiffer": make([]flowKeyPair, n),
|
|
||||||
}
|
|
||||||
for i := range n {
|
|
||||||
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
|
||||||
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
|
||||||
cases["sportDiffers"][i] = flowKeyPair{
|
|
||||||
a: base,
|
|
||||||
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
|
||||||
}
|
|
||||||
cases["dstDiffers"][i] = flowKeyPair{
|
|
||||||
a: base,
|
|
||||||
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
|
||||||
}
|
|
||||||
cases["allDiffer"][i] = flowKeyPair{
|
|
||||||
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
|
||||||
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return cases
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
|
||||||
// the sort step actually sees. Use this to compare reorderings.
|
|
||||||
func BenchmarkFlowKeyCompare(b *testing.B) {
|
|
||||||
for name, pairs := range flowKeyCases() {
|
|
||||||
b.Run(name, func(b *testing.B) {
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
var sink int
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
p := pairs[i&(len(pairs)-1)]
|
|
||||||
sink += flowKeyCompare(p.a, p.b)
|
|
||||||
}
|
|
||||||
runtime.KeepAlive(sink)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+472
-330
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@ const SendBatchCap = 128
|
|||||||
|
|
||||||
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||||
type batchWriter interface {
|
type batchWriter interface {
|
||||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||||
@@ -16,7 +16,6 @@ type SendBatch struct {
|
|||||||
out batchWriter
|
out batchWriter
|
||||||
bufs [][]byte
|
bufs [][]byte
|
||||||
dsts []netip.AddrPort
|
dsts []netip.AddrPort
|
||||||
ecns []byte
|
|
||||||
arena *Arena
|
arena *Arena
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -26,7 +25,6 @@ func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
|||||||
out: out,
|
out: out,
|
||||||
bufs: make([][]byte, 0, batchCap),
|
bufs: make([][]byte, 0, batchCap),
|
||||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||||
ecns: make([]byte, 0, batchCap),
|
|
||||||
arena: NewArena(arenaSize),
|
arena: NewArena(arenaSize),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -40,21 +38,22 @@ func (b *SendBatch) Reserve(sz int) []byte {
|
|||||||
// bounding how long the first packet of a large read batch waits.
|
// bounding how long the first packet of a large read batch waits.
|
||||||
func (b *SendBatch) Len() int { return len(b.bufs) }
|
func (b *SendBatch) Len() int { return len(b.bufs) }
|
||||||
|
|
||||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
||||||
b.bufs = append(b.bufs, pkt)
|
b.bufs = append(b.bufs, pkt)
|
||||||
b.dsts = append(b.dsts, dst)
|
b.dsts = append(b.dsts, dst)
|
||||||
b.ecns = append(b.ecns, outerECN)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *SendBatch) Flush() error {
|
// 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
|
var err error
|
||||||
|
written := 0
|
||||||
if len(b.bufs) > 0 {
|
if len(b.bufs) > 0 {
|
||||||
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
||||||
}
|
}
|
||||||
clear(b.bufs)
|
clear(b.bufs)
|
||||||
b.bufs = b.bufs[:0]
|
b.bufs = b.bufs[:0]
|
||||||
b.dsts = b.dsts[:0]
|
b.dsts = b.dsts[:0]
|
||||||
b.ecns = b.ecns[:0]
|
|
||||||
b.arena.Reset()
|
b.arena.Reset()
|
||||||
return err
|
return written, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,10 +8,9 @@ import (
|
|||||||
type fakeBatchWriter struct {
|
type fakeBatchWriter struct {
|
||||||
bufs [][]byte
|
bufs [][]byte
|
||||||
addrs []netip.AddrPort
|
addrs []netip.AddrPort
|
||||||
ecns []byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||||
// returns, so tests must capture data before that happens.
|
// returns, so tests must capture data before that happens.
|
||||||
w.bufs = make([][]byte, len(bufs))
|
w.bufs = make([][]byte, len(bufs))
|
||||||
@@ -21,8 +20,7 @@ func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns
|
|||||||
w.bufs[i] = cp
|
w.bufs[i] = cp
|
||||||
}
|
}
|
||||||
w.addrs = append(w.addrs[:0], addrs...)
|
w.addrs = append(w.addrs[:0], addrs...)
|
||||||
w.ecns = append(w.ecns[:0], ecns...)
|
return len(bufs), nil
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||||
@@ -36,9 +34,9 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
|||||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||||
}
|
}
|
||||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||||
b.Commit(pkt, ap, 0)
|
b.Commit(pkt, ap)
|
||||||
}
|
}
|
||||||
if err := b.Flush(); err != nil {
|
if _, err := b.Flush(); err != nil {
|
||||||
t.Fatalf("Flush: %v", err)
|
t.Fatalf("Flush: %v", err)
|
||||||
}
|
}
|
||||||
if len(fw.bufs) != 4 {
|
if len(fw.bufs) != 4 {
|
||||||
@@ -55,7 +53,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
|||||||
|
|
||||||
// Flush again with nothing committed — should be a no-op.
|
// Flush again with nothing committed — should be a no-op.
|
||||||
fw.bufs = nil
|
fw.bufs = nil
|
||||||
if err := b.Flush(); err != nil {
|
if _, err := b.Flush(); err != nil {
|
||||||
t.Fatalf("empty Flush: %v", err)
|
t.Fatalf("empty Flush: %v", err)
|
||||||
}
|
}
|
||||||
if fw.bufs != nil {
|
if fw.bufs != nil {
|
||||||
@@ -77,9 +75,9 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
|||||||
for i := 0; i < 3; i++ {
|
for i := 0; i < 3; i++ {
|
||||||
s := b.Reserve(8)
|
s := b.Reserve(8)
|
||||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||||
b.Commit(pkt, ap, 0)
|
b.Commit(pkt, ap)
|
||||||
}
|
}
|
||||||
if err := b.Flush(); err != nil {
|
if _, err := b.Flush(); err != nil {
|
||||||
t.Fatalf("Flush: %v", err)
|
t.Fatalf("Flush: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -98,18 +96,18 @@ func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
|||||||
|
|
||||||
s1 := b.Reserve(4)
|
s1 := b.Reserve(4)
|
||||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||||
b.Commit(pkt1, ap, 0)
|
b.Commit(pkt1, ap)
|
||||||
|
|
||||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||||
b.Commit(pkt2, ap, 0)
|
b.Commit(pkt2, ap)
|
||||||
|
|
||||||
// pkt1 must still be intact even though backing reallocated.
|
// pkt1 must still be intact even though backing reallocated.
|
||||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := b.Flush(); err != nil {
|
if _, err := b.Flush(); err != nil {
|
||||||
t.Fatalf("Flush: %v", err)
|
t.Fatalf("Flush: %v", err)
|
||||||
}
|
}
|
||||||
if len(fw.bufs) != 2 {
|
if len(fw.bufs) != 2 {
|
||||||
|
|||||||
+142
-152
@@ -1,6 +1,7 @@
|
|||||||
package batch
|
package batch
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
|
|
||||||
@@ -18,72 +19,56 @@ const udpCoalesceBufSize = 65535
|
|||||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||||
const udpCoalesceMaxSegs = 64
|
const udpCoalesceMaxSegs = 64
|
||||||
|
|
||||||
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
|
||||||
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
|
||||||
const udpCoalesceHdrCap = 64
|
|
||||||
|
|
||||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||||
type udpSlot struct {
|
type udpSlot struct {
|
||||||
passthrough bool
|
verbatim bool
|
||||||
rawPkt []byte // borrowed when passthrough
|
// 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
|
fk flowKey
|
||||||
hdrBuf [udpCoalesceHdrCap]byte
|
|
||||||
hdrLen int
|
hdrLen int
|
||||||
ipHdrLen int
|
ipHdrLen int
|
||||||
isV6 bool
|
isV6 bool
|
||||||
gsoSize int // per-segment UDP payload length
|
gsoSize int // per-segment UDP payload length
|
||||||
numSeg int
|
numSeg int
|
||||||
totalPay int
|
totalPay int
|
||||||
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
payIovs [][]byte
|
||||||
// (kernel UDP-GSO requires every segment but the last to be exactly
|
|
||||||
// gsoSize) or when limits are hit. No further appends after.
|
|
||||||
sealed bool
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
|
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
|
||||||
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
|
// Preserves the in-flow order of packets as they are Commit-ed
|
||||||
// underlying writer doesn't support USO.
|
|
||||||
//
|
|
||||||
// All output — coalesced or not — is deferred until Flush so per-flow
|
|
||||||
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
|
||||||
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
|
||||||
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
|
||||||
// single 5-tuple only ever lands in one lane and each lane preserves its
|
|
||||||
// own slot order.
|
|
||||||
//
|
//
|
||||||
// Owns no locks; one coalescer per TUN write queue.
|
// Owns no locks; one coalescer per TUN write queue.
|
||||||
type UDPCoalescer struct {
|
type UDPCoalescer struct {
|
||||||
plainW io.Writer
|
w tio.GSOWriter
|
||||||
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
|
||||||
|
|
||||||
slots []*udpSlot
|
slots []*udpSlot
|
||||||
openSlots map[flowKey]*udpSlot
|
openSlots map[flowKey]*udpSlot
|
||||||
pool []*udpSlot
|
// lastSlot caches the most recently touched open slot; see the
|
||||||
reserver Reserver
|
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
|
||||||
resetter Resetter
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
||||||
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
|
||||||
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
if !ok {
|
||||||
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
return nil
|
||||||
// plain Writes — same defensive shape as the TCP coalescer.
|
}
|
||||||
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer {
|
return &UDPCoalescer{
|
||||||
c := &UDPCoalescer{
|
w: gw,
|
||||||
plainW: w,
|
|
||||||
slots: make([]*udpSlot, 0, initialSlots),
|
slots: make([]*udpSlot, 0, initialSlots),
|
||||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||||
pool: make([]*udpSlot, 0, initialSlots),
|
pool: make([]*udpSlot, 0, initialSlots),
|
||||||
reserver: reserver,
|
|
||||||
resetter: resetter,
|
|
||||||
}
|
}
|
||||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
|
||||||
c.gsoW = gw
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||||
@@ -95,104 +80,111 @@ type parsedUDP struct {
|
|||||||
payLen int
|
payLen int
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
|
||||||
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||||
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
// entry and is filled in place. Returns false for malformed input or any shape that must not
|
||||||
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||||
var p parsedUDP
|
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||||
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||||
if !ok {
|
if !ok {
|
||||||
return p, false
|
return false
|
||||||
}
|
}
|
||||||
pkt = ip.pkt
|
return p.parseTail(trimmed, ipHdrLen)
|
||||||
p.fk = ip.fk
|
}
|
||||||
p.ipHdrLen = ip.ipHdrLen
|
|
||||||
|
|
||||||
if len(pkt) < p.ipHdrLen+8 {
|
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||||
return p, false
|
// fk's addresses are already filled.
|
||||||
|
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||||
|
if len(pkt) < ipHdrLen+8 {
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
p.hdrLen = p.ipHdrLen + 8
|
|
||||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||||
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
|
||||||
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
|
||||||
return p, false
|
return false
|
||||||
}
|
}
|
||||||
|
p.ipHdrLen = ipHdrLen
|
||||||
|
p.hdrLen = ipHdrLen + 8
|
||||||
p.payLen = udpLen - 8
|
p.payLen = udpLen - 8
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||||
return p, true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||||
return c.reserver(sz)
|
// 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
|
||||||
if c.gsoW == nil {
|
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||||
c.addPassthrough(pkt)
|
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
|
||||||
|
if sp.fragAny {
|
||||||
|
c.sealAllOpen()
|
||||||
|
c.addVerbatim(sp.pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
info, ok := parseUDP(pkt)
|
var info parsedUDP
|
||||||
if !ok {
|
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||||
c.addPassthrough(pkt)
|
c.sealAllOpen()
|
||||||
|
c.addVerbatim(sp.pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return c.commitParsed(pkt, info)
|
return c.commitParsed(sp.pkt, &info)
|
||||||
}
|
}
|
||||||
|
|
||||||
// commitParsed is the post-parse half of Commit. The caller must have
|
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
|
||||||
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
// valid parse so the header is not re-walked here.
|
||||||
// avoid re-walking the IP/UDP header.
|
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
|
||||||
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
|
||||||
if c.gsoW == nil {
|
// coalesced.
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
|
|
||||||
// reach the TUN, but it can't be coalesced: a GSO slot would store an
|
|
||||||
// empty payload iovec and the kernel has nothing to segment. Seal any
|
|
||||||
// open chain for this flow (so a later, non-empty datagram seeds fresh
|
|
||||||
// *after* this one and per-flow arrival order is preserved) and deliver
|
|
||||||
// it as a plain single datagram.
|
|
||||||
if info.payLen == 0 {
|
if info.payLen == 0 {
|
||||||
delete(c.openSlots, info.fk)
|
c.sealFlow(info.fk)
|
||||||
c.addPassthrough(pkt)
|
c.addVerbatim(pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if open := c.openSlots[info.fk]; open != 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.canAppend(open, pkt, info) {
|
||||||
c.appendPayload(open, pkt, info)
|
if c.appendPayload(open, pkt, info) {
|
||||||
if open.sealed {
|
// Chain closed (short segment): stop extending it.
|
||||||
delete(c.openSlots, info.fk)
|
c.sealFlow(info.fk)
|
||||||
|
} else {
|
||||||
|
c.lastSlot = open
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
// Can't extend: evict it from openSlots and fall through to seed a
|
||||||
delete(c.openSlots, info.fk)
|
// fresh slot.
|
||||||
|
c.sealFlow(info.fk)
|
||||||
}
|
}
|
||||||
c.seed(pkt, info)
|
c.seed(pkt, info)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush drains every queued slot and calls the configured Resetter.
|
|
||||||
func (c *UDPCoalescer) Flush() error {
|
func (c *UDPCoalescer) Flush() error {
|
||||||
first := c.drain()
|
|
||||||
if c.resetter != nil {
|
|
||||||
c.resetter()
|
|
||||||
}
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
// drain emits every queued slot in arrival order and clears the slot state.
|
|
||||||
// It does NOT reset the arena: borrowed payload slices stay valid until the
|
|
||||||
// arena's owner recycles it.
|
|
||||||
func (c *UDPCoalescer) drain() error {
|
|
||||||
var first error
|
var first error
|
||||||
for _, s := range c.slots {
|
for _, s := range c.slots {
|
||||||
var err error
|
var err error
|
||||||
if s.passthrough {
|
if s.verbatim || s.numSeg == 1 {
|
||||||
_, err = c.plainW.Write(s.rawPkt)
|
// 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 {
|
} else {
|
||||||
err = c.flushSlot(s)
|
err = c.flushSlot(s)
|
||||||
}
|
}
|
||||||
@@ -204,25 +196,38 @@ func (c *UDPCoalescer) drain() error {
|
|||||||
clear(c.slots)
|
clear(c.slots)
|
||||||
c.slots = c.slots[:0]
|
c.slots = c.slots[:0]
|
||||||
clear(c.openSlots)
|
clear(c.openSlots)
|
||||||
|
c.lastSlot = nil
|
||||||
return first
|
return first
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
// 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 := c.take()
|
||||||
s.passthrough = true
|
s.verbatim = true
|
||||||
s.rawPkt = pkt
|
s.rawPkt = pkt
|
||||||
c.slots = append(c.slots, s)
|
c.slots = append(c.slots, s)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
||||||
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
if info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||||
c.addPassthrough(pkt)
|
// 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
|
return
|
||||||
}
|
}
|
||||||
s := c.take()
|
s := c.take()
|
||||||
s.passthrough = false
|
s.verbatim = false
|
||||||
s.rawPkt = nil
|
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
// the superpacket header flushSlot patches in place.
|
||||||
|
s.rawPkt = pkt
|
||||||
s.hdrLen = info.hdrLen
|
s.hdrLen = info.hdrLen
|
||||||
s.ipHdrLen = info.ipHdrLen
|
s.ipHdrLen = info.ipHdrLen
|
||||||
s.isV6 = info.fk.isV6
|
s.isV6 = info.fk.isV6
|
||||||
@@ -230,19 +235,16 @@ func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
|||||||
s.gsoSize = info.payLen
|
s.gsoSize = info.payLen
|
||||||
s.numSeg = 1
|
s.numSeg = 1
|
||||||
s.totalPay = info.payLen
|
s.totalPay = info.payLen
|
||||||
s.sealed = false
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
c.slots = append(c.slots, s)
|
c.slots = append(c.slots, s)
|
||||||
c.openSlots[info.fk] = s
|
c.openSlots[info.fk] = s
|
||||||
|
c.lastSlot = s
|
||||||
}
|
}
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed.
|
// canAppend reports whether info's packet extends the slot's seed.
|
||||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
// Kernel UDP-GSO requires every segment except possibly the last to be
|
||||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
||||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||||
if s.sealed {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.hdrLen != s.hdrLen {
|
if info.hdrLen != s.hdrLen {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -255,20 +257,25 @@ func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
|||||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
// 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 false
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
// 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.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
s.numSeg++
|
s.numSeg++
|
||||||
s.totalPay += info.payLen
|
s.totalPay += info.payLen
|
||||||
if info.payLen < s.gsoSize {
|
return info.payLen < s.gsoSize
|
||||||
// Last-segment-can-be-shorter: this seals the chain.
|
|
||||||
s.sealed = true
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) take() *udpSlot {
|
func (c *UDPCoalescer) take() *udpSlot {
|
||||||
@@ -282,30 +289,19 @@ func (c *UDPCoalescer) take() *udpSlot {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||||
s.passthrough = false
|
// Reset every field, identity ones included; see TCPCoalescer.release.
|
||||||
s.rawPkt = nil
|
|
||||||
clear(s.payIovs)
|
clear(s.payIovs)
|
||||||
s.payIovs = s.payIovs[:0]
|
*s = udpSlot{payIovs: s.payIovs[:0]}
|
||||||
s.numSeg = 0
|
|
||||||
s.totalPay = 0
|
|
||||||
s.sealed = false
|
|
||||||
c.pool = append(c.pool, s)
|
c.pool = append(c.pool, s)
|
||||||
}
|
}
|
||||||
|
|
||||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
// 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 length to the *total* across all coalesced segments, then seeds
|
||||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||||
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
|
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
|
||||||
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
// slot is released right after, so nothing re-reads the patched header.
|
||||||
// values would silently drop everything but the first segment. The kernel
|
|
||||||
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
|
||||||
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
|
||||||
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
|
||||||
// our seed's uh->check must be consistent with the seed's uh->len, which
|
|
||||||
// is what passing the total to both pseudoSum and the UDP length field
|
|
||||||
// guarantees.
|
|
||||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||||
hdr := s.hdrBuf[:s.hdrLen]
|
hdr := s.rawPkt[:s.hdrLen]
|
||||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
||||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
||||||
|
|
||||||
@@ -330,14 +326,11 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
|||||||
udpCsumOff := s.ipHdrLen + 6
|
udpCsumOff := s.ipHdrLen + 6
|
||||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
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
|
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||||
// every field that must be identical across coalesced segments. Length
|
// every field that must be identical across coalesced segments
|
||||||
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
|
|
||||||
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
|
|
||||||
// don't coalesce, matching kernel GRO.
|
|
||||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||||
if len(a) != len(b) {
|
if len(a) != len(b) {
|
||||||
return false
|
return false
|
||||||
@@ -345,11 +338,8 @@ func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|||||||
if !ipHeadersMatch(a, b, isV6) {
|
if !ipHeadersMatch(a, b, isV6) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] —
|
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
|
||||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
// length varies (we rewrite at flush) and the checksum will be redone.
|
||||||
udp := ipHdrLen
|
udp := ipHdrLen
|
||||||
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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))
|
||||||
|
}
|
||||||
@@ -1,7 +1,9 @@
|
|||||||
package batch
|
package batch
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -58,29 +60,31 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
|||||||
return pkt
|
return pkt
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
|
||||||
w := &fakeTunWriter{gsoEnabled: false}
|
// do USO. See newTestTCPCoalescer.
|
||||||
arena := NewArena(0)
|
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
tb.Helper()
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
c := NewUDPCoalescer(w)
|
||||||
if err := c.Commit(pkt); err != nil {
|
if c == nil {
|
||||||
t.Fatal(err)
|
tb.Fatal("NewUDPCoalescer: writer does not support USO")
|
||||||
}
|
}
|
||||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
return c
|
||||||
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
}
|
||||||
|
|
||||||
|
// 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 err := c.Flush(); err != nil {
|
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
|
||||||
t.Fatal(err)
|
t.Fatalf("want nil for a plain writer, got %v", c)
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
// ICMP packet
|
// ICMP packet
|
||||||
pkt := make([]byte, 28)
|
pkt := make([]byte, 28)
|
||||||
pkt[0] = 0x45
|
pkt[0] = 0x45
|
||||||
@@ -101,8 +105,7 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
|||||||
|
|
||||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||||
if err := c.Commit(pkt); err != nil {
|
if err := c.Commit(pkt); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -110,17 +113,21 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
// A slot that never grew past one datagram flushes as a plain Write of
|
||||||
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
// the original packet bytes: the original (already valid) checksum
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
// 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))
|
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) {
|
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 1200)
|
pay := make([]byte, 1200)
|
||||||
for i := 0; i < 3; i++ {
|
for i := 0; i < 3; i++ {
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
@@ -160,8 +167,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
|||||||
// Last segment may be shorter, sealing the chain.
|
// Last segment may be shorter, sealing the chain.
|
||||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
full := make([]byte, 1200)
|
full := make([]byte, 1200)
|
||||||
tail := make([]byte, 600)
|
tail := make([]byte, 600)
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
@@ -180,22 +186,23 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites) != 2 {
|
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
|
||||||
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
// 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 {
|
if len(w.gsoWrites[0].pays) != 3 {
|
||||||
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites[1].pays) != 1 {
|
if got, want := len(w.writes[0]), 20+8+1200; got != want {
|
||||||
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
|
t.Errorf("re-seed plain write len=%d want %d", got, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
@@ -205,16 +212,21 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites) != 2 {
|
// Both seeds stay single-segment → two plain writes in arrival order.
|
||||||
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
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.
|
// Different 5-tuples must not coalesce.
|
||||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 800)
|
pay := make([]byte, 800)
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -245,8 +257,7 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
|||||||
// Caps at udpCoalesceMaxSegs.
|
// Caps at udpCoalesceMaxSegs.
|
||||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 100)
|
pay := make([]byte, 100)
|
||||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
@@ -271,12 +282,12 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
|||||||
|
|
||||||
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
||||||
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
||||||
// seals the Not-ECT chain and seeds a fresh superpacket that keeps CE; the
|
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
|
||||||
// trailing Not-ECT datagram seeds another.
|
// 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) {
|
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 800)
|
pay := make([]byte, 800)
|
||||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
pkt1 := buildUDPv4(1000, 53, pay)
|
||||||
@@ -290,16 +301,13 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites) != 3 {
|
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||||
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
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}
|
wantECN := []byte{0x00, 0x03, 0x00}
|
||||||
for i, g := range w.gsoWrites {
|
for i, p := range w.writes {
|
||||||
if len(g.pays) != 1 {
|
if got := p[1] & 0x03; got != wantECN[i] {
|
||||||
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||||
}
|
|
||||||
if got := g.hdr[1] & 0x03; got != wantECN[i] {
|
|
||||||
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -307,8 +315,7 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
|||||||
// IPv6 path: same flow, equal-sized → coalesced.
|
// IPv6 path: same flow, equal-sized → coalesced.
|
||||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 1200)
|
pay := make([]byte, 1200)
|
||||||
for i := 0; i < 3; i++ {
|
for i := 0; i < 3; i++ {
|
||||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||||
@@ -344,8 +351,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
|||||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
||||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pay := make([]byte, 800)
|
pay := make([]byte, 800)
|
||||||
pkt0 := buildUDPv4(1000, 53, pay)
|
pkt0 := buildUDPv4(1000, 53, pay)
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
pkt1 := buildUDPv4(1000, 53, pay)
|
||||||
@@ -359,16 +365,16 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites) != 2 {
|
// Both seeds stay single-segment → two plain writes, no gso.
|
||||||
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
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.
|
// Fragmented IPv4 must not be coalesced.
|
||||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||||
if err := c.Commit(pkt); err != nil {
|
if err := c.Commit(pkt); err != nil {
|
||||||
@@ -389,8 +395,7 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
|||||||
// reach the GSO path. Regression: must not panic and must be written.
|
// reach the GSO path. Regression: must not panic and must be written.
|
||||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||||
if err := c.Commit(pkt); err != nil {
|
if err := c.Commit(pkt); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -406,11 +411,10 @@ func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// IPv6 zero-length UDP datagram: same passthrough contract as v4.
|
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
||||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||||
if err := c.Commit(pkt); err != nil {
|
if err := c.Commit(pkt); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -431,8 +435,7 @@ func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
|||||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
||||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
full := make([]byte, 800)
|
full := make([]byte, 800)
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -447,17 +450,23 @@ func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// The empty datagram sealed the first slot, so the trailing full packet
|
// The empty datagram sealed the first slot, so the trailing full packet
|
||||||
// can't join it: two single-segment superpackets bracket one plain write.
|
// can't join it. All three emit as plain writes (the two full datagrams
|
||||||
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
|
// stayed single-segment; the empty one is verbatim) in per-flow
|
||||||
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
// 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).
|
// IPv4 with options is not admissible (we require IHL=5).
|
||||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
arena := NewArena(0)
|
c := newTestUDPCoalescer(t, w)
|
||||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||||
if err := c.Commit(pkt); err != nil {
|
if err := c.Commit(pkt); err != nil {
|
||||||
@@ -470,3 +479,58 @@ func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
|||||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -8,37 +8,69 @@ import (
|
|||||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
// archImpl names one checksum function under test. The per-arch
|
||||||
// seeds and a handful of starting alignments, asserting that our local
|
// export_*_test.go files enumerate the hand-written implementations so the
|
||||||
// Checksum matches gvisor's reference bit-for-bit.
|
// suite compares each one against gvisor directly, regardless of which one
|
||||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
// the public Checksum dispatches to on the running CPU. Testing only the
|
||||||
rng := rand.New(rand.NewPCG(1, 2))
|
// dispatcher was tautological wherever it resolved to the gvisor fallback
|
||||||
const padFront = 16
|
// (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
|
||||||
|
}
|
||||||
|
|
||||||
// Random pool large enough for the longest case + alignment slop.
|
// implsUnderTest is the public dispatcher plus every arch implementation.
|
||||||
pool := make([]byte, 4096+padFront)
|
func implsUnderTest() []archImpl {
|
||||||
for i := range pool {
|
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
|
||||||
pool[i] = byte(rng.Uint32())
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
// 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
|
||||||
|
|
||||||
for length := 0; length <= 4096; length++ {
|
// Random pool large enough for the longest case + alignment slop.
|
||||||
for _, seed := range seeds {
|
pool := make([]byte, 4096+padFront)
|
||||||
for _, off := range offsets {
|
for i := range pool {
|
||||||
if off+length > len(pool) {
|
pool[i] = byte(rng.Uint32())
|
||||||
continue
|
}
|
||||||
}
|
|
||||||
buf := pool[off : off+length]
|
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||||
got := Checksum(buf, seed)
|
|
||||||
if got != want {
|
for length := 0; length <= 4096; length++ {
|
||||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
for _, seed := range seeds {
|
||||||
length, off, seed, got, want)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -46,23 +78,28 @@ func TestChecksumMatchesGvisor(t *testing.T) {
|
|||||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||||
// alternating, and ascending sequences.
|
// alternating, and ascending sequences.
|
||||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||||
for length := 0; length <= 256; length++ {
|
for _, impl := range implsUnderTest() {
|
||||||
patterns := map[string][]byte{
|
t.Run(impl.name, func(t *testing.T) {
|
||||||
"zeros": make([]byte, length),
|
requireAvailable(t, impl)
|
||||||
"ones": bytes(length, 0xff),
|
for length := 0; length <= 256; length++ {
|
||||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
patterns := map[string][]byte{
|
||||||
"ascending": ascending(length),
|
"zeros": make([]byte, length),
|
||||||
}
|
"ones": bytes(length, 0xff),
|
||||||
for name, buf := range patterns {
|
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
"ascending": ascending(length),
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
}
|
||||||
got := Checksum(buf, seed)
|
for name, buf := range patterns {
|
||||||
if got != want {
|
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
name, length, seed, got, want)
|
got := impl.fn(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
name, length, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -98,36 +135,41 @@ func ascending(n int) []byte {
|
|||||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
// 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.
|
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
||||||
func TestChecksumTailPaths(t *testing.T) {
|
func TestChecksumTailPaths(t *testing.T) {
|
||||||
rng := rand.New(rand.NewPCG(42, 17))
|
for _, impl := range implsUnderTest() {
|
||||||
const padFront = 16
|
t.Run(impl.name, func(t *testing.T) {
|
||||||
const maxK = 8
|
requireAvailable(t, impl)
|
||||||
|
rng := rand.New(rand.NewPCG(42, 17))
|
||||||
|
const padFront = 16
|
||||||
|
const maxK = 8
|
||||||
|
|
||||||
pool := make([]byte, 64*maxK+padFront+64)
|
pool := make([]byte, 64*maxK+padFront+64)
|
||||||
for i := range pool {
|
for i := range pool {
|
||||||
pool[i] = byte(rng.Uint32())
|
pool[i] = byte(rng.Uint32())
|
||||||
}
|
}
|
||||||
|
|
||||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||||
|
|
||||||
for k := 0; k <= maxK; k++ {
|
for k := 0; k <= maxK; k++ {
|
||||||
for tail := 0; tail < 64; tail++ {
|
for tail := 0; tail < 64; tail++ {
|
||||||
length := 64*k + tail
|
length := 64*k + tail
|
||||||
for _, seed := range seeds {
|
for _, seed := range seeds {
|
||||||
for _, off := range offsets {
|
for _, off := range offsets {
|
||||||
if off+length > len(pool) {
|
if off+length > len(pool) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
buf := pool[off : off+length]
|
buf := pool[off : off+length]
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
got := Checksum(buf, seed)
|
got := impl.fn(buf, seed)
|
||||||
if got != want {
|
if got != want {
|
||||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||||
k, tail, length, off, seed, got, want)
|
k, tail, length, off, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -9,10 +9,9 @@ import (
|
|||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for
|
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
||||||
// reads, POLLOUT for writes) or shutdownFd signals teardown. It builds the
|
// (events is POLLIN for reads, POLLOUT for writes)
|
||||||
// pollfd array on the stack every call, so concurrent callers on the same
|
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
||||||
// Queue never share Revents storage.
|
|
||||||
//
|
//
|
||||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
||||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
@@ -17,18 +18,17 @@ type offloadQueueSet struct {
|
|||||||
// pqi is exactly the same as pq, but stored as the interface type
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
pqi []Queue
|
pqi []Queue
|
||||||
shutdownFd int
|
shutdownFd int
|
||||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6
|
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||||
// with the kernel. Queues created by Add inherit this and surface it
|
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||||
// via Offload.USOSupported so coalescers can gate USO emission.
|
|
||||||
usoEnabled bool
|
usoEnabled bool
|
||||||
closed atomic.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
|
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||||
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||||
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||||
// should fall back to per-packet writes when this is false.
|
|
||||||
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
|
||||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
@@ -39,6 +39,7 @@ func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
|||||||
pqi: []Queue{},
|
pqi: []Queue{},
|
||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
usoEnabled: usoEnabled,
|
usoEnabled: usoEnabled,
|
||||||
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
@@ -49,7 +50,10 @@ func (c *offloadQueueSet) Queues() []Queue {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *offloadQueueSet) Add(fd int) error {
|
func (c *offloadQueueSet) Add(fd int) error {
|
||||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
if c.closed.Load() {
|
||||||
|
return errors.New("queue set already closed")
|
||||||
|
}
|
||||||
|
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -73,23 +77,21 @@ func (c *offloadQueueSet) Close() error {
|
|||||||
|
|
||||||
errs := []error{}
|
errs := []error{}
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit. They observe
|
// Signal all readers blocked in poll to wake up and exit.
|
||||||
// POLLIN on the shutdown eventfd and return os.ErrClosed.
|
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||||
if err := c.wakeForShutdown(); err != nil {
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||||
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
|
||||||
// to this container.
|
|
||||||
for _, x := range c.pq {
|
for _, x := range c.pq {
|
||||||
if err := x.Close(); err != nil {
|
if err := x.Close(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close the shutdown eventfd last: every reader's pollfd set references
|
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||||
// it, so it must outlive the wake + per-queue teardown above.
|
// so it must outlive the wake + per-queue teardown above.
|
||||||
if err := unix.Close(c.shutdownFd); err != nil {
|
if err := unix.Close(c.shutdownFd); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -40,6 +40,9 @@ func (c *pollQueueSet) Queues() []Queue {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *pollQueueSet) Add(fd int) error {
|
func (c *pollQueueSet) Add(fd int) error {
|
||||||
|
if c.closed.Load() {
|
||||||
|
return errors.New("queue set already closed")
|
||||||
|
}
|
||||||
x, err := newPoll(fd, c.shutdownFd)
|
x, err := newPoll(fd, c.shutdownFd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -64,23 +67,21 @@ func (c *pollQueueSet) Close() error {
|
|||||||
|
|
||||||
errs := []error{}
|
errs := []error{}
|
||||||
|
|
||||||
// Wake any reader blocked in poll so it observes POLLIN on the shutdown
|
// Signal all readers blocked in poll to wake up and exit.
|
||||||
// eventfd and returns os.ErrClosed.
|
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||||
if err := c.wakeForShutdown(); err != nil {
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||||
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
|
||||||
// to this container.
|
|
||||||
for _, x := range c.pq {
|
for _, x := range c.pq {
|
||||||
if err := x.Close(); err != nil {
|
if err := x.Close(); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close the shutdown eventfd last: every reader's pollfd set references
|
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||||
// it, so it must outlive the wake + per-queue teardown above.
|
// so it must outlive the wake + per-queue teardown above.
|
||||||
if err := unix.Close(c.shutdownFd); err != nil {
|
if err := unix.Close(c.shutdownFd); err != nil {
|
||||||
errs = append(errs, err)
|
errs = append(errs, err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
//go:build !linux || android || e2e_testing
|
//go:build !linux || android
|
||||||
|
|
||||||
package tio
|
package tio
|
||||||
|
|
||||||
@@ -8,12 +8,6 @@ func protoFromGSOType(_ uint8) (GSOProto, error) {
|
|||||||
return 0, fmt.Errorf("GSO unsupported")
|
return 0, fmt.Errorf("GSO unsupported")
|
||||||
}
|
}
|
||||||
|
|
||||||
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
|
|
||||||
// builds (and Android/e2e_testing) this package does not provide a Queue
|
|
||||||
// implementation, so any caller that does construct a Packet here can only
|
|
||||||
// be operating on non-superpacket bytes and the stub forwards them
|
|
||||||
// directly. A non-zero GSO field is a programming error from the caller
|
|
||||||
// and returns an explicit error rather than silently misbehaving.
|
|
||||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||||
if pkt.GSO.IsSuperpacket() {
|
if pkt.GSO.IsSuperpacket() {
|
||||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||||
|
|||||||
@@ -4,9 +4,8 @@ import "io"
|
|||||||
|
|
||||||
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
||||||
// Read fills a private scratch buffer and returns exactly one Packet whose
|
// 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
|
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
|
||||||
// Queue contract. Single-reader like every Queue; Write is exactly as safe
|
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
|
||||||
// for concurrent use as the underlying source's Write.
|
|
||||||
type singleQueue struct {
|
type singleQueue struct {
|
||||||
rw io.ReadWriter
|
rw io.ReadWriter
|
||||||
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
||||||
@@ -14,9 +13,9 @@ type singleQueue struct {
|
|||||||
ret [1]Packet
|
ret [1]Packet
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun
|
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
|
||||||
// device) into a Queue. bufSize is the per-queue read scratch size and must
|
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
|
||||||
// be at least the largest datagram the source can return. Close closes rwc.
|
// Close closes rwc.
|
||||||
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
||||||
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
||||||
}
|
}
|
||||||
|
|||||||
+60
-85
@@ -13,79 +13,72 @@ type QueueSet interface {
|
|||||||
Add(fd int) error
|
Add(fd int) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// Capabilities advertises which kernel offload features a Queue
|
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||||
// successfully negotiated. Callers consult this to decide which coalescers
|
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||||
// to wire onto the write path — a Queue without TSO can't usefully accept a
|
|
||||||
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
|
|
||||||
type Capabilities struct {
|
type Capabilities struct {
|
||||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed
|
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
|
||||||
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe.
|
// and WriteGSO with GSOProtoTCP is safe.
|
||||||
TSO bool
|
TSO bool
|
||||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so
|
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
|
||||||
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||||
USO bool
|
USO bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// Queue is a readable/writable Poll queue. Concurrency contract: a single
|
// Queue is a readable/writable Poll queue.
|
||||||
// read goroutine drives Read; plain Write is safe for concurrent callers;
|
// 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.
|
// 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 {
|
type Queue interface {
|
||||||
io.Closer
|
io.Closer
|
||||||
|
|
||||||
// Read returns one or more packets. The returned Packet.Bytes slices
|
// Read returns one or more packets.
|
||||||
// are borrowed from the Queue's internal buffer and are only valid
|
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
|
||||||
// until the next Read or Close on this Queue - callers must encrypt
|
// until the next Read or Close on this Queue.
|
||||||
// or copy each slice before the next call. A Packet may carry a
|
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
|
||||||
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
|
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
||||||
// true the caller must segment Bytes before treating it as a single
|
|
||||||
// IP datagram. Single-reader only: not safe for concurrent Reads (it
|
|
||||||
// reuses per-queue rx scratch each call).
|
|
||||||
Read() ([]Packet, error)
|
Read() ([]Packet, error)
|
||||||
|
|
||||||
// Write emits a single packet on the plaintext (outside→inside)
|
// Write emits a single packet on the plaintext (outside→inside) delivery path.
|
||||||
// delivery path. Safe for concurrent use.
|
// Safe for concurrent use.
|
||||||
Write(p []byte) (int, error)
|
Write(p []byte) (int, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Packet is the unit Queue.Read returns. Bytes points into the queue's
|
// Packet is the unit Queue.Read returns.
|
||||||
// internal buffer and is only valid until the next Read or Close on the
|
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
|
||||||
// queue that produced it. GSO is the zero value for an already-segmented
|
// GSO is the zero value for an already-segmented IP datagram;
|
||||||
// IP datagram; when non-zero it describes a kernel-supplied TSO/USO
|
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
|
||||||
// superpacket the caller must segment before consuming.
|
|
||||||
type Packet struct {
|
type Packet struct {
|
||||||
Bytes []byte
|
Bytes []byte
|
||||||
GSO GSOInfo
|
GSO GSOInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||||
// The zero value means "not a superpacket" — Bytes is one regular IP
|
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
|
||||||
// datagram and no segmentation is required.
|
|
||||||
type GSOInfo struct {
|
type GSOInfo struct {
|
||||||
// Size is the GSO segment size: max payload bytes per segment
|
// Size is the GSO segment size: max payload bytes per segment
|
||||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means
|
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
|
||||||
// not a superpacket.
|
|
||||||
Size uint16
|
Size uint16
|
||||||
// HdrLen is the total L3+L4 header length within Bytes (already
|
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
|
||||||
// corrected via correctHdrLen, so safe to slice on).
|
|
||||||
HdrLen uint16
|
HdrLen uint16
|
||||||
// CsumStart is the L4 header offset inside Bytes (== L3 header
|
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
|
||||||
// length).
|
|
||||||
CsumStart uint16
|
CsumStart uint16
|
||||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows
|
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
|
||||||
// which checksum/header layout to apply.
|
|
||||||
Proto GSOProto
|
Proto GSOProto
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
||||||
// superpacket that needs segmentation before its bytes can be encrypted
|
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
|
||||||
// and sent on the wire.
|
|
||||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
||||||
|
|
||||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
// 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.
|
// safe to retain past the next Read or Close on the originating Queue.
|
||||||
// GSO metadata is copied verbatim. Use this only when a caller genuinely
|
// GSO metadata is copied verbatim.
|
||||||
// needs to outlive the borrowed-slice contract — the hot path reads should
|
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
|
||||||
// continue to consume the borrow synchronously to avoid the allocation.
|
|
||||||
func (p Packet) Clone() Packet {
|
func (p Packet) Clone() Packet {
|
||||||
if p.Bytes == nil {
|
if p.Bytes == nil {
|
||||||
return p
|
return p
|
||||||
@@ -95,78 +88,60 @@ func (p Packet) Clone() Packet {
|
|||||||
return Packet{Bytes: cp, GSO: p.GSO}
|
return Packet{Bytes: cp, GSO: p.GSO}
|
||||||
}
|
}
|
||||||
|
|
||||||
// CapsProvider is an optional interface implemented by Queues that
|
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
|
||||||
// successfully negotiated kernel offload features at open time. Callers
|
// Callers pick a write-path coalescer based on the result.
|
||||||
// pick a write-path coalescer based on the result. Queues that don't
|
// Queues that don't implement it are treated as having no offload capability.
|
||||||
// implement it are treated as having no offload capability — callers must
|
|
||||||
// fall back to plain per-packet writes.
|
|
||||||
type CapsProvider interface {
|
type CapsProvider interface {
|
||||||
Capabilities() Capabilities
|
Capabilities() Capabilities
|
||||||
}
|
}
|
||||||
|
|
||||||
// QueueCapabilities returns q's negotiated offload capabilities, or the
|
// GSOProto selects the L4 protocol for a GSO superpacket.
|
||||||
// zero value when q does not advertise any.
|
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||||
func QueueCapabilities(q Queue) Capabilities {
|
|
||||||
if cp, ok := q.(CapsProvider); ok {
|
|
||||||
return cp.Capabilities()
|
|
||||||
}
|
|
||||||
return Capabilities{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
|
|
||||||
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
|
||||||
// inside the transport header virtio NEEDS_CSUM expects.
|
// inside the transport header virtio NEEDS_CSUM expects.
|
||||||
type GSOProto uint8
|
type GSOProto uint8
|
||||||
|
|
||||||
const (
|
const (
|
||||||
GSOProtoTCP GSOProto = iota
|
GSOProtoUnknown GSOProto = iota
|
||||||
|
GSOProtoTCP
|
||||||
GSOProtoUDP
|
GSOProtoUDP
|
||||||
)
|
)
|
||||||
|
|
||||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||||
// assembled from a header prefix plus one or more borrowed payload
|
// assembled from a header prefix plus one or more borrowed payload fragments,
|
||||||
// fragments, in a single vectored write (writev with a leading
|
// in a single vectored write (writev with a leading virtio_net_hdr).
|
||||||
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
|
||||||
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
// Backends without GSO support do not implement this interface and coalescing is skipped.
|
||||||
// support do not implement this interface and coalescing is skipped.
|
|
||||||
//
|
//
|
||||||
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
|
||||||
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
// transportHdr is the TCP or UDP header
|
||||||
// header (mutable - the L4 checksum field must hold the pseudo-header
|
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||||
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
|
||||||
// pays are non-overlapping payload fragments whose concatenation is the
|
// They are read-only from the writer's perspective and must remain valid until the call returns.
|
||||||
// full superpacket payload; they are read-only from the writer's
|
// Every segment in pays except possibly the last must be exactly the same size.
|
||||||
// perspective and must remain valid until the call returns. Every segment
|
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
|
||||||
// in pays except possibly the last is exactly the same size. proto picks
|
|
||||||
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
|
|
||||||
//
|
//
|
||||||
// Callers should also consult CapsProvider (via SupportsGSO or
|
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
|
||||||
// QueueCapabilities) for the per-protocol negotiated capability; an
|
// USO may not have been negotiated even when TSO was.
|
||||||
// implementation of GSOWriter is necessary but not sufficient since USO
|
|
||||||
// may not have been negotiated even when TSO was.
|
|
||||||
type GSOWriter interface {
|
type GSOWriter interface {
|
||||||
|
io.Writer
|
||||||
|
CapsProvider
|
||||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||||
// queue advertises the negotiated capability for `want`. A writer that
|
// queue advertises the negotiated capability for `want`.
|
||||||
// implements GSOWriter but not CapsProvider is treated as permissive
|
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
|
||||||
// (used by tests and fakes that don't negotiate).
|
|
||||||
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
|
|
||||||
gw, ok := w.(GSOWriter)
|
gw, ok := w.(GSOWriter)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
cp, ok := w.(CapsProvider)
|
caps := gw.Capabilities()
|
||||||
if !ok {
|
|
||||||
return gw, true
|
|
||||||
}
|
|
||||||
caps := cp.Capabilities()
|
|
||||||
switch want {
|
switch want {
|
||||||
case GSOProtoTCP:
|
case GSOProtoTCP:
|
||||||
return gw, caps.TSO
|
return gw, caps.TSO
|
||||||
case GSOProtoUDP:
|
case GSOProtoUDP:
|
||||||
return gw, caps.USO
|
return gw, caps.USO
|
||||||
|
default:
|
||||||
|
return gw, false
|
||||||
}
|
}
|
||||||
return gw, false
|
|
||||||
}
|
}
|
||||||
|
|||||||
+153
-150
@@ -4,6 +4,7 @@
|
|||||||
package tio
|
package tio
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
@@ -17,71 +18,56 @@ import (
|
|||||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||||
)
|
)
|
||||||
|
|
||||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one
|
const maxSuperpacketLen = 65535
|
||||||
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
|
|
||||||
|
// 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
|
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||||
// We round up to give comfortable margin for the drain headroom check
|
// We round up to give margin for the drain headroom check below.
|
||||||
// below.
|
|
||||||
const tunRxBufSize = 64 * 1024
|
const tunRxBufSize = 64 * 1024
|
||||||
|
|
||||||
// tunRxBufCap is the total size we allocate for the per-reader rx
|
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
|
||||||
// buffer. With reads landing directly in rxBuf, each drain iteration
|
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||||
// 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,
|
||||||
// Sized to eight such iterations so a single poll wake can drain several
|
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
|
||||||
// TSO/USO superpackets under bulk load, amortizing the wake and giving
|
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
|
||||||
// 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
|
const tunRxBufCap = tunRxBufSize * 8
|
||||||
|
|
||||||
// tunDrainCap caps how many packets a single Read will accumulate via
|
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
|
||||||
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
|
||||||
// bounding how much work a single caller holds before handing off.
|
|
||||||
const tunDrainCap = 64
|
const tunDrainCap = 64
|
||||||
|
|
||||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call: 3 fixed
|
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
|
||||||
// entries (virtio_net_hdr, IP hdr, transport hdr) plus up to gsoMaxIovs-3
|
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
|
||||||
// payload fragments. Sized comfortably above the typical kernel GSO
|
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
|
||||||
// segment cap (Linux UDP_GRO is 64) so realistic coalesced bursts never
|
// so realistic coalesced bursts never touch the limit.
|
||||||
// touch the limit. iovecs are tiny (16 bytes), so the entire scratch is
|
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
|
||||||
// 4 KiB — fine to keep resident on every queue. WriteGSO returns an error
|
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
|
||||||
// rather than reallocating when a caller exceeds this budget.
|
|
||||||
const gsoMaxIovs = 256
|
const gsoMaxIovs = 256
|
||||||
|
|
||||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
|
||||||
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
|
||||||
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
|
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
|
||||||
// verification. All packets that reach the plain Write paths already carry
|
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
|
||||||
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we
|
// L4 checksum anyway. What matters here is what the header does NOT say:
|
||||||
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during
|
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
|
||||||
// superpacket segmentation, or built locally by CreateRejectPacket), so
|
// All packets that reach the plain Write paths already carry a valid L4 checksum.
|
||||||
// trusting them is safe.
|
|
||||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
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.
|
// 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.
|
// 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 {
|
type Offload struct {
|
||||||
fd int
|
fd int
|
||||||
shutdownFd int
|
shutdownFd int
|
||||||
closed atomic.Bool
|
|
||||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
|
||||||
rxOff int // cursor into rxBuf for the current Read drain
|
|
||||||
pending []Packet // packets returned from the most recent Read
|
|
||||||
|
|
||||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
|
||||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
|
||||||
// lets us read the body directly into rxBuf at the current rxOff with
|
|
||||||
// no userspace copy on the GSO_NONE fast path.
|
|
||||||
readVnetScratch [virtio.Size]byte
|
|
||||||
// readIovs is the readv(2) iovec scratch wired once at construction —
|
|
||||||
// iovec[0] points at readVnetScratch; iovec[1].Base/Len is updated per
|
|
||||||
// read to address the current rxBuf slot.
|
|
||||||
readIovs [2]unix.Iovec
|
|
||||||
|
|
||||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
// 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.
|
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||||
usoEnabled bool
|
usoEnabled bool
|
||||||
|
closed atomic.Bool
|
||||||
|
|
||||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
// 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
|
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
||||||
@@ -92,9 +78,27 @@ type Offload struct {
|
|||||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
||||||
// (and drops the call) if a caller hands it more fragments than fit.
|
// (and drops the call) if a caller hands it more fragments than fit.
|
||||||
gsoIovs []unix.Iovec
|
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) (*Offload, error) {
|
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||||
}
|
}
|
||||||
@@ -104,6 +108,7 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
|||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
usoEnabled: usoEnabled,
|
usoEnabled: usoEnabled,
|
||||||
closed: atomic.Bool{},
|
closed: atomic.Bool{},
|
||||||
|
l: l,
|
||||||
|
|
||||||
rxBuf: make([]byte, tunRxBufCap),
|
rxBuf: make([]byte, tunRxBufCap),
|
||||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
||||||
@@ -128,26 +133,18 @@ func (r *Offload) blockOnWrite() error {
|
|||||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
||||||
}
|
}
|
||||||
|
|
||||||
// readPacket issues a single readv(2) splitting the virtio_net_hdr off
|
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
|
||||||
// into readVnetScratch and reading the packet body directly into rxBuf at
|
// and reading the packet body directly into rxBuf at the current rxOff.
|
||||||
// the current rxOff. Returns the body length (zero virtio header bytes,
|
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
|
||||||
// just the IP packet/superpacket). block controls whether EAGAIN is
|
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
|
||||||
// retried via poll: the initial read of a drain blocks; subsequent drain
|
|
||||||
// reads do not.
|
|
||||||
//
|
|
||||||
// The body iovec capacity is always tunReadBufSize; callers (the Read
|
|
||||||
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
|
|
||||||
// hold one worst-case kernel-supplied packet body. Without that gate the
|
|
||||||
// body iovec could be smaller than the next inbound packet and the
|
|
||||||
// kernel would truncate.
|
|
||||||
func (r *Offload) readPacket(block bool) (int, error) {
|
func (r *Offload) readPacket(block bool) (int, error) {
|
||||||
for {
|
for {
|
||||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||||
r.readIovs[1].SetLen(tunReadBufSize)
|
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)))
|
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||||
if errno == 0 {
|
if errno == 0 {
|
||||||
if int(n) < virtio.Size {
|
if int(n) < virtio.Size {
|
||||||
return 0, io.ErrShortWrite
|
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
|
||||||
}
|
}
|
||||||
return int(n) - virtio.Size, nil
|
return int(n) - virtio.Size, nil
|
||||||
}
|
}
|
||||||
@@ -170,29 +167,30 @@ func (r *Offload) readPacket(block bool) (int, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Read returns one or more packets from the tun. Each Packet either
|
// Read returns one or more packets from the tun.
|
||||||
// carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO
|
// 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).
|
||||||
// superpacket plus the GSOInfo a caller needs to segment it (see
|
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
|
||||||
// SegmentSuperpacket). The first read blocks via poll; once the fd is
|
// - the kernel queue is empty (EAGAIN)
|
||||||
// known readable we drain additional packets non-blocking until the
|
// - we've collected tunDrainCap packets,
|
||||||
// kernel queue is empty (EAGAIN), we've collected tunDrainCap packets,
|
// - or we're out of rxBuf headroom.
|
||||||
// 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
|
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
|
||||||
// into the Offload's internal buffer and are only valid until the next
|
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
|
||||||
// Read or Close on this Queue.
|
|
||||||
func (r *Offload) Read() ([]Packet, error) {
|
func (r *Offload) Read() ([]Packet, error) {
|
||||||
r.pending = r.pending[:0]
|
r.pending = r.pending[:0]
|
||||||
r.rxOff = 0
|
r.rxOff = 0
|
||||||
|
|
||||||
// Initial (blocking) read. Retry on decode errors so a single bad
|
// Initial (blocking) read.
|
||||||
// packet does not stall the reader.
|
// Retry on decode errors so a single bad packet does not stall the reader.
|
||||||
for {
|
for {
|
||||||
n, err := r.readPacket(true)
|
n, err := r.readPacket(true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := r.decodeRead(n); err != nil {
|
if err := r.decodeRead(n); err != nil {
|
||||||
// Drop and read again — a bad packet should not kill the reader.
|
// 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
|
continue
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
@@ -214,6 +212,7 @@ func (r *Offload) Read() ([]Packet, error) {
|
|||||||
if err := r.decodeRead(n); err != nil {
|
if err := r.decodeRead(n); err != nil {
|
||||||
// Drop this packet and stop the drain; we'd rather hand off
|
// Drop this packet and stop the drain; we'd rather hand off
|
||||||
// what we have than keep spinning here.
|
// what we have than keep spinning here.
|
||||||
|
r.logDroppedRead(err)
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -221,13 +220,20 @@ func (r *Offload) Read() ([]Packet, error) {
|
|||||||
return r.pending, nil
|
return r.pending, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length
|
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
|
||||||
// pktLen). The bytes stay in rxBuf — for GSO_NONE we slice them as a
|
// header. Debug-gated so the happy path never pays for attribute assembly.
|
||||||
// regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
func (r *Offload) logDroppedRead(err error) {
|
||||||
// for TSO/USO superpackets we attach the corrected GSO metadata so the
|
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
// caller can segment lazily at encrypt time. rxOff advances past the
|
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
|
||||||
// kernel-supplied body and nothing else, since segmentation no longer
|
}
|
||||||
// writes back into rxBuf.
|
}
|
||||||
|
|
||||||
|
// 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 {
|
func (r *Offload) decodeRead(pktLen int) error {
|
||||||
if pktLen <= 0 {
|
if pktLen <= 0 {
|
||||||
return fmt.Errorf("short tun read: %d", pktLen)
|
return fmt.Errorf("short tun read: %d", pktLen)
|
||||||
@@ -237,7 +243,7 @@ func (r *Offload) decodeRead(pktLen int) error {
|
|||||||
|
|
||||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
||||||
|
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -248,17 +254,13 @@ func (r *Offload) decodeRead(pktLen int) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
|
|
||||||
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
|
|
||||||
// the metadata. The bytes stay in rxBuf untouched, segmentation
|
|
||||||
// happens in SegmentSuperpacket at encrypt time.
|
|
||||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
proto, err := protoFromGSOType(hdr.GSOType)
|
proto, err := protoFromGSOType(hdr.GSOType())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -276,22 +278,16 @@ func (r *Offload) decodeRead(pktLen int) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *Offload) Write(buf []byte) (int, error) {
|
func (r *Offload) Write(buf []byte) (int, error) {
|
||||||
|
if len(buf) == 0 {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
iovs := [2]unix.Iovec{
|
iovs := [2]unix.Iovec{
|
||||||
{Base: &validVnetHdr[0]},
|
{Base: &validVnetHdr[0]},
|
||||||
{Base: &buf[0]},
|
{Base: &buf[0]},
|
||||||
}
|
}
|
||||||
iovs[0].SetLen(virtio.Size)
|
iovs[0].SetLen(virtio.Size)
|
||||||
iovs[1].SetLen(len(buf))
|
iovs[1].SetLen(len(buf))
|
||||||
return r.writeWithScratch(buf, &iovs)
|
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
|
|
||||||
if len(buf) == 0 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
iovs[1].Base = &buf[0]
|
|
||||||
iovs[1].SetLen(len(buf))
|
|
||||||
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||||
@@ -321,57 +317,34 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
|||||||
|
|
||||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
// Capabilities reports the offload features negotiated for this Queue. TSO
|
||||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
// 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
|
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
|
||||||
// (Linux ≥ 6.2).
|
|
||||||
func (r *Offload) Capabilities() Capabilities {
|
func (r *Offload) Capabilities() Capabilities {
|
||||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
return Capabilities{TSO: true, USO: r.usoEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
||||||
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
|
if len(pays) == 0 {
|
||||||
|
// There are no payload fragments. There is nothing to send.
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after
|
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
|
||||||
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
|
|
||||||
var csumOff uint16
|
|
||||||
switch proto {
|
switch proto {
|
||||||
case GSOProtoUDP:
|
case GSOProtoUDP:
|
||||||
csumOff = 6
|
csumOff = 6
|
||||||
default:
|
case GSOProtoTCP:
|
||||||
csumOff = 16
|
csumOff = 16
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unknown GSO proto: %d", proto)
|
||||||
}
|
}
|
||||||
vhdr := virtio.Hdr{
|
// Incorrect geometry must cause an error, not a silent drop.
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
// No sane packet should ever make it inside this branch.
|
||||||
HdrLen: uint16(len(hdr) + len(transportHdr)),
|
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
|
||||||
GSOSize: uint16(len(pays[0])),
|
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
|
||||||
CsumStart: uint16(len(hdr)),
|
|
||||||
CsumOffset: csumOff,
|
|
||||||
}
|
}
|
||||||
if len(pays) > 1 {
|
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
|
||||||
ipVer := hdr[0] >> 4
|
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
|
||||||
switch {
|
|
||||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
|
||||||
case ipVer == 6:
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
|
||||||
case ipVer == 4:
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
|
||||||
default:
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
|
||||||
vhdr.GSOSize = 0
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
|
||||||
vhdr.GSOSize = 0
|
|
||||||
}
|
|
||||||
vhdr.Encode(r.gsoHdrBuf[:])
|
|
||||||
|
|
||||||
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
|
|
||||||
// wired to gsoHdrBuf at construction and never changes.
|
|
||||||
need := 3 + len(pays)
|
need := 3 + len(pays)
|
||||||
if need > cap(r.gsoIovs) {
|
if need > cap(r.gsoIovs) {
|
||||||
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
|
|
||||||
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
|
|
||||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||||
}
|
}
|
||||||
r.gsoIovs = r.gsoIovs[:need]
|
r.gsoIovs = r.gsoIovs[:need]
|
||||||
@@ -379,22 +352,52 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
|
|||||||
r.gsoIovs[1].SetLen(len(hdr))
|
r.gsoIovs[1].SetLen(len(hdr))
|
||||||
r.gsoIovs[2].Base = &transportHdr[0]
|
r.gsoIovs[2].Base = &transportHdr[0]
|
||||||
r.gsoIovs[2].SetLen(len(transportHdr))
|
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||||
// Defense in depth: an empty payload fragment can't be a valid GSO
|
|
||||||
// segment and &p[0] would panic on it. Callers route zero-length
|
segSize := len(pays[0])
|
||||||
// datagrams through the plain path (see UDPCoalescer.commitParsed), so
|
total := len(hdr) + len(transportHdr)
|
||||||
// this should never fire, but skip empties rather than index into one.
|
for i, p := range pays {
|
||||||
// `n` tracks where the next payload iovec lands, since skips make it
|
|
||||||
// drift from 3+i.
|
|
||||||
n := 3
|
|
||||||
for _, p := range pays {
|
|
||||||
if len(p) == 0 {
|
if len(p) == 0 {
|
||||||
continue
|
// 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)
|
||||||
}
|
}
|
||||||
r.gsoIovs[n].Base = &p[0]
|
total += len(p)
|
||||||
r.gsoIovs[n].SetLen(len(p))
|
r.gsoIovs[3+i].Base = &p[0]
|
||||||
n++
|
r.gsoIovs[3+i].SetLen(len(p))
|
||||||
}
|
}
|
||||||
r.gsoIovs = r.gsoIovs[:n]
|
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
|
||||||
|
if total > maxSuperpacketLen {
|
||||||
|
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
|
||||||
|
// Multiple segments carry the real GSO type and segSize, which the loop
|
||||||
|
// above verified is the size of every fragment except possibly the last.
|
||||||
|
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||||
|
if len(pays) > 1 {
|
||||||
|
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
|
||||||
|
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
|
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
|
||||||
|
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
|
||||||
|
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var gsoSize uint16
|
||||||
|
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
|
gsoSize = uint16(segSize)
|
||||||
|
}
|
||||||
|
virtio.EncodeHeader(
|
||||||
|
r.gsoHdrBuf[:],
|
||||||
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
|
gsoType, /*gsoType*/
|
||||||
|
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
|
||||||
|
gsoSize, /*gsoSize*/
|
||||||
|
uint16(len(hdr)), /*csumStart*/
|
||||||
|
csumOff, /*csumOffset*/
|
||||||
|
)
|
||||||
|
|
||||||
_, err := r.rawWrite(r.gsoIovs)
|
_, err := r.rawWrite(r.gsoIovs)
|
||||||
return err
|
return err
|
||||||
@@ -405,10 +408,10 @@ func (r *Offload) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
//shutdownFd is owned by the container, so we should not close it
|
// 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
|
// 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.
|
||||||
// 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
|
||||||
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
|
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||||
// ppoll first). closed.Swap already guarantees we only close once.
|
// closed.Swap already guarantees we only close once.
|
||||||
return unix.Close(r.fd)
|
return unix.Close(r.fd)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,11 +11,6 @@ import (
|
|||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
|
|
||||||
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
|
|
||||||
// prefix plus the virtio header.
|
|
||||||
const tunReadBufSize = 65535
|
|
||||||
|
|
||||||
type Poll struct {
|
type Poll struct {
|
||||||
fd int
|
fd int
|
||||||
shutdownFd int
|
shutdownFd int
|
||||||
@@ -25,10 +20,10 @@ type Poll struct {
|
|||||||
batchRet [1]Packet
|
batchRet [1]Packet
|
||||||
}
|
}
|
||||||
|
|
||||||
// newPoll wraps an existing tun fd. On failure it does NOT close fd: the
|
// newPoll wraps an existing tun fd.
|
||||||
// caller owns fd and is the sole closer (see pollQueueSet.Add callers in
|
// On failure it does NOT close fd: the caller owns fd and is the sole closer
|
||||||
// overlay/tun_linux.go, which unix.Close on Add error). This matches the
|
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
|
||||||
// newOffload convention and keeps closes at exactly one on every path.
|
// This matches the newOffload convention and keeps closes at exactly one on every path.
|
||||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||||
@@ -37,7 +32,7 @@ func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
|||||||
out := &Poll{
|
out := &Poll{
|
||||||
fd: fd,
|
fd: fd,
|
||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
readBuf: make([]byte, tunReadBufSize),
|
readBuf: make([]byte, 65535), // largest possible size Linux permits
|
||||||
}
|
}
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
@@ -52,6 +47,11 @@ func (t *Poll) blockOnWrite() error {
|
|||||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
|
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) {
|
func (t *Poll) Read() ([]Packet, error) {
|
||||||
n, err := t.readOne(t.readBuf)
|
n, err := t.readOne(t.readBuf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -108,10 +108,11 @@ func (t *Poll) Close() error {
|
|||||||
if t.closed.Swap(true) {
|
if t.closed.Swap(true) {
|
||||||
return nil
|
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
|
// shutdownFd is owned by the container, so we should not close it
|
||||||
// loading it in readOne, and mutating the field would race that load.
|
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
|
||||||
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
|
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||||
// ppoll first). closed.Swap already guarantees we only close once.
|
// 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)
|
return unix.Close(t.fd)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package tio
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"log/slog"
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -210,7 +211,7 @@ func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
|||||||
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
||||||
// for the GSO/offload queueset.
|
// for the GSO/offload queueset.
|
||||||
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||||
qs, err := NewOffloadQueueSet(false)
|
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
c, ok := qs.(*offloadQueueSet)
|
c, ok := qs.(*offloadQueueSet)
|
||||||
require.True(t, ok)
|
require.True(t, ok)
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
//go:build linux && !android
|
||||||
// +build linux,!android,!e2e_testing
|
// +build linux,!android
|
||||||
|
|
||||||
package tio
|
package tio
|
||||||
|
|
||||||
@@ -11,11 +11,11 @@ import (
|
|||||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||||
)
|
)
|
||||||
|
|
||||||
// protoFromGSOType maps a virtio_net_hdr GSOType to the GSOProto value the
|
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
|
||||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
||||||
// value — the caller should only invoke this on a confirmed superpacket.
|
// value. The caller should only invoke this on a confirmed superpacket.
|
||||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||||
switch t {
|
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||||
return GSOProtoTCP, nil
|
return GSOProtoTCP, nil
|
||||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||||
@@ -25,17 +25,26 @@ func protoFromGSOType(t uint8) (GSOProto, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts
|
// gsoTypeFromProto is the reverse of protoFromGSOType
|
||||||
// fn is called once with pkt.Bytes (no segmentation, no copy). For GSO/USO
|
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
|
||||||
// superpackets fn is called once per segment with a slice of pkt.Bytes
|
switch {
|
||||||
// holding that segment's plaintext (a freshly-patched L3+L4 header sliced
|
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||||
// in front of the original payload chunk). The slide is destructive: pkt is
|
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||||
// consumed by this call and its bytes are in an undefined state when
|
case ipVer == 6:
|
||||||
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier
|
return unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||||
// seg slice past fn's return for that segment. The scratch parameter is
|
case ipVer == 4:
|
||||||
// unused on the destructive path and kept only for cross-platform
|
return unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||||
// signature compatibility. Aborts and returns the first error from fn or
|
default:
|
||||||
// from per-segment construction.
|
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 {
|
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||||
if !pkt.GSO.IsSuperpacket() {
|
if !pkt.GSO.IsSuperpacket() {
|
||||||
return fn(pkt.Bytes)
|
return fn(pkt.Bytes)
|
||||||
|
|||||||
@@ -19,6 +19,36 @@ import (
|
|||||||
// worst-case 64 KiB superpacket plus replicated per-segment headers).
|
// worst-case 64 KiB superpacket plus replicated per-segment headers).
|
||||||
const testSegScratchSize = 192 * 1024
|
const testSegScratchSize = 192 * 1024
|
||||||
|
|
||||||
|
// TestProtoFromGSOTypeMasksECN guards the CWR-superpacket drop bug: the
|
||||||
|
// kernel qualifies a TSO superpacket whose TCP header carries CWR with
|
||||||
|
// VIRTIO_NET_HDR_GSO_ECN (we negotiate TUN_F_TSO_ECN, so it WILL send
|
||||||
|
// them once ECN feedback flows), and the decoder must mask that bit
|
||||||
|
// rather than reject the packet as an unknown type.
|
||||||
|
func TestProtoFromGSOTypeMasksECN(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
typ uint8
|
||||||
|
want GSOProto
|
||||||
|
}{
|
||||||
|
{unix.VIRTIO_NET_HDR_GSO_TCPV4, GSOProtoTCP},
|
||||||
|
{unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
|
||||||
|
{unix.VIRTIO_NET_HDR_GSO_TCPV6, GSOProtoTCP},
|
||||||
|
{unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
|
||||||
|
{unix.VIRTIO_NET_HDR_GSO_UDP_L4, GSOProtoUDP},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got, err := protoFromGSOType(c.typ)
|
||||||
|
if err != nil || got != c.want {
|
||||||
|
t.Errorf("protoFromGSOType(%#x) = (%v, %v), want (%v, nil)", c.typ, got, err, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_NONE); err == nil {
|
||||||
|
t.Error("GSO_NONE must still be rejected")
|
||||||
|
}
|
||||||
|
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_ECN); err == nil {
|
||||||
|
t.Error("a bare ECN bit with no base type must still be rejected")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
||||||
// with a folded pseudo-header sum, equals all-ones (valid).
|
// with a folded pseudo-header sum, equals all-ones (valid).
|
||||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||||
@@ -33,7 +63,7 @@ func verifyChecksum(b []byte, pseudo uint16) bool {
|
|||||||
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
|
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
|
||||||
// invoked here.
|
// invoked here.
|
||||||
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
|
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
cp := append([]byte(nil), pkt...)
|
cp := append([]byte(nil), pkt...)
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||||
if err := virtio.FinishChecksum(cp, hdr); err != nil {
|
if err := virtio.FinishChecksum(cp, hdr); err != nil {
|
||||||
@@ -43,7 +73,7 @@ func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) e
|
|||||||
*out = append(*out, cp)
|
*out = append(*out, cp)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
proto, err := protoFromGSOType(hdr.GSOType)
|
proto, err := protoFromGSOType(hdr.GSOType())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -110,15 +140,14 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
|
|||||||
for i := 0; i < payLen; i++ {
|
for i := 0; i < payLen; i++ {
|
||||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||||
}
|
}
|
||||||
|
return pkt, virtio.NewHeader(
|
||||||
return pkt, virtio.Hdr{
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
uint16(mss), /*gsoSize*/
|
||||||
GSOSize: uint16(mss),
|
uint16(ipLen), /*csumStart*/
|
||||||
CsumStart: uint16(ipLen),
|
16, /*csumOffset*/
|
||||||
CsumOffset: 16,
|
)
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSegmentTCPv4(t *testing.T) {
|
func TestSegmentTCPv4(t *testing.T) {
|
||||||
@@ -232,14 +261,14 @@ func TestSegmentTCPv6(t *testing.T) {
|
|||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
}
|
}
|
||||||
|
|
||||||
hdr := virtio.Hdr{
|
hdr := virtio.NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||||
GSOSize: uint16(mss),
|
uint16(mss), /*gsoSize*/
|
||||||
CsumStart: uint16(ipLen),
|
uint16(ipLen), /*csumStart*/
|
||||||
CsumOffset: 16,
|
16, /*csumOffset*/
|
||||||
}
|
)
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, testSegScratchSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
@@ -281,7 +310,7 @@ func TestSegmentTCPv6(t *testing.T) {
|
|||||||
|
|
||||||
func TestSegmentGSONonePassesThrough(t *testing.T) {
|
func TestSegmentGSONonePassesThrough(t *testing.T) {
|
||||||
pkt, hdr := buildTSOv4(t, 100, 100)
|
pkt, hdr := buildTSOv4(t, 100, 100)
|
||||||
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
hdr.SetGSOType(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||||
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, testSegScratchSize)
|
||||||
@@ -300,7 +329,7 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
|
|||||||
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
|
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
|
||||||
// still rejected; only modern GSO_UDP_L4 (USO) is supported.
|
// still rejected; only modern GSO_UDP_L4 (USO) is supported.
|
||||||
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
|
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
|
||||||
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
hdr := virtio.NewHeader(0, unix.VIRTIO_NET_HDR_GSO_UDP, 0, 0, 0, 0)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
|
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
|
||||||
t.Fatalf("expected rejection for legacy UDP GSO")
|
t.Fatalf("expected rejection for legacy UDP GSO")
|
||||||
@@ -324,22 +353,26 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
|
|||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
|
||||||
// UDP header (length + checksum filled in per segment by segmentUDPYield)
|
// UDP header. The kernel hands us a USO superpacket whose length field
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
// covers the WHOLE superpacket; the segmenter overwrites it per segment.
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
// Populating it here matters: leaving it zero makes the base-checksum path
|
||||||
|
// that must exclude it untestable, since excluding zero is a no-op.
|
||||||
|
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||||
|
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||||
|
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
for i := 0; i < payLen; i++ {
|
||||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||||
}
|
}
|
||||||
|
|
||||||
return pkt, virtio.Hdr{
|
return pkt, virtio.NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||||
HdrLen: uint16(ipLen + udpLen),
|
uint16(ipLen+udpLen), /*hdrLen*/
|
||||||
GSOSize: uint16(gsoSize),
|
uint16(gsoSize), /*gsoSize*/
|
||||||
CsumStart: uint16(ipLen),
|
uint16(ipLen), /*csumStart*/
|
||||||
CsumOffset: 6,
|
6, /*csumOffset*/
|
||||||
}
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSegmentUDPv4(t *testing.T) {
|
func TestSegmentUDPv4(t *testing.T) {
|
||||||
@@ -364,11 +397,12 @@ func TestSegmentUDPv4(t *testing.T) {
|
|||||||
if totalLen != uint16(28+gso) {
|
if totalLen != uint16(28+gso) {
|
||||||
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
|
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
|
||||||
}
|
}
|
||||||
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every
|
// Software UDP GSO bumps the IPv4 ID per segment exactly like TSO
|
||||||
// segment carries the same ID as the seed.
|
// (inet_gso_segment's fixed-ID case is TCP-only); wireguard-go's
|
||||||
|
// gsoSplit increments unconditionally too.
|
||||||
id := binary.BigEndian.Uint16(seg[4:6])
|
id := binary.BigEndian.Uint16(seg[4:6])
|
||||||
if id != 0x4242 {
|
if id != 0x4242+uint16(i) {
|
||||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242)
|
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||||
}
|
}
|
||||||
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
||||||
if udpLen != uint16(8+gso) {
|
if udpLen != uint16(8+gso) {
|
||||||
@@ -436,19 +470,21 @@ func TestSegmentUDPv6(t *testing.T) {
|
|||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
||||||
|
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
|
||||||
|
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
for i := 0; i < payLen; i++ {
|
||||||
pkt[ipLen+udpLen+i] = byte(i)
|
pkt[ipLen+udpLen+i] = byte(i)
|
||||||
}
|
}
|
||||||
|
|
||||||
hdr := virtio.Hdr{
|
hdr := virtio.NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||||
HdrLen: uint16(ipLen + udpLen),
|
uint16(ipLen+udpLen), /*hdrLen*/
|
||||||
GSOSize: uint16(gso),
|
uint16(gso), /*gsoSize*/
|
||||||
CsumStart: uint16(ipLen),
|
uint16(ipLen), /*csumStart*/
|
||||||
CsumOffset: 6,
|
6, /*csumOffset*/
|
||||||
}
|
)
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, testSegScratchSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
@@ -580,14 +616,14 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
|
|||||||
for i := 0; i < sz.payLen; i++ {
|
for i := 0; i < sz.payLen; i++ {
|
||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
}
|
}
|
||||||
hdr := virtio.Hdr{
|
hdr := virtio.NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||||
GSOSize: uint16(sz.mss),
|
uint16(sz.mss), /*gsoSize*/
|
||||||
CsumStart: uint16(ipLen),
|
uint16(ipLen), /*csumStart*/
|
||||||
CsumOffset: 16,
|
16, /*csumOffset*/
|
||||||
}
|
)
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, testSegScratchSize)
|
||||||
out := make([][]byte, 0, 64)
|
out := make([][]byte, 0, 64)
|
||||||
@@ -640,35 +676,72 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the
|
// TestSegmentSuperpacketNoAlloc pins the segmenters' zero-allocation
|
||||||
// zero-length UDP DoS: a payload fragment of length zero would make &p[0]
|
// contract. Both SegmentTCP and SegmentUDP derive their per-superpacket
|
||||||
// panic (index-out-of-range) when building the iovec array. WriteGSO must
|
// constants into fixed-size arrays (tmp/ipTmp/savedHdr) that must stay on
|
||||||
// skip empties instead. We write to /dev/null so the writev always succeeds
|
// the stack, and both take a yield closure that must not escape. Any of
|
||||||
// synchronously; the point is simply that neither call panics.
|
// those escaping turns one allocation into one-per-superpacket on the
|
||||||
func TestWriteGSOSkipsEmptyPayloads(t *testing.T) {
|
// hottest path in the reader, which BenchmarkSegmentSuperpacketAllocsTSO
|
||||||
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
// reports but nothing fails on. This does.
|
||||||
if err != nil {
|
//
|
||||||
t.Fatalf("open /dev/null: %v", err)
|
// The yield closure here only touches captured scalars: appending segments
|
||||||
|
// to a slice would allocate in the test itself and mask the measurement.
|
||||||
|
func TestSegmentSuperpacketNoAlloc(t *testing.T) {
|
||||||
|
const mss = 1400
|
||||||
|
const numSeg = 8
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
build func() ([]byte, virtio.Hdr)
|
||||||
|
}{
|
||||||
|
{"tso-v4", func() ([]byte, virtio.Hdr) { return buildTSOv4(t, mss*numSeg, mss) }},
|
||||||
|
{"uso-v4", func() ([]byte, virtio.Hdr) { return buildUSOv4(t, mss*numSeg, mss) }},
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
|
||||||
|
|
||||||
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
for _, tc := range cases {
|
||||||
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
o.gsoIovs[0].SetLen(virtio.Size)
|
master, hdr := tc.build()
|
||||||
|
proto, err := protoFromGSOType(hdr.GSOType())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("protoFromGSOType: %v", err)
|
||||||
|
}
|
||||||
|
work := make([]byte, len(master))
|
||||||
|
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||||
|
Size: hdr.GSOSize,
|
||||||
|
HdrLen: hdr.HdrLen,
|
||||||
|
CsumStart: hdr.CsumStart,
|
||||||
|
Proto: proto,
|
||||||
|
}}
|
||||||
|
|
||||||
ipHdr := make([]byte, 20)
|
// Segmentation consumes its input destructively, so restore from
|
||||||
ipHdr[0] = 0x45 // IPv4, IHL 5
|
// the master copy each run; copy(2) into an existing slice does
|
||||||
udpHdr := make([]byte, 8)
|
// not allocate. seen/bytes keep the closure from being optimized
|
||||||
|
// away and double as a sanity check that work actually happened.
|
||||||
|
var seen, bytes int
|
||||||
|
run := func() {
|
||||||
|
copy(work, master)
|
||||||
|
seen, bytes = 0, 0
|
||||||
|
if err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||||
|
seen++
|
||||||
|
bytes += len(seg)
|
||||||
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("SegmentSuperpacket: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Sole payload empty: exercises the all-empty skip (n stays at 3).
|
run() // warm up: absorb any one-time allocation elsewhere
|
||||||
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil {
|
if seen != numSeg {
|
||||||
t.Fatalf("WriteGSO with a single empty payload: %v", err)
|
t.Fatalf("yielded %d segments, want %d", seen, numSeg)
|
||||||
}
|
}
|
||||||
// Empty mixed with a real fragment: exercises the index-drift skip so a
|
|
||||||
// later non-empty payload still lands in the right iovec slot.
|
if allocs := testing.AllocsPerRun(200, run); allocs != 0 {
|
||||||
real := make([]byte, 1200)
|
t.Fatalf("SegmentSuperpacket allocated %.1f times per call, want 0", allocs)
|
||||||
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil {
|
}
|
||||||
t.Fatalf("WriteGSO with a trailing empty payload: %v", err)
|
if seen != numSeg || bytes == 0 {
|
||||||
|
t.Fatalf("post-measure sanity: seen=%d bytes=%d", seen, bytes)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -721,9 +794,11 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
|||||||
const ipv6HdrLen = 40
|
const ipv6HdrLen = 40
|
||||||
const tcpHdrLen = 20
|
const tcpHdrLen = 20
|
||||||
const headerLen = ipv6HdrLen + tcpHdrLen
|
const headerLen = ipv6HdrLen + tcpHdrLen
|
||||||
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec
|
// Maximum TUN read body at the drain threshold. readv bounds the body
|
||||||
// is what bounds the kernel's superpacket length.
|
// iovec by the space actually left in rxBuf, and the drain gate keeps that
|
||||||
pktLen := tunReadBufSize
|
// at >= tunRxBufSize, so that is the largest superpacket the kernel can
|
||||||
|
// hand back on the last permitted drain read.
|
||||||
|
pktLen := tunRxBufSize
|
||||||
payLen := pktLen - headerLen
|
payLen := pktLen - headerLen
|
||||||
const targetSegs = 64
|
const targetSegs = 64
|
||||||
gsoSize := (payLen + targetSegs - 1) / targetSegs
|
gsoSize := (payLen + targetSegs - 1) / targetSegs
|
||||||
@@ -745,14 +820,14 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
|||||||
copy(o.rxBuf[o.rxOff:], pkt)
|
copy(o.rxBuf[o.rxOff:], pkt)
|
||||||
|
|
||||||
// Encode the matching virtio_net_hdr.
|
// Encode the matching virtio_net_hdr.
|
||||||
hdr := virtio.Hdr{
|
hdr := virtio.NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
|
||||||
HdrLen: uint16(headerLen),
|
uint16(headerLen), /*hdrLen*/
|
||||||
GSOSize: uint16(gsoSize),
|
uint16(gsoSize), /*gsoSize*/
|
||||||
CsumStart: uint16(ipv6HdrLen),
|
uint16(ipv6HdrLen), /*csumStart*/
|
||||||
CsumOffset: 16,
|
16, /*csumOffset*/
|
||||||
}
|
)
|
||||||
hdr.Encode(o.readVnetScratch[:])
|
hdr.Encode(o.readVnetScratch[:])
|
||||||
|
|
||||||
startRxOff := o.rxOff
|
startRxOff := o.rxOff
|
||||||
@@ -824,3 +899,224 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
|||||||
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
|
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestOffloadWriteZeroLength: a zero-length Write must be a no-op, not a
|
||||||
|
// panic. The guard used to live below the &buf[0] that tripped on it.
|
||||||
|
func TestOffloadWriteZeroLength(t *testing.T) {
|
||||||
|
tf := &Offload{fd: -1} // any write reaching the fd would fail loudly
|
||||||
|
for _, buf := range [][]byte{nil, {}} {
|
||||||
|
n, err := tf.Write(buf)
|
||||||
|
if n != 0 || err != nil {
|
||||||
|
t.Errorf("Write(len=0) = (%d, %v), want (0, nil)", n, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteGSOSuperpacketGeometry decodes the vnet header the kernel would see for a multi-segment write:
|
||||||
|
// the GSO type must match the proto and IP version
|
||||||
|
// gso_size must be the per-segment size (the kernel rejects a superpacket with gso_size == 0),
|
||||||
|
// and the csum fields must point at the transport header's checksum slot.
|
||||||
|
// Write through a pipe so the bytes can be read back and decoded.
|
||||||
|
func TestWriteGSOSuperpacketGeometry(t *testing.T) {
|
||||||
|
var pfds [2]int
|
||||||
|
if err := unix.Pipe(pfds[:]); err != nil {
|
||||||
|
t.Fatalf("pipe: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { unix.Close(pfds[0]); unix.Close(pfds[1]) })
|
||||||
|
|
||||||
|
o := &Offload{fd: pfds[1], gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||||
|
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||||
|
o.gsoIovs[0].SetLen(virtio.Size)
|
||||||
|
|
||||||
|
ipHdr := make([]byte, 20)
|
||||||
|
ipHdr[0] = 0x45
|
||||||
|
udpHdr := make([]byte, 8)
|
||||||
|
seg := make([]byte, 1200)
|
||||||
|
|
||||||
|
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP); err != nil {
|
||||||
|
t.Fatalf("WriteGSO: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := make([]byte, virtio.Size+len(ipHdr)+len(udpHdr)+2*len(seg)+64)
|
||||||
|
n, err := unix.Read(pfds[0], buf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pipe: %v", err)
|
||||||
|
}
|
||||||
|
var vhdr virtio.Hdr
|
||||||
|
vhdr.Decode(buf[:virtio.Size])
|
||||||
|
if vhdr.GSOType() != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||||
|
t.Errorf("gsoType=%d want UDP_L4", vhdr.GSOType())
|
||||||
|
}
|
||||||
|
if vhdr.GSOSize != 1200 {
|
||||||
|
t.Errorf("GSOSize=%d want 1200 (per-segment size from pays[0])", vhdr.GSOSize)
|
||||||
|
}
|
||||||
|
if vhdr.HdrLen != uint16(len(ipHdr)+len(udpHdr)) {
|
||||||
|
t.Errorf("HdrLen=%d want %d", vhdr.HdrLen, len(ipHdr)+len(udpHdr))
|
||||||
|
}
|
||||||
|
if vhdr.CsumStart != uint16(len(ipHdr)) || vhdr.CsumOffset != 6 {
|
||||||
|
t.Errorf("csum start/offset = %d/%d want %d/6", vhdr.CsumStart, vhdr.CsumOffset, len(ipHdr))
|
||||||
|
}
|
||||||
|
if want := virtio.Size + len(ipHdr) + len(udpHdr) + 2*len(seg); n != want {
|
||||||
|
t.Errorf("wrote %d bytes want %d", n, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteGSORejectsBadGeometry pins the length-check contracts
|
||||||
|
func TestWriteGSORejectsBadGeometry(t *testing.T) {
|
||||||
|
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open /dev/null: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||||
|
|
||||||
|
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||||
|
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||||
|
o.gsoIovs[0].SetLen(virtio.Size)
|
||||||
|
|
||||||
|
ipHdr := make([]byte, 20)
|
||||||
|
ipHdr[0] = 0x45
|
||||||
|
udpHdr := make([]byte, 8)
|
||||||
|
tcpHdr := make([]byte, 20)
|
||||||
|
seg := make([]byte, 1200)
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
hdr, thdr []byte
|
||||||
|
pays [][]byte
|
||||||
|
proto GSOProto
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"empty-ip-hdr-with-payload", nil, udpHdr, [][]byte{seg}, GSOProtoUDP, true},
|
||||||
|
{"udp-transport-too-short-for-csum", ipHdr, udpHdr[:6], [][]byte{seg}, GSOProtoUDP, true},
|
||||||
|
{"tcp-transport-too-short-for-csum", ipHdr, tcpHdr[:16], [][]byte{seg}, GSOProtoTCP, true},
|
||||||
|
{"superpacket-over-65535", ipHdr, tcpHdr, [][]byte{make([]byte, 40000), make([]byte, 40000)}, GSOProtoTCP, true},
|
||||||
|
{"sole-payload-empty", ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP, true},
|
||||||
|
{"leading-empty-fragment", ipHdr, udpHdr, [][]byte{{}, seg, seg}, GSOProtoUDP, true},
|
||||||
|
{"trailing-empty-fragment", ipHdr, tcpHdr, [][]byte{seg, {}}, GSOProtoTCP, true},
|
||||||
|
{"oversize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 1201), seg}, GSOProtoUDP, true},
|
||||||
|
{"undersize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 100), seg}, GSOProtoUDP, true},
|
||||||
|
{"oversize-last-fragment", ipHdr, tcpHdr, [][]byte{seg, make([]byte, 1201)}, GSOProtoTCP, true},
|
||||||
|
{"short-last-fragment-ok", ipHdr, udpHdr, [][]byte{seg, seg, make([]byte, 100)}, GSOProtoUDP, false},
|
||||||
|
{"multi-segment-bad-ip-version", []byte{0x05}, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, true},
|
||||||
|
{"single-segment-bad-ip-version-ok", []byte{0x05}, udpHdr, [][]byte{seg}, GSOProtoUDP, false},
|
||||||
|
{"no-pays-noop", ipHdr, udpHdr, nil, GSOProtoUDP, false},
|
||||||
|
{"valid-udp", ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, false},
|
||||||
|
{"valid-tcp", ipHdr, tcpHdr, [][]byte{seg, seg}, GSOProtoTCP, false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
err := o.WriteGSO(tc.hdr, tc.thdr, tc.pays, tc.proto)
|
||||||
|
if tc.wantErr && err == nil {
|
||||||
|
t.Errorf("WriteGSO = nil, want error")
|
||||||
|
}
|
||||||
|
if !tc.wantErr && err != nil {
|
||||||
|
t.Errorf("WriteGSO = %v, want nil", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkSegmentUDPv4 is the USO counterpart to BenchmarkSegmentTCPv4. The
|
||||||
|
// yield is a no-op so the measurement is segmentation plus checksum work only.
|
||||||
|
func BenchmarkSegmentUDPv4(b *testing.B) {
|
||||||
|
sizes := []struct {
|
||||||
|
name string
|
||||||
|
payLen int
|
||||||
|
gsoSize int
|
||||||
|
}{
|
||||||
|
{"64KiB_GSO1400", 64000, 1400},
|
||||||
|
{"16KiB_GSO1400", 16384, 1400},
|
||||||
|
{"4KiB_GSO1400", 4096, 1400},
|
||||||
|
}
|
||||||
|
for _, sz := range sizes {
|
||||||
|
b.Run(sz.name, func(b *testing.B) {
|
||||||
|
const ipLen = 20
|
||||||
|
const udpLen = 8
|
||||||
|
pkt := make([]byte, ipLen+udpLen+sz.payLen)
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+sz.payLen))
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||||
|
pkt[8] = 64
|
||||||
|
pkt[9] = unix.IPPROTO_UDP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
binary.BigEndian.PutUint16(pkt[20:22], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[22:24], 53)
|
||||||
|
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+sz.payLen))
|
||||||
|
for i := 0; i < sz.payLen; i++ {
|
||||||
|
pkt[ipLen+udpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
|
||||||
|
master := append([]byte(nil), pkt...)
|
||||||
|
work := make([]byte, len(pkt))
|
||||||
|
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||||
|
Size: uint16(sz.gsoSize),
|
||||||
|
HdrLen: ipLen + udpLen,
|
||||||
|
CsumStart: ipLen,
|
||||||
|
Proto: GSOProtoUDP,
|
||||||
|
}}
|
||||||
|
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
copy(work, master)
|
||||||
|
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkSegmentUDPv6 mirrors BenchmarkSegmentUDPv4 for IPv6, where the
|
||||||
|
// pseudo-header address sum is 32 bytes rather than 8.
|
||||||
|
func BenchmarkSegmentUDPv6(b *testing.B) {
|
||||||
|
sizes := []struct {
|
||||||
|
name string
|
||||||
|
payLen int
|
||||||
|
gsoSize int
|
||||||
|
}{
|
||||||
|
{"64KiB_GSO1400", 64000, 1400},
|
||||||
|
{"16KiB_GSO1400", 16384, 1400},
|
||||||
|
{"4KiB_GSO1400", 4096, 1400},
|
||||||
|
}
|
||||||
|
for _, sz := range sizes {
|
||||||
|
b.Run(sz.name, func(b *testing.B) {
|
||||||
|
const ipLen = 40
|
||||||
|
const udpLen = 8
|
||||||
|
pkt := make([]byte, ipLen+udpLen+sz.payLen)
|
||||||
|
pkt[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+sz.payLen))
|
||||||
|
pkt[6] = unix.IPPROTO_UDP
|
||||||
|
pkt[7] = 64
|
||||||
|
pkt[8], pkt[9], pkt[23] = 0xfe, 0x80, 1
|
||||||
|
pkt[24], pkt[25], pkt[39] = 0xfe, 0x80, 2
|
||||||
|
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
||||||
|
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+sz.payLen))
|
||||||
|
for i := 0; i < sz.payLen; i++ {
|
||||||
|
pkt[ipLen+udpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
|
||||||
|
master := append([]byte(nil), pkt...)
|
||||||
|
work := make([]byte, len(pkt))
|
||||||
|
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||||
|
Size: uint16(sz.gsoSize),
|
||||||
|
HdrLen: ipLen + udpLen,
|
||||||
|
CsumStart: ipLen,
|
||||||
|
Proto: GSOProtoUDP,
|
||||||
|
}}
|
||||||
|
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
copy(work, master)
|
||||||
|
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,7 +3,11 @@
|
|||||||
|
|
||||||
package virtio
|
package virtio
|
||||||
|
|
||||||
import "encoding/binary"
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||||
@@ -13,31 +17,59 @@ const Size = 10
|
|||||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||||
type Hdr struct {
|
type Hdr struct {
|
||||||
Flags uint8
|
Flags uint8
|
||||||
GSOType uint8
|
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
|
||||||
HdrLen uint16
|
HdrLen uint16
|
||||||
GSOSize uint16
|
GSOSize uint16
|
||||||
CsumStart uint16
|
CsumStart uint16
|
||||||
CsumOffset 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
|
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||||
func (h *Hdr) Decode(b []byte) {
|
func (h *Hdr) Decode(b []byte) {
|
||||||
h.Flags = b[0]
|
h.Flags = b[0]
|
||||||
h.GSOType = b[1]
|
h.gsoType = b[1]
|
||||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
||||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
||||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
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
|
// 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.
|
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
||||||
func (h *Hdr) Encode(b []byte) {
|
func (h *Hdr) Encode(b []byte) {
|
||||||
b[0] = h.Flags
|
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
|
||||||
b[1] = h.GSOType
|
}
|
||||||
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
|
||||||
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize)
|
// GSOType returns gsoType with the ECN-flag masked out
|
||||||
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart)
|
func (h *Hdr) GSOType() uint8 {
|
||||||
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset)
|
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
|
||||||
}
|
}
|
||||||
|
|||||||
+138
-131
@@ -27,11 +27,8 @@ const (
|
|||||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||||
)
|
)
|
||||||
|
|
||||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each
|
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
|
||||||
// segment. The largest header the segmenter supports is IPv4 (max IHL 60)
|
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
|
||||||
// plus TCP (max data-offset 60) = 120 bytes; the array is sized to that
|
|
||||||
// worst case so the snapshot lives on the stack with no per-call heap
|
|
||||||
// allocation.
|
|
||||||
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||||
|
|
||||||
// Byte offsets inside an IPv4 header.
|
// Byte offsets inside an IPv4 header.
|
||||||
@@ -65,63 +62,72 @@ const (
|
|||||||
udpChecksumOff = 6
|
udpChecksumOff = 6
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var errPacketTooShort = errors.New("packet too short")
|
||||||
|
|
||||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||||
|
|
||||||
// tcpCwrFlag is cleared on every segment except the first. Per RFC 3168
|
// tcpCwrFlag is cleared on every segment except the first.
|
||||||
// §6.1.2 the CWR bit signals a one-shot transition (the sender just halved
|
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
|
||||||
// its window) and must appear on the first segment of a TSO burst only.
|
// and must appear on the first segment of a TSO burst only.
|
||||||
const tcpCwrFlag = 0x80
|
const tcpCwrFlag = 0x80
|
||||||
|
|
||||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
||||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
||||||
// the GSO type must agree with the IP version nibble.
|
// the GSO type must agree with the IP version nibble.
|
||||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
func CheckValid(pkt []byte, hdr Hdr) error {
|
||||||
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
|
||||||
// carry coalescing info rather than checksum offsets. A TUN writing via
|
|
||||||
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
|
||||||
// miscompute the segment checksums — refuse the packet instead.
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||||
}
|
}
|
||||||
if len(pkt) < ipv4HeaderMinLen {
|
if len(pkt) < ipv4HeaderMinLen {
|
||||||
return fmt.Errorf("packet too short")
|
return errPacketTooShort
|
||||||
}
|
}
|
||||||
ipVersion := pkt[0] >> 4
|
ipVersion := pkt[0] >> 4
|
||||||
switch hdr.GSOType {
|
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:
|
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||||
if ipVersion != 4 {
|
if ipVersion != 4 {
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||||
}
|
}
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||||
if ipVersion != 6 {
|
if ipVersion != 6 {
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||||
}
|
}
|
||||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
// USO carries either v4 or v6; the leading nibble disambiguates.
|
||||||
if !(ipVersion == 4 || ipVersion == 6) {
|
if !(ipVersion == 4 || ipVersion == 6) {
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
if !(ipVersion == 6 || ipVersion == 4) {
|
if !(ipVersion == 6 || ipVersion == 4) {
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header
|
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
|
||||||
// length read out of pkt. The kernel's hdr.HdrLen on the FORWARD path can
|
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
|
||||||
// be the length of the entire first packet, so we don't trust it.
|
|
||||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||||
// Thank you wireguard-go for documenting these edge-cases
|
// 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
|
// 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
|
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
|
||||||
// FORWARD path. Instead, parse the transport header length and add it onto
|
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
|
||||||
// csumStart, which is synonymous for IP header length.
|
|
||||||
|
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||||
hdr.HdrLen = hdr.CsumStart + 8
|
hdr.HdrLen = hdr.CsumStart + 8
|
||||||
} else {
|
} else {
|
||||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
||||||
@@ -129,8 +135,7 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
||||||
if tcpHLen < 20 || tcpHLen > 60 {
|
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
|
||||||
// A TCP header must be between 20 and 60 bytes in length.
|
|
||||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||||
}
|
}
|
||||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||||
@@ -150,19 +155,62 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
|
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
|
||||||
// slice into pkt itself. Per-segment plaintext is laid out by stamping a
|
// with a floor of one so a header-only superpacket still yields a single segment.
|
||||||
// copy of the original L3+L4 header into pkt at offset i*gsoSize, where it
|
func segCount(payLen, gsoSize int) int {
|
||||||
// sits immediately before that segment's payload chunk in the original
|
n := (payLen + gsoSize - 1) / gsoSize
|
||||||
// buffer. The stamp is destructive but harmless: iter i's header write lands
|
if n == 0 {
|
||||||
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already
|
return 1
|
||||||
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers
|
}
|
||||||
// live payload — this holds even when gsoSize < hdrLen. The header bytes are
|
return n
|
||||||
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from
|
}
|
||||||
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
|
|
||||||
// overwrite the leading header in place and every stamp after the first would
|
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
|
||||||
// copy corrupted bytes. pkt is consumed by this call and must not be inspected
|
// for every segment: the source and destination addresses plus the protocol
|
||||||
// by the caller after the final yield.
|
// 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 {
|
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||||
if gsoSizeU == 0 {
|
if gsoSizeU == 0 {
|
||||||
return fmt.Errorf("gso_size is zero")
|
return fmt.Errorf("gso_size is zero")
|
||||||
@@ -181,49 +229,28 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||||
payLen := len(pkt) - headerLen
|
payLen := len(pkt) - headerLen
|
||||||
gsoSize := int(gsoSizeU)
|
gsoSize := int(gsoSizeU)
|
||||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
numSeg := segCount(payLen, gsoSize)
|
||||||
if numSeg == 0 {
|
|
||||||
numSeg = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||||
|
|
||||||
var tmp [tcpHeaderMaxLen]byte
|
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
|
||||||
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
|
||||||
tmp[tcpFlagsOff] = 0
|
|
||||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
|
||||||
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
|
||||||
|
|
||||||
var baseProtoSum uint32
|
|
||||||
if isV4 {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
|
||||||
} else {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
|
||||||
}
|
|
||||||
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
|
||||||
|
|
||||||
var origIPID uint16
|
var origIPID uint16
|
||||||
var baseIPHdrSum uint32
|
var baseIPHdrSum uint32
|
||||||
if isV4 {
|
if isV4 {
|
||||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
var err error
|
||||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
// TSO bumps the ID per segment, so it stays out of the base sum.
|
||||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
var ipTmp [ipv4HeaderMaxLen]byte
|
|
||||||
copy(ipTmp[:ihl], pkt[:ihl])
|
|
||||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
|
||||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Snapshot the pristine L3+L4 header once. Every segment's header is
|
// Snapshot the pristine L3+L4 header once. '
|
||||||
// stamped from this copy, so overlapping stamps (gsoSize < headerLen)
|
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
|
||||||
// can never corrupt the source. The variable fields (seq/flags/cksum/
|
|
||||||
// totalLen/id) captured here are stale but are overwritten per segment.
|
|
||||||
var savedHdr [maxSegHdrLen]byte
|
var savedHdr [maxSegHdrLen]byte
|
||||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||||
|
|
||||||
@@ -237,12 +264,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
segLen := headerLen + segPayLen
|
segLen := headerLen + segPayLen
|
||||||
headerOff := i * gsoSize
|
headerOff := i * gsoSize
|
||||||
|
|
||||||
// Stamp the header into place immediately before this segment's
|
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
|
||||||
// payload, sourced from the pristine snapshot. Iter 0's header is
|
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
|
||||||
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1
|
|
||||||
// needs the stamp. The per-segment patches below overwrite the
|
|
||||||
// variable fields.
|
|
||||||
if i > 0 {
|
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])
|
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||||
}
|
}
|
||||||
seg := pkt[headerOff : headerOff+segLen]
|
seg := pkt[headerOff : headerOff+segLen]
|
||||||
@@ -271,10 +296,9 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
seg[csumStart+tcpFlagsOff] = segFlags
|
seg[csumStart+tcpFlagsOff] = segFlags
|
||||||
|
|
||||||
tcpLen := tcpHdrLen + segPayLen
|
tcpLen := tcpHdrLen + segPayLen
|
||||||
// Payload bytes still live at their original offset in pkt. The
|
// Payload bytes still live at their original offset in pkt.
|
||||||
// header slide above only writes into pkt[i*G : i*G+H], which is
|
// 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)
|
||||||
// the tail of seg_{i-1}'s payload (already consumed) and never
|
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
|
||||||
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
|
|
||||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||||
@@ -290,17 +314,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of
|
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
|
||||||
// the original L3+L4 header into pkt at offset i*gsoSize and yielding
|
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
|
||||||
// pkt[i*G:i*G+segLen] to the caller. Per-segment patches are total_len +
|
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
|
||||||
// IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. pkt is
|
// pkt is consumed destructively.
|
||||||
// consumed destructively; see SegmentTCP for the layout reasoning, including
|
|
||||||
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
|
|
||||||
// (correctness when gsoSize < hdrLen).
|
|
||||||
//
|
|
||||||
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
|
|
||||||
// bump it), which is why the IP-level per-segment work is limited to
|
|
||||||
// total_len + IPv4 header checksum (v4) or payload_len (v6).
|
|
||||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||||
if gsoSizeU == 0 {
|
if gsoSizeU == 0 {
|
||||||
return fmt.Errorf("gso_size is zero")
|
return fmt.Errorf("gso_size is zero")
|
||||||
@@ -321,41 +338,24 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
|
|
||||||
payLen := len(pkt) - headerLen
|
payLen := len(pkt) - headerLen
|
||||||
gsoSize := int(gsoSizeU)
|
gsoSize := int(gsoSizeU)
|
||||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
numSeg := segCount(payLen, gsoSize)
|
||||||
if numSeg == 0 {
|
|
||||||
numSeg = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
var udpTmp [udpHeaderLen]byte
|
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
|
||||||
copy(udpTmp[:], pkt[csumStart:headerLen])
|
|
||||||
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
|
|
||||||
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
|
|
||||||
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
|
|
||||||
|
|
||||||
var baseProtoSum uint32
|
|
||||||
if isV4 {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
|
||||||
} else {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
|
||||||
}
|
|
||||||
baseProtoSum += uint32(unix.IPPROTO_UDP)
|
|
||||||
|
|
||||||
|
var origIPID uint16
|
||||||
var baseIPHdrSum uint32
|
var baseIPHdrSum uint32
|
||||||
if isV4 {
|
if isV4 {
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
var err error
|
||||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
// 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
|
||||||
}
|
}
|
||||||
var ipTmp [ipv4HeaderMaxLen]byte
|
|
||||||
copy(ipTmp[:ihl], pkt[:ihl])
|
|
||||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
|
||||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Snapshot the pristine L3+L4 header once and stamp every segment from
|
// Snapshot the pristine L3+L4 header once and stamp every segment from it
|
||||||
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
|
|
||||||
// segments when gsoSize < headerLen.
|
|
||||||
var savedHdr [maxSegHdrLen]byte
|
var savedHdr [maxSegHdrLen]byte
|
||||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||||
|
|
||||||
@@ -378,8 +378,10 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
udpLen := udpHeaderLen + segPayLen
|
udpLen := udpHeaderLen + segPayLen
|
||||||
|
|
||||||
if isV4 {
|
if isV4 {
|
||||||
|
segID := origIPID + uint16(i)
|
||||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||||
ipSum := baseIPHdrSum + uint32(totalLen)
|
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||||
|
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||||
} else {
|
} else {
|
||||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||||
@@ -387,12 +389,13 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
|
|
||||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||||
|
|
||||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
// Sum the UDP header (length just written, checksum zeroed) together with
|
||||||
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
// this segment's payload in one pass, seeded with the pseudo-header sum.
|
||||||
wide += uint64(udpLen) + uint64(udpLen)
|
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
pseudo := baseProtoSum + uint32(udpLen)
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||||
csum := foldComplement(uint32(wide))
|
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||||
|
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
|
||||||
if csum == 0 {
|
if csum == 0 {
|
||||||
csum = 0xffff
|
csum = 0xffff
|
||||||
}
|
}
|
||||||
@@ -406,10 +409,9 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel
|
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
|
||||||
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
|
// CsumStart / CsumOffset point at the 16-bit checksum field.
|
||||||
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
|
||||||
// the pseudo-header partial sum by the kernel), and store the result.
|
|
||||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||||
cs := int(hdr.CsumStart)
|
cs := int(hdr.CsumStart)
|
||||||
co := int(hdr.CsumOffset)
|
co := int(hdr.CsumOffset)
|
||||||
@@ -421,7 +423,12 @@ func FinishChecksum(seg []byte, hdr Hdr) error {
|
|||||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||||
seg[cs+co] = 0
|
seg[cs+co] = 0
|
||||||
seg[cs+co+1] = 0
|
seg[cs+co+1] = 0
|
||||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -226,13 +226,14 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
|||||||
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
||||||
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
||||||
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
||||||
hdr := Hdr{
|
hdr := NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||||
GSOSize: 6, // two 6-byte segments
|
0, /*hdrLen*/
|
||||||
CsumStart: csumStart,
|
6, /*gsoSize: two 6-byte segments*/
|
||||||
CsumOffset: 6,
|
csumStart, /*csumStart*/
|
||||||
}
|
6, /*csumOffset*/
|
||||||
|
)
|
||||||
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||||
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||||
}
|
}
|
||||||
@@ -247,13 +248,14 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
|||||||
t.Run("too-short-rejected", func(t *testing.T) {
|
t.Run("too-short-rejected", func(t *testing.T) {
|
||||||
pkt := make([]byte, 25)
|
pkt := make([]byte, 25)
|
||||||
pkt[0] = 0x45 // IPv4, IHL 5
|
pkt[0] = 0x45 // IPv4, IHL 5
|
||||||
hdr := Hdr{
|
hdr := NewHeader(
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||||
GSOSize: 6,
|
0, /*hdrLen*/
|
||||||
CsumStart: 20,
|
6, /*gsoSize*/
|
||||||
CsumOffset: 6,
|
20, /*csumStart*/
|
||||||
}
|
6, /*csumOffset*/
|
||||||
|
)
|
||||||
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||||
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||||
}
|
}
|
||||||
@@ -303,9 +305,10 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
|||||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
||||||
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
||||||
}
|
}
|
||||||
// UDP-GSO keeps the same IPv4 ID across every segment.
|
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
|
||||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 {
|
// (inet_gso_segment's fixed-ID case is TCP-only).
|
||||||
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id)
|
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)
|
segPayLen := len(seg) - int(hdrLen)
|
||||||
@@ -333,3 +336,267 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
|
||||||
|
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
|
||||||
|
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
|
||||||
|
const ipLen, udpLen = 20, 8
|
||||||
|
pkt = make([]byte, ipLen+udpLen+len(payload))
|
||||||
|
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||||
|
pkt[8] = 64
|
||||||
|
pkt[9] = unix.IPPROTO_UDP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
|
||||||
|
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
|
||||||
|
copy(pkt[ipLen+udpLen:], payload)
|
||||||
|
|
||||||
|
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
|
||||||
|
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
|
||||||
|
|
||||||
|
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
|
||||||
|
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
|
||||||
|
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
|
||||||
|
var payload []byte
|
||||||
|
for i := 0; i < 0x10000; i++ {
|
||||||
|
p := []byte{byte(i >> 8), byte(i)}
|
||||||
|
pkt, hdr := buildUDPv4Single(p)
|
||||||
|
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
|
||||||
|
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
|
||||||
|
pkt[cs+co], pkt[cs+co+1] = 0, 0
|
||||||
|
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
|
||||||
|
payload = p
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if payload == nil {
|
||||||
|
t.Fatal("no 2-byte payload produced a zero checksum")
|
||||||
|
}
|
||||||
|
|
||||||
|
pkt, hdr := buildUDPv4Single(payload)
|
||||||
|
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||||
|
t.Fatalf("FinishChecksum: %v", err)
|
||||||
|
}
|
||||||
|
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
|
||||||
|
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
|
||||||
|
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
|
||||||
|
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
|
||||||
|
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
|
||||||
|
const cs, co = 20, tcpChecksumOff
|
||||||
|
|
||||||
|
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
|
||||||
|
seg := make([]byte, cs+co+2)
|
||||||
|
for i := range seg[cs:] {
|
||||||
|
seg[cs+i] = byte(i * 7)
|
||||||
|
}
|
||||||
|
var partial uint16
|
||||||
|
for i := 0; i <= 0xffff; i++ {
|
||||||
|
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
|
||||||
|
probe := append([]byte(nil), seg...)
|
||||||
|
probe[cs+co], probe[cs+co+1] = 0, 0
|
||||||
|
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
|
||||||
|
partial = uint16(i)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
|
||||||
|
|
||||||
|
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
|
||||||
|
if err := FinishChecksum(seg, hdr); err != nil {
|
||||||
|
t.Fatalf("FinishChecksum: %v", err)
|
||||||
|
}
|
||||||
|
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
|
||||||
|
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
|
||||||
|
func TestFinishChecksumUDPValidates(t *testing.T) {
|
||||||
|
payload := []byte("the definitive tun offloads branch")
|
||||||
|
pkt, hdr := buildUDPv4Single(payload)
|
||||||
|
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||||
|
t.Fatalf("FinishChecksum: %v", err)
|
||||||
|
}
|
||||||
|
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
|
||||||
|
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
|
||||||
|
t.Fatal("completed UDP checksum does not validate")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
|
||||||
|
// into gso_type for TSO superpackets with CWR set. CheckValid must
|
||||||
|
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
|
||||||
|
// fell into the default case and skipped the IP-version agreement check.
|
||||||
|
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
|
||||||
|
func TestCheckValidMasksGSOECN(t *testing.T) {
|
||||||
|
v4pkt, _, _ := buildTCPv4Super(100)
|
||||||
|
v6pkt := make([]byte, len(v4pkt))
|
||||||
|
copy(v6pkt, v4pkt)
|
||||||
|
v6pkt[0] = 0x60 // claim IPv6
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
pkt []byte
|
||||||
|
gsoType uint8
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
|
||||||
|
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||||
|
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||||
|
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||||
|
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
|
||||||
|
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
|
||||||
|
if tc.wantErr && err == nil {
|
||||||
|
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
|
||||||
|
}
|
||||||
|
if !tc.wantErr && err != nil {
|
||||||
|
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
|
||||||
|
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
|
||||||
|
// dodging both segmentation and FinishChecksum on its way downstream.
|
||||||
|
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
|
||||||
|
v4pkt, _, _ := buildTCPv4Super(100)
|
||||||
|
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
|
||||||
|
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
|
||||||
|
// against an independent RFC 1071 reference fold, hitting the carry edge
|
||||||
|
// cases (values whose first fold produces another carry).
|
||||||
|
func TestFoldComplementMatchesReference(t *testing.T) {
|
||||||
|
refFold := func(s uint64) uint16 {
|
||||||
|
for s>>16 != 0 {
|
||||||
|
s = s&0xffff + s>>16
|
||||||
|
}
|
||||||
|
return uint16(s)
|
||||||
|
}
|
||||||
|
cases := []uint32{
|
||||||
|
0, 1, 0xffff,
|
||||||
|
0x10000, // single carry
|
||||||
|
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||||
|
0xffff0000, // high half only
|
||||||
|
0xfffeffff, // first fold yields another carry
|
||||||
|
0xffffffff, // worst case
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
|
||||||
|
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
|
||||||
|
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
|
||||||
|
// into scratch, zero the fields the segment loop rewrites, sum. The production
|
||||||
|
// versions instead sum in place and subtract those fields via one's-complement
|
||||||
|
// arithmetic, which is faster but far less obvious — particularly for the TCP
|
||||||
|
// flags byte, which is only half of a 16-bit word. These references exist so
|
||||||
|
// that trade is checked rather than asserted.
|
||||||
|
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
|
||||||
|
var ipTmp [ipv4HeaderMaxLen]byte
|
||||||
|
copy(ipTmp[:ihl], pkt[:ihl])
|
||||||
|
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||||
|
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||||
|
tcpLen := headerLen - csumStart
|
||||||
|
var tmp [tcpHeaderMaxLen]byte
|
||||||
|
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
|
||||||
|
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||||
|
tmp[tcpFlagsOff] = 0
|
||||||
|
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||||
|
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
|
||||||
|
// what the file already has and reproduces identically on every run.
|
||||||
|
func randByte(state *uint32) byte {
|
||||||
|
*state = *state*1664525 + 1013904223
|
||||||
|
return byte(*state >> 24)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBaseSumsMatchZeroingReference(t *testing.T) {
|
||||||
|
state := uint32(12345)
|
||||||
|
|
||||||
|
t.Run("ipv4", func(t *testing.T) {
|
||||||
|
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
|
||||||
|
for iter := 0; iter < 5000; iter++ {
|
||||||
|
pkt := make([]byte, ihl)
|
||||||
|
for i := range pkt {
|
||||||
|
pkt[i] = randByte(&state)
|
||||||
|
}
|
||||||
|
pkt[0] = byte(0x40 | (ihl / 4))
|
||||||
|
|
||||||
|
want := referenceBaseIPv4HdrSum(pkt, ihl)
|
||||||
|
got, err := baseIPv4HdrSum(pkt, ihl)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ihl=%d: %v", ihl, err)
|
||||||
|
}
|
||||||
|
// Compare the value that reaches the wire: the raw partial
|
||||||
|
// sums may legally differ by one's-complement -0 vs +0.
|
||||||
|
for _, tl := range []uint32{20, 1500, 65535} {
|
||||||
|
for _, id := range []uint32{0, 0x4242, 0xffff} {
|
||||||
|
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
|
||||||
|
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("tcp", func(t *testing.T) {
|
||||||
|
const csumStart = 20
|
||||||
|
for dataOff := 5; dataOff <= 15; dataOff++ {
|
||||||
|
tcpLen := dataOff * 4
|
||||||
|
headerLen := csumStart + tcpLen
|
||||||
|
for iter := 0; iter < 5000; iter++ {
|
||||||
|
pkt := make([]byte, headerLen+64)
|
||||||
|
for i := range pkt {
|
||||||
|
pkt[i] = randByte(&state)
|
||||||
|
}
|
||||||
|
pkt[0] = 0x45
|
||||||
|
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
|
||||||
|
|
||||||
|
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
|
||||||
|
got := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||||
|
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
|
||||||
|
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
|
||||||
|
for _, l4 := range []uint32{20, 1460, 65535} {
|
||||||
|
a := foldComplement(want + seq + fl + l4)
|
||||||
|
b := foldComplement(got + seq + fl + l4)
|
||||||
|
if a != b {
|
||||||
|
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
|
||||||
|
dataOff, seq, fl, l4, a, b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -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. Only valid for single threaded use.
|
// 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
|
||||||
|
|||||||
+43
-92
@@ -25,32 +25,17 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
readers tio.QueueSet
|
readers tio.QueueSet
|
||||||
closeLock sync.Mutex
|
closeLock sync.Mutex
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
MaxMTU int
|
MaxMTU int
|
||||||
DefaultMTU int
|
DefaultMTU int
|
||||||
TXQueueLen int
|
TXQueueLen int
|
||||||
deviceIndex int
|
deviceIndex int
|
||||||
ioctlFd uintptr
|
ioctlFd uintptr
|
||||||
vnetHdr bool
|
vnetHdr bool
|
||||||
// offloadFlags is the exact TUN_F_* offload mask newTun negotiated with
|
|
||||||
// the kernel: usoOffloadFlags when USO was accepted, tsoOffloadFlags on
|
|
||||||
// the TSO-only fallback, or 0 when vnetHdr is off. TUNSETOFFLOAD is
|
|
||||||
// device-wide (drivers/net/tun.c set_offload updates tun->set_features
|
|
||||||
// for the whole netdev), so addQueue must replay this exact
|
|
||||||
// mask on every added queue — issuing a narrower mask there would
|
|
||||||
// silently downgrade offloads (e.g. disable USO) for all queues while
|
|
||||||
// they still advertise the stale capability.
|
|
||||||
offloadFlags uint
|
offloadFlags uint
|
||||||
// routeFeatureECN, when true, sets RTAX_FEATURE_ECN on every route we
|
|
||||||
// install for the tun. The kernel then actively negotiates ECN for
|
|
||||||
// connections destined to those prefixes (equivalent to `ip route
|
|
||||||
// change ... features ecn`) regardless of net.ipv4.tcp_ecn, so flows
|
|
||||||
// across the nebula mesh use ECN even when the host default is the
|
|
||||||
// passive setting (=2). Disable via tunnels.ecn=false.
|
|
||||||
routeFeatureECN bool
|
|
||||||
|
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
@@ -91,14 +76,7 @@ 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) {
|
||||||
// We don't know what flags the caller opened this fd with and can't turn
|
// We don't know what flags the caller opened this fd with and can't turn
|
||||||
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||||
t, err := newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks)
|
return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Device = "tun0"
|
|
||||||
|
|
||||||
return t, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||||
@@ -124,8 +102,7 @@ func openTunDev() (int, error) {
|
|||||||
return fd, nil
|
return fd, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen
|
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
|
||||||
// device name on success.
|
|
||||||
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||||
var req ifReq
|
var req ifReq
|
||||||
req.Flags = flags
|
req.Flags = flags
|
||||||
@@ -136,57 +113,45 @@ func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
|||||||
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a
|
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
|
||||||
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
|
|
||||||
// TSO_ECN tells the kernel we propagate ECN correctly through coalesce and
|
|
||||||
// segmentation, so it can deliver superpackets whose seed has CWR/ECE set
|
|
||||||
// or whose IP-level codepoint is CE.
|
|
||||||
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
||||||
|
|
||||||
// usoOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. Requires
|
// usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags.
|
||||||
// Linux ≥ 6.2; older kernels reject it and we fall back to TCP-only TSO via
|
// Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO
|
||||||
// tsoOffloadFlags.
|
const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||||
const usoOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
|
||||||
|
|
||||||
// offloadUSOEnabled reports whether the negotiated offload mask includes UDP
|
|
||||||
// Segmentation Offload. It is the single source of truth for the usoEnabled
|
|
||||||
// capability surfaced by each queue, so the mask stored on the tun and the USO
|
|
||||||
// bit reported to coalescers can never drift apart.
|
|
||||||
func offloadUSOEnabled(offloadFlags uint) bool {
|
func offloadUSOEnabled(offloadFlags uint) bool {
|
||||||
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
|
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) {
|
||||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
// IFF_TUN_EXCL prevents us from attaching to an already-running tun
|
||||||
|
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL)
|
||||||
if multiqueue {
|
if multiqueue {
|
||||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||||
}
|
}
|
||||||
nameStr := c.GetString("tun.dev", "")
|
nameStr := c.GetString("tun.dev", "")
|
||||||
|
|
||||||
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_*
|
|
||||||
// offloads via TUNSETOFFLOAD so we can receive TSO/USO superpackets.
|
|
||||||
// We try TSO+USO first, fall back to TSO-only on kernels without USO
|
|
||||||
// (Linux < 6.2), and finally give up on virtio headers entirely and
|
|
||||||
// reopen as a plain TUN if neither offload mask is accepted.
|
|
||||||
fd, err := openTunDev()
|
fd, err := openTunDev()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
vnetHdr := true
|
vnetHdr := true
|
||||||
// offloadFlags is the exact TUN_F_* mask the kernel accepted. We remember
|
|
||||||
// it (rather than a plain bool) so addQueue can replay the
|
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_* offloads
|
||||||
// identical device-wide mask on added queues instead of downgrading them.
|
// 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
|
var offloadFlags uint
|
||||||
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
vnetHdr = false
|
vnetHdr = false
|
||||||
} else {
|
} else {
|
||||||
// Try TSO+USO first. On kernels without USO support (Linux < 6.2)
|
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
|
||||||
// the ioctl returns EINVAL; fall back to the TCP-only mask before
|
offloadFlags = usoAndTSOOffloadFlags
|
||||||
// giving up on VNET_HDR entirely.
|
|
||||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoOffloadFlags)); err == nil {
|
|
||||||
offloadFlags = usoOffloadFlags
|
|
||||||
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
||||||
offloadFlags = tsoOffloadFlags
|
offloadFlags = tsoOffloadFlags
|
||||||
} else {
|
} else {
|
||||||
@@ -212,7 +177,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
|
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
|
||||||
}
|
}
|
||||||
|
|
||||||
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks)
|
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks, name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -222,16 +187,14 @@ 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
|
// newTunGeneric does all the stuff common to different tun initialization paths.
|
||||||
// paths. It will close your files on error. offloadFlags is the TUN_F_* mask
|
// It will close your files on error.
|
||||||
// newTun negotiated (0 when vnetHdr is off); the queues' USO capability is
|
// offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false)
|
||||||
// derived from it so it can never disagree with the mask we replay on added
|
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) {
|
||||||
// multiqueue readers.
|
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix) (*tun, error) {
|
|
||||||
var qs tio.QueueSet
|
var qs tio.QueueSet
|
||||||
var err error
|
var err error
|
||||||
if vnetHdr {
|
if vnetHdr {
|
||||||
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags))
|
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
|
||||||
} else {
|
} else {
|
||||||
qs, err = tio.NewPollQueueSet()
|
qs, err = tio.NewPollQueueSet()
|
||||||
}
|
}
|
||||||
@@ -242,11 +205,15 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
|
|||||||
}
|
}
|
||||||
err = qs.Add(fd)
|
err = qs.Add(fd)
|
||||||
if err != nil {
|
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)
|
_ = unix.Close(fd)
|
||||||
|
_ = qs.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
|
Device: name,
|
||||||
readers: qs,
|
readers: qs,
|
||||||
closeLock: sync.Mutex{},
|
closeLock: sync.Mutex{},
|
||||||
vnetHdr: vnetHdr,
|
vnetHdr: vnetHdr,
|
||||||
@@ -255,7 +222,6 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
|
|||||||
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),
|
||||||
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
|
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
|
||||||
routeFeatureECN: c.GetBool("tunnels.ecn", true),
|
|
||||||
routesFromSystem: map[netip.Prefix]routing.Gateways{},
|
routesFromSystem: map[netip.Prefix]routing.Gateways{},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
@@ -349,9 +315,7 @@ func (t *tun) reload(c *config.C, initial bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Queues opens additional kernel multiqueue fds until the device has n
|
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
|
||||||
// queues, then returns them all. The first queue was opened by newTun; each
|
|
||||||
// extra fd replays the negotiated offload state (see addQueue).
|
|
||||||
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||||
for len(t.readers.Queues()) < n {
|
for len(t.readers.Queues()) < n {
|
||||||
if err := t.addQueue(); err != nil {
|
if err := t.addQueue(); err != nil {
|
||||||
@@ -361,8 +325,7 @@ func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
|||||||
return t.readers.Queues(), nil
|
return t.readers.Queues(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to
|
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
|
||||||
// the queue set.
|
|
||||||
func (t *tun) addQueue() error {
|
func (t *tun) addQueue() error {
|
||||||
t.closeLock.Lock()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
@@ -382,10 +345,6 @@ func (t *tun) addQueue() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if t.vnetHdr {
|
if t.vnetHdr {
|
||||||
// Replay the exact mask newTun negotiated. TUNSETOFFLOAD is
|
|
||||||
// device-wide, so issuing the TSO-only mask here would disable USO
|
|
||||||
// for every queue (including queue 0) on kernels where newTun
|
|
||||||
// successfully enabled it, while the queues keep advertising USO.
|
|
||||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
||||||
@@ -566,18 +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
|
// Match the metric the kernel uses for its auto-installed connected route,
|
||||||
// route, so RouteReplace overwrites it in place instead of adding a
|
// so RouteReplace overwrites it in place instead of adding a second route at a worse metric.
|
||||||
// second route at a worse metric. IPv6 connected routes are installed
|
// IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0.
|
||||||
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the
|
// Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6.
|
||||||
// kernel route wins lookups and our MTU / AdvMSS / Features never
|
|
||||||
// apply on v6.
|
|
||||||
if cidr.Addr().Is6() {
|
if cidr.Addr().Is6() {
|
||||||
nr.Priority = 256
|
nr.Priority = 256
|
||||||
}
|
}
|
||||||
if t.routeFeatureECN {
|
|
||||||
nr.Features |= unix.RTAX_FEATURE_ECN
|
|
||||||
}
|
|
||||||
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)
|
||||||
@@ -627,9 +581,6 @@ func (t *tun) addRoutes(logErrors bool) error {
|
|||||||
if r.Metric > 0 {
|
if r.Metric > 0 {
|
||||||
nr.Priority = r.Metric
|
nr.Priority = r.Metric
|
||||||
}
|
}
|
||||||
if t.routeFeatureECN {
|
|
||||||
nr.Features |= unix.RTAX_FEATURE_ECN
|
|
||||||
}
|
|
||||||
|
|
||||||
err := netlink.RouteReplace(&nr)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -39,14 +39,14 @@ func TestTunAdvMSS(t *testing.T) {
|
|||||||
// capability: it is derived from the negotiated offload mask, so the mask
|
// capability: it is derived from the negotiated offload mask, so the mask
|
||||||
// stored on the tun and the capability reported to coalescers cannot drift.
|
// stored on the tun and the capability reported to coalescers cannot drift.
|
||||||
func TestOffloadUSOEnabled(t *testing.T) {
|
func TestOffloadUSOEnabled(t *testing.T) {
|
||||||
// usoOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
// usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||||
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
||||||
// addQueue) would not actually be a downgrade.
|
// addQueue) would not actually be a downgrade.
|
||||||
if usoOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||||
t.Fatalf("usoOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoOffloadFlags, tsoOffloadFlags)
|
t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags)
|
||||||
}
|
}
|
||||||
if usoOffloadFlags == tsoOffloadFlags {
|
if usoAndTSOOffloadFlags == tsoOffloadFlags {
|
||||||
t.Fatal("usoOffloadFlags must add bits beyond tsoOffloadFlags")
|
t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||||
}
|
}
|
||||||
|
|
||||||
cases := []struct {
|
cases := []struct {
|
||||||
@@ -54,7 +54,7 @@ func TestOffloadUSOEnabled(t *testing.T) {
|
|||||||
offloadFlags uint
|
offloadFlags uint
|
||||||
wantUSO bool
|
wantUSO bool
|
||||||
}{
|
}{
|
||||||
{"uso-negotiated", usoOffloadFlags, true},
|
{"uso-negotiated", usoAndTSOOffloadFlags, true},
|
||||||
{"tso-fallback", tsoOffloadFlags, false},
|
{"tso-fallback", tsoOffloadFlags, false},
|
||||||
{"no-vnet-hdr", 0, false},
|
{"no-vnet-hdr", 0, false},
|
||||||
}
|
}
|
||||||
@@ -78,12 +78,12 @@ func TestOffloadUSOEnabled(t *testing.T) {
|
|||||||
// TUNSETOFFLOAD argument is read from.
|
// TUNSETOFFLOAD argument is read from.
|
||||||
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
||||||
t.Run("uso-negotiated", func(t *testing.T) {
|
t.Run("uso-negotiated", func(t *testing.T) {
|
||||||
tn := &tun{vnetHdr: true, offloadFlags: usoOffloadFlags}
|
tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
|
||||||
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
||||||
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
||||||
// mask (the original bug).
|
// mask (the original bug).
|
||||||
if tn.offloadFlags != usoOffloadFlags {
|
if tn.offloadFlags != usoAndTSOOffloadFlags {
|
||||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoOffloadFlags)
|
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags)
|
||||||
}
|
}
|
||||||
if tn.offloadFlags == tsoOffloadFlags {
|
if tn.offloadFlags == tsoOffloadFlags {
|
||||||
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
||||||
|
|||||||
+6
-4
@@ -10,11 +10,13 @@ import (
|
|||||||
_ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux
|
_ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux
|
||||||
)
|
)
|
||||||
|
|
||||||
// startPprofServer serves net/http/pprof on :6060 for the life of ctx. It is
|
// startPprofServer serves net/http/pprof on localhost:6060 for the life of
|
||||||
// only compiled into debug builds (`-tags debug`, `make debug`), so a debug
|
// ctx. It is only compiled into debug builds (`-tags debug`, `make debug`),
|
||||||
// build announces itself with the Info line below.
|
// so a debug build announces itself with the Info line below. Loopback only:
|
||||||
|
// a wildcard bind would expose profiles (peer addresses, config-derived
|
||||||
|
// state) to anything that can reach the host, the overlay included.
|
||||||
func startPprofServer(ctx context.Context, l *slog.Logger) {
|
func startPprofServer(ctx context.Context, l *slog.Logger) {
|
||||||
server := &http.Server{Addr: ":6060", Handler: nil}
|
server := &http.Server{Addr: "localhost:6060", Handler: nil}
|
||||||
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
|
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
|
|||||||
+1
-1
@@ -161,7 +161,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
switch existingRelay.State {
|
switch existingRelay.State {
|
||||||
case Established:
|
case Established:
|
||||||
hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String())
|
hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String())
|
||||||
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false)
|
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false, 0)
|
||||||
case Disestablished:
|
case Disestablished:
|
||||||
// Mark this relay as 'requested'
|
// Mark this relay as 'requested'
|
||||||
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||||
|
|||||||
+13
-28
@@ -14,43 +14,28 @@ const MTU = 9001
|
|||||||
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||||
const MaxWriteBatch = 128
|
const MaxWriteBatch = 128
|
||||||
|
|
||||||
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
|
|
||||||
// data, kernel offload state, etc.) and passed to EncReader callbacks.
|
|
||||||
// Backends that do not produce a particular signal leave its zero value.
|
|
||||||
//
|
|
||||||
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
|
|
||||||
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
|
|
||||||
// means Not-ECT, which is also the value backends without ECN RX support
|
|
||||||
// supply on every packet.
|
|
||||||
type RxMeta struct {
|
|
||||||
OuterECN byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type EncReader func(
|
type EncReader func(
|
||||||
addr netip.AddrPort,
|
addr netip.AddrPort,
|
||||||
payload []byte,
|
payload []byte,
|
||||||
meta RxMeta,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type Conn interface {
|
type Conn interface {
|
||||||
Rebind() error
|
Rebind() error
|
||||||
LocalAddr() (netip.AddrPort, error)
|
LocalAddr() (netip.AddrPort, error)
|
||||||
// ListenOut invokes r for each received packet. On batch-capable
|
// ListenOut invokes r for each received packet.
|
||||||
// backends (recvmmsg), flush is called after each batch is fully
|
// On batch-capable backends (recvmmsg), flush is called after each batch is fully delivered.
|
||||||
// delivered — callers use it to flush per-batch accumulators such as
|
// Callers use it to flush per-batch accumulators such as TUN write coalescers.
|
||||||
// TUN write coalescers. Single-packet backends call flush after each
|
// Single-packet backends call flush after each packet. flush must not be nil.
|
||||||
// packet. flush must not be nil.
|
|
||||||
ListenOut(r EncReader, flush func()) error
|
ListenOut(r EncReader, flush func()) error
|
||||||
WriteTo(b []byte, addr netip.AddrPort) error
|
WriteTo(b []byte, addr netip.AddrPort) error
|
||||||
// WriteBatch sends a contiguous batch of packets, each with its own
|
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||||
// destination. bufs and addrs must have the same length. outerECNs may
|
// destination. bufs and addrs must have the same length. Linux uses
|
||||||
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
|
// sendmmsg(2) for a single syscall.
|
||||||
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
|
//
|
||||||
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
|
// Returns the number of packets successfully written. A destination the kernel rejects costs only
|
||||||
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
|
// its own packet, so a short count means some peers were undeliverable, not that the batch failed.
|
||||||
// cmsg; other backends ignore it. Returns on the first error; callers
|
// Not safe for concurrent use on the same Conn.
|
||||||
// may observe a partial send if some packets went out before the error.
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
|
||||||
ReloadConfig(c *config.C)
|
ReloadConfig(c *config.C)
|
||||||
SupportsMultipleReaders() bool
|
SupportsMultipleReaders() bool
|
||||||
Close() error
|
Close() error
|
||||||
@@ -73,8 +58,8 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
|||||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
func (NoopConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||||
return nil
|
return len(bufs), nil
|
||||||
}
|
}
|
||||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
|
||||||
|
//
|
||||||
|
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
|
||||||
|
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
|
||||||
|
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
|
||||||
|
// keeps this package from referencing whatever owns the rebind.
|
||||||
|
//
|
||||||
|
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
|
||||||
|
type NetworkChangeMonitor struct {
|
||||||
|
l *slog.Logger
|
||||||
|
ctx context.Context
|
||||||
|
enabled bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
|
||||||
|
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
|
||||||
|
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
|
||||||
|
return &NetworkChangeMonitor{
|
||||||
|
l: l,
|
||||||
|
ctx: ctx,
|
||||||
|
enabled: c.GetBool("listen.rebind_on_network_change", true),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
|
||||||
|
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
|
||||||
|
func (m *NetworkChangeMonitor) Start(rebind func()) {
|
||||||
|
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
changes, err := watchNetworkChanges(m.ctx, m.l)
|
||||||
|
if err != nil {
|
||||||
|
// Not fatal. Everything else still works, we just won't notice a network change on our own.
|
||||||
|
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if changes == nil {
|
||||||
|
// This platform's sockets don't go stale, so there is nothing to watch for.
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
m.l.Info("Watching for network changes to rebind the udp listener")
|
||||||
|
|
||||||
|
for range changes {
|
||||||
|
m.l.Info("Local network changed, rebinding the udp listener")
|
||||||
|
rebind()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
//go:build darwin && !ios && !e2e_testing
|
||||||
|
// +build darwin,!ios,!e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"os"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
|
||||||
|
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
|
||||||
|
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
|
||||||
|
netChangeSettleWindow = time.Second
|
||||||
|
|
||||||
|
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
|
||||||
|
// the kernel, so being generous here is how we avoid missing a message.
|
||||||
|
netChangeReadBuffer = 4096
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
|
||||||
|
//
|
||||||
|
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
|
||||||
|
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
|
||||||
|
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
|
||||||
|
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
|
||||||
|
//
|
||||||
|
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
|
||||||
|
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
|
||||||
|
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
|
||||||
|
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
|
||||||
|
sock, err := openRouteSocket()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
changes := make(chan struct{}, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer close(changes)
|
||||||
|
defer func() { _ = sock.Close() }()
|
||||||
|
|
||||||
|
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
|
||||||
|
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
|
||||||
|
done := make(chan struct{})
|
||||||
|
defer close(done)
|
||||||
|
go func() {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
_ = sock.Close()
|
||||||
|
case <-done:
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
watchRouteSocket(l, sock, changes)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return changes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
|
||||||
|
// the socket is closed, which is how cancellation gets us out of here.
|
||||||
|
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
|
||||||
|
buf := make([]byte, netChangeReadBuffer)
|
||||||
|
|
||||||
|
for {
|
||||||
|
n, err := sock.Read(buf)
|
||||||
|
if err != nil {
|
||||||
|
logRouteSocketError(l, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isNetworkChange(buf[:n]) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
|
||||||
|
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
|
||||||
|
// simply produce another report, which is the correct outcome anyway.
|
||||||
|
deadline := time.Now().Add(netChangeSettleWindow)
|
||||||
|
for {
|
||||||
|
if err = sock.SetReadDeadline(deadline); err != nil {
|
||||||
|
logRouteSocketError(l, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err = sock.Read(buf); err != nil {
|
||||||
|
if os.IsTimeout(err) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
logRouteSocketError(l, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = sock.SetReadDeadline(time.Time{}); err != nil {
|
||||||
|
logRouteSocketError(l, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case changes <- struct{}{}:
|
||||||
|
default:
|
||||||
|
// One already pending, and a second "the network moved" tells the reader nothing new.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
|
||||||
|
func logRouteSocketError(l *slog.Logger, err error) {
|
||||||
|
if errors.Is(err, os.ErrClosed) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
|
||||||
|
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
|
||||||
|
func openRouteSocket() (*os.File, error) {
|
||||||
|
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = unix.SetNonblock(fd, true); err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return os.NewFile(uintptr(fd), "route"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
|
||||||
|
//
|
||||||
|
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
|
||||||
|
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
|
||||||
|
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
|
||||||
|
func isNetworkChange(msg []byte) bool {
|
||||||
|
if len(msg) < 4 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// u_short msglen, u_char version, u_char type
|
||||||
|
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
switch msg[3] {
|
||||||
|
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
|
||||||
|
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
|
||||||
|
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,244 @@
|
|||||||
|
//go:build darwin && !ios && !e2e_testing
|
||||||
|
// +build darwin,!ios,!e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.uber.org/goleak"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
|
||||||
|
func routeMsg(msgType uint8, extra int) []byte {
|
||||||
|
msg := make([]byte, 4+extra)
|
||||||
|
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
|
||||||
|
msg[2] = unix.RTM_VERSION
|
||||||
|
msg[3] = msgType
|
||||||
|
return msg
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsNetworkChange(t *testing.T) {
|
||||||
|
// The three that mean our addressing may have moved
|
||||||
|
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
|
||||||
|
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
|
||||||
|
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
|
||||||
|
|
||||||
|
// Route churn is not something a rebind helps with
|
||||||
|
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
|
||||||
|
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
|
||||||
|
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
|
||||||
|
|
||||||
|
// Garbage must not be mistaken for a change
|
||||||
|
assert.False(t, isNetworkChange(nil), "empty")
|
||||||
|
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
|
||||||
|
|
||||||
|
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
|
||||||
|
wrongVersion[2] = unix.RTM_VERSION + 1
|
||||||
|
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
|
||||||
|
|
||||||
|
lying := routeMsg(unix.RTM_NEWADDR, 0)
|
||||||
|
binary.NativeEndian.PutUint16(lying[0:2], 512)
|
||||||
|
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
|
||||||
|
}
|
||||||
|
|
||||||
|
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
|
||||||
|
// stands in for the kernel so the watch loop can be driven with synthetic messages.
|
||||||
|
func socketPair(t *testing.T) (*os.File, int) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, unix.SetNonblock(fds[0], true))
|
||||||
|
|
||||||
|
f := os.NewFile(uintptr(fds[0]), "route")
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = f.Close()
|
||||||
|
_ = unix.Close(fds[1])
|
||||||
|
})
|
||||||
|
|
||||||
|
return f, fds[1]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
|
||||||
|
sock, kernel := socketPair(t)
|
||||||
|
changes := make(chan struct{}, 1)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
|
||||||
|
// exactly one report rather than one apiece.
|
||||||
|
for range 5 {
|
||||||
|
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
// Uninteresting messages in the middle of a burst must not add a report of their own either.
|
||||||
|
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-changes:
|
||||||
|
case <-time.After(netChangeSettleWindow * 4):
|
||||||
|
t.Fatal("a burst should have reported a change")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nothing more from that burst
|
||||||
|
select {
|
||||||
|
case <-changes:
|
||||||
|
t.Fatal("a burst should report exactly once")
|
||||||
|
case <-time.After(netChangeSettleWindow):
|
||||||
|
}
|
||||||
|
|
||||||
|
// A change after the window has closed is a separate event and gets its own report.
|
||||||
|
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
|
||||||
|
require.NoError(t, err)
|
||||||
|
select {
|
||||||
|
case <-changes:
|
||||||
|
case <-time.After(netChangeSettleWindow * 4):
|
||||||
|
t.Fatal("a later change should report again")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Closing the socket is how the real thing shuts down
|
||||||
|
require.NoError(t, sock.Close())
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
|
||||||
|
sock, kernel := socketPair(t)
|
||||||
|
changes := make(chan struct{}, 1)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
|
||||||
|
_, err := unix.Write(kernel, routeMsg(msgType, 8))
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-changes:
|
||||||
|
t.Fatal("route churn alone must not report a change")
|
||||||
|
case <-time.After(netChangeSettleWindow * 2):
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, sock.Close())
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
|
||||||
|
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
|
||||||
|
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
|
||||||
|
sock, kernel := socketPair(t)
|
||||||
|
changes := make(chan struct{}, 1)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
|
||||||
|
for range 3 {
|
||||||
|
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||||
|
require.NoError(t, err)
|
||||||
|
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The watcher must still be alive and responsive to a close
|
||||||
|
require.NoError(t, sock.Close())
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("watchRouteSocket wedged on a full channel")
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
|
||||||
|
// cancelling the context closes the channel so a ranging caller falls out of its loop.
|
||||||
|
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
changes, err := watchNetworkChanges(ctx, test.NewLogger())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, changes, "darwin should support watching")
|
||||||
|
|
||||||
|
drained := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
for range changes {
|
||||||
|
}
|
||||||
|
close(drained)
|
||||||
|
}()
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
select {
|
||||||
|
case <-drained:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("cancelling the context should close the changes channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
|
||||||
|
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
|
||||||
|
// must return it and clean up the watch goroutines.
|
||||||
|
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
|
||||||
|
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
|
||||||
|
// leaks itself.
|
||||||
|
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
l := test.NewLogger()
|
||||||
|
c := config.NewC(l)
|
||||||
|
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
|
||||||
|
m := NewNetworkChangeMonitor(ctx, l, c)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
m.Start(func() {})
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
|
||||||
|
// watch or no-op'd, both of which we want to catch.
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
t.Fatal("Start returned instead of watching")
|
||||||
|
case <-time.After(time.Millisecond * 250):
|
||||||
|
}
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("Start did not return after the context was cancelled")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Starting again after the context is dead must not open anything.
|
||||||
|
m.Start(func() {})
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
//go:build !darwin || ios || e2e_testing
|
||||||
|
// +build !darwin ios e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchNetworkChanges is a no-op outside of darwin.
|
||||||
|
//
|
||||||
|
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
|
||||||
|
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
|
||||||
|
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
|
||||||
|
// and two things racing to rebind the same socket is worse than one.
|
||||||
|
//
|
||||||
|
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
|
||||||
|
// selecting on it, since a receive from a nil channel blocks forever.
|
||||||
|
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
|
||||||
|
t.Helper()
|
||||||
|
l := test.NewLogger()
|
||||||
|
c := config.NewC(l)
|
||||||
|
require.NoError(t, c.LoadString(cfg))
|
||||||
|
return NewNetworkChangeMonitor(ctx, l, c)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
|
||||||
|
// Says nothing about rebinding, so this covers the default.
|
||||||
|
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
|
||||||
|
assert.True(t, m.enabled, "should default to on")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
|
||||||
|
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
|
||||||
|
require.False(t, m.enabled)
|
||||||
|
|
||||||
|
// Must return without opening a socket. If it watched anything this would block.
|
||||||
|
m.Start(func() {})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
|
||||||
|
// Nothing to rebind, so there is no point watching, on any platform.
|
||||||
|
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
|
||||||
|
m.Start(nil)
|
||||||
|
}
|
||||||
@@ -1,62 +0,0 @@
|
|||||||
//go:build !android && !e2e_testing
|
|
||||||
// +build !android,!e2e_testing
|
|
||||||
|
|
||||||
package udp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
"syscall"
|
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// rawSendmmsg performs sendmmsg(2) over a syscall.RawConn without
|
|
||||||
// allocating a closure per call. The struct holds preallocated in/out
|
|
||||||
// scratch (chunk/sent/errno) and a method-value bound at construction so
|
|
||||||
// rawConn.Write receives a stable function pointer instead of a fresh
|
|
||||||
// closure on every send.
|
|
||||||
type rawSendmmsg struct {
|
|
||||||
msgs []rawMessage
|
|
||||||
chunk int
|
|
||||||
sent int
|
|
||||||
errno syscall.Errno
|
|
||||||
callback func(fd uintptr) bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// bind wires r.callback to r.run. Must be called once after r.msgs is set;
|
|
||||||
// subsequent send calls invoke r.callback without rebinding.
|
|
||||||
func (r *rawSendmmsg) bind() { r.callback = r.run }
|
|
||||||
|
|
||||||
// run is the preallocated callback rawConn.Write invokes. It reads its
|
|
||||||
// input (r.chunk) and writes its outputs (r.sent, r.errno) through the
|
|
||||||
// rawSendmmsg fields so the method value does not capture per-call locals
|
|
||||||
// and therefore does not heap-allocate.
|
|
||||||
func (r *rawSendmmsg) run(fd uintptr) bool {
|
|
||||||
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
|
|
||||||
uintptr(unsafe.Pointer(&r.msgs[0])), uintptr(r.chunk),
|
|
||||||
0, 0, 0,
|
|
||||||
)
|
|
||||||
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
r.sent = int(r1)
|
|
||||||
r.errno = errno
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// send issues sendmmsg over rc against the first n entries of r.msgs.
|
|
||||||
// Returns the number of entries the kernel processed and any error;
|
|
||||||
// matches the original sendmmsg helper's contract.
|
|
||||||
func (r *rawSendmmsg) send(rc syscall.RawConn, n int) (int, error) {
|
|
||||||
r.chunk = n
|
|
||||||
r.sent = 0
|
|
||||||
r.errno = 0
|
|
||||||
if err := rc.Write(r.callback); err != nil {
|
|
||||||
return r.sent, err
|
|
||||||
}
|
|
||||||
if r.errno != 0 {
|
|
||||||
return r.sent, &net.OpError{Op: "sendmmsg", Err: r.errno}
|
|
||||||
}
|
|
||||||
return r.sent, nil
|
|
||||||
}
|
|
||||||
+16
-10
@@ -140,13 +140,20 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||||
|
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||||
|
// TODO: WriteTo maps EWOULDBLOCK to an error, so a full send buffer
|
||||||
|
// silently drops the rest of a burst (linux blocks instead). Poll for
|
||||||
|
// writability on EAGAIN before giving up on the remainder.
|
||||||
|
written := 0
|
||||||
for i, b := range bufs {
|
for i, b := range bufs {
|
||||||
if err := u.WriteTo(b, addrs[i]); err != nil {
|
if err := u.WriteTo(b, addrs[i]); err == nil {
|
||||||
return err
|
written++
|
||||||
|
} else {
|
||||||
|
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return written, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
@@ -188,7 +195,7 @@ func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||||
flush()
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -197,6 +204,9 @@ func (u *StdConn) SupportsMultipleReaders() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
|
||||||
|
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
|
||||||
|
// this way on its own, which is what strands us after the underlying network changes.
|
||||||
func (u *StdConn) Rebind() error {
|
func (u *StdConn) Rebind() error {
|
||||||
var err error
|
var err error
|
||||||
if u.isV4 {
|
if u.isV4 {
|
||||||
@@ -205,9 +215,5 @@ func (u *StdConn) Rebind() error {
|
|||||||
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
|
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
return err
|
||||||
u.l.Error("Failed to rebind udp socket", "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,61 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
|
|
||||||
package udp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestPlanRunBreaksOnECNChange confirms that two same-destination, same-size
|
|
||||||
// packets with different outer ECN end up in separate sendmmsg entries (the
|
|
||||||
// kernel stamps one outer codepoint per entry, so a run that straddled the
|
|
||||||
// boundary would silently lose information).
|
|
||||||
func TestPlanRunBreaksOnECNChange(t *testing.T) {
|
|
||||||
u := &StdConn{gsoSupported: true, maxGSOSegments: 63}
|
|
||||||
dst := netip.MustParseAddrPort("10.0.0.1:4242")
|
|
||||||
|
|
||||||
bufs := [][]byte{
|
|
||||||
make([]byte, 1200),
|
|
||||||
make([]byte, 1200),
|
|
||||||
make([]byte, 1200),
|
|
||||||
}
|
|
||||||
addrs := []netip.AddrPort{dst, dst, dst}
|
|
||||||
|
|
||||||
t.Run("uniform_ecn_runs_together", func(t *testing.T) {
|
|
||||||
ecns := []byte{0x02, 0x02, 0x02}
|
|
||||||
runLen, segSize := u.planRun(bufs, addrs, ecns, 0, 64)
|
|
||||||
if runLen != 3 {
|
|
||||||
t.Errorf("runLen=%d want 3 (uniform ECT(0))", runLen)
|
|
||||||
}
|
|
||||||
if segSize != 1200 {
|
|
||||||
t.Errorf("segSize=%d want 1200", segSize)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("ecn_change_truncates_run", func(t *testing.T) {
|
|
||||||
// 0,0,3: first two run together, CE seeds a fresh entry.
|
|
||||||
ecns := []byte{0x00, 0x00, 0x03}
|
|
||||||
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
|
||||||
if runLen != 2 {
|
|
||||||
t.Errorf("runLen=%d want 2 (ECN changes at index 2)", runLen)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("nil_ecns_runs_full", func(t *testing.T) {
|
|
||||||
runLen, _ := u.planRun(bufs, addrs, nil, 0, 64)
|
|
||||||
if runLen != 3 {
|
|
||||||
t.Errorf("runLen=%d want 3 (nil ecns means no break)", runLen)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("first_ecn_is_singleton", func(t *testing.T) {
|
|
||||||
// Second packet has different ECN from the first → run halts at 1
|
|
||||||
// (the first packet alone forms the run).
|
|
||||||
ecns := []byte{0x00, 0x03, 0x03}
|
|
||||||
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
|
||||||
if runLen != 1 {
|
|
||||||
t.Errorf("runLen=%d want 1 (different ECN immediately)", runLen)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
+9
-5
@@ -44,13 +44,17 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||||
|
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||||
|
written := 0
|
||||||
for i, b := range bufs {
|
for i, b := range bufs {
|
||||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err == nil {
|
||||||
return err
|
written++
|
||||||
|
} else {
|
||||||
|
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return written, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
@@ -102,7 +106,7 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||||
flush()
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+197
-717
File diff suppressed because it is too large
Load Diff
@@ -30,39 +30,6 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
|
||||||
msgs := make([]rawMessage, n)
|
|
||||||
buffers := make([][]byte, n)
|
|
||||||
names := make([][]byte, n)
|
|
||||||
|
|
||||||
var cmsgs []byte
|
|
||||||
if cmsgSpace > 0 {
|
|
||||||
cmsgs = make([]byte, n*cmsgSpace)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range msgs {
|
|
||||||
buffers[i] = make([]byte, bufSize)
|
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
|
||||||
|
|
||||||
vs := []iovec{
|
|
||||||
{Base: &buffers[i][0], Len: uint32(len(buffers[i]))},
|
|
||||||
}
|
|
||||||
|
|
||||||
msgs[i].Hdr.Iov = &vs[0]
|
|
||||||
msgs[i].Hdr.Iovlen = uint32(len(vs))
|
|
||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
|
||||||
|
|
||||||
if cmsgSpace > 0 {
|
|
||||||
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
|
||||||
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return msgs, buffers, names, cmsgs
|
|
||||||
}
|
|
||||||
|
|
||||||
func setIovLen(v *iovec, n int) {
|
func setIovLen(v *iovec, n int) {
|
||||||
v.Len = uint32(n)
|
v.Len = uint32(n)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -33,39 +33,6 @@ type rawMessage struct {
|
|||||||
Pad0 [4]byte
|
Pad0 [4]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
|
||||||
msgs := make([]rawMessage, n)
|
|
||||||
buffers := make([][]byte, n)
|
|
||||||
names := make([][]byte, n)
|
|
||||||
|
|
||||||
var cmsgs []byte
|
|
||||||
if cmsgSpace > 0 {
|
|
||||||
cmsgs = make([]byte, n*cmsgSpace)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range msgs {
|
|
||||||
buffers[i] = make([]byte, bufSize)
|
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
|
||||||
|
|
||||||
vs := []iovec{
|
|
||||||
{Base: &buffers[i][0], Len: uint64(len(buffers[i]))},
|
|
||||||
}
|
|
||||||
|
|
||||||
msgs[i].Hdr.Iov = &vs[0]
|
|
||||||
msgs[i].Hdr.Iovlen = uint64(len(vs))
|
|
||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
|
||||||
|
|
||||||
if cmsgSpace > 0 {
|
|
||||||
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
|
||||||
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return msgs, buffers, names, cmsgs
|
|
||||||
}
|
|
||||||
|
|
||||||
func setIovLen(v *iovec, n int) {
|
func setIovLen(v *iovec, n int) {
|
||||||
v.Len = uint64(n)
|
v.Len = uint64(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+559
-104
@@ -3,11 +3,11 @@
|
|||||||
package udp
|
package udp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"syscall"
|
"slices"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
@@ -54,39 +54,6 @@ func buildCmsg(level, typ int32, data []byte) []byte {
|
|||||||
return buf
|
return buf
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestParseRecvCmsgOuterECNFamily is the RX half of the dual-stack ECN fix:
|
|
||||||
// parseRecvCmsg must read the outer ECN from whichever family the kernel
|
|
||||||
// delivered, not from the socket family. On the default `::` dual-stack bind
|
|
||||||
// a v4 peer's outer ECN arrives as an IP_TOS cmsg, which the old socket-family
|
|
||||||
// gate ignored entirely.
|
|
||||||
func TestParseRecvCmsgOuterECNFamily(t *testing.T) {
|
|
||||||
tc := make([]byte, 4)
|
|
||||||
binary.NativeEndian.PutUint32(tc, 0x02)
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
ctrl []byte
|
|
||||||
want byte
|
|
||||||
}{
|
|
||||||
{"ip_tos_ce", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x03}), 0x03},
|
|
||||||
{"ip_tos_ect0", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x02}), 0x02},
|
|
||||||
{"ipv6_tclass_ect0", buildCmsg(int32(unix.IPPROTO_IPV6), int32(unix.IPV6_TCLASS), tc), 0x02},
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
t.Run(c.name, func(t *testing.T) {
|
|
||||||
hdr := &msghdr{Control: &c.ctrl[0]}
|
|
||||||
setMsgControllen(hdr, len(c.ctrl))
|
|
||||||
gso, ecn := parseRecvCmsg(hdr, false, true)
|
|
||||||
if gso != 0 {
|
|
||||||
t.Errorf("gso = %d, want 0 (no UDP_GRO cmsg present)", gso)
|
|
||||||
}
|
|
||||||
if ecn != c.want {
|
|
||||||
t.Errorf("ecn = 0x%02x, want 0x%02x", ecn, c.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func testLogger() *slog.Logger {
|
func testLogger() *slog.Logger {
|
||||||
return slog.New(slog.DiscardHandler)
|
return slog.New(slog.DiscardHandler)
|
||||||
}
|
}
|
||||||
@@ -123,9 +90,13 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
|||||||
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
|
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
|
||||||
addrs := []netip.AddrPort{good, bad, good}
|
addrs := []netip.AddrPort{good, bad, good}
|
||||||
|
|
||||||
if err := sender.WriteBatch(bufs, addrs, nil); err != nil {
|
n, err := sender.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
|
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
|
||||||
}
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Errorf("WriteBatch wrote %d packets, want 2 of 3 (the bad-family dest is the only casualty)", n)
|
||||||
|
}
|
||||||
|
|
||||||
got := map[string]bool{}
|
got := map[string]bool{}
|
||||||
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
@@ -145,12 +116,11 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWriteBatchOuterTOSToV4Mapped is the TX half of the dual-stack ECN fix,
|
// TestWriteBatchUnreachableDestDeliversOthers is the kernel-rejection twin of
|
||||||
// verified against a live kernel: WriteBatch on the default `::` dual-stack
|
// TestWriteBatchBadFamilyDeliversOthers. A destination the kernel refuses outright (240.0.0.0/4 is reserved, so
|
||||||
// socket, sending to a v4-mapped destination, must stamp the outer ECN via an
|
// the send returns EINVAL) fails its sendmmsg entry; WriteBatch must drop only that entry and still deliver
|
||||||
// IP_TOS cmsg (not IPV6_TCLASS, which the kernel's v4 path ignores) so a v4
|
// every other packet rather than abandoning the batch at the first failure.
|
||||||
// receiver actually sees it.
|
func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
|
||||||
func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
|
|
||||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||||
@@ -158,76 +128,561 @@ func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
|
|||||||
defer rx.Close()
|
defer rx.Close()
|
||||||
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||||
|
|
||||||
// Ask the kernel to deliver the received outer TOS as ancillary data.
|
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 1)
|
||||||
rxRaw, err := rx.SyscallConn()
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("SyscallConn: %v", err)
|
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
|
||||||
}
|
|
||||||
var soErr error
|
|
||||||
if err := rxRaw.Control(func(fd uintptr) {
|
|
||||||
soErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
|
|
||||||
}); err != nil || soErr != nil {
|
|
||||||
t.Skipf("cannot enable IP_RECVTOS (sandbox/kernel?): ctrl=%v so=%v", err, soErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := NewListener(testLogger(), netip.IPv6Unspecified(), 0, false, 1)
|
|
||||||
if err != nil {
|
|
||||||
t.Skipf("cannot open dual-stack sender (sandbox?): %v", err)
|
|
||||||
}
|
}
|
||||||
defer c.Close()
|
defer c.Close()
|
||||||
sender := c.(*StdConn)
|
sender := c.(*StdConn)
|
||||||
if sender.isV4 {
|
|
||||||
t.Skipf("sender came up v4-only; need a dual-stack v6 socket for this test")
|
good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||||
|
bad := netip.MustParseAddrPort("240.0.0.1:9999") // reserved space, the kernel refuses it
|
||||||
|
|
||||||
|
bufs := [][]byte{[]byte("P0"), []byte("P1"), []byte("BAD"), []byte("P3"), []byte("P4")}
|
||||||
|
addrs := []netip.AddrPort{good, good, bad, good, good}
|
||||||
|
|
||||||
|
// The bad destination is reported, but only after every other packet has been attempted.
|
||||||
|
if _, err := sender.WriteBatch(bufs, addrs); err == nil {
|
||||||
|
t.Log("WriteBatch returned nil; kernel accepted the reserved address, delivery assertions still apply")
|
||||||
}
|
}
|
||||||
|
|
||||||
// v4-mapped-in-v6 destination: routed through the kernel's IPv4 path.
|
got := map[string]bool{}
|
||||||
dst := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
const wantECN = byte(0x02) // ECT(0)
|
buf := make([]byte, 64)
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
if err := sender.WriteBatch([][]byte{[]byte("tos-probe")}, []netip.AddrPort{dst}, []byte{wantECN}); err != nil {
|
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
|
||||||
t.Fatalf("WriteBatch: %v", err)
|
if rerr != nil {
|
||||||
}
|
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
|
||||||
|
|
||||||
// Read the datagram plus its ancillary TOS.
|
|
||||||
rx.SetReadDeadline(time.Now().Add(3 * time.Second))
|
|
||||||
payload := make([]byte, 128)
|
|
||||||
oob := make([]byte, 512)
|
|
||||||
var n, oobn int
|
|
||||||
var rerr error
|
|
||||||
if err := rxRaw.Read(func(fd uintptr) bool {
|
|
||||||
n, oobn, _, _, rerr = unix.Recvmsg(int(fd), payload, oob, 0)
|
|
||||||
if rerr == syscall.EAGAIN || rerr == syscall.EWOULDBLOCK {
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
return true
|
got[string(buf[:n])] = true
|
||||||
}); err != nil {
|
|
||||||
t.Fatalf("waiting for datagram failed (no delivery?): %v", err)
|
|
||||||
}
|
}
|
||||||
if rerr != nil {
|
for _, want := range []string{"P0", "P1", "P3", "P4"} {
|
||||||
t.Fatalf("Recvmsg: %v", rerr)
|
if !got[want] {
|
||||||
}
|
t.Errorf("packet %s was not delivered; delivered set = %v", want, got)
|
||||||
if string(payload[:n]) != "tos-probe" {
|
}
|
||||||
t.Fatalf("payload = %q, want %q", string(payload[:n]), "tos-probe")
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
cmsgs, err := unix.ParseSocketControlMessage(oob[:oobn])
|
// TestParseRecvCmsgCorruptLenNoPanic: a cmsg Len near max-int used to wrap
|
||||||
if err != nil {
|
// off+clen negative, slip past the bounds check, and drive the walk offset
|
||||||
t.Fatalf("ParseSocketControlMessage: %v", err)
|
// negative -- a panic on the next ctrl[off]. The guard must compare Len
|
||||||
}
|
// against the remaining bytes instead. Also pins the plain truncated-Len
|
||||||
found := false
|
// cases (too small, larger than the buffer) to a clean early return.
|
||||||
var gotTOS byte
|
func TestParseRecvCmsgCorruptLenNoPanic(t *testing.T) {
|
||||||
for _, m := range cmsgs {
|
// First cmsg: a valid empty one so the walk advances past off=0
|
||||||
if m.Header.Level == unix.IPPROTO_IP && m.Header.Type == unix.IP_TOS && len(m.Data) >= 1 {
|
// (off+clen can't overflow while off is still zero).
|
||||||
found = true
|
valid := buildCmsg(int32(unix.SOL_UDP), int32(unix.UDP_GRO), make([]byte, 4))
|
||||||
gotTOS = m.Data[0]
|
|
||||||
|
corrupt := func(lenVal int) []byte {
|
||||||
|
buf := make([]byte, len(valid)+unix.CmsgSpace(4))
|
||||||
|
copy(buf, valid)
|
||||||
|
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[len(valid)]))
|
||||||
|
h.Level = int32(unix.IPPROTO_IP)
|
||||||
|
h.Type = int32(unix.IP_TOS)
|
||||||
|
setCmsgLen(h, lenVal)
|
||||||
|
return buf
|
||||||
|
}
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
ctrl []byte
|
||||||
|
}{
|
||||||
|
{"len_near_max_int", corrupt(int(^uint(0)>>1) - 8)},
|
||||||
|
{"len_too_small", corrupt(unix.SizeofCmsghdr - 1)},
|
||||||
|
{"len_past_buffer", corrupt(1 << 20)},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
hdr := &msghdr{Control: &c.ctrl[0]}
|
||||||
|
setMsgControllen(hdr, len(c.ctrl))
|
||||||
|
gso := parseRecvCmsg(hdr)
|
||||||
|
// The valid leading UDP_GRO cmsg (payload 0) must still parse;
|
||||||
|
// the corrupt trailer just ends the walk.
|
||||||
|
if gso != 0 {
|
||||||
|
t.Errorf("parseRecvCmsg = %d, want 0", gso)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDeliverSegments pins the GRO RX splitting: a kernel-coalesced buffer
|
||||||
|
// must come back out as the exact pre-coalesce packets -- every boundary
|
||||||
|
// error here shreds encrypted packets and every decrypt downstream fails.
|
||||||
|
func TestDeliverSegments(t *testing.T) {
|
||||||
|
from := netip.MustParseAddrPort("192.0.2.1:4242")
|
||||||
|
// Spare backing capacity mimics the recvmmsg row a real payload sits in;
|
||||||
|
// the cap checks below prove none of it leaks to a delivered segment.
|
||||||
|
pay := func(n int) []byte {
|
||||||
|
b := make([]byte, n, n+512)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = byte(i)
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
payload []byte
|
||||||
|
segSize int
|
||||||
|
wantLens []int
|
||||||
|
}{
|
||||||
|
{"no-gro", pay(1400), 0, []int{1400}},
|
||||||
|
{"negative-segsize", pay(1400), -5, []int{1400}},
|
||||||
|
{"segsize-equals-payload", pay(1400), 1400, []int{1400}},
|
||||||
|
{"segsize-past-payload", pay(1400), 2000, []int{1400}},
|
||||||
|
{"even-split", pay(4200), 1400, []int{1400, 1400, 1400}},
|
||||||
|
{"short-tail", pay(3000), 1400, []int{1400, 1400, 200}},
|
||||||
|
{"single-byte-tail", pay(2801), 1400, []int{1400, 1400, 1}},
|
||||||
|
{"segsize-one", pay(3), 1, []int{1, 1, 1}},
|
||||||
|
{"empty-payload", pay(0), 1400, []int{0}},
|
||||||
|
{"max-coalesce", pay(65500), 1372, nil}, // lens derived below
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
wantLens := c.wantLens
|
||||||
|
if wantLens == nil {
|
||||||
|
for rem := len(c.payload); rem > 0; rem -= c.segSize {
|
||||||
|
wantLens = append(wantLens, min(c.segSize, rem))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var got [][]byte
|
||||||
|
deliverSegments(func(a netip.AddrPort, seg []byte) {
|
||||||
|
if a != from {
|
||||||
|
t.Errorf("from = %v, want %v", a, from)
|
||||||
|
}
|
||||||
|
got = append(got, seg)
|
||||||
|
}, from, c.payload, c.segSize)
|
||||||
|
|
||||||
|
if len(got) != len(wantLens) {
|
||||||
|
t.Fatalf("delivered %d segments, want %d", len(got), len(wantLens))
|
||||||
|
}
|
||||||
|
// Segments must tile the payload in order with no gap, overlap,
|
||||||
|
// or copy: each must alias the payload at the right offset.
|
||||||
|
off := 0
|
||||||
|
for i, seg := range got {
|
||||||
|
if len(seg) != wantLens[i] {
|
||||||
|
t.Fatalf("segment %d len=%d want %d", i, len(seg), wantLens[i])
|
||||||
|
}
|
||||||
|
if cap(seg) != len(seg) {
|
||||||
|
// EncReader contract: an append into spare capacity would
|
||||||
|
// scribble into the next segment of the shared row.
|
||||||
|
t.Errorf("segment %d cap=%d, want %d (capacity must not reach into the row)", i, cap(seg), len(seg))
|
||||||
|
}
|
||||||
|
if len(seg) > 0 && &seg[0] != &c.payload[off] {
|
||||||
|
t.Errorf("segment %d does not alias payload at offset %d", i, off)
|
||||||
|
}
|
||||||
|
off += len(seg)
|
||||||
|
}
|
||||||
|
if off != len(c.payload) {
|
||||||
|
t.Errorf("segments cover %d bytes, payload has %d", off, len(c.payload))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newRewindTestWriter builds a batchWriter with no socket: GSO planning on,
|
||||||
|
// sendFn left for the test to script. fd is invalid on purpose -- any path
|
||||||
|
// that actually hits the kernel fails loudly.
|
||||||
|
func newRewindTestWriter() *batchWriter {
|
||||||
|
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
|
||||||
|
w.prepareWriteMessages(MaxWriteBatch)
|
||||||
|
w.gsoSupported = true
|
||||||
|
w.maxGSOSegments = 63
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// capturePrepared decodes n prepared mmsghdr entries beginning at start
|
||||||
|
// straight from their iovecs -- ground truth, deliberately not the entryEnd
|
||||||
|
// bookkeeping the resume logic itself relies on. Returns one []byte per
|
||||||
|
// packed packet, in entry order.
|
||||||
|
func capturePrepared(w *batchWriter, start, n int) [][]byte {
|
||||||
|
var out [][]byte
|
||||||
|
for e := start; e < start+n; e++ {
|
||||||
|
hdr := &w.msgs[e].Hdr
|
||||||
|
iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen))
|
||||||
|
for _, iov := range iovs {
|
||||||
|
b := make([]byte, int(iov.Len))
|
||||||
|
if iov.Len > 0 {
|
||||||
|
copy(b, unsafe.Slice(iov.Base, int(iov.Len)))
|
||||||
|
}
|
||||||
|
out = append(out, b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchPartialSendRewind drives WriteBatch through scripted
|
||||||
|
// partial sendmmsg results and asserts the rewind resumes exactly where
|
||||||
|
// the kernel stopped: every packet on the wire exactly once, in order,
|
||||||
|
// no duplicate, no loss. This is the hairiest logic in the write path
|
||||||
|
// and a rewind bug means silent packet duplication or loss under EAGAIN-
|
||||||
|
// style backpressure.
|
||||||
|
func TestWriteBatchPartialSendRewind(t *testing.T) {
|
||||||
|
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||||
|
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||||
|
|
||||||
|
mkBuf := func(tag byte, n int) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = tag
|
||||||
|
}
|
||||||
|
b[0] = tag // tag identifies the packet uniquely below
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run
|
||||||
|
// tail), then two to B. The planner packs this as multiple entries with
|
||||||
|
// multi-iovec runs, which is what makes the rewind arithmetic hairy.
|
||||||
|
bufs := [][]byte{
|
||||||
|
mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A
|
||||||
|
mkBuf(4, 600), // short tail to A
|
||||||
|
mkBuf(5, 900), mkBuf(6, 900), // run to B
|
||||||
|
}
|
||||||
|
addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB}
|
||||||
|
|
||||||
|
scripts := [][]int{
|
||||||
|
{99}, // accept everything first call
|
||||||
|
{1, 99}, // one entry per call, then the rest
|
||||||
|
{1, 1, 1, 99}, // strictly one entry per call
|
||||||
|
{2, 99}, // two entries, then the rest
|
||||||
|
}
|
||||||
|
for si, script := range scripts {
|
||||||
|
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
var wire [][]byte
|
||||||
|
call := 0
|
||||||
|
w.sendFn = func(start, n int) (int, error) {
|
||||||
|
accept := n
|
||||||
|
if call < len(script) && script[call] < n {
|
||||||
|
accept = script[call]
|
||||||
|
}
|
||||||
|
call++
|
||||||
|
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||||
|
return accept, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := w.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if written != len(bufs) {
|
||||||
|
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||||
|
}
|
||||||
|
if len(wire) != len(bufs) {
|
||||||
|
t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs))
|
||||||
|
}
|
||||||
|
for i, b := range wire {
|
||||||
|
if len(b) != len(bufs[i]) || b[0] != bufs[i][0] {
|
||||||
|
t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)",
|
||||||
|
i, b[0], len(b), bufs[i][0], len(bufs[i]))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchSkipUnroutableRunAccounting: an unroutable destination mid-
|
||||||
|
// batch is skipped without committing an entry, leaving a hole in the bufs
|
||||||
|
// index space. The written count must tally packets per sent entry -- the
|
||||||
|
// index span would count the hole -- across both full and partial sendmmsg
|
||||||
|
// success.
|
||||||
|
func TestWriteBatchSkipUnroutableRunAccounting(t *testing.T) {
|
||||||
|
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||||
|
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||||
|
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // v6 dest, v4 writer
|
||||||
|
|
||||||
|
mk := func(tag byte, n int) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
b[0] = tag
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 500), mk(4, 900), mk(5, 900)}
|
||||||
|
addrs := []netip.AddrPort{dstA, dstA, bad, dstB, dstB}
|
||||||
|
|
||||||
|
for si, script := range [][]int{{99}, {1, 99}} {
|
||||||
|
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
var wire [][]byte
|
||||||
|
call := 0
|
||||||
|
w.sendFn = func(start, n int) (int, error) {
|
||||||
|
accept := n
|
||||||
|
if call < len(script) && script[call] < n {
|
||||||
|
accept = script[call]
|
||||||
|
}
|
||||||
|
call++
|
||||||
|
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||||
|
return accept, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := w.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if written != 4 {
|
||||||
|
t.Errorf("written = %d, want 4 (the unroutable run is the only casualty)", written)
|
||||||
|
}
|
||||||
|
wantTags := []byte{1, 2, 4, 5}
|
||||||
|
if len(wire) != len(wantTags) {
|
||||||
|
t.Fatalf("wire got %d packets, want %d (dup or loss around the skip)", len(wire), len(wantTags))
|
||||||
|
}
|
||||||
|
for i, b := range wire {
|
||||||
|
if b[0] != wantTags[i] {
|
||||||
|
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchMidChunkRejectResumes: after a partial success, a zero-sent
|
||||||
|
// error on the FIRST REMAINING entry (done > 0) must drop only that entry's
|
||||||
|
// run and resume the rest of the chunk in place -- no repacking, no packets
|
||||||
|
// lost from entries before or after the rejected one.
|
||||||
|
func TestWriteBatchMidChunkRejectResumes(t *testing.T) {
|
||||||
|
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||||
|
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||||
|
dstC := netip.MustParseAddrPort("127.0.0.3:4242")
|
||||||
|
|
||||||
|
mk := func(tag byte, n int) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
b[0] = tag
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
// Three entries: a 2-packet GSO run to A, a 2-packet run to B, one to C.
|
||||||
|
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 900), mk(4, 900), mk(5, 600)}
|
||||||
|
addrs := []netip.AddrPort{dstA, dstA, dstB, dstB, dstC}
|
||||||
|
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
var wire [][]byte
|
||||||
|
var starts []int
|
||||||
|
call := 0
|
||||||
|
w.sendFn = func(start, n int) (int, error) {
|
||||||
|
starts = append(starts, start)
|
||||||
|
call++
|
||||||
|
switch call {
|
||||||
|
case 1: // accept only entry 0 (the run to A)
|
||||||
|
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||||
|
return 1, nil
|
||||||
|
case 2: // reject entry 1 (the run to B) outright
|
||||||
|
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EPERM}
|
||||||
|
default: // accept the rest
|
||||||
|
wire = append(wire, capturePrepared(w, start, n)...)
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := w.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if written != 3 {
|
||||||
|
t.Errorf("written = %d, want 3 (B's rejected run is the only casualty)", written)
|
||||||
|
}
|
||||||
|
wantTags := []byte{1, 2, 5}
|
||||||
|
if len(wire) != len(wantTags) {
|
||||||
|
t.Fatalf("wire got %d packets, want %d (dup or loss around the mid-chunk reject)", len(wire), len(wantTags))
|
||||||
|
}
|
||||||
|
for i, b := range wire {
|
||||||
|
if b[0] != wantTags[i] {
|
||||||
|
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// The resume must reuse the prepared entries: same chunk, advancing
|
||||||
|
// start offsets, no repack (which would restart at 0 with fresh entries).
|
||||||
|
if want := []int{0, 1, 2}; !slices.Equal(starts, want) {
|
||||||
|
t.Errorf("sendFn start offsets = %v, want %v", starts, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchMidChunkEIODisablesGSOWithoutDup: an EIO on a GSO entry
|
||||||
|
// after earlier entries in the chunk already went out must replay ONLY from
|
||||||
|
// the failed run (replanned as single-packet entries) -- the already-sent
|
||||||
|
// entries must not be duplicated.
|
||||||
|
func TestWriteBatchMidChunkEIODisablesGSOWithoutDup(t *testing.T) {
|
||||||
|
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||||
|
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||||
|
|
||||||
|
mk := func(tag byte, n int) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
b[0] = tag
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
// Entry 0: single packet to A. Entry 1: 2-packet GSO run to B.
|
||||||
|
bufs := [][]byte{mk(1, 600), mk(2, 1200), mk(3, 1200)}
|
||||||
|
addrs := []netip.AddrPort{dstA, dstB, dstB}
|
||||||
|
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
var wire [][]byte
|
||||||
|
call := 0
|
||||||
|
w.sendFn = func(start, n int) (int, error) {
|
||||||
|
call++
|
||||||
|
switch call {
|
||||||
|
case 1: // accept entry 0 only
|
||||||
|
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||||
|
return 1, nil
|
||||||
|
case 2: // EIO on the GSO run to B
|
||||||
|
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||||
|
default: // replanned single-packet replay
|
||||||
|
wire = append(wire, capturePrepared(w, start, n)...)
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := w.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if w.gsoSupported {
|
||||||
|
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||||
|
}
|
||||||
|
if written != len(bufs) {
|
||||||
|
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||||
|
}
|
||||||
|
wantTags := []byte{1, 2, 3}
|
||||||
|
if len(wire) != len(wantTags) {
|
||||||
|
t.Fatalf("wire got %d packets, want %d (packet 1 duplicated, or B's run lost)", len(wire), len(wantTags))
|
||||||
|
}
|
||||||
|
for i, b := range wire {
|
||||||
|
if b[0] != wantTags[i] {
|
||||||
|
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an
|
||||||
|
// error rather than spin forever replaying the same chunk.
|
||||||
|
func TestWriteBatchZeroProgress(t *testing.T) {
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
w.sendFn = func(start, n int) (int, error) { return 0, nil }
|
||||||
|
bufs := [][]byte{make([]byte, 100)}
|
||||||
|
addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")}
|
||||||
|
if _, err := w.WriteBatch(bufs, addrs); err == nil {
|
||||||
|
t.Fatal("WriteBatch = nil error on zero progress, want error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a
|
||||||
|
// sendmmsg rejected with EIO on a GSO superpacket entry must clear
|
||||||
|
// gsoSupported and replay the same packets as per-packet entries through
|
||||||
|
// sendmmsg (keeping batching), not fall back to per-packet sendto.
|
||||||
|
func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) {
|
||||||
|
dst := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||||
|
bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)}
|
||||||
|
addrs := []netip.AddrPort{dst, dst, dst}
|
||||||
|
|
||||||
|
w := newRewindTestWriter()
|
||||||
|
var entryCounts []int
|
||||||
|
call := 0
|
||||||
|
w.sendFn = func(start, n int) (int, error) {
|
||||||
|
entryCounts = append(entryCounts, n)
|
||||||
|
call++
|
||||||
|
if call == 1 {
|
||||||
|
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||||
|
}
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := w.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if w.gsoSupported {
|
||||||
|
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||||
|
}
|
||||||
|
if written != len(bufs) {
|
||||||
|
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||||
|
}
|
||||||
|
// First call: one GSO entry carrying the whole run. Replay: one entry
|
||||||
|
// per packet, still via sendmmsg.
|
||||||
|
want := []int{1, 3}
|
||||||
|
if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] {
|
||||||
|
t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGSOEngagesOnLoopback is the offload smoke test: real sockets, real
|
||||||
|
// UDP_SEGMENT cmsg, real kernel segmentation over loopback. It asserts
|
||||||
|
// both that GSO *engaged* (the whole batch left in a single sendmmsg
|
||||||
|
// entry -- a silent fallback to per-packet entries fails the test) and
|
||||||
|
// that the kernel carved the superpacket back into the exact original
|
||||||
|
// datagrams on the receive side. Runs in CI (make test on ubuntu-latest),
|
||||||
|
// which is what guards against the offload path silently degrading.
|
||||||
|
func TestGSOEngagesOnLoopback(t *testing.T) {
|
||||||
|
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen rx: %v", err)
|
||||||
|
}
|
||||||
|
defer rx.Close()
|
||||||
|
dst := rx.LocalAddr().(*net.UDPAddr).AddrPort()
|
||||||
|
|
||||||
|
uc, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 8)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewListener: %v", err)
|
||||||
|
}
|
||||||
|
sc := uc.(*StdConn)
|
||||||
|
defer sc.Close()
|
||||||
|
|
||||||
|
if !sc.bw.gsoSupported {
|
||||||
|
var un unix.Utsname
|
||||||
|
_ = unix.Uname(&un)
|
||||||
|
release := string(un.Release[:])
|
||||||
|
if major, minor := parseRelease(release); major > 4 || (major == 4 && minor >= 18) {
|
||||||
|
t.Fatalf("kernel %q supports UDP_SEGMENT but the GSO probe failed", release)
|
||||||
|
}
|
||||||
|
t.Skipf("kernel %q predates UDP_SEGMENT (4.18)", release)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Spy on the real syscall to count entries per sendmmsg without
|
||||||
|
// changing what hits the kernel.
|
||||||
|
var entryCounts []int
|
||||||
|
real := sc.bw.sendFn
|
||||||
|
sc.bw.sendFn = func(start, n int) (int, error) {
|
||||||
|
entryCounts = append(entryCounts, n)
|
||||||
|
return real(start, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
const numPkts = 8
|
||||||
|
const pktLen = 1200
|
||||||
|
bufs := make([][]byte, numPkts)
|
||||||
|
addrs := make([]netip.AddrPort, numPkts)
|
||||||
|
for i := range bufs {
|
||||||
|
bufs[i] = make([]byte, pktLen)
|
||||||
|
for j := range bufs[i] {
|
||||||
|
bufs[i][j] = byte(i)
|
||||||
|
}
|
||||||
|
addrs[i] = dst
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := sc.WriteBatch(bufs, addrs)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
if written != numPkts {
|
||||||
|
t.Fatalf("written = %d, want %d", written, numPkts)
|
||||||
|
}
|
||||||
|
// GSO engaged means the run went out as ONE sendmmsg entry carrying a
|
||||||
|
// UDP_SEGMENT superpacket. Per-packet entries mean it silently fell
|
||||||
|
// back -- exactly the regression this test exists to catch.
|
||||||
|
if len(entryCounts) != 1 || entryCounts[0] != 1 {
|
||||||
|
t.Fatalf("sendmmsg entry counts = %v, want [1]: GSO did not engage", entryCounts)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The kernel must deliver the original datagram boundaries and bytes.
|
||||||
|
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||||
|
got := make([]byte, pktLen+1)
|
||||||
|
for i := 0; i < numPkts; i++ {
|
||||||
|
n, _, err := rx.ReadFromUDP(got)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("rx read %d: %v", i, err)
|
||||||
|
}
|
||||||
|
if n != pktLen {
|
||||||
|
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
|
||||||
|
}
|
||||||
|
for j := 0; j < n; j++ {
|
||||||
|
if got[j] != byte(i) {
|
||||||
|
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !found {
|
|
||||||
t.Fatalf("no IP_TOS cmsg delivered to v4 receiver — outer ECN did not land (%d cmsgs)", len(cmsgs))
|
|
||||||
}
|
|
||||||
if gotTOS&0x03 != wantECN {
|
|
||||||
t.Errorf("received outer TOS = 0x%02x, want low-2-bits = 0x%02x", gotTOS, wantECN)
|
|
||||||
} else {
|
|
||||||
t.Logf("verified: v4 receiver saw outer TOS 0x%02x (ECN=0x%02x) from dual-stack sender", gotTOS, gotTOS&0x03)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user